| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041 |
- package analyzer
- import (
- "fmt"
- "strings"
- "github.com/danfragoso/pizzasql-next/pkg/lexer"
- "github.com/danfragoso/pizzasql-next/pkg/parser"
- )
- // ErrorType categorizes analysis errors.
- type ErrorType int
- const (
- ErrUnknown ErrorType = iota
- ErrTableNotFound
- ErrTableExists
- ErrColumnNotFound
- ErrColumnAmbiguous
- ErrTypeMismatch
- ErrInvalidFunction
- ErrInvalidArgCount
- ErrAggregateInWhere
- ErrNonAggregateInSelect
- ErrInvalidGroupBy
- )
- // AnalysisError represents a semantic analysis error.
- type AnalysisError struct {
- Type ErrorType
- Message string
- Line int
- Column int
- Context string
- }
- func (e *AnalysisError) Error() string {
- if e.Line > 0 {
- return fmt.Sprintf("analysis error at line %d, column %d: %s", e.Line, e.Column, e.Message)
- }
- return fmt.Sprintf("analysis error: %s", e.Message)
- }
- // Analyzer performs semantic analysis on parsed SQL statements.
- type Analyzer struct {
- catalog *Catalog
- scope *Scope
- errors []*AnalysisError
- }
- // New creates a new Analyzer with the given catalog.
- func New(catalog *Catalog) *Analyzer {
- if catalog == nil {
- catalog = NewCatalog()
- }
- return &Analyzer{
- catalog: catalog,
- }
- }
- // Analyze performs semantic analysis on a statement.
- func (a *Analyzer) Analyze(stmt parser.Statement) error {
- a.errors = nil
- a.scope = NewScope(nil)
- switch s := stmt.(type) {
- case *parser.SelectStmt:
- return a.analyzeSelect(s)
- case *parser.InsertStmt:
- return a.analyzeInsert(s)
- case *parser.UpdateStmt:
- return a.analyzeUpdate(s)
- case *parser.DeleteStmt:
- return a.analyzeDelete(s)
- case *parser.CreateTableStmt:
- return a.analyzeCreateTable(s)
- case *parser.DropTableStmt:
- return a.analyzeDropTable(s)
- case *parser.AlterTableStmt:
- // ALTER TABLE is handled directly by executor, no semantic analysis needed
- return nil
- case *parser.AttachStmt:
- // ATTACH DATABASE is handled directly by executor
- return nil
- case *parser.DetachStmt:
- // DETACH DATABASE is handled directly by executor
- return nil
- case *parser.BeginStmt, *parser.CommitStmt, *parser.RollbackStmt,
- *parser.SavepointStmt, *parser.ReleaseStmt:
- // Transaction statements don't need semantic analysis
- return nil
- case *parser.CreateIndexStmt, *parser.DropIndexStmt:
- // Index statements don't need semantic analysis
- return nil
- default:
- return &AnalysisError{
- Type: ErrUnknown,
- Message: fmt.Sprintf("unknown statement type: %T", stmt),
- }
- }
- }
- // GetCatalog returns the analyzer's catalog.
- func (a *Analyzer) GetCatalog() *Catalog {
- return a.catalog
- }
- // analyzeSelect analyzes a SELECT statement.
- func (a *Analyzer) analyzeSelect(stmt *parser.SelectStmt) error {
- // First, resolve tables in FROM clause
- if err := a.resolveFromClause(stmt.From); err != nil {
- return err
- }
- // Analyze WHERE clause
- if stmt.Where != nil {
- info, err := a.analyzeExpr(stmt.Where)
- if err != nil {
- return err
- }
- // WHERE clause cannot contain aggregates
- if info.IsAggregate {
- return &AnalysisError{
- Type: ErrAggregateInWhere,
- Message: "aggregate functions not allowed in WHERE clause",
- }
- }
- }
- // Determine if this is an aggregate query
- hasAggregate := false
- hasGroupBy := len(stmt.GroupBy) > 0
- // Analyze GROUP BY expressions first
- for _, expr := range stmt.GroupBy {
- if _, err := a.analyzeExpr(expr); err != nil {
- return err
- }
- }
- // Analyze SELECT columns and collect aliases for ORDER BY/HAVING reference
- selectAliases := make(map[string]*ExprInfo)
- for _, col := range stmt.Columns {
- if col.Star {
- // SELECT * - all columns from all tables
- continue
- }
- info, err := a.analyzeExpr(col.Expr)
- if err != nil {
- return err
- }
- if info.IsAggregate {
- hasAggregate = true
- }
- // Track column aliases so ORDER BY and HAVING can reference them
- if col.Alias != "" {
- selectAliases[strings.ToUpper(col.Alias)] = info
- }
- }
- // Register SELECT aliases as virtual columns for ORDER BY/HAVING reference
- for alias, info := range selectAliases {
- a.scope.DefineSelectAlias(alias, info.Type)
- }
- // Validate GROUP BY semantics
- if hasAggregate && !hasGroupBy {
- // Aggregate query without GROUP BY - all non-aggregate columns must be constants
- for _, col := range stmt.Columns {
- if col.Star {
- return &AnalysisError{
- Type: ErrNonAggregateInSelect,
- Message: "SELECT * not allowed with aggregate functions without GROUP BY",
- }
- }
- info, exprErr := a.analyzeExpr(col.Expr)
- if exprErr != nil || info == nil {
- continue
- }
- if !info.IsAggregate && !info.IsConstant {
- // Check if it's a simple column reference
- if ref, ok := col.Expr.(*parser.ColumnRef); ok {
- return &AnalysisError{
- Type: ErrNonAggregateInSelect,
- Message: fmt.Sprintf("column %q must appear in GROUP BY clause or be in an aggregate function", ref.Column),
- }
- }
- }
- }
- }
- // Analyze HAVING clause
- if stmt.Having != nil {
- info, err := a.analyzeExpr(stmt.Having)
- if err != nil {
- return err
- }
- // HAVING without GROUP BY requires aggregates
- if !hasGroupBy && !info.IsAggregate {
- return &AnalysisError{
- Type: ErrInvalidGroupBy,
- Message: "HAVING clause requires GROUP BY or aggregate function",
- }
- }
- }
- // Analyze ORDER BY
- for _, item := range stmt.OrderBy {
- if _, err := a.analyzeExpr(item.Expr); err != nil {
- return err
- }
- }
- // Analyze LIMIT/OFFSET
- if stmt.Limit != nil {
- info, err := a.analyzeExpr(stmt.Limit)
- if err != nil {
- return err
- }
- if !info.Type.IsNumeric() && info.Type != TypeNull {
- return &AnalysisError{
- Type: ErrTypeMismatch,
- Message: "LIMIT must be numeric",
- }
- }
- }
- if stmt.Offset != nil {
- info, err := a.analyzeExpr(stmt.Offset)
- if err != nil {
- return err
- }
- if !info.Type.IsNumeric() && info.Type != TypeNull {
- return &AnalysisError{
- Type: ErrTypeMismatch,
- Message: "OFFSET must be numeric",
- }
- }
- }
- return nil
- }
- // resolveFromClause adds tables from FROM clause to scope.
- func (a *Analyzer) resolveFromClause(tables []parser.TableRef) error {
- for _, ref := range tables {
- // Handle subquery (derived table)
- if ref.Subquery != nil {
- // Analyze the subquery
- if err := a.analyzeSelect(ref.Subquery); err != nil {
- return err
- }
- // Create a table info from subquery columns
- // For now, we'll use a simplified approach - just mark it as a derived table
- tableInfo := &TableInfo{
- Name: ref.Alias, // Derived tables MUST have an alias
- Columns: []ColumnInfo{},
- Alias: ref.Alias,
- }
- // Add columns from SELECT list
- for _, col := range ref.Subquery.Columns {
- colName := ""
- if col.Alias != "" {
- colName = col.Alias
- } else if colRef, ok := col.Expr.(*parser.ColumnRef); ok {
- colName = colRef.Column
- } else {
- // For expressions without alias, use a generated name
- colName = fmt.Sprintf("col_%d", len(tableInfo.Columns))
- }
- tableInfo.Columns = append(tableInfo.Columns, ColumnInfo{
- Name: colName,
- TableName: ref.Alias,
- Type: TypeAny, // We'd need type inference for proper typing
- })
- }
- a.scope.DefineTable(tableInfo)
- } else {
- // Regular table reference
- table, ok := a.catalog.GetTable(ref.Name)
- if !ok {
- return &AnalysisError{
- Type: ErrTableNotFound,
- Message: fmt.Sprintf("table not found: %s", ref.Name),
- }
- }
- // Create a copy with alias if specified
- tableInfo := &TableInfo{
- Name: table.Name,
- Columns: table.Columns,
- Alias: ref.Alias,
- IsView: table.IsView,
- }
- a.scope.DefineTable(tableInfo)
- }
- // Handle JOINs
- if ref.Join != nil {
- if err := a.resolveJoin(ref.Join); err != nil {
- return err
- }
- }
- }
- return nil
- }
- // resolveJoin resolves a JOIN clause.
- func (a *Analyzer) resolveJoin(join *parser.JoinClause) error {
- if join.Table == nil {
- return nil
- }
- table, ok := a.catalog.GetTable(join.Table.Name)
- if !ok {
- return &AnalysisError{
- Type: ErrTableNotFound,
- Message: fmt.Sprintf("table not found: %s", join.Table.Name),
- }
- }
- tableInfo := &TableInfo{
- Name: table.Name,
- Columns: table.Columns,
- Alias: join.Table.Alias,
- IsView: table.IsView,
- }
- a.scope.DefineTable(tableInfo)
- // Analyze ON condition
- if join.Condition != nil {
- if _, err := a.analyzeExpr(join.Condition); err != nil {
- return err
- }
- }
- // Handle USING clause
- for _, colName := range join.Using {
- _, _, ok := a.scope.LookupColumn("", colName)
- if !ok {
- return &AnalysisError{
- Type: ErrColumnNotFound,
- Message: fmt.Sprintf("column not found in USING clause: %s", colName),
- }
- }
- }
- // Recursively handle chained JOINs
- if join.Table.Join != nil {
- if err := a.resolveJoin(join.Table.Join); err != nil {
- return err
- }
- }
- return nil
- }
- // analyzeInsert analyzes an INSERT statement.
- func (a *Analyzer) analyzeInsert(stmt *parser.InsertStmt) error {
- table, ok := a.catalog.GetTable(stmt.Table.Name)
- if !ok {
- return &AnalysisError{
- Type: ErrTableNotFound,
- Message: fmt.Sprintf("table not found: %s", stmt.Table.Name),
- }
- }
- // Validate column list if specified
- var targetCols []ColumnInfo
- if len(stmt.Columns) > 0 {
- for _, colName := range stmt.Columns {
- col, ok := table.GetColumn(colName)
- if !ok {
- return &AnalysisError{
- Type: ErrColumnNotFound,
- Message: fmt.Sprintf("column not found: %s", colName),
- }
- }
- targetCols = append(targetCols, *col)
- }
- } else {
- targetCols = table.Columns
- }
- // Validate VALUES
- for _, row := range stmt.Values {
- if len(row) != len(targetCols) {
- return &AnalysisError{
- Type: ErrTypeMismatch,
- Message: fmt.Sprintf("INSERT has %d columns but %d values", len(targetCols), len(row)),
- }
- }
- for i, expr := range row {
- info, err := a.analyzeExpr(expr)
- if err != nil {
- return err
- }
- // Check type compatibility
- if !info.Type.IsComparable(targetCols[i].Type) && info.Type != TypeNull {
- return &AnalysisError{
- Type: ErrTypeMismatch,
- Message: fmt.Sprintf("type mismatch for column %s: expected %s, got %s",
- targetCols[i].Name, targetCols[i].Type, info.Type),
- }
- }
- }
- }
- // Analyze INSERT ... SELECT
- if stmt.Select != nil {
- a.scope.DefineTable(&TableInfo{Name: table.Name, Columns: table.Columns, IsView: table.IsView})
- if err := a.analyzeSelect(stmt.Select); err != nil {
- return err
- }
- }
- return nil
- }
- // analyzeUpdate analyzes an UPDATE statement.
- func (a *Analyzer) analyzeUpdate(stmt *parser.UpdateStmt) error {
- table, ok := a.catalog.GetTable(stmt.Table.Name)
- if !ok {
- return &AnalysisError{
- Type: ErrTableNotFound,
- Message: fmt.Sprintf("table not found: %s", stmt.Table.Name),
- }
- }
- a.scope.DefineTable(&TableInfo{Name: table.Name, Columns: table.Columns, Alias: stmt.Table.Alias, IsView: table.IsView})
- // Validate SET assignments
- for _, assign := range stmt.Set {
- col, ok := table.GetColumn(assign.Column)
- if !ok {
- return &AnalysisError{
- Type: ErrColumnNotFound,
- Message: fmt.Sprintf("column not found: %s", assign.Column),
- }
- }
- info, err := a.analyzeExpr(assign.Value)
- if err != nil {
- return err
- }
- if !info.Type.IsComparable(col.Type) && info.Type != TypeNull {
- return &AnalysisError{
- Type: ErrTypeMismatch,
- Message: fmt.Sprintf("type mismatch for column %s: expected %s, got %s",
- col.Name, col.Type, info.Type),
- }
- }
- }
- // Analyze WHERE clause
- if stmt.Where != nil {
- info, err := a.analyzeExpr(stmt.Where)
- if err != nil {
- return err
- }
- if info.IsAggregate {
- return &AnalysisError{
- Type: ErrAggregateInWhere,
- Message: "aggregate functions not allowed in WHERE clause",
- }
- }
- }
- return nil
- }
- // analyzeDelete analyzes a DELETE statement.
- func (a *Analyzer) analyzeDelete(stmt *parser.DeleteStmt) error {
- table, ok := a.catalog.GetTable(stmt.Table.Name)
- if !ok {
- return &AnalysisError{
- Type: ErrTableNotFound,
- Message: fmt.Sprintf("table not found: %s", stmt.Table.Name),
- }
- }
- a.scope.DefineTable(&TableInfo{Name: table.Name, Columns: table.Columns, IsView: table.IsView})
- // Analyze WHERE clause
- if stmt.Where != nil {
- info, err := a.analyzeExpr(stmt.Where)
- if err != nil {
- return err
- }
- if info.IsAggregate {
- return &AnalysisError{
- Type: ErrAggregateInWhere,
- Message: "aggregate functions not allowed in WHERE clause",
- }
- }
- }
- return nil
- }
- // analyzeCreateTable analyzes a CREATE TABLE statement.
- func (a *Analyzer) analyzeCreateTable(stmt *parser.CreateTableStmt) error {
- // Check if table already exists
- if a.catalog.TableExists(stmt.Table.Name) {
- if stmt.IfNotExists {
- return nil // Silently succeed
- }
- return &AnalysisError{
- Type: ErrTableExists,
- Message: fmt.Sprintf("table already exists: %s", stmt.Table.Name),
- }
- }
- // Build table info
- tableInfo := &TableInfo{
- Name: stmt.Table.Name,
- }
- columnNames := make(map[string]bool)
- for _, colDef := range stmt.Columns {
- upperName := strings.ToUpper(colDef.Name)
- if columnNames[upperName] {
- return &AnalysisError{
- Type: ErrColumnAmbiguous,
- Message: fmt.Sprintf("duplicate column name: %s", colDef.Name),
- }
- }
- columnNames[upperName] = true
- colInfo := ColumnInfo{
- Name: colDef.Name,
- Type: TypeFromName(colDef.Type.Name),
- Nullable: true,
- TableName: stmt.Table.Name,
- }
- // Process constraints
- for _, constraint := range colDef.Constraints {
- switch constraint.Type {
- case parser.ConstraintPrimaryKey:
- colInfo.PrimaryKey = true
- colInfo.Nullable = false
- case parser.ConstraintNotNull:
- colInfo.Nullable = false
- case parser.ConstraintDefault:
- // Store default value (not evaluated here)
- colInfo.Default = constraint.Default
- }
- }
- tableInfo.Columns = append(tableInfo.Columns, colInfo)
- }
- // Process table-level constraints
- for _, constraint := range stmt.Constraints {
- switch constraint.Type {
- case parser.ConstraintPrimaryKey:
- for _, colName := range constraint.Columns {
- for i := range tableInfo.Columns {
- if strings.EqualFold(tableInfo.Columns[i].Name, colName) {
- tableInfo.Columns[i].PrimaryKey = true
- tableInfo.Columns[i].Nullable = false
- }
- }
- }
- }
- }
- // Add to catalog
- return a.catalog.CreateTable(tableInfo)
- }
- // analyzeDropTable analyzes a DROP TABLE statement.
- func (a *Analyzer) analyzeDropTable(stmt *parser.DropTableStmt) error {
- for _, tableRef := range stmt.Tables {
- if !a.catalog.TableExists(tableRef.Name) {
- if stmt.IfExists {
- continue // Silently succeed
- }
- return &AnalysisError{
- Type: ErrTableNotFound,
- Message: fmt.Sprintf("table not found: %s", tableRef.Name),
- }
- }
- if err := a.catalog.DropTable(tableRef.Name); err != nil {
- return err
- }
- }
- return nil
- }
- // analyzeExpr analyzes an expression and returns type information.
- func (a *Analyzer) analyzeExpr(expr parser.Expr) (*ExprInfo, error) {
- switch e := expr.(type) {
- case *parser.LiteralExpr:
- return a.analyzeLiteral(e)
- case *parser.ColumnRef:
- return a.analyzeColumnRef(e)
- case *parser.BinaryExpr:
- return a.analyzeBinaryExpr(e)
- case *parser.UnaryExpr:
- return a.analyzeUnaryExpr(e)
- case *parser.FunctionCall:
- return a.analyzeFunctionCall(e)
- case *parser.ParenExpr:
- return a.analyzeExpr(e.Expr)
- case *parser.CaseExpr:
- return a.analyzeCaseExpr(e)
- case *parser.CastExpr:
- return a.analyzeCastExpr(e)
- case *parser.InExpr:
- return a.analyzeInExpr(e)
- case *parser.BetweenExpr:
- return a.analyzeBetweenExpr(e)
- case *parser.LikeExpr:
- return a.analyzeLikeExpr(e)
- case *parser.IsNullExpr:
- return a.analyzeIsNullExpr(e)
- case *parser.ExistsExpr:
- return a.analyzeExistsExpr(e)
- case *parser.SubqueryExpr:
- return a.analyzeSubqueryExpr(e)
- default:
- return &ExprInfo{Type: TypeUnknown}, nil
- }
- }
- func (a *Analyzer) analyzeLiteral(e *parser.LiteralExpr) (*ExprInfo, error) {
- info := &ExprInfo{IsConstant: true}
- switch e.Type {
- case lexer.TokenNumber:
- if strings.Contains(e.Value, ".") || strings.Contains(strings.ToLower(e.Value), "e") {
- info.Type = TypeReal
- } else {
- info.Type = TypeInteger
- }
- case lexer.TokenString:
- info.Type = TypeText
- case lexer.TokenNULL:
- info.Type = TypeNull
- info.Nullable = true
- case lexer.TokenTRUE, lexer.TokenFALSE:
- info.Type = TypeBoolean
- case lexer.TokenStar:
- info.Type = TypeAny
- default:
- info.Type = TypeUnknown
- }
- return info, nil
- }
- func (a *Analyzer) analyzeColumnRef(e *parser.ColumnRef) (*ExprInfo, error) {
- col, _, ok := a.scope.LookupColumn(e.Table, e.Column)
- if !ok {
- // If no tables are in scope, treat as unknown (for standalone expressions)
- if len(a.scope.GetTables()) == 0 {
- return &ExprInfo{Type: TypeUnknown}, nil
- }
- return nil, &AnalysisError{
- Type: ErrColumnNotFound,
- Message: fmt.Sprintf("column not found: %s", formatColumnRef(e)),
- }
- }
- return &ExprInfo{
- Type: col.Type,
- Nullable: col.Nullable,
- }, nil
- }
- func formatColumnRef(e *parser.ColumnRef) string {
- if e.Table != "" {
- return e.Table + "." + e.Column
- }
- return e.Column
- }
- func (a *Analyzer) analyzeBinaryExpr(e *parser.BinaryExpr) (*ExprInfo, error) {
- left, err := a.analyzeExpr(e.Left)
- if err != nil {
- return nil, err
- }
- right, err := a.analyzeExpr(e.Right)
- if err != nil {
- return nil, err
- }
- info := &ExprInfo{
- IsAggregate: left.IsAggregate || right.IsAggregate,
- IsConstant: left.IsConstant && right.IsConstant,
- Nullable: left.Nullable || right.Nullable,
- }
- switch e.Op {
- case lexer.TokenPlus, lexer.TokenMinus, lexer.TokenStar, lexer.TokenSlash, lexer.TokenPercent:
- // Arithmetic operators
- info.Type = CommonType(left.Type, right.Type)
- if !left.Type.IsNumeric() && left.Type != TypeNull && left.Type != TypeUnknown {
- return nil, &AnalysisError{
- Type: ErrTypeMismatch,
- Message: fmt.Sprintf("arithmetic operator requires numeric type, got %s", left.Type),
- }
- }
- case lexer.TokenEq, lexer.TokenNeq, lexer.TokenLt, lexer.TokenLte, lexer.TokenGt, lexer.TokenGte:
- // Comparison operators
- info.Type = TypeBoolean
- if !left.Type.IsComparable(right.Type) {
- return nil, &AnalysisError{
- Type: ErrTypeMismatch,
- Message: fmt.Sprintf("cannot compare %s with %s", left.Type, right.Type),
- }
- }
- case lexer.TokenAND, lexer.TokenOR:
- // Logical operators
- info.Type = TypeBoolean
- case lexer.TokenConcat:
- // String concatenation
- info.Type = TypeText
- default:
- info.Type = TypeUnknown
- }
- return info, nil
- }
- func (a *Analyzer) analyzeUnaryExpr(e *parser.UnaryExpr) (*ExprInfo, error) {
- operand, err := a.analyzeExpr(e.Operand)
- if err != nil {
- return nil, err
- }
- info := &ExprInfo{
- IsAggregate: operand.IsAggregate,
- IsConstant: operand.IsConstant,
- Nullable: operand.Nullable,
- }
- switch e.Op {
- case lexer.TokenPlus:
- // Unary + is a no-op in SQLite — passes any type through unchanged.
- info.Type = operand.Type
- case lexer.TokenMinus:
- info.Type = operand.Type
- if !operand.Type.IsNumeric() && operand.Type != TypeNull && operand.Type != TypeUnknown {
- return nil, &AnalysisError{
- Type: ErrTypeMismatch,
- Message: fmt.Sprintf("unary - requires numeric type, got %s", operand.Type),
- }
- }
- case lexer.TokenNOT:
- info.Type = TypeBoolean
- default:
- info.Type = operand.Type
- }
- return info, nil
- }
- func (a *Analyzer) analyzeFunctionCall(e *parser.FunctionCall) (*ExprInfo, error) {
- sig, ok := LookupFunction(e.Name)
- if !ok {
- return nil, &AnalysisError{
- Type: ErrInvalidFunction,
- Message: fmt.Sprintf("unknown function: %s", e.Name),
- }
- }
- // Handle COUNT(*)
- argCount := len(e.Args)
- if e.Star {
- argCount = 0 // COUNT(*) has 0 real args
- }
- // Check argument count
- if argCount < sig.MinArgs {
- return nil, &AnalysisError{
- Type: ErrInvalidArgCount,
- Message: fmt.Sprintf("function %s requires at least %d arguments, got %d", e.Name, sig.MinArgs, argCount),
- }
- }
- if sig.MaxArgs >= 0 && argCount > sig.MaxArgs {
- return nil, &AnalysisError{
- Type: ErrInvalidArgCount,
- Message: fmt.Sprintf("function %s accepts at most %d arguments, got %d", e.Name, sig.MaxArgs, argCount),
- }
- }
- // Analyze arguments
- info := &ExprInfo{
- Type: sig.ReturnType,
- IsAggregate: sig.IsAggregate,
- }
- for _, arg := range e.Args {
- argInfo, err := a.analyzeExpr(arg)
- if err != nil {
- return nil, err
- }
- if argInfo.Nullable {
- info.Nullable = true
- }
- // Propagate aggregate status from arguments
- if argInfo.IsAggregate && !sig.IsAggregate {
- info.IsAggregate = true
- }
- }
- // Special case: MIN/MAX/COALESCE return type depends on argument
- if sig.ReturnType == TypeAny && len(e.Args) > 0 {
- argInfo, _ := a.analyzeExpr(e.Args[0])
- if argInfo != nil {
- info.Type = argInfo.Type
- }
- }
- return info, nil
- }
- func (a *Analyzer) analyzeCaseExpr(e *parser.CaseExpr) (*ExprInfo, error) {
- info := &ExprInfo{
- Nullable: true, // CASE can return NULL
- }
- // Analyze operand if present (simple CASE)
- if e.Operand != nil {
- opInfo, err := a.analyzeExpr(e.Operand)
- if err != nil {
- return nil, err
- }
- if opInfo.IsAggregate {
- info.IsAggregate = true
- }
- }
- // Analyze WHEN clauses
- var resultType Type
- for _, when := range e.Whens {
- condInfo, err := a.analyzeExpr(when.Condition)
- if err != nil {
- return nil, err
- }
- if condInfo.IsAggregate {
- info.IsAggregate = true
- }
- resInfo, err := a.analyzeExpr(when.Result)
- if err != nil {
- return nil, err
- }
- if resInfo.IsAggregate {
- info.IsAggregate = true
- }
- if resultType == TypeUnknown {
- resultType = resInfo.Type
- } else {
- resultType = CommonType(resultType, resInfo.Type)
- }
- }
- // Analyze ELSE clause
- if e.Else != nil {
- elseInfo, err := a.analyzeExpr(e.Else)
- if err != nil {
- return nil, err
- }
- if elseInfo.IsAggregate {
- info.IsAggregate = true
- }
- resultType = CommonType(resultType, elseInfo.Type)
- }
- info.Type = resultType
- return info, nil
- }
- func (a *Analyzer) analyzeCastExpr(e *parser.CastExpr) (*ExprInfo, error) {
- exprInfo, err := a.analyzeExpr(e.Expr)
- if err != nil {
- return nil, err
- }
- return &ExprInfo{
- Type: TypeFromName(e.Type.Name),
- IsAggregate: exprInfo.IsAggregate,
- IsConstant: exprInfo.IsConstant,
- Nullable: exprInfo.Nullable,
- }, nil
- }
- func (a *Analyzer) analyzeInExpr(e *parser.InExpr) (*ExprInfo, error) {
- leftInfo, err := a.analyzeExpr(e.Left)
- if err != nil {
- return nil, err
- }
- info := &ExprInfo{
- Type: TypeBoolean,
- IsAggregate: leftInfo.IsAggregate,
- }
- // Analyze value list
- for _, val := range e.Values {
- valInfo, err := a.analyzeExpr(val)
- if err != nil {
- return nil, err
- }
- if valInfo.IsAggregate {
- info.IsAggregate = true
- }
- }
- // Analyze subquery
- if e.Subquery != nil {
- subScope := NewScope(a.scope)
- oldScope := a.scope
- a.scope = subScope
- err := a.analyzeSelect(e.Subquery)
- a.scope = oldScope
- if err != nil {
- return nil, err
- }
- }
- return info, nil
- }
- func (a *Analyzer) analyzeBetweenExpr(e *parser.BetweenExpr) (*ExprInfo, error) {
- leftInfo, err := a.analyzeExpr(e.Left)
- if err != nil {
- return nil, err
- }
- lowInfo, err := a.analyzeExpr(e.Low)
- if err != nil {
- return nil, err
- }
- highInfo, err := a.analyzeExpr(e.High)
- if err != nil {
- return nil, err
- }
- return &ExprInfo{
- Type: TypeBoolean,
- IsAggregate: leftInfo.IsAggregate || lowInfo.IsAggregate || highInfo.IsAggregate,
- Nullable: leftInfo.Nullable || lowInfo.Nullable || highInfo.Nullable,
- }, nil
- }
- func (a *Analyzer) analyzeLikeExpr(e *parser.LikeExpr) (*ExprInfo, error) {
- leftInfo, err := a.analyzeExpr(e.Left)
- if err != nil {
- return nil, err
- }
- patternInfo, err := a.analyzeExpr(e.Pattern)
- if err != nil {
- return nil, err
- }
- info := &ExprInfo{
- Type: TypeBoolean,
- IsAggregate: leftInfo.IsAggregate || patternInfo.IsAggregate,
- Nullable: leftInfo.Nullable || patternInfo.Nullable,
- }
- if e.Escape != nil {
- escInfo, err := a.analyzeExpr(e.Escape)
- if err != nil {
- return nil, err
- }
- if escInfo.IsAggregate {
- info.IsAggregate = true
- }
- }
- return info, nil
- }
- func (a *Analyzer) analyzeIsNullExpr(e *parser.IsNullExpr) (*ExprInfo, error) {
- leftInfo, err := a.analyzeExpr(e.Left)
- if err != nil {
- return nil, err
- }
- return &ExprInfo{
- Type: TypeBoolean,
- IsAggregate: leftInfo.IsAggregate,
- IsConstant: leftInfo.IsConstant,
- }, nil
- }
- func (a *Analyzer) analyzeExistsExpr(e *parser.ExistsExpr) (*ExprInfo, error) {
- // Analyze subquery in its own scope
- subScope := NewScope(a.scope)
- oldScope := a.scope
- a.scope = subScope
- err := a.analyzeSelect(e.Subquery)
- a.scope = oldScope
- if err != nil {
- return nil, err
- }
- return &ExprInfo{
- Type: TypeBoolean,
- }, nil
- }
- func (a *Analyzer) analyzeSubqueryExpr(e *parser.SubqueryExpr) (*ExprInfo, error) {
- // Analyze subquery in its own scope
- subScope := NewScope(a.scope)
- oldScope := a.scope
- a.scope = subScope
- err := a.analyzeSelect(e.Query)
- a.scope = oldScope
- if err != nil {
- return nil, err
- }
- // Scalar subquery - return type of first column
- // For simplicity, return TypeAny
- return &ExprInfo{
- Type: TypeAny,
- }, nil
- }
|