2
0

analyzer.go 26 KB

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