2
0

analyzer.go 24 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041
  1. package analyzer
  2. import (
  3. "fmt"
  4. "strings"
  5. "github.com/danfragoso/pizzasql-next/pkg/lexer"
  6. "github.com/danfragoso/pizzasql-next/pkg/parser"
  7. )
  8. // ErrorType categorizes analysis errors.
  9. type ErrorType int
  10. const (
  11. ErrUnknown ErrorType = iota
  12. ErrTableNotFound
  13. ErrTableExists
  14. ErrColumnNotFound
  15. ErrColumnAmbiguous
  16. ErrTypeMismatch
  17. ErrInvalidFunction
  18. ErrInvalidArgCount
  19. ErrAggregateInWhere
  20. ErrNonAggregateInSelect
  21. ErrInvalidGroupBy
  22. )
  23. // AnalysisError represents a semantic analysis error.
  24. type AnalysisError struct {
  25. Type ErrorType
  26. Message string
  27. Line int
  28. Column int
  29. Context string
  30. }
  31. func (e *AnalysisError) Error() string {
  32. if e.Line > 0 {
  33. return fmt.Sprintf("analysis error at line %d, column %d: %s", e.Line, e.Column, e.Message)
  34. }
  35. return fmt.Sprintf("analysis error: %s", e.Message)
  36. }
  37. // Analyzer performs semantic analysis on parsed SQL statements.
  38. type Analyzer struct {
  39. catalog *Catalog
  40. scope *Scope
  41. errors []*AnalysisError
  42. }
  43. // New creates a new Analyzer with the given catalog.
  44. func New(catalog *Catalog) *Analyzer {
  45. if catalog == nil {
  46. catalog = NewCatalog()
  47. }
  48. return &Analyzer{
  49. catalog: catalog,
  50. }
  51. }
  52. // Analyze performs semantic analysis on a statement.
  53. func (a *Analyzer) Analyze(stmt parser.Statement) error {
  54. a.errors = nil
  55. a.scope = NewScope(nil)
  56. switch s := stmt.(type) {
  57. case *parser.SelectStmt:
  58. return a.analyzeSelect(s)
  59. case *parser.InsertStmt:
  60. return a.analyzeInsert(s)
  61. case *parser.UpdateStmt:
  62. return a.analyzeUpdate(s)
  63. case *parser.DeleteStmt:
  64. return a.analyzeDelete(s)
  65. case *parser.CreateTableStmt:
  66. return a.analyzeCreateTable(s)
  67. case *parser.DropTableStmt:
  68. return a.analyzeDropTable(s)
  69. case *parser.AlterTableStmt:
  70. // ALTER TABLE is handled directly by executor, no semantic analysis needed
  71. return nil
  72. case *parser.AttachStmt:
  73. // ATTACH DATABASE is handled directly by executor
  74. return nil
  75. case *parser.DetachStmt:
  76. // DETACH DATABASE is handled directly by executor
  77. return nil
  78. case *parser.BeginStmt, *parser.CommitStmt, *parser.RollbackStmt,
  79. *parser.SavepointStmt, *parser.ReleaseStmt:
  80. // Transaction statements don't need semantic analysis
  81. return nil
  82. case *parser.CreateIndexStmt, *parser.DropIndexStmt:
  83. // Index statements don't need semantic analysis
  84. return nil
  85. default:
  86. return &AnalysisError{
  87. Type: ErrUnknown,
  88. Message: fmt.Sprintf("unknown statement type: %T", stmt),
  89. }
  90. }
  91. }
  92. // GetCatalog returns the analyzer's catalog.
  93. func (a *Analyzer) GetCatalog() *Catalog {
  94. return a.catalog
  95. }
  96. // analyzeSelect analyzes a SELECT statement.
  97. func (a *Analyzer) analyzeSelect(stmt *parser.SelectStmt) error {
  98. // First, resolve tables in FROM clause
  99. if err := a.resolveFromClause(stmt.From); err != nil {
  100. return err
  101. }
  102. // Analyze WHERE clause
  103. if stmt.Where != nil {
  104. info, err := a.analyzeExpr(stmt.Where)
  105. if err != nil {
  106. return err
  107. }
  108. // WHERE clause cannot contain aggregates
  109. if info.IsAggregate {
  110. return &AnalysisError{
  111. Type: ErrAggregateInWhere,
  112. Message: "aggregate functions not allowed in WHERE clause",
  113. }
  114. }
  115. }
  116. // Determine if this is an aggregate query
  117. hasAggregate := false
  118. hasGroupBy := len(stmt.GroupBy) > 0
  119. // Analyze GROUP BY expressions first
  120. for _, expr := range stmt.GroupBy {
  121. if _, err := a.analyzeExpr(expr); err != nil {
  122. return err
  123. }
  124. }
  125. // Analyze SELECT columns and collect aliases for ORDER BY/HAVING reference
  126. selectAliases := make(map[string]*ExprInfo)
  127. for _, col := range stmt.Columns {
  128. if col.Star {
  129. // SELECT * - all columns from all tables
  130. continue
  131. }
  132. info, err := a.analyzeExpr(col.Expr)
  133. if err != nil {
  134. return err
  135. }
  136. if info.IsAggregate {
  137. hasAggregate = true
  138. }
  139. // Track column aliases so ORDER BY and HAVING can reference them
  140. if col.Alias != "" {
  141. selectAliases[strings.ToUpper(col.Alias)] = info
  142. }
  143. }
  144. // Register SELECT aliases as virtual columns for ORDER BY/HAVING reference
  145. for alias, info := range selectAliases {
  146. a.scope.DefineSelectAlias(alias, info.Type)
  147. }
  148. // Validate GROUP BY semantics
  149. if hasAggregate && !hasGroupBy {
  150. // Aggregate query without GROUP BY - all non-aggregate columns must be constants
  151. for _, col := range stmt.Columns {
  152. if col.Star {
  153. return &AnalysisError{
  154. Type: ErrNonAggregateInSelect,
  155. Message: "SELECT * not allowed with aggregate functions without GROUP BY",
  156. }
  157. }
  158. info, exprErr := a.analyzeExpr(col.Expr)
  159. if exprErr != nil || info == nil {
  160. continue
  161. }
  162. if !info.IsAggregate && !info.IsConstant {
  163. // Check if it's a simple column reference
  164. if ref, ok := col.Expr.(*parser.ColumnRef); ok {
  165. return &AnalysisError{
  166. Type: ErrNonAggregateInSelect,
  167. Message: fmt.Sprintf("column %q must appear in GROUP BY clause or be in an aggregate function", ref.Column),
  168. }
  169. }
  170. }
  171. }
  172. }
  173. // Analyze HAVING clause
  174. if stmt.Having != nil {
  175. info, err := a.analyzeExpr(stmt.Having)
  176. if err != nil {
  177. return err
  178. }
  179. // HAVING without GROUP BY requires aggregates
  180. if !hasGroupBy && !info.IsAggregate {
  181. return &AnalysisError{
  182. Type: ErrInvalidGroupBy,
  183. Message: "HAVING clause requires GROUP BY or aggregate function",
  184. }
  185. }
  186. }
  187. // Analyze ORDER BY
  188. for _, item := range stmt.OrderBy {
  189. if _, err := a.analyzeExpr(item.Expr); err != nil {
  190. return err
  191. }
  192. }
  193. // Analyze LIMIT/OFFSET
  194. if stmt.Limit != nil {
  195. info, err := a.analyzeExpr(stmt.Limit)
  196. if err != nil {
  197. return err
  198. }
  199. if !info.Type.IsNumeric() && info.Type != TypeNull {
  200. return &AnalysisError{
  201. Type: ErrTypeMismatch,
  202. Message: "LIMIT must be numeric",
  203. }
  204. }
  205. }
  206. if stmt.Offset != nil {
  207. info, err := a.analyzeExpr(stmt.Offset)
  208. if err != nil {
  209. return err
  210. }
  211. if !info.Type.IsNumeric() && info.Type != TypeNull {
  212. return &AnalysisError{
  213. Type: ErrTypeMismatch,
  214. Message: "OFFSET must be numeric",
  215. }
  216. }
  217. }
  218. return nil
  219. }
  220. // resolveFromClause adds tables from FROM clause to scope.
  221. func (a *Analyzer) resolveFromClause(tables []parser.TableRef) error {
  222. for _, ref := range tables {
  223. // Handle subquery (derived table)
  224. if ref.Subquery != nil {
  225. // Analyze the subquery
  226. if err := a.analyzeSelect(ref.Subquery); err != nil {
  227. return err
  228. }
  229. // Create a table info from subquery columns
  230. // For now, we'll use a simplified approach - just mark it as a derived table
  231. tableInfo := &TableInfo{
  232. Name: ref.Alias, // Derived tables MUST have an alias
  233. Columns: []ColumnInfo{},
  234. Alias: ref.Alias,
  235. }
  236. // Add columns from SELECT list
  237. for _, col := range ref.Subquery.Columns {
  238. colName := ""
  239. if col.Alias != "" {
  240. colName = col.Alias
  241. } else if colRef, ok := col.Expr.(*parser.ColumnRef); ok {
  242. colName = colRef.Column
  243. } else {
  244. // For expressions without alias, use a generated name
  245. colName = fmt.Sprintf("col_%d", len(tableInfo.Columns))
  246. }
  247. tableInfo.Columns = append(tableInfo.Columns, ColumnInfo{
  248. Name: colName,
  249. TableName: ref.Alias,
  250. Type: TypeAny, // We'd need type inference for proper typing
  251. })
  252. }
  253. a.scope.DefineTable(tableInfo)
  254. } else {
  255. // Regular table reference
  256. table, ok := a.catalog.GetTable(ref.Name)
  257. if !ok {
  258. return &AnalysisError{
  259. Type: ErrTableNotFound,
  260. Message: fmt.Sprintf("table not found: %s", ref.Name),
  261. }
  262. }
  263. // Create a copy with alias if specified
  264. tableInfo := &TableInfo{
  265. Name: table.Name,
  266. Columns: table.Columns,
  267. Alias: ref.Alias,
  268. IsView: table.IsView,
  269. }
  270. a.scope.DefineTable(tableInfo)
  271. }
  272. // Handle JOINs
  273. if ref.Join != nil {
  274. if err := a.resolveJoin(ref.Join); err != nil {
  275. return err
  276. }
  277. }
  278. }
  279. return nil
  280. }
  281. // resolveJoin resolves a JOIN clause.
  282. func (a *Analyzer) resolveJoin(join *parser.JoinClause) error {
  283. if join.Table == nil {
  284. return nil
  285. }
  286. table, ok := a.catalog.GetTable(join.Table.Name)
  287. if !ok {
  288. return &AnalysisError{
  289. Type: ErrTableNotFound,
  290. Message: fmt.Sprintf("table not found: %s", join.Table.Name),
  291. }
  292. }
  293. tableInfo := &TableInfo{
  294. Name: table.Name,
  295. Columns: table.Columns,
  296. Alias: join.Table.Alias,
  297. IsView: table.IsView,
  298. }
  299. a.scope.DefineTable(tableInfo)
  300. // Analyze ON condition
  301. if join.Condition != nil {
  302. if _, err := a.analyzeExpr(join.Condition); err != nil {
  303. return err
  304. }
  305. }
  306. // Handle USING clause
  307. for _, colName := range join.Using {
  308. _, _, ok := a.scope.LookupColumn("", colName)
  309. if !ok {
  310. return &AnalysisError{
  311. Type: ErrColumnNotFound,
  312. Message: fmt.Sprintf("column not found in USING clause: %s", colName),
  313. }
  314. }
  315. }
  316. // Recursively handle chained JOINs
  317. if join.Table.Join != nil {
  318. if err := a.resolveJoin(join.Table.Join); err != nil {
  319. return err
  320. }
  321. }
  322. return nil
  323. }
  324. // analyzeInsert analyzes an INSERT statement.
  325. func (a *Analyzer) analyzeInsert(stmt *parser.InsertStmt) error {
  326. table, ok := a.catalog.GetTable(stmt.Table.Name)
  327. if !ok {
  328. return &AnalysisError{
  329. Type: ErrTableNotFound,
  330. Message: fmt.Sprintf("table not found: %s", stmt.Table.Name),
  331. }
  332. }
  333. // Validate column list if specified
  334. var targetCols []ColumnInfo
  335. if len(stmt.Columns) > 0 {
  336. for _, colName := range stmt.Columns {
  337. col, ok := table.GetColumn(colName)
  338. if !ok {
  339. return &AnalysisError{
  340. Type: ErrColumnNotFound,
  341. Message: fmt.Sprintf("column not found: %s", colName),
  342. }
  343. }
  344. targetCols = append(targetCols, *col)
  345. }
  346. } else {
  347. targetCols = table.Columns
  348. }
  349. // Validate VALUES
  350. for _, row := range stmt.Values {
  351. if len(row) != len(targetCols) {
  352. return &AnalysisError{
  353. Type: ErrTypeMismatch,
  354. Message: fmt.Sprintf("INSERT has %d columns but %d values", len(targetCols), len(row)),
  355. }
  356. }
  357. for i, expr := range row {
  358. info, err := a.analyzeExpr(expr)
  359. if err != nil {
  360. return err
  361. }
  362. // Check type compatibility
  363. if !info.Type.IsComparable(targetCols[i].Type) && info.Type != TypeNull {
  364. return &AnalysisError{
  365. Type: ErrTypeMismatch,
  366. Message: fmt.Sprintf("type mismatch for column %s: expected %s, got %s",
  367. targetCols[i].Name, targetCols[i].Type, info.Type),
  368. }
  369. }
  370. }
  371. }
  372. // Analyze INSERT ... SELECT
  373. if stmt.Select != nil {
  374. a.scope.DefineTable(&TableInfo{Name: table.Name, Columns: table.Columns, IsView: table.IsView})
  375. if err := a.analyzeSelect(stmt.Select); err != nil {
  376. return err
  377. }
  378. }
  379. return nil
  380. }
  381. // analyzeUpdate analyzes an UPDATE statement.
  382. func (a *Analyzer) analyzeUpdate(stmt *parser.UpdateStmt) error {
  383. table, ok := a.catalog.GetTable(stmt.Table.Name)
  384. if !ok {
  385. return &AnalysisError{
  386. Type: ErrTableNotFound,
  387. Message: fmt.Sprintf("table not found: %s", stmt.Table.Name),
  388. }
  389. }
  390. a.scope.DefineTable(&TableInfo{Name: table.Name, Columns: table.Columns, Alias: stmt.Table.Alias, IsView: table.IsView})
  391. // Validate SET assignments
  392. for _, assign := range stmt.Set {
  393. col, ok := table.GetColumn(assign.Column)
  394. if !ok {
  395. return &AnalysisError{
  396. Type: ErrColumnNotFound,
  397. Message: fmt.Sprintf("column not found: %s", assign.Column),
  398. }
  399. }
  400. info, err := a.analyzeExpr(assign.Value)
  401. if err != nil {
  402. return err
  403. }
  404. if !info.Type.IsComparable(col.Type) && info.Type != TypeNull {
  405. return &AnalysisError{
  406. Type: ErrTypeMismatch,
  407. Message: fmt.Sprintf("type mismatch for column %s: expected %s, got %s",
  408. col.Name, col.Type, info.Type),
  409. }
  410. }
  411. }
  412. // Analyze WHERE clause
  413. if stmt.Where != nil {
  414. info, err := a.analyzeExpr(stmt.Where)
  415. if err != nil {
  416. return err
  417. }
  418. if info.IsAggregate {
  419. return &AnalysisError{
  420. Type: ErrAggregateInWhere,
  421. Message: "aggregate functions not allowed in WHERE clause",
  422. }
  423. }
  424. }
  425. return nil
  426. }
  427. // analyzeDelete analyzes a DELETE statement.
  428. func (a *Analyzer) analyzeDelete(stmt *parser.DeleteStmt) error {
  429. table, ok := a.catalog.GetTable(stmt.Table.Name)
  430. if !ok {
  431. return &AnalysisError{
  432. Type: ErrTableNotFound,
  433. Message: fmt.Sprintf("table not found: %s", stmt.Table.Name),
  434. }
  435. }
  436. a.scope.DefineTable(&TableInfo{Name: table.Name, Columns: table.Columns, IsView: table.IsView})
  437. // Analyze WHERE clause
  438. if stmt.Where != nil {
  439. info, err := a.analyzeExpr(stmt.Where)
  440. if err != nil {
  441. return err
  442. }
  443. if info.IsAggregate {
  444. return &AnalysisError{
  445. Type: ErrAggregateInWhere,
  446. Message: "aggregate functions not allowed in WHERE clause",
  447. }
  448. }
  449. }
  450. return nil
  451. }
  452. // analyzeCreateTable analyzes a CREATE TABLE statement.
  453. func (a *Analyzer) analyzeCreateTable(stmt *parser.CreateTableStmt) error {
  454. // Check if table already exists
  455. if a.catalog.TableExists(stmt.Table.Name) {
  456. if stmt.IfNotExists {
  457. return nil // Silently succeed
  458. }
  459. return &AnalysisError{
  460. Type: ErrTableExists,
  461. Message: fmt.Sprintf("table already exists: %s", stmt.Table.Name),
  462. }
  463. }
  464. // Build table info
  465. tableInfo := &TableInfo{
  466. Name: stmt.Table.Name,
  467. }
  468. columnNames := make(map[string]bool)
  469. for _, colDef := range stmt.Columns {
  470. upperName := strings.ToUpper(colDef.Name)
  471. if columnNames[upperName] {
  472. return &AnalysisError{
  473. Type: ErrColumnAmbiguous,
  474. Message: fmt.Sprintf("duplicate column name: %s", colDef.Name),
  475. }
  476. }
  477. columnNames[upperName] = true
  478. colInfo := ColumnInfo{
  479. Name: colDef.Name,
  480. Type: TypeFromName(colDef.Type.Name),
  481. Nullable: true,
  482. TableName: stmt.Table.Name,
  483. }
  484. // Process constraints
  485. for _, constraint := range colDef.Constraints {
  486. switch constraint.Type {
  487. case parser.ConstraintPrimaryKey:
  488. colInfo.PrimaryKey = true
  489. colInfo.Nullable = false
  490. case parser.ConstraintNotNull:
  491. colInfo.Nullable = false
  492. case parser.ConstraintDefault:
  493. // Store default value (not evaluated here)
  494. colInfo.Default = constraint.Default
  495. }
  496. }
  497. tableInfo.Columns = append(tableInfo.Columns, colInfo)
  498. }
  499. // Process table-level constraints
  500. for _, constraint := range stmt.Constraints {
  501. switch constraint.Type {
  502. case parser.ConstraintPrimaryKey:
  503. for _, colName := range constraint.Columns {
  504. for i := range tableInfo.Columns {
  505. if strings.EqualFold(tableInfo.Columns[i].Name, colName) {
  506. tableInfo.Columns[i].PrimaryKey = true
  507. tableInfo.Columns[i].Nullable = false
  508. }
  509. }
  510. }
  511. }
  512. }
  513. // Add to catalog
  514. return a.catalog.CreateTable(tableInfo)
  515. }
  516. // analyzeDropTable analyzes a DROP TABLE statement.
  517. func (a *Analyzer) analyzeDropTable(stmt *parser.DropTableStmt) error {
  518. for _, tableRef := range stmt.Tables {
  519. if !a.catalog.TableExists(tableRef.Name) {
  520. if stmt.IfExists {
  521. continue // Silently succeed
  522. }
  523. return &AnalysisError{
  524. Type: ErrTableNotFound,
  525. Message: fmt.Sprintf("table not found: %s", tableRef.Name),
  526. }
  527. }
  528. if err := a.catalog.DropTable(tableRef.Name); err != nil {
  529. return err
  530. }
  531. }
  532. return nil
  533. }
  534. // analyzeExpr analyzes an expression and returns type information.
  535. func (a *Analyzer) analyzeExpr(expr parser.Expr) (*ExprInfo, error) {
  536. switch e := expr.(type) {
  537. case *parser.LiteralExpr:
  538. return a.analyzeLiteral(e)
  539. case *parser.ColumnRef:
  540. return a.analyzeColumnRef(e)
  541. case *parser.BinaryExpr:
  542. return a.analyzeBinaryExpr(e)
  543. case *parser.UnaryExpr:
  544. return a.analyzeUnaryExpr(e)
  545. case *parser.FunctionCall:
  546. return a.analyzeFunctionCall(e)
  547. case *parser.ParenExpr:
  548. return a.analyzeExpr(e.Expr)
  549. case *parser.CaseExpr:
  550. return a.analyzeCaseExpr(e)
  551. case *parser.CastExpr:
  552. return a.analyzeCastExpr(e)
  553. case *parser.InExpr:
  554. return a.analyzeInExpr(e)
  555. case *parser.BetweenExpr:
  556. return a.analyzeBetweenExpr(e)
  557. case *parser.LikeExpr:
  558. return a.analyzeLikeExpr(e)
  559. case *parser.IsNullExpr:
  560. return a.analyzeIsNullExpr(e)
  561. case *parser.ExistsExpr:
  562. return a.analyzeExistsExpr(e)
  563. case *parser.SubqueryExpr:
  564. return a.analyzeSubqueryExpr(e)
  565. default:
  566. return &ExprInfo{Type: TypeUnknown}, nil
  567. }
  568. }
  569. func (a *Analyzer) analyzeLiteral(e *parser.LiteralExpr) (*ExprInfo, error) {
  570. info := &ExprInfo{IsConstant: true}
  571. switch e.Type {
  572. case lexer.TokenNumber:
  573. if strings.Contains(e.Value, ".") || strings.Contains(strings.ToLower(e.Value), "e") {
  574. info.Type = TypeReal
  575. } else {
  576. info.Type = TypeInteger
  577. }
  578. case lexer.TokenString:
  579. info.Type = TypeText
  580. case lexer.TokenNULL:
  581. info.Type = TypeNull
  582. info.Nullable = true
  583. case lexer.TokenTRUE, lexer.TokenFALSE:
  584. info.Type = TypeBoolean
  585. case lexer.TokenStar:
  586. info.Type = TypeAny
  587. default:
  588. info.Type = TypeUnknown
  589. }
  590. return info, nil
  591. }
  592. func (a *Analyzer) analyzeColumnRef(e *parser.ColumnRef) (*ExprInfo, error) {
  593. col, _, ok := a.scope.LookupColumn(e.Table, e.Column)
  594. if !ok {
  595. // If no tables are in scope, treat as unknown (for standalone expressions)
  596. if len(a.scope.GetTables()) == 0 {
  597. return &ExprInfo{Type: TypeUnknown}, nil
  598. }
  599. return nil, &AnalysisError{
  600. Type: ErrColumnNotFound,
  601. Message: fmt.Sprintf("column not found: %s", formatColumnRef(e)),
  602. }
  603. }
  604. return &ExprInfo{
  605. Type: col.Type,
  606. Nullable: col.Nullable,
  607. }, nil
  608. }
  609. func formatColumnRef(e *parser.ColumnRef) string {
  610. if e.Table != "" {
  611. return e.Table + "." + e.Column
  612. }
  613. return e.Column
  614. }
  615. func (a *Analyzer) analyzeBinaryExpr(e *parser.BinaryExpr) (*ExprInfo, error) {
  616. left, err := a.analyzeExpr(e.Left)
  617. if err != nil {
  618. return nil, err
  619. }
  620. right, err := a.analyzeExpr(e.Right)
  621. if err != nil {
  622. return nil, err
  623. }
  624. info := &ExprInfo{
  625. IsAggregate: left.IsAggregate || right.IsAggregate,
  626. IsConstant: left.IsConstant && right.IsConstant,
  627. Nullable: left.Nullable || right.Nullable,
  628. }
  629. switch e.Op {
  630. case lexer.TokenPlus, lexer.TokenMinus, lexer.TokenStar, lexer.TokenSlash, lexer.TokenPercent:
  631. // Arithmetic operators
  632. info.Type = CommonType(left.Type, right.Type)
  633. if !left.Type.IsNumeric() && left.Type != TypeNull && left.Type != TypeUnknown {
  634. return nil, &AnalysisError{
  635. Type: ErrTypeMismatch,
  636. Message: fmt.Sprintf("arithmetic operator requires numeric type, got %s", left.Type),
  637. }
  638. }
  639. case lexer.TokenEq, lexer.TokenNeq, lexer.TokenLt, lexer.TokenLte, lexer.TokenGt, lexer.TokenGte:
  640. // Comparison operators
  641. info.Type = TypeBoolean
  642. if !left.Type.IsComparable(right.Type) {
  643. return nil, &AnalysisError{
  644. Type: ErrTypeMismatch,
  645. Message: fmt.Sprintf("cannot compare %s with %s", left.Type, right.Type),
  646. }
  647. }
  648. case lexer.TokenAND, lexer.TokenOR:
  649. // Logical operators
  650. info.Type = TypeBoolean
  651. case lexer.TokenConcat:
  652. // String concatenation
  653. info.Type = TypeText
  654. default:
  655. info.Type = TypeUnknown
  656. }
  657. return info, nil
  658. }
  659. func (a *Analyzer) analyzeUnaryExpr(e *parser.UnaryExpr) (*ExprInfo, error) {
  660. operand, err := a.analyzeExpr(e.Operand)
  661. if err != nil {
  662. return nil, err
  663. }
  664. info := &ExprInfo{
  665. IsAggregate: operand.IsAggregate,
  666. IsConstant: operand.IsConstant,
  667. Nullable: operand.Nullable,
  668. }
  669. switch e.Op {
  670. case lexer.TokenPlus:
  671. // Unary + is a no-op in SQLite — passes any type through unchanged.
  672. info.Type = operand.Type
  673. case lexer.TokenMinus:
  674. info.Type = operand.Type
  675. if !operand.Type.IsNumeric() && operand.Type != TypeNull && operand.Type != TypeUnknown {
  676. return nil, &AnalysisError{
  677. Type: ErrTypeMismatch,
  678. Message: fmt.Sprintf("unary - requires numeric type, got %s", operand.Type),
  679. }
  680. }
  681. case lexer.TokenNOT:
  682. info.Type = TypeBoolean
  683. default:
  684. info.Type = operand.Type
  685. }
  686. return info, nil
  687. }
  688. func (a *Analyzer) analyzeFunctionCall(e *parser.FunctionCall) (*ExprInfo, error) {
  689. sig, ok := LookupFunction(e.Name)
  690. if !ok {
  691. return nil, &AnalysisError{
  692. Type: ErrInvalidFunction,
  693. Message: fmt.Sprintf("unknown function: %s", e.Name),
  694. }
  695. }
  696. // Handle COUNT(*)
  697. argCount := len(e.Args)
  698. if e.Star {
  699. argCount = 0 // COUNT(*) has 0 real args
  700. }
  701. // Check argument count
  702. if argCount < sig.MinArgs {
  703. return nil, &AnalysisError{
  704. Type: ErrInvalidArgCount,
  705. Message: fmt.Sprintf("function %s requires at least %d arguments, got %d", e.Name, sig.MinArgs, argCount),
  706. }
  707. }
  708. if sig.MaxArgs >= 0 && argCount > sig.MaxArgs {
  709. return nil, &AnalysisError{
  710. Type: ErrInvalidArgCount,
  711. Message: fmt.Sprintf("function %s accepts at most %d arguments, got %d", e.Name, sig.MaxArgs, argCount),
  712. }
  713. }
  714. // Analyze arguments
  715. info := &ExprInfo{
  716. Type: sig.ReturnType,
  717. IsAggregate: sig.IsAggregate,
  718. }
  719. for _, arg := range e.Args {
  720. argInfo, err := a.analyzeExpr(arg)
  721. if err != nil {
  722. return nil, err
  723. }
  724. if argInfo.Nullable {
  725. info.Nullable = true
  726. }
  727. // Propagate aggregate status from arguments
  728. if argInfo.IsAggregate && !sig.IsAggregate {
  729. info.IsAggregate = true
  730. }
  731. }
  732. // Special case: MIN/MAX/COALESCE return type depends on argument
  733. if sig.ReturnType == TypeAny && len(e.Args) > 0 {
  734. argInfo, _ := a.analyzeExpr(e.Args[0])
  735. if argInfo != nil {
  736. info.Type = argInfo.Type
  737. }
  738. }
  739. return info, nil
  740. }
  741. func (a *Analyzer) analyzeCaseExpr(e *parser.CaseExpr) (*ExprInfo, error) {
  742. info := &ExprInfo{
  743. Nullable: true, // CASE can return NULL
  744. }
  745. // Analyze operand if present (simple CASE)
  746. if e.Operand != nil {
  747. opInfo, err := a.analyzeExpr(e.Operand)
  748. if err != nil {
  749. return nil, err
  750. }
  751. if opInfo.IsAggregate {
  752. info.IsAggregate = true
  753. }
  754. }
  755. // Analyze WHEN clauses
  756. var resultType Type
  757. for _, when := range e.Whens {
  758. condInfo, err := a.analyzeExpr(when.Condition)
  759. if err != nil {
  760. return nil, err
  761. }
  762. if condInfo.IsAggregate {
  763. info.IsAggregate = true
  764. }
  765. resInfo, err := a.analyzeExpr(when.Result)
  766. if err != nil {
  767. return nil, err
  768. }
  769. if resInfo.IsAggregate {
  770. info.IsAggregate = true
  771. }
  772. if resultType == TypeUnknown {
  773. resultType = resInfo.Type
  774. } else {
  775. resultType = CommonType(resultType, resInfo.Type)
  776. }
  777. }
  778. // Analyze ELSE clause
  779. if e.Else != nil {
  780. elseInfo, err := a.analyzeExpr(e.Else)
  781. if err != nil {
  782. return nil, err
  783. }
  784. if elseInfo.IsAggregate {
  785. info.IsAggregate = true
  786. }
  787. resultType = CommonType(resultType, elseInfo.Type)
  788. }
  789. info.Type = resultType
  790. return info, nil
  791. }
  792. func (a *Analyzer) analyzeCastExpr(e *parser.CastExpr) (*ExprInfo, error) {
  793. exprInfo, err := a.analyzeExpr(e.Expr)
  794. if err != nil {
  795. return nil, err
  796. }
  797. return &ExprInfo{
  798. Type: TypeFromName(e.Type.Name),
  799. IsAggregate: exprInfo.IsAggregate,
  800. IsConstant: exprInfo.IsConstant,
  801. Nullable: exprInfo.Nullable,
  802. }, nil
  803. }
  804. func (a *Analyzer) analyzeInExpr(e *parser.InExpr) (*ExprInfo, error) {
  805. leftInfo, err := a.analyzeExpr(e.Left)
  806. if err != nil {
  807. return nil, err
  808. }
  809. info := &ExprInfo{
  810. Type: TypeBoolean,
  811. IsAggregate: leftInfo.IsAggregate,
  812. }
  813. // Analyze value list
  814. for _, val := range e.Values {
  815. valInfo, err := a.analyzeExpr(val)
  816. if err != nil {
  817. return nil, err
  818. }
  819. if valInfo.IsAggregate {
  820. info.IsAggregate = true
  821. }
  822. }
  823. // Analyze subquery
  824. if e.Subquery != nil {
  825. subScope := NewScope(a.scope)
  826. oldScope := a.scope
  827. a.scope = subScope
  828. err := a.analyzeSelect(e.Subquery)
  829. a.scope = oldScope
  830. if err != nil {
  831. return nil, err
  832. }
  833. }
  834. return info, nil
  835. }
  836. func (a *Analyzer) analyzeBetweenExpr(e *parser.BetweenExpr) (*ExprInfo, error) {
  837. leftInfo, err := a.analyzeExpr(e.Left)
  838. if err != nil {
  839. return nil, err
  840. }
  841. lowInfo, err := a.analyzeExpr(e.Low)
  842. if err != nil {
  843. return nil, err
  844. }
  845. highInfo, err := a.analyzeExpr(e.High)
  846. if err != nil {
  847. return nil, err
  848. }
  849. return &ExprInfo{
  850. Type: TypeBoolean,
  851. IsAggregate: leftInfo.IsAggregate || lowInfo.IsAggregate || highInfo.IsAggregate,
  852. Nullable: leftInfo.Nullable || lowInfo.Nullable || highInfo.Nullable,
  853. }, nil
  854. }
  855. func (a *Analyzer) analyzeLikeExpr(e *parser.LikeExpr) (*ExprInfo, error) {
  856. leftInfo, err := a.analyzeExpr(e.Left)
  857. if err != nil {
  858. return nil, err
  859. }
  860. patternInfo, err := a.analyzeExpr(e.Pattern)
  861. if err != nil {
  862. return nil, err
  863. }
  864. info := &ExprInfo{
  865. Type: TypeBoolean,
  866. IsAggregate: leftInfo.IsAggregate || patternInfo.IsAggregate,
  867. Nullable: leftInfo.Nullable || patternInfo.Nullable,
  868. }
  869. if e.Escape != nil {
  870. escInfo, err := a.analyzeExpr(e.Escape)
  871. if err != nil {
  872. return nil, err
  873. }
  874. if escInfo.IsAggregate {
  875. info.IsAggregate = true
  876. }
  877. }
  878. return info, nil
  879. }
  880. func (a *Analyzer) analyzeIsNullExpr(e *parser.IsNullExpr) (*ExprInfo, error) {
  881. leftInfo, err := a.analyzeExpr(e.Left)
  882. if err != nil {
  883. return nil, err
  884. }
  885. return &ExprInfo{
  886. Type: TypeBoolean,
  887. IsAggregate: leftInfo.IsAggregate,
  888. IsConstant: leftInfo.IsConstant,
  889. }, nil
  890. }
  891. func (a *Analyzer) analyzeExistsExpr(e *parser.ExistsExpr) (*ExprInfo, error) {
  892. // Analyze subquery in its own scope
  893. subScope := NewScope(a.scope)
  894. oldScope := a.scope
  895. a.scope = subScope
  896. err := a.analyzeSelect(e.Subquery)
  897. a.scope = oldScope
  898. if err != nil {
  899. return nil, err
  900. }
  901. return &ExprInfo{
  902. Type: TypeBoolean,
  903. }, nil
  904. }
  905. func (a *Analyzer) analyzeSubqueryExpr(e *parser.SubqueryExpr) (*ExprInfo, error) {
  906. // Analyze subquery in its own scope
  907. subScope := NewScope(a.scope)
  908. oldScope := a.scope
  909. a.scope = subScope
  910. err := a.analyzeSelect(e.Query)
  911. a.scope = oldScope
  912. if err != nil {
  913. return nil, err
  914. }
  915. // Scalar subquery - return type of first column
  916. // For simplicity, return TypeAny
  917. return &ExprInfo{
  918. Type: TypeAny,
  919. }, nil
  920. }