analyzer.go 25 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021102210231024102510261027102810291030103110321033103410351036103710381039104010411042
  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. // Analysis is validation-only. Publishing the table before the durable
  514. // schema write can leave a phantom catalog entry when that write fails.
  515. return nil
  516. }
  517. // analyzeDropTable analyzes a DROP TABLE statement.
  518. func (a *Analyzer) analyzeDropTable(stmt *parser.DropTableStmt) error {
  519. for _, tableRef := range stmt.Tables {
  520. if !a.catalog.TableExists(tableRef.Name) {
  521. if stmt.IfExists {
  522. continue // Silently succeed
  523. }
  524. return &AnalysisError{
  525. Type: ErrTableNotFound,
  526. Message: fmt.Sprintf("table not found: %s", tableRef.Name),
  527. }
  528. }
  529. if err := a.catalog.DropTable(tableRef.Name); err != nil {
  530. return err
  531. }
  532. }
  533. return nil
  534. }
  535. // analyzeExpr analyzes an expression and returns type information.
  536. func (a *Analyzer) analyzeExpr(expr parser.Expr) (*ExprInfo, error) {
  537. switch e := expr.(type) {
  538. case *parser.LiteralExpr:
  539. return a.analyzeLiteral(e)
  540. case *parser.ColumnRef:
  541. return a.analyzeColumnRef(e)
  542. case *parser.BinaryExpr:
  543. return a.analyzeBinaryExpr(e)
  544. case *parser.UnaryExpr:
  545. return a.analyzeUnaryExpr(e)
  546. case *parser.FunctionCall:
  547. return a.analyzeFunctionCall(e)
  548. case *parser.ParenExpr:
  549. return a.analyzeExpr(e.Expr)
  550. case *parser.CaseExpr:
  551. return a.analyzeCaseExpr(e)
  552. case *parser.CastExpr:
  553. return a.analyzeCastExpr(e)
  554. case *parser.InExpr:
  555. return a.analyzeInExpr(e)
  556. case *parser.BetweenExpr:
  557. return a.analyzeBetweenExpr(e)
  558. case *parser.LikeExpr:
  559. return a.analyzeLikeExpr(e)
  560. case *parser.IsNullExpr:
  561. return a.analyzeIsNullExpr(e)
  562. case *parser.ExistsExpr:
  563. return a.analyzeExistsExpr(e)
  564. case *parser.SubqueryExpr:
  565. return a.analyzeSubqueryExpr(e)
  566. default:
  567. return &ExprInfo{Type: TypeUnknown}, nil
  568. }
  569. }
  570. func (a *Analyzer) analyzeLiteral(e *parser.LiteralExpr) (*ExprInfo, error) {
  571. info := &ExprInfo{IsConstant: true}
  572. switch e.Type {
  573. case lexer.TokenNumber:
  574. if strings.Contains(e.Value, ".") || strings.Contains(strings.ToLower(e.Value), "e") {
  575. info.Type = TypeReal
  576. } else {
  577. info.Type = TypeInteger
  578. }
  579. case lexer.TokenString:
  580. info.Type = TypeText
  581. case lexer.TokenNULL:
  582. info.Type = TypeNull
  583. info.Nullable = true
  584. case lexer.TokenTRUE, lexer.TokenFALSE:
  585. info.Type = TypeBoolean
  586. case lexer.TokenStar:
  587. info.Type = TypeAny
  588. default:
  589. info.Type = TypeUnknown
  590. }
  591. return info, nil
  592. }
  593. func (a *Analyzer) analyzeColumnRef(e *parser.ColumnRef) (*ExprInfo, error) {
  594. col, _, ok := a.scope.LookupColumn(e.Table, e.Column)
  595. if !ok {
  596. // If no tables are in scope, treat as unknown (for standalone expressions)
  597. if len(a.scope.GetTables()) == 0 {
  598. return &ExprInfo{Type: TypeUnknown}, nil
  599. }
  600. return nil, &AnalysisError{
  601. Type: ErrColumnNotFound,
  602. Message: fmt.Sprintf("column not found: %s", formatColumnRef(e)),
  603. }
  604. }
  605. return &ExprInfo{
  606. Type: col.Type,
  607. Nullable: col.Nullable,
  608. }, nil
  609. }
  610. func formatColumnRef(e *parser.ColumnRef) string {
  611. if e.Table != "" {
  612. return e.Table + "." + e.Column
  613. }
  614. return e.Column
  615. }
  616. func (a *Analyzer) analyzeBinaryExpr(e *parser.BinaryExpr) (*ExprInfo, error) {
  617. left, err := a.analyzeExpr(e.Left)
  618. if err != nil {
  619. return nil, err
  620. }
  621. right, err := a.analyzeExpr(e.Right)
  622. if err != nil {
  623. return nil, err
  624. }
  625. info := &ExprInfo{
  626. IsAggregate: left.IsAggregate || right.IsAggregate,
  627. IsConstant: left.IsConstant && right.IsConstant,
  628. Nullable: left.Nullable || right.Nullable,
  629. }
  630. switch e.Op {
  631. case lexer.TokenPlus, lexer.TokenMinus, lexer.TokenStar, lexer.TokenSlash, lexer.TokenPercent:
  632. // Arithmetic operators
  633. info.Type = CommonType(left.Type, right.Type)
  634. if !left.Type.IsNumeric() && left.Type != TypeNull && left.Type != TypeUnknown {
  635. return nil, &AnalysisError{
  636. Type: ErrTypeMismatch,
  637. Message: fmt.Sprintf("arithmetic operator requires numeric type, got %s", left.Type),
  638. }
  639. }
  640. case lexer.TokenEq, lexer.TokenNeq, lexer.TokenLt, lexer.TokenLte, lexer.TokenGt, lexer.TokenGte:
  641. // Comparison operators
  642. info.Type = TypeBoolean
  643. if !left.Type.IsComparable(right.Type) {
  644. return nil, &AnalysisError{
  645. Type: ErrTypeMismatch,
  646. Message: fmt.Sprintf("cannot compare %s with %s", left.Type, right.Type),
  647. }
  648. }
  649. case lexer.TokenAND, lexer.TokenOR:
  650. // Logical operators
  651. info.Type = TypeBoolean
  652. case lexer.TokenConcat:
  653. // String concatenation
  654. info.Type = TypeText
  655. default:
  656. info.Type = TypeUnknown
  657. }
  658. return info, nil
  659. }
  660. func (a *Analyzer) analyzeUnaryExpr(e *parser.UnaryExpr) (*ExprInfo, error) {
  661. operand, err := a.analyzeExpr(e.Operand)
  662. if err != nil {
  663. return nil, err
  664. }
  665. info := &ExprInfo{
  666. IsAggregate: operand.IsAggregate,
  667. IsConstant: operand.IsConstant,
  668. Nullable: operand.Nullable,
  669. }
  670. switch e.Op {
  671. case lexer.TokenPlus:
  672. // Unary + is a no-op in SQLite — passes any type through unchanged.
  673. info.Type = operand.Type
  674. case lexer.TokenMinus:
  675. info.Type = operand.Type
  676. if !operand.Type.IsNumeric() && operand.Type != TypeNull && operand.Type != TypeUnknown {
  677. return nil, &AnalysisError{
  678. Type: ErrTypeMismatch,
  679. Message: fmt.Sprintf("unary - requires numeric type, got %s", operand.Type),
  680. }
  681. }
  682. case lexer.TokenNOT:
  683. info.Type = TypeBoolean
  684. default:
  685. info.Type = operand.Type
  686. }
  687. return info, nil
  688. }
  689. func (a *Analyzer) analyzeFunctionCall(e *parser.FunctionCall) (*ExprInfo, error) {
  690. sig, ok := LookupFunction(e.Name)
  691. if !ok {
  692. return nil, &AnalysisError{
  693. Type: ErrInvalidFunction,
  694. Message: fmt.Sprintf("unknown function: %s", e.Name),
  695. }
  696. }
  697. // Handle COUNT(*)
  698. argCount := len(e.Args)
  699. if e.Star {
  700. argCount = 0 // COUNT(*) has 0 real args
  701. }
  702. // Check argument count
  703. if argCount < sig.MinArgs {
  704. return nil, &AnalysisError{
  705. Type: ErrInvalidArgCount,
  706. Message: fmt.Sprintf("function %s requires at least %d arguments, got %d", e.Name, sig.MinArgs, argCount),
  707. }
  708. }
  709. if sig.MaxArgs >= 0 && argCount > sig.MaxArgs {
  710. return nil, &AnalysisError{
  711. Type: ErrInvalidArgCount,
  712. Message: fmt.Sprintf("function %s accepts at most %d arguments, got %d", e.Name, sig.MaxArgs, argCount),
  713. }
  714. }
  715. // Analyze arguments
  716. info := &ExprInfo{
  717. Type: sig.ReturnType,
  718. IsAggregate: sig.IsAggregate,
  719. }
  720. for _, arg := range e.Args {
  721. argInfo, err := a.analyzeExpr(arg)
  722. if err != nil {
  723. return nil, err
  724. }
  725. if argInfo.Nullable {
  726. info.Nullable = true
  727. }
  728. // Propagate aggregate status from arguments
  729. if argInfo.IsAggregate && !sig.IsAggregate {
  730. info.IsAggregate = true
  731. }
  732. }
  733. // Special case: MIN/MAX/COALESCE return type depends on argument
  734. if sig.ReturnType == TypeAny && len(e.Args) > 0 {
  735. argInfo, _ := a.analyzeExpr(e.Args[0])
  736. if argInfo != nil {
  737. info.Type = argInfo.Type
  738. }
  739. }
  740. return info, nil
  741. }
  742. func (a *Analyzer) analyzeCaseExpr(e *parser.CaseExpr) (*ExprInfo, error) {
  743. info := &ExprInfo{
  744. Nullable: true, // CASE can return NULL
  745. }
  746. // Analyze operand if present (simple CASE)
  747. if e.Operand != nil {
  748. opInfo, err := a.analyzeExpr(e.Operand)
  749. if err != nil {
  750. return nil, err
  751. }
  752. if opInfo.IsAggregate {
  753. info.IsAggregate = true
  754. }
  755. }
  756. // Analyze WHEN clauses
  757. var resultType Type
  758. for _, when := range e.Whens {
  759. condInfo, err := a.analyzeExpr(when.Condition)
  760. if err != nil {
  761. return nil, err
  762. }
  763. if condInfo.IsAggregate {
  764. info.IsAggregate = true
  765. }
  766. resInfo, err := a.analyzeExpr(when.Result)
  767. if err != nil {
  768. return nil, err
  769. }
  770. if resInfo.IsAggregate {
  771. info.IsAggregate = true
  772. }
  773. if resultType == TypeUnknown {
  774. resultType = resInfo.Type
  775. } else {
  776. resultType = CommonType(resultType, resInfo.Type)
  777. }
  778. }
  779. // Analyze ELSE clause
  780. if e.Else != nil {
  781. elseInfo, err := a.analyzeExpr(e.Else)
  782. if err != nil {
  783. return nil, err
  784. }
  785. if elseInfo.IsAggregate {
  786. info.IsAggregate = true
  787. }
  788. resultType = CommonType(resultType, elseInfo.Type)
  789. }
  790. info.Type = resultType
  791. return info, nil
  792. }
  793. func (a *Analyzer) analyzeCastExpr(e *parser.CastExpr) (*ExprInfo, error) {
  794. exprInfo, err := a.analyzeExpr(e.Expr)
  795. if err != nil {
  796. return nil, err
  797. }
  798. return &ExprInfo{
  799. Type: TypeFromName(e.Type.Name),
  800. IsAggregate: exprInfo.IsAggregate,
  801. IsConstant: exprInfo.IsConstant,
  802. Nullable: exprInfo.Nullable,
  803. }, nil
  804. }
  805. func (a *Analyzer) analyzeInExpr(e *parser.InExpr) (*ExprInfo, error) {
  806. leftInfo, err := a.analyzeExpr(e.Left)
  807. if err != nil {
  808. return nil, err
  809. }
  810. info := &ExprInfo{
  811. Type: TypeBoolean,
  812. IsAggregate: leftInfo.IsAggregate,
  813. }
  814. // Analyze value list
  815. for _, val := range e.Values {
  816. valInfo, err := a.analyzeExpr(val)
  817. if err != nil {
  818. return nil, err
  819. }
  820. if valInfo.IsAggregate {
  821. info.IsAggregate = true
  822. }
  823. }
  824. // Analyze subquery
  825. if e.Subquery != nil {
  826. subScope := NewScope(a.scope)
  827. oldScope := a.scope
  828. a.scope = subScope
  829. err := a.analyzeSelect(e.Subquery)
  830. a.scope = oldScope
  831. if err != nil {
  832. return nil, err
  833. }
  834. }
  835. return info, nil
  836. }
  837. func (a *Analyzer) analyzeBetweenExpr(e *parser.BetweenExpr) (*ExprInfo, error) {
  838. leftInfo, err := a.analyzeExpr(e.Left)
  839. if err != nil {
  840. return nil, err
  841. }
  842. lowInfo, err := a.analyzeExpr(e.Low)
  843. if err != nil {
  844. return nil, err
  845. }
  846. highInfo, err := a.analyzeExpr(e.High)
  847. if err != nil {
  848. return nil, err
  849. }
  850. return &ExprInfo{
  851. Type: TypeBoolean,
  852. IsAggregate: leftInfo.IsAggregate || lowInfo.IsAggregate || highInfo.IsAggregate,
  853. Nullable: leftInfo.Nullable || lowInfo.Nullable || highInfo.Nullable,
  854. }, nil
  855. }
  856. func (a *Analyzer) analyzeLikeExpr(e *parser.LikeExpr) (*ExprInfo, error) {
  857. leftInfo, err := a.analyzeExpr(e.Left)
  858. if err != nil {
  859. return nil, err
  860. }
  861. patternInfo, err := a.analyzeExpr(e.Pattern)
  862. if err != nil {
  863. return nil, err
  864. }
  865. info := &ExprInfo{
  866. Type: TypeBoolean,
  867. IsAggregate: leftInfo.IsAggregate || patternInfo.IsAggregate,
  868. Nullable: leftInfo.Nullable || patternInfo.Nullable,
  869. }
  870. if e.Escape != nil {
  871. escInfo, err := a.analyzeExpr(e.Escape)
  872. if err != nil {
  873. return nil, err
  874. }
  875. if escInfo.IsAggregate {
  876. info.IsAggregate = true
  877. }
  878. }
  879. return info, nil
  880. }
  881. func (a *Analyzer) analyzeIsNullExpr(e *parser.IsNullExpr) (*ExprInfo, error) {
  882. leftInfo, err := a.analyzeExpr(e.Left)
  883. if err != nil {
  884. return nil, err
  885. }
  886. return &ExprInfo{
  887. Type: TypeBoolean,
  888. IsAggregate: leftInfo.IsAggregate,
  889. IsConstant: leftInfo.IsConstant,
  890. }, nil
  891. }
  892. func (a *Analyzer) analyzeExistsExpr(e *parser.ExistsExpr) (*ExprInfo, error) {
  893. // Analyze subquery in its own scope
  894. subScope := NewScope(a.scope)
  895. oldScope := a.scope
  896. a.scope = subScope
  897. err := a.analyzeSelect(e.Subquery)
  898. a.scope = oldScope
  899. if err != nil {
  900. return nil, err
  901. }
  902. return &ExprInfo{
  903. Type: TypeBoolean,
  904. }, nil
  905. }
  906. func (a *Analyzer) analyzeSubqueryExpr(e *parser.SubqueryExpr) (*ExprInfo, error) {
  907. // Analyze subquery in its own scope
  908. subScope := NewScope(a.scope)
  909. oldScope := a.scope
  910. a.scope = subScope
  911. err := a.analyzeSelect(e.Query)
  912. a.scope = oldScope
  913. if err != nil {
  914. return nil, err
  915. }
  916. // Scalar subquery - return type of first column
  917. // For simplicity, return TypeAny
  918. return &ExprInfo{
  919. Type: TypeAny,
  920. }, nil
  921. }