2
0

analyzer.go 24 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021102210231024102510261027102810291030103110321033103410351036103710381039
  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. }
  269. a.scope.DefineTable(tableInfo)
  270. }
  271. // Handle JOINs
  272. if ref.Join != nil {
  273. if err := a.resolveJoin(ref.Join); err != nil {
  274. return err
  275. }
  276. }
  277. }
  278. return nil
  279. }
  280. // resolveJoin resolves a JOIN clause.
  281. func (a *Analyzer) resolveJoin(join *parser.JoinClause) error {
  282. if join.Table == nil {
  283. return nil
  284. }
  285. table, ok := a.catalog.GetTable(join.Table.Name)
  286. if !ok {
  287. return &AnalysisError{
  288. Type: ErrTableNotFound,
  289. Message: fmt.Sprintf("table not found: %s", join.Table.Name),
  290. }
  291. }
  292. tableInfo := &TableInfo{
  293. Name: table.Name,
  294. Columns: table.Columns,
  295. Alias: join.Table.Alias,
  296. }
  297. a.scope.DefineTable(tableInfo)
  298. // Analyze ON condition
  299. if join.Condition != nil {
  300. if _, err := a.analyzeExpr(join.Condition); err != nil {
  301. return err
  302. }
  303. }
  304. // Handle USING clause
  305. for _, colName := range join.Using {
  306. _, _, ok := a.scope.LookupColumn("", colName)
  307. if !ok {
  308. return &AnalysisError{
  309. Type: ErrColumnNotFound,
  310. Message: fmt.Sprintf("column not found in USING clause: %s", colName),
  311. }
  312. }
  313. }
  314. // Recursively handle chained JOINs
  315. if join.Table.Join != nil {
  316. if err := a.resolveJoin(join.Table.Join); err != nil {
  317. return err
  318. }
  319. }
  320. return nil
  321. }
  322. // analyzeInsert analyzes an INSERT statement.
  323. func (a *Analyzer) analyzeInsert(stmt *parser.InsertStmt) error {
  324. table, ok := a.catalog.GetTable(stmt.Table.Name)
  325. if !ok {
  326. return &AnalysisError{
  327. Type: ErrTableNotFound,
  328. Message: fmt.Sprintf("table not found: %s", stmt.Table.Name),
  329. }
  330. }
  331. // Validate column list if specified
  332. var targetCols []ColumnInfo
  333. if len(stmt.Columns) > 0 {
  334. for _, colName := range stmt.Columns {
  335. col, ok := table.GetColumn(colName)
  336. if !ok {
  337. return &AnalysisError{
  338. Type: ErrColumnNotFound,
  339. Message: fmt.Sprintf("column not found: %s", colName),
  340. }
  341. }
  342. targetCols = append(targetCols, *col)
  343. }
  344. } else {
  345. targetCols = table.Columns
  346. }
  347. // Validate VALUES
  348. for _, row := range stmt.Values {
  349. if len(row) != len(targetCols) {
  350. return &AnalysisError{
  351. Type: ErrTypeMismatch,
  352. Message: fmt.Sprintf("INSERT has %d columns but %d values", len(targetCols), len(row)),
  353. }
  354. }
  355. for i, expr := range row {
  356. info, err := a.analyzeExpr(expr)
  357. if err != nil {
  358. return err
  359. }
  360. // Check type compatibility
  361. if !info.Type.IsComparable(targetCols[i].Type) && info.Type != TypeNull {
  362. return &AnalysisError{
  363. Type: ErrTypeMismatch,
  364. Message: fmt.Sprintf("type mismatch for column %s: expected %s, got %s",
  365. targetCols[i].Name, targetCols[i].Type, info.Type),
  366. }
  367. }
  368. }
  369. }
  370. // Analyze INSERT ... SELECT
  371. if stmt.Select != nil {
  372. a.scope.DefineTable(&TableInfo{Name: table.Name, Columns: table.Columns, IsView: table.IsView})
  373. if err := a.analyzeSelect(stmt.Select); err != nil {
  374. return err
  375. }
  376. }
  377. return nil
  378. }
  379. // analyzeUpdate analyzes an UPDATE statement.
  380. func (a *Analyzer) analyzeUpdate(stmt *parser.UpdateStmt) error {
  381. table, ok := a.catalog.GetTable(stmt.Table.Name)
  382. if !ok {
  383. return &AnalysisError{
  384. Type: ErrTableNotFound,
  385. Message: fmt.Sprintf("table not found: %s", stmt.Table.Name),
  386. }
  387. }
  388. a.scope.DefineTable(&TableInfo{Name: table.Name, Columns: table.Columns, Alias: stmt.Table.Alias, IsView: table.IsView})
  389. // Validate SET assignments
  390. for _, assign := range stmt.Set {
  391. col, ok := table.GetColumn(assign.Column)
  392. if !ok {
  393. return &AnalysisError{
  394. Type: ErrColumnNotFound,
  395. Message: fmt.Sprintf("column not found: %s", assign.Column),
  396. }
  397. }
  398. info, err := a.analyzeExpr(assign.Value)
  399. if err != nil {
  400. return err
  401. }
  402. if !info.Type.IsComparable(col.Type) && info.Type != TypeNull {
  403. return &AnalysisError{
  404. Type: ErrTypeMismatch,
  405. Message: fmt.Sprintf("type mismatch for column %s: expected %s, got %s",
  406. col.Name, col.Type, info.Type),
  407. }
  408. }
  409. }
  410. // Analyze WHERE clause
  411. if stmt.Where != nil {
  412. info, err := a.analyzeExpr(stmt.Where)
  413. if err != nil {
  414. return err
  415. }
  416. if info.IsAggregate {
  417. return &AnalysisError{
  418. Type: ErrAggregateInWhere,
  419. Message: "aggregate functions not allowed in WHERE clause",
  420. }
  421. }
  422. }
  423. return nil
  424. }
  425. // analyzeDelete analyzes a DELETE statement.
  426. func (a *Analyzer) analyzeDelete(stmt *parser.DeleteStmt) error {
  427. table, ok := a.catalog.GetTable(stmt.Table.Name)
  428. if !ok {
  429. return &AnalysisError{
  430. Type: ErrTableNotFound,
  431. Message: fmt.Sprintf("table not found: %s", stmt.Table.Name),
  432. }
  433. }
  434. a.scope.DefineTable(&TableInfo{Name: table.Name, Columns: table.Columns, IsView: table.IsView})
  435. // Analyze WHERE clause
  436. if stmt.Where != nil {
  437. info, err := a.analyzeExpr(stmt.Where)
  438. if err != nil {
  439. return err
  440. }
  441. if info.IsAggregate {
  442. return &AnalysisError{
  443. Type: ErrAggregateInWhere,
  444. Message: "aggregate functions not allowed in WHERE clause",
  445. }
  446. }
  447. }
  448. return nil
  449. }
  450. // analyzeCreateTable analyzes a CREATE TABLE statement.
  451. func (a *Analyzer) analyzeCreateTable(stmt *parser.CreateTableStmt) error {
  452. // Check if table already exists
  453. if a.catalog.TableExists(stmt.Table.Name) {
  454. if stmt.IfNotExists {
  455. return nil // Silently succeed
  456. }
  457. return &AnalysisError{
  458. Type: ErrTableExists,
  459. Message: fmt.Sprintf("table already exists: %s", stmt.Table.Name),
  460. }
  461. }
  462. // Build table info
  463. tableInfo := &TableInfo{
  464. Name: stmt.Table.Name,
  465. }
  466. columnNames := make(map[string]bool)
  467. for _, colDef := range stmt.Columns {
  468. upperName := strings.ToUpper(colDef.Name)
  469. if columnNames[upperName] {
  470. return &AnalysisError{
  471. Type: ErrColumnAmbiguous,
  472. Message: fmt.Sprintf("duplicate column name: %s", colDef.Name),
  473. }
  474. }
  475. columnNames[upperName] = true
  476. colInfo := ColumnInfo{
  477. Name: colDef.Name,
  478. Type: TypeFromName(colDef.Type.Name),
  479. Nullable: true,
  480. TableName: stmt.Table.Name,
  481. }
  482. // Process constraints
  483. for _, constraint := range colDef.Constraints {
  484. switch constraint.Type {
  485. case parser.ConstraintPrimaryKey:
  486. colInfo.PrimaryKey = true
  487. colInfo.Nullable = false
  488. case parser.ConstraintNotNull:
  489. colInfo.Nullable = false
  490. case parser.ConstraintDefault:
  491. // Store default value (not evaluated here)
  492. colInfo.Default = constraint.Default
  493. }
  494. }
  495. tableInfo.Columns = append(tableInfo.Columns, colInfo)
  496. }
  497. // Process table-level constraints
  498. for _, constraint := range stmt.Constraints {
  499. switch constraint.Type {
  500. case parser.ConstraintPrimaryKey:
  501. for _, colName := range constraint.Columns {
  502. for i := range tableInfo.Columns {
  503. if strings.EqualFold(tableInfo.Columns[i].Name, colName) {
  504. tableInfo.Columns[i].PrimaryKey = true
  505. tableInfo.Columns[i].Nullable = false
  506. }
  507. }
  508. }
  509. }
  510. }
  511. // Add to catalog
  512. return a.catalog.CreateTable(tableInfo)
  513. }
  514. // analyzeDropTable analyzes a DROP TABLE statement.
  515. func (a *Analyzer) analyzeDropTable(stmt *parser.DropTableStmt) error {
  516. for _, tableRef := range stmt.Tables {
  517. if !a.catalog.TableExists(tableRef.Name) {
  518. if stmt.IfExists {
  519. continue // Silently succeed
  520. }
  521. return &AnalysisError{
  522. Type: ErrTableNotFound,
  523. Message: fmt.Sprintf("table not found: %s", tableRef.Name),
  524. }
  525. }
  526. if err := a.catalog.DropTable(tableRef.Name); err != nil {
  527. return err
  528. }
  529. }
  530. return nil
  531. }
  532. // analyzeExpr analyzes an expression and returns type information.
  533. func (a *Analyzer) analyzeExpr(expr parser.Expr) (*ExprInfo, error) {
  534. switch e := expr.(type) {
  535. case *parser.LiteralExpr:
  536. return a.analyzeLiteral(e)
  537. case *parser.ColumnRef:
  538. return a.analyzeColumnRef(e)
  539. case *parser.BinaryExpr:
  540. return a.analyzeBinaryExpr(e)
  541. case *parser.UnaryExpr:
  542. return a.analyzeUnaryExpr(e)
  543. case *parser.FunctionCall:
  544. return a.analyzeFunctionCall(e)
  545. case *parser.ParenExpr:
  546. return a.analyzeExpr(e.Expr)
  547. case *parser.CaseExpr:
  548. return a.analyzeCaseExpr(e)
  549. case *parser.CastExpr:
  550. return a.analyzeCastExpr(e)
  551. case *parser.InExpr:
  552. return a.analyzeInExpr(e)
  553. case *parser.BetweenExpr:
  554. return a.analyzeBetweenExpr(e)
  555. case *parser.LikeExpr:
  556. return a.analyzeLikeExpr(e)
  557. case *parser.IsNullExpr:
  558. return a.analyzeIsNullExpr(e)
  559. case *parser.ExistsExpr:
  560. return a.analyzeExistsExpr(e)
  561. case *parser.SubqueryExpr:
  562. return a.analyzeSubqueryExpr(e)
  563. default:
  564. return &ExprInfo{Type: TypeUnknown}, nil
  565. }
  566. }
  567. func (a *Analyzer) analyzeLiteral(e *parser.LiteralExpr) (*ExprInfo, error) {
  568. info := &ExprInfo{IsConstant: true}
  569. switch e.Type {
  570. case lexer.TokenNumber:
  571. if strings.Contains(e.Value, ".") || strings.Contains(strings.ToLower(e.Value), "e") {
  572. info.Type = TypeReal
  573. } else {
  574. info.Type = TypeInteger
  575. }
  576. case lexer.TokenString:
  577. info.Type = TypeText
  578. case lexer.TokenNULL:
  579. info.Type = TypeNull
  580. info.Nullable = true
  581. case lexer.TokenTRUE, lexer.TokenFALSE:
  582. info.Type = TypeBoolean
  583. case lexer.TokenStar:
  584. info.Type = TypeAny
  585. default:
  586. info.Type = TypeUnknown
  587. }
  588. return info, nil
  589. }
  590. func (a *Analyzer) analyzeColumnRef(e *parser.ColumnRef) (*ExprInfo, error) {
  591. col, _, ok := a.scope.LookupColumn(e.Table, e.Column)
  592. if !ok {
  593. // If no tables are in scope, treat as unknown (for standalone expressions)
  594. if len(a.scope.GetTables()) == 0 {
  595. return &ExprInfo{Type: TypeUnknown}, nil
  596. }
  597. return nil, &AnalysisError{
  598. Type: ErrColumnNotFound,
  599. Message: fmt.Sprintf("column not found: %s", formatColumnRef(e)),
  600. }
  601. }
  602. return &ExprInfo{
  603. Type: col.Type,
  604. Nullable: col.Nullable,
  605. }, nil
  606. }
  607. func formatColumnRef(e *parser.ColumnRef) string {
  608. if e.Table != "" {
  609. return e.Table + "." + e.Column
  610. }
  611. return e.Column
  612. }
  613. func (a *Analyzer) analyzeBinaryExpr(e *parser.BinaryExpr) (*ExprInfo, error) {
  614. left, err := a.analyzeExpr(e.Left)
  615. if err != nil {
  616. return nil, err
  617. }
  618. right, err := a.analyzeExpr(e.Right)
  619. if err != nil {
  620. return nil, err
  621. }
  622. info := &ExprInfo{
  623. IsAggregate: left.IsAggregate || right.IsAggregate,
  624. IsConstant: left.IsConstant && right.IsConstant,
  625. Nullable: left.Nullable || right.Nullable,
  626. }
  627. switch e.Op {
  628. case lexer.TokenPlus, lexer.TokenMinus, lexer.TokenStar, lexer.TokenSlash, lexer.TokenPercent:
  629. // Arithmetic operators
  630. info.Type = CommonType(left.Type, right.Type)
  631. if !left.Type.IsNumeric() && left.Type != TypeNull && left.Type != TypeUnknown {
  632. return nil, &AnalysisError{
  633. Type: ErrTypeMismatch,
  634. Message: fmt.Sprintf("arithmetic operator requires numeric type, got %s", left.Type),
  635. }
  636. }
  637. case lexer.TokenEq, lexer.TokenNeq, lexer.TokenLt, lexer.TokenLte, lexer.TokenGt, lexer.TokenGte:
  638. // Comparison operators
  639. info.Type = TypeBoolean
  640. if !left.Type.IsComparable(right.Type) {
  641. return nil, &AnalysisError{
  642. Type: ErrTypeMismatch,
  643. Message: fmt.Sprintf("cannot compare %s with %s", left.Type, right.Type),
  644. }
  645. }
  646. case lexer.TokenAND, lexer.TokenOR:
  647. // Logical operators
  648. info.Type = TypeBoolean
  649. case lexer.TokenConcat:
  650. // String concatenation
  651. info.Type = TypeText
  652. default:
  653. info.Type = TypeUnknown
  654. }
  655. return info, nil
  656. }
  657. func (a *Analyzer) analyzeUnaryExpr(e *parser.UnaryExpr) (*ExprInfo, error) {
  658. operand, err := a.analyzeExpr(e.Operand)
  659. if err != nil {
  660. return nil, err
  661. }
  662. info := &ExprInfo{
  663. IsAggregate: operand.IsAggregate,
  664. IsConstant: operand.IsConstant,
  665. Nullable: operand.Nullable,
  666. }
  667. switch e.Op {
  668. case lexer.TokenPlus:
  669. // Unary + is a no-op in SQLite — passes any type through unchanged.
  670. info.Type = operand.Type
  671. case lexer.TokenMinus:
  672. info.Type = operand.Type
  673. if !operand.Type.IsNumeric() && operand.Type != TypeNull && operand.Type != TypeUnknown {
  674. return nil, &AnalysisError{
  675. Type: ErrTypeMismatch,
  676. Message: fmt.Sprintf("unary - requires numeric type, got %s", operand.Type),
  677. }
  678. }
  679. case lexer.TokenNOT:
  680. info.Type = TypeBoolean
  681. default:
  682. info.Type = operand.Type
  683. }
  684. return info, nil
  685. }
  686. func (a *Analyzer) analyzeFunctionCall(e *parser.FunctionCall) (*ExprInfo, error) {
  687. sig, ok := LookupFunction(e.Name)
  688. if !ok {
  689. return nil, &AnalysisError{
  690. Type: ErrInvalidFunction,
  691. Message: fmt.Sprintf("unknown function: %s", e.Name),
  692. }
  693. }
  694. // Handle COUNT(*)
  695. argCount := len(e.Args)
  696. if e.Star {
  697. argCount = 0 // COUNT(*) has 0 real args
  698. }
  699. // Check argument count
  700. if argCount < sig.MinArgs {
  701. return nil, &AnalysisError{
  702. Type: ErrInvalidArgCount,
  703. Message: fmt.Sprintf("function %s requires at least %d arguments, got %d", e.Name, sig.MinArgs, argCount),
  704. }
  705. }
  706. if sig.MaxArgs >= 0 && argCount > sig.MaxArgs {
  707. return nil, &AnalysisError{
  708. Type: ErrInvalidArgCount,
  709. Message: fmt.Sprintf("function %s accepts at most %d arguments, got %d", e.Name, sig.MaxArgs, argCount),
  710. }
  711. }
  712. // Analyze arguments
  713. info := &ExprInfo{
  714. Type: sig.ReturnType,
  715. IsAggregate: sig.IsAggregate,
  716. }
  717. for _, arg := range e.Args {
  718. argInfo, err := a.analyzeExpr(arg)
  719. if err != nil {
  720. return nil, err
  721. }
  722. if argInfo.Nullable {
  723. info.Nullable = true
  724. }
  725. // Propagate aggregate status from arguments
  726. if argInfo.IsAggregate && !sig.IsAggregate {
  727. info.IsAggregate = true
  728. }
  729. }
  730. // Special case: MIN/MAX/COALESCE return type depends on argument
  731. if sig.ReturnType == TypeAny && len(e.Args) > 0 {
  732. argInfo, _ := a.analyzeExpr(e.Args[0])
  733. if argInfo != nil {
  734. info.Type = argInfo.Type
  735. }
  736. }
  737. return info, nil
  738. }
  739. func (a *Analyzer) analyzeCaseExpr(e *parser.CaseExpr) (*ExprInfo, error) {
  740. info := &ExprInfo{
  741. Nullable: true, // CASE can return NULL
  742. }
  743. // Analyze operand if present (simple CASE)
  744. if e.Operand != nil {
  745. opInfo, err := a.analyzeExpr(e.Operand)
  746. if err != nil {
  747. return nil, err
  748. }
  749. if opInfo.IsAggregate {
  750. info.IsAggregate = true
  751. }
  752. }
  753. // Analyze WHEN clauses
  754. var resultType Type
  755. for _, when := range e.Whens {
  756. condInfo, err := a.analyzeExpr(when.Condition)
  757. if err != nil {
  758. return nil, err
  759. }
  760. if condInfo.IsAggregate {
  761. info.IsAggregate = true
  762. }
  763. resInfo, err := a.analyzeExpr(when.Result)
  764. if err != nil {
  765. return nil, err
  766. }
  767. if resInfo.IsAggregate {
  768. info.IsAggregate = true
  769. }
  770. if resultType == TypeUnknown {
  771. resultType = resInfo.Type
  772. } else {
  773. resultType = CommonType(resultType, resInfo.Type)
  774. }
  775. }
  776. // Analyze ELSE clause
  777. if e.Else != nil {
  778. elseInfo, err := a.analyzeExpr(e.Else)
  779. if err != nil {
  780. return nil, err
  781. }
  782. if elseInfo.IsAggregate {
  783. info.IsAggregate = true
  784. }
  785. resultType = CommonType(resultType, elseInfo.Type)
  786. }
  787. info.Type = resultType
  788. return info, nil
  789. }
  790. func (a *Analyzer) analyzeCastExpr(e *parser.CastExpr) (*ExprInfo, error) {
  791. exprInfo, err := a.analyzeExpr(e.Expr)
  792. if err != nil {
  793. return nil, err
  794. }
  795. return &ExprInfo{
  796. Type: TypeFromName(e.Type.Name),
  797. IsAggregate: exprInfo.IsAggregate,
  798. IsConstant: exprInfo.IsConstant,
  799. Nullable: exprInfo.Nullable,
  800. }, nil
  801. }
  802. func (a *Analyzer) analyzeInExpr(e *parser.InExpr) (*ExprInfo, error) {
  803. leftInfo, err := a.analyzeExpr(e.Left)
  804. if err != nil {
  805. return nil, err
  806. }
  807. info := &ExprInfo{
  808. Type: TypeBoolean,
  809. IsAggregate: leftInfo.IsAggregate,
  810. }
  811. // Analyze value list
  812. for _, val := range e.Values {
  813. valInfo, err := a.analyzeExpr(val)
  814. if err != nil {
  815. return nil, err
  816. }
  817. if valInfo.IsAggregate {
  818. info.IsAggregate = true
  819. }
  820. }
  821. // Analyze subquery
  822. if e.Subquery != nil {
  823. subScope := NewScope(a.scope)
  824. oldScope := a.scope
  825. a.scope = subScope
  826. err := a.analyzeSelect(e.Subquery)
  827. a.scope = oldScope
  828. if err != nil {
  829. return nil, err
  830. }
  831. }
  832. return info, nil
  833. }
  834. func (a *Analyzer) analyzeBetweenExpr(e *parser.BetweenExpr) (*ExprInfo, error) {
  835. leftInfo, err := a.analyzeExpr(e.Left)
  836. if err != nil {
  837. return nil, err
  838. }
  839. lowInfo, err := a.analyzeExpr(e.Low)
  840. if err != nil {
  841. return nil, err
  842. }
  843. highInfo, err := a.analyzeExpr(e.High)
  844. if err != nil {
  845. return nil, err
  846. }
  847. return &ExprInfo{
  848. Type: TypeBoolean,
  849. IsAggregate: leftInfo.IsAggregate || lowInfo.IsAggregate || highInfo.IsAggregate,
  850. Nullable: leftInfo.Nullable || lowInfo.Nullable || highInfo.Nullable,
  851. }, nil
  852. }
  853. func (a *Analyzer) analyzeLikeExpr(e *parser.LikeExpr) (*ExprInfo, error) {
  854. leftInfo, err := a.analyzeExpr(e.Left)
  855. if err != nil {
  856. return nil, err
  857. }
  858. patternInfo, err := a.analyzeExpr(e.Pattern)
  859. if err != nil {
  860. return nil, err
  861. }
  862. info := &ExprInfo{
  863. Type: TypeBoolean,
  864. IsAggregate: leftInfo.IsAggregate || patternInfo.IsAggregate,
  865. Nullable: leftInfo.Nullable || patternInfo.Nullable,
  866. }
  867. if e.Escape != nil {
  868. escInfo, err := a.analyzeExpr(e.Escape)
  869. if err != nil {
  870. return nil, err
  871. }
  872. if escInfo.IsAggregate {
  873. info.IsAggregate = true
  874. }
  875. }
  876. return info, nil
  877. }
  878. func (a *Analyzer) analyzeIsNullExpr(e *parser.IsNullExpr) (*ExprInfo, error) {
  879. leftInfo, err := a.analyzeExpr(e.Left)
  880. if err != nil {
  881. return nil, err
  882. }
  883. return &ExprInfo{
  884. Type: TypeBoolean,
  885. IsAggregate: leftInfo.IsAggregate,
  886. IsConstant: leftInfo.IsConstant,
  887. }, nil
  888. }
  889. func (a *Analyzer) analyzeExistsExpr(e *parser.ExistsExpr) (*ExprInfo, error) {
  890. // Analyze subquery in its own scope
  891. subScope := NewScope(a.scope)
  892. oldScope := a.scope
  893. a.scope = subScope
  894. err := a.analyzeSelect(e.Subquery)
  895. a.scope = oldScope
  896. if err != nil {
  897. return nil, err
  898. }
  899. return &ExprInfo{
  900. Type: TypeBoolean,
  901. }, nil
  902. }
  903. func (a *Analyzer) analyzeSubqueryExpr(e *parser.SubqueryExpr) (*ExprInfo, error) {
  904. // Analyze subquery in its own scope
  905. subScope := NewScope(a.scope)
  906. oldScope := a.scope
  907. a.scope = subScope
  908. err := a.analyzeSelect(e.Query)
  909. a.scope = oldScope
  910. if err != nil {
  911. return nil, err
  912. }
  913. // Scalar subquery - return type of first column
  914. // For simplicity, return TypeAny
  915. return &ExprInfo{
  916. Type: TypeAny,
  917. }, nil
  918. }