2
0

analyzer.go 24 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021102210231024102510261027102810291030103110321033103410351036
  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})
  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})
  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})
  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.TokenMinus, lexer.TokenPlus:
  669. info.Type = operand.Type
  670. if !operand.Type.IsNumeric() && operand.Type != TypeNull && operand.Type != TypeUnknown {
  671. return nil, &AnalysisError{
  672. Type: ErrTypeMismatch,
  673. Message: fmt.Sprintf("unary %s requires numeric type, got %s", e.Op, operand.Type),
  674. }
  675. }
  676. case lexer.TokenNOT:
  677. info.Type = TypeBoolean
  678. default:
  679. info.Type = operand.Type
  680. }
  681. return info, nil
  682. }
  683. func (a *Analyzer) analyzeFunctionCall(e *parser.FunctionCall) (*ExprInfo, error) {
  684. sig, ok := LookupFunction(e.Name)
  685. if !ok {
  686. return nil, &AnalysisError{
  687. Type: ErrInvalidFunction,
  688. Message: fmt.Sprintf("unknown function: %s", e.Name),
  689. }
  690. }
  691. // Handle COUNT(*)
  692. argCount := len(e.Args)
  693. if e.Star {
  694. argCount = 0 // COUNT(*) has 0 real args
  695. }
  696. // Check argument count
  697. if argCount < sig.MinArgs {
  698. return nil, &AnalysisError{
  699. Type: ErrInvalidArgCount,
  700. Message: fmt.Sprintf("function %s requires at least %d arguments, got %d", e.Name, sig.MinArgs, argCount),
  701. }
  702. }
  703. if sig.MaxArgs >= 0 && argCount > sig.MaxArgs {
  704. return nil, &AnalysisError{
  705. Type: ErrInvalidArgCount,
  706. Message: fmt.Sprintf("function %s accepts at most %d arguments, got %d", e.Name, sig.MaxArgs, argCount),
  707. }
  708. }
  709. // Analyze arguments
  710. info := &ExprInfo{
  711. Type: sig.ReturnType,
  712. IsAggregate: sig.IsAggregate,
  713. }
  714. for _, arg := range e.Args {
  715. argInfo, err := a.analyzeExpr(arg)
  716. if err != nil {
  717. return nil, err
  718. }
  719. if argInfo.Nullable {
  720. info.Nullable = true
  721. }
  722. // Propagate aggregate status from arguments
  723. if argInfo.IsAggregate && !sig.IsAggregate {
  724. info.IsAggregate = true
  725. }
  726. }
  727. // Special case: MIN/MAX/COALESCE return type depends on argument
  728. if sig.ReturnType == TypeAny && len(e.Args) > 0 {
  729. argInfo, _ := a.analyzeExpr(e.Args[0])
  730. if argInfo != nil {
  731. info.Type = argInfo.Type
  732. }
  733. }
  734. return info, nil
  735. }
  736. func (a *Analyzer) analyzeCaseExpr(e *parser.CaseExpr) (*ExprInfo, error) {
  737. info := &ExprInfo{
  738. Nullable: true, // CASE can return NULL
  739. }
  740. // Analyze operand if present (simple CASE)
  741. if e.Operand != nil {
  742. opInfo, err := a.analyzeExpr(e.Operand)
  743. if err != nil {
  744. return nil, err
  745. }
  746. if opInfo.IsAggregate {
  747. info.IsAggregate = true
  748. }
  749. }
  750. // Analyze WHEN clauses
  751. var resultType Type
  752. for _, when := range e.Whens {
  753. condInfo, err := a.analyzeExpr(when.Condition)
  754. if err != nil {
  755. return nil, err
  756. }
  757. if condInfo.IsAggregate {
  758. info.IsAggregate = true
  759. }
  760. resInfo, err := a.analyzeExpr(when.Result)
  761. if err != nil {
  762. return nil, err
  763. }
  764. if resInfo.IsAggregate {
  765. info.IsAggregate = true
  766. }
  767. if resultType == TypeUnknown {
  768. resultType = resInfo.Type
  769. } else {
  770. resultType = CommonType(resultType, resInfo.Type)
  771. }
  772. }
  773. // Analyze ELSE clause
  774. if e.Else != nil {
  775. elseInfo, err := a.analyzeExpr(e.Else)
  776. if err != nil {
  777. return nil, err
  778. }
  779. if elseInfo.IsAggregate {
  780. info.IsAggregate = true
  781. }
  782. resultType = CommonType(resultType, elseInfo.Type)
  783. }
  784. info.Type = resultType
  785. return info, nil
  786. }
  787. func (a *Analyzer) analyzeCastExpr(e *parser.CastExpr) (*ExprInfo, error) {
  788. exprInfo, err := a.analyzeExpr(e.Expr)
  789. if err != nil {
  790. return nil, err
  791. }
  792. return &ExprInfo{
  793. Type: TypeFromName(e.Type.Name),
  794. IsAggregate: exprInfo.IsAggregate,
  795. IsConstant: exprInfo.IsConstant,
  796. Nullable: exprInfo.Nullable,
  797. }, nil
  798. }
  799. func (a *Analyzer) analyzeInExpr(e *parser.InExpr) (*ExprInfo, error) {
  800. leftInfo, err := a.analyzeExpr(e.Left)
  801. if err != nil {
  802. return nil, err
  803. }
  804. info := &ExprInfo{
  805. Type: TypeBoolean,
  806. IsAggregate: leftInfo.IsAggregate,
  807. }
  808. // Analyze value list
  809. for _, val := range e.Values {
  810. valInfo, err := a.analyzeExpr(val)
  811. if err != nil {
  812. return nil, err
  813. }
  814. if valInfo.IsAggregate {
  815. info.IsAggregate = true
  816. }
  817. }
  818. // Analyze subquery
  819. if e.Subquery != nil {
  820. subScope := NewScope(a.scope)
  821. oldScope := a.scope
  822. a.scope = subScope
  823. err := a.analyzeSelect(e.Subquery)
  824. a.scope = oldScope
  825. if err != nil {
  826. return nil, err
  827. }
  828. }
  829. return info, nil
  830. }
  831. func (a *Analyzer) analyzeBetweenExpr(e *parser.BetweenExpr) (*ExprInfo, error) {
  832. leftInfo, err := a.analyzeExpr(e.Left)
  833. if err != nil {
  834. return nil, err
  835. }
  836. lowInfo, err := a.analyzeExpr(e.Low)
  837. if err != nil {
  838. return nil, err
  839. }
  840. highInfo, err := a.analyzeExpr(e.High)
  841. if err != nil {
  842. return nil, err
  843. }
  844. return &ExprInfo{
  845. Type: TypeBoolean,
  846. IsAggregate: leftInfo.IsAggregate || lowInfo.IsAggregate || highInfo.IsAggregate,
  847. Nullable: leftInfo.Nullable || lowInfo.Nullable || highInfo.Nullable,
  848. }, nil
  849. }
  850. func (a *Analyzer) analyzeLikeExpr(e *parser.LikeExpr) (*ExprInfo, error) {
  851. leftInfo, err := a.analyzeExpr(e.Left)
  852. if err != nil {
  853. return nil, err
  854. }
  855. patternInfo, err := a.analyzeExpr(e.Pattern)
  856. if err != nil {
  857. return nil, err
  858. }
  859. info := &ExprInfo{
  860. Type: TypeBoolean,
  861. IsAggregate: leftInfo.IsAggregate || patternInfo.IsAggregate,
  862. Nullable: leftInfo.Nullable || patternInfo.Nullable,
  863. }
  864. if e.Escape != nil {
  865. escInfo, err := a.analyzeExpr(e.Escape)
  866. if err != nil {
  867. return nil, err
  868. }
  869. if escInfo.IsAggregate {
  870. info.IsAggregate = true
  871. }
  872. }
  873. return info, nil
  874. }
  875. func (a *Analyzer) analyzeIsNullExpr(e *parser.IsNullExpr) (*ExprInfo, error) {
  876. leftInfo, err := a.analyzeExpr(e.Left)
  877. if err != nil {
  878. return nil, err
  879. }
  880. return &ExprInfo{
  881. Type: TypeBoolean,
  882. IsAggregate: leftInfo.IsAggregate,
  883. IsConstant: leftInfo.IsConstant,
  884. }, nil
  885. }
  886. func (a *Analyzer) analyzeExistsExpr(e *parser.ExistsExpr) (*ExprInfo, error) {
  887. // Analyze subquery in its own scope
  888. subScope := NewScope(a.scope)
  889. oldScope := a.scope
  890. a.scope = subScope
  891. err := a.analyzeSelect(e.Subquery)
  892. a.scope = oldScope
  893. if err != nil {
  894. return nil, err
  895. }
  896. return &ExprInfo{
  897. Type: TypeBoolean,
  898. }, nil
  899. }
  900. func (a *Analyzer) analyzeSubqueryExpr(e *parser.SubqueryExpr) (*ExprInfo, error) {
  901. // Analyze subquery in its own scope
  902. subScope := NewScope(a.scope)
  903. oldScope := a.scope
  904. a.scope = subScope
  905. err := a.analyzeSelect(e.Query)
  906. a.scope = oldScope
  907. if err != nil {
  908. return nil, err
  909. }
  910. // Scalar subquery - return type of first column
  911. // For simplicity, return TypeAny
  912. return &ExprInfo{
  913. Type: TypeAny,
  914. }, nil
  915. }