analyzer_test.go 18 KB

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