analyzer.go 24 KB

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