analyzer_test.go 20 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786
  1. package analyzer
  2. import (
  3. "testing"
  4. "github.com/danfragoso/pizzasql-next/pkg/lexer"
  5. "github.com/danfragoso/pizzasql-next/pkg/parser"
  6. )
  7. func parse(t *testing.T, sql string) parser.Statement {
  8. t.Helper()
  9. l := lexer.New(sql)
  10. p := parser.New(l)
  11. stmt, err := p.Parse()
  12. if err != nil {
  13. t.Fatalf("parse error: %v", err)
  14. }
  15. return stmt
  16. }
  17. func setupCatalog() *Catalog {
  18. catalog := NewCatalog()
  19. // Create users table
  20. catalog.CreateTable(&TableInfo{
  21. Name: "users",
  22. Columns: []ColumnInfo{
  23. {Name: "id", Type: TypeInteger, PrimaryKey: true},
  24. {Name: "name", Type: TypeText, Nullable: false},
  25. {Name: "email", Type: TypeText, Nullable: true},
  26. {Name: "age", Type: TypeInteger, Nullable: true},
  27. {Name: "active", Type: TypeBoolean, Nullable: false},
  28. {Name: "balance", Type: TypeReal, Nullable: true},
  29. },
  30. })
  31. // Create orders table
  32. catalog.CreateTable(&TableInfo{
  33. Name: "orders",
  34. Columns: []ColumnInfo{
  35. {Name: "id", Type: TypeInteger, PrimaryKey: true},
  36. {Name: "user_id", Type: TypeInteger, Nullable: false},
  37. {Name: "amount", Type: TypeReal, Nullable: false},
  38. {Name: "status", Type: TypeText, Nullable: false},
  39. {Name: "created_at", Type: TypeText, Nullable: false},
  40. },
  41. })
  42. // Create products table
  43. catalog.CreateTable(&TableInfo{
  44. Name: "products",
  45. Columns: []ColumnInfo{
  46. {Name: "id", Type: TypeInteger, PrimaryKey: true},
  47. {Name: "name", Type: TypeText, Nullable: false},
  48. {Name: "price", Type: TypeReal, Nullable: false},
  49. {Name: "stock", Type: TypeInteger, Nullable: false},
  50. },
  51. })
  52. return catalog
  53. }
  54. // Type system tests
  55. func TestTypeFromName(t *testing.T) {
  56. tests := []struct {
  57. name string
  58. expected Type
  59. }{
  60. {"INTEGER", TypeInteger},
  61. {"INT", TypeInteger},
  62. {"SMALLINT", TypeInteger},
  63. {"BIGINT", TypeInteger},
  64. {"TINYINT", TypeInteger},
  65. {"REAL", TypeReal},
  66. {"FLOAT", TypeReal},
  67. {"DOUBLE", TypeReal},
  68. {"TEXT", TypeText},
  69. {"VARCHAR", TypeText},
  70. {"CHAR", TypeText},
  71. {"CHARACTER", TypeText},
  72. {"CLOB", TypeText},
  73. {"BLOB", TypeBlob},
  74. {"BOOLEAN", TypeBoolean},
  75. {"NUMERIC", TypeNumeric},
  76. {"DECIMAL", TypeNumeric},
  77. {"UUID", TypeText},
  78. {"", TypeBlob}, // Empty type -> BLOB (SQLite rule)
  79. }
  80. for _, tt := range tests {
  81. t.Run(tt.name, func(t *testing.T) {
  82. got := TypeFromName(tt.name)
  83. if got != tt.expected {
  84. t.Errorf("TypeFromName(%q) = %v, want %v", tt.name, got, tt.expected)
  85. }
  86. })
  87. }
  88. }
  89. func TestTypeComparable(t *testing.T) {
  90. tests := []struct {
  91. a, b Type
  92. expected bool
  93. }{
  94. {TypeInteger, TypeInteger, true},
  95. {TypeInteger, TypeReal, true},
  96. {TypeInteger, TypeNumeric, true},
  97. {TypeReal, TypeNumeric, true},
  98. {TypeText, TypeText, true},
  99. {TypeText, TypeBlob, true},
  100. {TypeNull, TypeInteger, true},
  101. {TypeNull, TypeText, true},
  102. {TypeAny, TypeInteger, true},
  103. {TypeInteger, TypeText, false},
  104. {TypeReal, TypeBlob, false},
  105. }
  106. for _, tt := range tests {
  107. t.Run(tt.a.String()+"_"+tt.b.String(), func(t *testing.T) {
  108. got := tt.a.IsComparable(tt.b)
  109. if got != tt.expected {
  110. t.Errorf("%v.IsComparable(%v) = %v, want %v", tt.a, tt.b, got, tt.expected)
  111. }
  112. })
  113. }
  114. }
  115. func TestCommonType(t *testing.T) {
  116. tests := []struct {
  117. a, b Type
  118. expected Type
  119. }{
  120. {TypeInteger, TypeInteger, TypeInteger},
  121. {TypeInteger, TypeReal, TypeReal},
  122. {TypeReal, TypeInteger, TypeReal},
  123. {TypeInteger, TypeNumeric, TypeNumeric},
  124. {TypeNull, TypeInteger, TypeInteger},
  125. {TypeText, TypeText, TypeText},
  126. {TypeText, TypeBlob, TypeText},
  127. }
  128. for _, tt := range tests {
  129. t.Run(tt.a.String()+"_"+tt.b.String(), func(t *testing.T) {
  130. got := CommonType(tt.a, tt.b)
  131. if got != tt.expected {
  132. t.Errorf("CommonType(%v, %v) = %v, want %v", tt.a, tt.b, got, tt.expected)
  133. }
  134. })
  135. }
  136. }
  137. // Function lookup tests
  138. func TestLookupFunction(t *testing.T) {
  139. tests := []struct {
  140. name string
  141. exists bool
  142. isAggregate bool
  143. }{
  144. {"COUNT", true, true},
  145. {"SUM", true, true},
  146. {"AVG", true, true},
  147. {"MIN", true, true},
  148. {"MAX", true, true},
  149. {"UPPER", true, false},
  150. {"LOWER", true, false},
  151. {"LENGTH", true, false},
  152. {"COALESCE", true, false},
  153. {"UNKNOWN_FUNC", false, false},
  154. }
  155. for _, tt := range tests {
  156. t.Run(tt.name, func(t *testing.T) {
  157. sig, ok := LookupFunction(tt.name)
  158. if ok != tt.exists {
  159. t.Errorf("LookupFunction(%q) exists = %v, want %v", tt.name, ok, tt.exists)
  160. }
  161. if ok && sig.IsAggregate != tt.isAggregate {
  162. t.Errorf("LookupFunction(%q).IsAggregate = %v, want %v", tt.name, sig.IsAggregate, tt.isAggregate)
  163. }
  164. })
  165. }
  166. }
  167. // SELECT analysis tests
  168. func TestAnalyzeSelectBasic(t *testing.T) {
  169. catalog := setupCatalog()
  170. analyzer := New(catalog)
  171. tests := []struct {
  172. name string
  173. sql string
  174. }{
  175. {"select star", "SELECT * FROM users"},
  176. {"select columns", "SELECT id, name FROM users"},
  177. {"select with alias", "SELECT id AS user_id, name AS full_name FROM users"},
  178. {"select with where", "SELECT * FROM users WHERE id = 1"},
  179. {"select with complex where", "SELECT * FROM users WHERE id = 1 AND name = 'John'"},
  180. {"select with order by", "SELECT * FROM users ORDER BY name ASC"},
  181. {"select with limit", "SELECT * FROM users LIMIT 10"},
  182. {"select with limit offset", "SELECT * FROM users LIMIT 10 OFFSET 5"},
  183. {"select distinct", "SELECT DISTINCT name FROM users"},
  184. }
  185. for _, tt := range tests {
  186. t.Run(tt.name, func(t *testing.T) {
  187. stmt := parse(t, tt.sql)
  188. err := analyzer.Analyze(stmt)
  189. if err != nil {
  190. t.Errorf("Analyze(%q) error: %v", tt.sql, err)
  191. }
  192. })
  193. }
  194. }
  195. func TestAnalyzeSelectJoin(t *testing.T) {
  196. catalog := setupCatalog()
  197. analyzer := New(catalog)
  198. tests := []struct {
  199. name string
  200. sql string
  201. }{
  202. {"inner join", "SELECT * FROM users JOIN orders ON users.id = orders.user_id"},
  203. {"left join", "SELECT * FROM users LEFT JOIN orders ON users.id = orders.user_id"},
  204. {"join with alias", "SELECT u.id, o.amount FROM users u JOIN orders o ON u.id = o.user_id"},
  205. {"multiple joins", "SELECT * FROM users u JOIN orders o ON u.id = o.user_id JOIN products p ON p.id = 1"},
  206. }
  207. for _, tt := range tests {
  208. t.Run(tt.name, func(t *testing.T) {
  209. stmt := parse(t, tt.sql)
  210. err := analyzer.Analyze(stmt)
  211. if err != nil {
  212. t.Errorf("Analyze(%q) error: %v", tt.sql, err)
  213. }
  214. })
  215. }
  216. }
  217. func TestAnalyzeSelectQualifiedWildcard(t *testing.T) {
  218. catalog := setupCatalog()
  219. analyzer := New(catalog)
  220. valid := []struct {
  221. name string
  222. sql string
  223. }{
  224. {"alias wildcard", "SELECT u.* FROM users u"},
  225. {"table name wildcard", "SELECT users.* FROM users"},
  226. {"join alias wildcard", "SELECT u.* FROM users u JOIN orders o ON u.id = o.user_id"},
  227. {"mixed wildcard and column", "SELECT u.*, o.amount FROM users u JOIN orders o ON u.id = o.user_id"},
  228. }
  229. for _, tt := range valid {
  230. t.Run(tt.name, func(t *testing.T) {
  231. err := analyzer.Analyze(parse(t, tt.sql))
  232. if err != nil {
  233. t.Errorf("Analyze(%q) error: %v", tt.sql, err)
  234. }
  235. })
  236. }
  237. invalid := []struct {
  238. name string
  239. sql string
  240. errType ErrorType
  241. }{
  242. {"unknown qualifier", "SELECT nope.* FROM users u", ErrTableNotFound},
  243. {"aggregate without group by", "SELECT u.*, COUNT(*) FROM users u", ErrNonAggregateInSelect},
  244. }
  245. for _, tt := range invalid {
  246. t.Run(tt.name, func(t *testing.T) {
  247. err := analyzer.Analyze(parse(t, tt.sql))
  248. if err == nil {
  249. t.Fatalf("Analyze(%q) expected error, got nil", tt.sql)
  250. }
  251. analysisErr, ok := err.(*AnalysisError)
  252. if !ok {
  253. t.Fatalf("expected AnalysisError, got %T", err)
  254. }
  255. if analysisErr.Type != tt.errType {
  256. t.Errorf("expected error type %v, got %v (%s)", tt.errType, analysisErr.Type, analysisErr.Message)
  257. }
  258. })
  259. }
  260. }
  261. func TestAnalyzeSelectAggregate(t *testing.T) {
  262. catalog := setupCatalog()
  263. analyzer := New(catalog)
  264. tests := []struct {
  265. name string
  266. sql string
  267. }{
  268. {"count star", "SELECT COUNT(*) FROM users"},
  269. {"count column", "SELECT COUNT(id) FROM users"},
  270. {"sum", "SELECT SUM(age) FROM users"},
  271. {"avg", "SELECT AVG(balance) FROM users"},
  272. {"min max", "SELECT MIN(age), MAX(age) FROM users"},
  273. {"group by", "SELECT name, COUNT(*) FROM users GROUP BY name"},
  274. {"group by having", "SELECT name, COUNT(*) FROM users GROUP BY name HAVING COUNT(*) > 1"},
  275. }
  276. for _, tt := range tests {
  277. t.Run(tt.name, func(t *testing.T) {
  278. stmt := parse(t, tt.sql)
  279. err := analyzer.Analyze(stmt)
  280. if err != nil {
  281. t.Errorf("Analyze(%q) error: %v", tt.sql, err)
  282. }
  283. })
  284. }
  285. }
  286. func TestAnalyzeSelectErrors(t *testing.T) {
  287. catalog := setupCatalog()
  288. analyzer := New(catalog)
  289. tests := []struct {
  290. name string
  291. sql string
  292. errType ErrorType
  293. }{
  294. {"table not found", "SELECT * FROM nonexistent", ErrTableNotFound},
  295. {"column not found", "SELECT nonexistent FROM users", ErrColumnNotFound},
  296. {"aggregate in where", "SELECT * FROM users WHERE COUNT(*) > 0", ErrAggregateInWhere},
  297. }
  298. for _, tt := range tests {
  299. t.Run(tt.name, func(t *testing.T) {
  300. stmt := parse(t, tt.sql)
  301. err := analyzer.Analyze(stmt)
  302. if err == nil {
  303. t.Errorf("Analyze(%q) expected error, got nil", tt.sql)
  304. return
  305. }
  306. if ae, ok := err.(*AnalysisError); ok {
  307. if ae.Type != tt.errType {
  308. t.Errorf("Analyze(%q) error type = %v, want %v", tt.sql, ae.Type, tt.errType)
  309. }
  310. }
  311. })
  312. }
  313. }
  314. // INSERT analysis tests
  315. func TestAnalyzeInsert(t *testing.T) {
  316. catalog := setupCatalog()
  317. analyzer := New(catalog)
  318. tests := []struct {
  319. name string
  320. sql string
  321. }{
  322. {"insert all columns", "INSERT INTO users VALUES (1, 'John', 'john@example.com', 30, TRUE, 100.50)"},
  323. {"insert with columns", "INSERT INTO users (id, name, active) VALUES (1, 'John', TRUE)"},
  324. {"insert multiple rows", "INSERT INTO users (id, name, active) VALUES (1, 'John', TRUE), (2, 'Jane', FALSE)"},
  325. }
  326. for _, tt := range tests {
  327. t.Run(tt.name, func(t *testing.T) {
  328. stmt := parse(t, tt.sql)
  329. err := analyzer.Analyze(stmt)
  330. if err != nil {
  331. t.Errorf("Analyze(%q) error: %v", tt.sql, err)
  332. }
  333. })
  334. }
  335. }
  336. func TestAnalyzeInsertErrors(t *testing.T) {
  337. catalog := setupCatalog()
  338. analyzer := New(catalog)
  339. tests := []struct {
  340. name string
  341. sql string
  342. errType ErrorType
  343. }{
  344. {"table not found", "INSERT INTO nonexistent VALUES (1)", ErrTableNotFound},
  345. {"column not found", "INSERT INTO users (nonexistent) VALUES (1)", ErrColumnNotFound},
  346. {"wrong column count", "INSERT INTO users (id, name) VALUES (1)", ErrTypeMismatch},
  347. }
  348. for _, tt := range tests {
  349. t.Run(tt.name, func(t *testing.T) {
  350. stmt := parse(t, tt.sql)
  351. err := analyzer.Analyze(stmt)
  352. if err == nil {
  353. t.Errorf("Analyze(%q) expected error, got nil", tt.sql)
  354. return
  355. }
  356. if ae, ok := err.(*AnalysisError); ok {
  357. if ae.Type != tt.errType {
  358. t.Errorf("Analyze(%q) error type = %v, want %v", tt.sql, ae.Type, tt.errType)
  359. }
  360. }
  361. })
  362. }
  363. }
  364. // UPDATE analysis tests
  365. func TestAnalyzeUpdate(t *testing.T) {
  366. catalog := setupCatalog()
  367. analyzer := New(catalog)
  368. tests := []struct {
  369. name string
  370. sql string
  371. }{
  372. {"update single column", "UPDATE users SET name = 'John' WHERE id = 1"},
  373. {"update multiple columns", "UPDATE users SET name = 'John', age = 30 WHERE id = 1"},
  374. {"update with expression", "UPDATE users SET age = age + 1 WHERE active = TRUE"},
  375. }
  376. for _, tt := range tests {
  377. t.Run(tt.name, func(t *testing.T) {
  378. stmt := parse(t, tt.sql)
  379. err := analyzer.Analyze(stmt)
  380. if err != nil {
  381. t.Errorf("Analyze(%q) error: %v", tt.sql, err)
  382. }
  383. })
  384. }
  385. }
  386. func TestAnalyzeUpdateErrors(t *testing.T) {
  387. catalog := setupCatalog()
  388. analyzer := New(catalog)
  389. tests := []struct {
  390. name string
  391. sql string
  392. errType ErrorType
  393. }{
  394. {"table not found", "UPDATE nonexistent SET x = 1", ErrTableNotFound},
  395. {"column not found", "UPDATE users SET nonexistent = 1", ErrColumnNotFound},
  396. }
  397. for _, tt := range tests {
  398. t.Run(tt.name, func(t *testing.T) {
  399. stmt := parse(t, tt.sql)
  400. err := analyzer.Analyze(stmt)
  401. if err == nil {
  402. t.Errorf("Analyze(%q) expected error, got nil", tt.sql)
  403. return
  404. }
  405. if ae, ok := err.(*AnalysisError); ok {
  406. if ae.Type != tt.errType {
  407. t.Errorf("Analyze(%q) error type = %v, want %v", tt.sql, ae.Type, tt.errType)
  408. }
  409. }
  410. })
  411. }
  412. }
  413. // DELETE analysis tests
  414. func TestAnalyzeDelete(t *testing.T) {
  415. catalog := setupCatalog()
  416. analyzer := New(catalog)
  417. tests := []struct {
  418. name string
  419. sql string
  420. }{
  421. {"delete all", "DELETE FROM users"},
  422. {"delete with where", "DELETE FROM users WHERE id = 1"},
  423. {"delete with complex where", "DELETE FROM users WHERE active = FALSE AND age < 18"},
  424. }
  425. for _, tt := range tests {
  426. t.Run(tt.name, func(t *testing.T) {
  427. stmt := parse(t, tt.sql)
  428. err := analyzer.Analyze(stmt)
  429. if err != nil {
  430. t.Errorf("Analyze(%q) error: %v", tt.sql, err)
  431. }
  432. })
  433. }
  434. }
  435. // CREATE TABLE analysis tests
  436. func TestAnalyzeCreateTable(t *testing.T) {
  437. tests := []struct {
  438. name string
  439. sql string
  440. }{
  441. {"basic table", "CREATE TABLE test (id INTEGER PRIMARY KEY, name TEXT)"},
  442. {"with constraints", "CREATE TABLE test (id INTEGER PRIMARY KEY, name TEXT NOT NULL, email TEXT UNIQUE)"},
  443. {"with default", "CREATE TABLE test (id INTEGER PRIMARY KEY, active BOOLEAN DEFAULT TRUE)"},
  444. {"if not exists", "CREATE TABLE IF NOT EXISTS test (id INTEGER)"},
  445. }
  446. for _, tt := range tests {
  447. t.Run(tt.name, func(t *testing.T) {
  448. // Use fresh catalog for each test
  449. c := NewCatalog()
  450. a := New(c)
  451. stmt := parse(t, tt.sql)
  452. err := a.Analyze(stmt)
  453. if err != nil {
  454. t.Errorf("Analyze(%q) error: %v", tt.sql, err)
  455. }
  456. })
  457. }
  458. }
  459. func TestAnalyzeCreateTableErrors(t *testing.T) {
  460. catalog := setupCatalog()
  461. analyzer := New(catalog)
  462. tests := []struct {
  463. name string
  464. sql string
  465. errType ErrorType
  466. }{
  467. {"table exists", "CREATE TABLE users (id INTEGER)", ErrTableExists},
  468. }
  469. for _, tt := range tests {
  470. t.Run(tt.name, func(t *testing.T) {
  471. stmt := parse(t, tt.sql)
  472. err := analyzer.Analyze(stmt)
  473. if err == nil {
  474. t.Errorf("Analyze(%q) expected error, got nil", tt.sql)
  475. return
  476. }
  477. if ae, ok := err.(*AnalysisError); ok {
  478. if ae.Type != tt.errType {
  479. t.Errorf("Analyze(%q) error type = %v, want %v", tt.sql, ae.Type, tt.errType)
  480. }
  481. }
  482. })
  483. }
  484. }
  485. // DROP TABLE analysis tests
  486. func TestAnalyzeDropTable(t *testing.T) {
  487. tests := []struct {
  488. name string
  489. sql string
  490. }{
  491. {"drop existing", "DROP TABLE products"},
  492. {"drop if exists", "DROP TABLE IF EXISTS nonexistent"},
  493. }
  494. for _, tt := range tests {
  495. t.Run(tt.name, func(t *testing.T) {
  496. // Use fresh catalog for each test
  497. c := setupCatalog()
  498. a := New(c)
  499. stmt := parse(t, tt.sql)
  500. err := a.Analyze(stmt)
  501. if err != nil {
  502. t.Errorf("Analyze(%q) error: %v", tt.sql, err)
  503. }
  504. })
  505. }
  506. }
  507. func TestAnalyzeDropTableErrors(t *testing.T) {
  508. catalog := setupCatalog()
  509. analyzer := New(catalog)
  510. tests := []struct {
  511. name string
  512. sql string
  513. errType ErrorType
  514. }{
  515. {"table not found", "DROP TABLE nonexistent", ErrTableNotFound},
  516. }
  517. for _, tt := range tests {
  518. t.Run(tt.name, func(t *testing.T) {
  519. stmt := parse(t, tt.sql)
  520. err := analyzer.Analyze(stmt)
  521. if err == nil {
  522. t.Errorf("Analyze(%q) expected error, got nil", tt.sql)
  523. return
  524. }
  525. if ae, ok := err.(*AnalysisError); ok {
  526. if ae.Type != tt.errType {
  527. t.Errorf("Analyze(%q) error type = %v, want %v", tt.sql, ae.Type, tt.errType)
  528. }
  529. }
  530. })
  531. }
  532. }
  533. // Expression analysis tests
  534. func TestAnalyzeExpressions(t *testing.T) {
  535. catalog := setupCatalog()
  536. analyzer := New(catalog)
  537. tests := []struct {
  538. name string
  539. sql string
  540. }{
  541. {"arithmetic", "SELECT 1 + 2 * 3 FROM users"},
  542. {"comparison", "SELECT * FROM users WHERE age > 18"},
  543. {"logical", "SELECT * FROM users WHERE active = TRUE AND age >= 21"},
  544. {"is null", "SELECT * FROM users WHERE email IS NULL"},
  545. {"is not null", "SELECT * FROM users WHERE email IS NOT NULL"},
  546. {"in list", "SELECT * FROM users WHERE id IN (1, 2, 3)"},
  547. {"not in", "SELECT * FROM users WHERE id NOT IN (1, 2, 3)"},
  548. {"between", "SELECT * FROM users WHERE age BETWEEN 18 AND 65"},
  549. {"like", "SELECT * FROM users WHERE name LIKE 'J%'"},
  550. {"case when", "SELECT CASE WHEN age >= 18 THEN 'adult' ELSE 'minor' END FROM users"},
  551. {"cast", "SELECT CAST(age AS TEXT) FROM users"},
  552. {"coalesce", "SELECT COALESCE(email, 'no email') FROM users"},
  553. {"function", "SELECT UPPER(name), LENGTH(email) FROM users"},
  554. }
  555. for _, tt := range tests {
  556. t.Run(tt.name, func(t *testing.T) {
  557. stmt := parse(t, tt.sql)
  558. err := analyzer.Analyze(stmt)
  559. if err != nil {
  560. t.Errorf("Analyze(%q) error: %v", tt.sql, err)
  561. }
  562. })
  563. }
  564. }
  565. func TestAnalyzeFunctionErrors(t *testing.T) {
  566. catalog := setupCatalog()
  567. analyzer := New(catalog)
  568. tests := []struct {
  569. name string
  570. sql string
  571. errType ErrorType
  572. }{
  573. {"unknown function", "SELECT UNKNOWN_FUNC(id) FROM users", ErrInvalidFunction},
  574. {"wrong arg count", "SELECT UPPER() FROM users", ErrInvalidArgCount},
  575. {"too many args", "SELECT LENGTH(name, 1) FROM users", ErrInvalidArgCount},
  576. }
  577. for _, tt := range tests {
  578. t.Run(tt.name, func(t *testing.T) {
  579. stmt := parse(t, tt.sql)
  580. err := analyzer.Analyze(stmt)
  581. if err == nil {
  582. t.Errorf("Analyze(%q) expected error, got nil", tt.sql)
  583. return
  584. }
  585. if ae, ok := err.(*AnalysisError); ok {
  586. if ae.Type != tt.errType {
  587. t.Errorf("Analyze(%q) error type = %v, want %v", tt.sql, ae.Type, tt.errType)
  588. }
  589. }
  590. })
  591. }
  592. }
  593. // Subquery tests
  594. func TestAnalyzeSubqueries(t *testing.T) {
  595. catalog := setupCatalog()
  596. analyzer := New(catalog)
  597. tests := []struct {
  598. name string
  599. sql string
  600. }{
  601. {"in subquery", "SELECT * FROM users WHERE id IN (SELECT user_id FROM orders)"},
  602. {"exists subquery", "SELECT * FROM users WHERE EXISTS (SELECT 1 FROM orders WHERE orders.user_id = users.id)"},
  603. {"scalar subquery", "SELECT (SELECT COUNT(*) FROM orders) FROM users"},
  604. }
  605. for _, tt := range tests {
  606. t.Run(tt.name, func(t *testing.T) {
  607. stmt := parse(t, tt.sql)
  608. err := analyzer.Analyze(stmt)
  609. if err != nil {
  610. t.Errorf("Analyze(%q) error: %v", tt.sql, err)
  611. }
  612. })
  613. }
  614. }
  615. // Scope tests
  616. func TestScope(t *testing.T) {
  617. scope := NewScope(nil)
  618. table := &TableInfo{
  619. Name: "users",
  620. Columns: []ColumnInfo{
  621. {Name: "id", Type: TypeInteger},
  622. {Name: "name", Type: TypeText},
  623. },
  624. }
  625. scope.DefineTable(table)
  626. // Test table lookup
  627. if _, ok := scope.LookupTable("users"); !ok {
  628. t.Error("expected to find table 'users'")
  629. }
  630. if _, ok := scope.LookupTable("USERS"); !ok {
  631. t.Error("expected case-insensitive table lookup")
  632. }
  633. if _, ok := scope.LookupTable("nonexistent"); ok {
  634. t.Error("expected not to find table 'nonexistent'")
  635. }
  636. // Test column lookup
  637. if col, _, ok := scope.LookupColumn("", "id"); !ok || col.Type != TypeInteger {
  638. t.Error("expected to find column 'id' with type INTEGER")
  639. }
  640. if col, _, ok := scope.LookupColumn("users", "name"); !ok || col.Type != TypeText {
  641. t.Error("expected to find column 'users.name' with type TEXT")
  642. }
  643. if _, _, ok := scope.LookupColumn("", "nonexistent"); ok {
  644. t.Error("expected not to find column 'nonexistent'")
  645. }
  646. }
  647. func TestCatalog(t *testing.T) {
  648. catalog := NewCatalog()
  649. // Create table
  650. err := catalog.CreateTable(&TableInfo{
  651. Name: "test",
  652. Columns: []ColumnInfo{
  653. {Name: "id", Type: TypeInteger},
  654. },
  655. })
  656. if err != nil {
  657. t.Errorf("CreateTable error: %v", err)
  658. }
  659. // Check exists
  660. if !catalog.TableExists("test") {
  661. t.Error("expected table 'test' to exist")
  662. }
  663. // Duplicate create should fail
  664. err = catalog.CreateTable(&TableInfo{Name: "test"})
  665. if err == nil {
  666. t.Error("expected error for duplicate table")
  667. }
  668. // Drop table
  669. err = catalog.DropTable("test")
  670. if err != nil {
  671. t.Errorf("DropTable error: %v", err)
  672. }
  673. // Check not exists
  674. if catalog.TableExists("test") {
  675. t.Error("expected table 'test' to not exist after drop")
  676. }
  677. // Drop non-existent should fail
  678. err = catalog.DropTable("test")
  679. if err == nil {
  680. t.Error("expected error for dropping non-existent table")
  681. }
  682. }
  683. // Benchmark
  684. func BenchmarkAnalyzeSelect(b *testing.B) {
  685. catalog := setupCatalog()
  686. analyzer := New(catalog)
  687. sql := "SELECT u.id, u.name, COUNT(o.id) FROM users u LEFT JOIN orders o ON u.id = o.user_id WHERE u.active = TRUE GROUP BY u.id, u.name HAVING COUNT(o.id) > 0 ORDER BY u.name LIMIT 100"
  688. l := lexer.New(sql)
  689. p := parser.New(l)
  690. stmt, _ := p.Parse()
  691. b.ResetTimer()
  692. for i := 0; i < b.N; i++ {
  693. _ = analyzer.Analyze(stmt)
  694. }
  695. }