2
0

main.go 21 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881
  1. // sqllogictest runner — sends SQL via HTTP to a running PizzaSQL server and
  2. // compares results against the expected output in .test files.
  3. //
  4. // File format: https://www.sqlite.org/sqllogictest/doc/trunk/about.wiki
  5. //
  6. // Usage:
  7. //
  8. // go run ./cmd/sqllogictest -url http://localhost:8080 -dir testdata/sqllogictest
  9. package main
  10. import (
  11. "bufio"
  12. "bytes"
  13. "context"
  14. "crypto/md5"
  15. "database/sql"
  16. "flag"
  17. "fmt"
  18. "math"
  19. "net/http"
  20. "os"
  21. "path/filepath"
  22. "sort"
  23. "strconv"
  24. "strings"
  25. "time"
  26. "github.com/goccy/go-json"
  27. _ "github.com/lib/pq" // PostgreSQL driver
  28. )
  29. const engineName = "pizzasql"
  30. // ANSI color helpers
  31. const (
  32. colorReset = "\033[0m"
  33. colorRed = "\033[31m"
  34. colorGreen = "\033[32m"
  35. colorYellow = "\033[33m"
  36. colorCyan = "\033[36m"
  37. colorBold = "\033[1m"
  38. colorDim = "\033[2m"
  39. )
  40. // ── types ────────────────────────────────────────────────────────────────────
  41. type lineInfo struct {
  42. text string
  43. num int
  44. }
  45. type record struct {
  46. isStatement bool
  47. isQuery bool
  48. expectOK bool // statement: true → expect success
  49. typeStr string // query: column type chars (I/R/T)
  50. sortMode string // nosort | rowsort | valuesort
  51. label string
  52. sql string
  53. expected []string // flattened expected values, one per line
  54. skip bool
  55. file string
  56. line int
  57. }
  58. type queryRequest struct {
  59. SQL string `json:"sql"`
  60. }
  61. type queryResponse struct {
  62. Columns []struct {
  63. Name string `json:"name"`
  64. Type string `json:"type"`
  65. } `json:"columns"`
  66. Rows [][]interface{} `json:"rows"`
  67. Error *struct {
  68. Code string `json:"code"`
  69. Message string `json:"message"`
  70. } `json:"error"`
  71. }
  72. // ── runner ───────────────────────────────────────────────────────────────────
  73. type runner struct {
  74. baseURL string
  75. client *http.Client
  76. pgDB *sql.DB // PostgreSQL connection (if using -pg flag)
  77. usePG bool // Use PostgreSQL wire protocol instead of HTTP
  78. verbose bool
  79. stopOnFail bool
  80. passed int
  81. failed int
  82. skipped int
  83. total int // total files to run
  84. filesDone int // files completed
  85. logW *bufio.Writer
  86. logPath string
  87. }
  88. func main() {
  89. urlFlag := flag.String("url", "http://localhost:8080", "PizzaSQL server URL")
  90. pgFlag := flag.Bool("pg", false, "Use PostgreSQL wire protocol instead of HTTP")
  91. pgHostFlag := flag.String("pg-host", "localhost", "PostgreSQL server host")
  92. pgPortFlag := flag.Int("pg-port", 5432, "PostgreSQL server port")
  93. pgDBFlag := flag.String("pg-db", "pizzasql", "PostgreSQL database name")
  94. dirFlag := flag.String("dir", "testdata/sqllogictest", "Directory containing .test files")
  95. fileFlag := flag.String("file", "", "Single .test file to run (overrides -dir)")
  96. verboseFlag := flag.Bool("v", false, "Print each passing record")
  97. stopFlag := flag.Bool("stop", false, "Stop on first failure")
  98. logFlag := flag.String("log", "sqllogictest-failures.log", "File to write failures to ('' to disable)")
  99. flag.Parse()
  100. r := &runner{
  101. baseURL: strings.TrimRight(*urlFlag, "/"),
  102. client: &http.Client{Timeout: 120 * time.Second},
  103. usePG: *pgFlag,
  104. verbose: *verboseFlag,
  105. stopOnFail: *stopFlag,
  106. logPath: *logFlag,
  107. }
  108. // If using PostgreSQL wire protocol, establish connection
  109. if *pgFlag {
  110. connStr := fmt.Sprintf("host=%s port=%d dbname=%s sslmode=disable",
  111. *pgHostFlag, *pgPortFlag, *pgDBFlag)
  112. db, err := sql.Open("postgres", connStr)
  113. if err != nil {
  114. fmt.Fprintf(os.Stderr, "Failed to connect to PostgreSQL: %v\n", err)
  115. os.Exit(1)
  116. }
  117. defer db.Close()
  118. // Test connection
  119. ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
  120. defer cancel()
  121. if err := db.PingContext(ctx); err != nil {
  122. fmt.Fprintf(os.Stderr, "Failed to ping PostgreSQL server: %v\n", err)
  123. os.Exit(1)
  124. }
  125. r.pgDB = db
  126. fmt.Printf("Connected to PostgreSQL at %s:%d (database: %s)\n", *pgHostFlag, *pgPortFlag, *pgDBFlag)
  127. }
  128. if *logFlag != "" {
  129. lf, err := os.Create(*logFlag)
  130. if err != nil {
  131. fmt.Fprintf(os.Stderr, "cannot open log file: %v\n", err)
  132. os.Exit(1)
  133. }
  134. defer lf.Close()
  135. r.logW = bufio.NewWriter(lf)
  136. defer r.logW.Flush()
  137. }
  138. var files []string
  139. if *fileFlag != "" {
  140. files = []string{*fileFlag}
  141. } else {
  142. err := filepath.WalkDir(*dirFlag, func(path string, d os.DirEntry, err error) error {
  143. if err != nil {
  144. return err
  145. }
  146. if !d.IsDir() && strings.HasSuffix(path, ".test") {
  147. files = append(files, path)
  148. }
  149. return nil
  150. })
  151. if err != nil || len(files) == 0 {
  152. fmt.Fprintf(os.Stderr, "no .test files found in %s\n", *dirFlag)
  153. os.Exit(1)
  154. }
  155. sort.Strings(files)
  156. }
  157. r.total = len(files)
  158. start := time.Now()
  159. for _, f := range files {
  160. if err := r.runFile(f, start); err != nil {
  161. fmt.Fprintf(os.Stderr, "error in %s: %v\n", f, err)
  162. }
  163. if r.stopOnFail && r.failed > 0 {
  164. break
  165. }
  166. }
  167. // clear the progress line
  168. fmt.Print("\r\033[K")
  169. total := r.passed + r.failed
  170. elapsed := time.Since(start).Round(time.Millisecond)
  171. passColor, failColor := colorGreen, colorDim
  172. if r.failed > 0 {
  173. failColor = colorRed
  174. }
  175. pct := 0.0
  176. if total > 0 {
  177. pct = 100.0 * float64(r.passed) / float64(total)
  178. }
  179. fmt.Printf("%s--- Summary ---%s\n", colorBold, colorReset)
  180. var summaryQPS string
  181. if secs := elapsed.Seconds(); secs > 0 && total > 0 {
  182. qps := float64(total) / secs
  183. switch {
  184. case qps >= 1_000_000:
  185. summaryQPS = fmt.Sprintf("%.2fM q/s", qps/1_000_000)
  186. case qps >= 1_000:
  187. summaryQPS = fmt.Sprintf("%.2fk q/s", qps/1_000)
  188. default:
  189. summaryQPS = fmt.Sprintf("%.0f q/s", qps)
  190. }
  191. }
  192. fmt.Printf("passed: %s%d/%d (%.1f%%)%s\n", passColor, r.passed, total, pct, colorReset)
  193. fmt.Printf("failed: %s%d%s\n", failColor, r.failed, colorReset)
  194. fmt.Printf("skipped: %d\n", r.skipped)
  195. fmt.Printf("time: %s\n", elapsed)
  196. fmt.Printf("thru: %s%s%s\n", colorCyan, summaryQPS, colorReset)
  197. if r.failed > 0 && *logFlag != "" {
  198. fmt.Printf("log: %s%s%s\n", colorCyan, *logFlag, colorReset)
  199. }
  200. if r.failed > 0 {
  201. os.Exit(1)
  202. }
  203. }
  204. func (r *runner) runFile(path string, start time.Time) error {
  205. f, err := os.Open(path)
  206. if err != nil {
  207. return err
  208. }
  209. defer f.Close()
  210. records, err := parseFile(path, f)
  211. if err != nil {
  212. return err
  213. }
  214. // Drop any tables/views this file creates so it always runs against a clean state.
  215. for _, tbl := range collectCreatedTables(records) {
  216. r.execQuery("DROP TABLE IF EXISTS " + tbl) //nolint:errcheck
  217. }
  218. for _, v := range collectCreatedViews(records) {
  219. r.execQuery("DROP VIEW IF EXISTS " + v) //nolint:errcheck
  220. }
  221. labelCache := make(map[string][]string) // label → first result
  222. failsBefore := r.failed
  223. for _, rec := range records {
  224. if r.stopOnFail && r.failed > 0 {
  225. break
  226. }
  227. if rec.skip {
  228. r.skipped++
  229. continue
  230. }
  231. r.runRecord(rec, labelCache)
  232. r.printProgress(path, start)
  233. }
  234. r.filesDone++
  235. newFails := r.failed - failsBefore
  236. rel, _ := filepath.Rel("testdata/sqllogictest", path)
  237. if rel == "" {
  238. rel = filepath.Base(path)
  239. }
  240. var statusStr string
  241. if newFails == 0 {
  242. statusStr = colorGreen + "ok" + colorReset
  243. } else {
  244. statusStr = fmt.Sprintf("%s%d FAILED%s", colorRed, newFails, colorReset)
  245. }
  246. fmt.Printf("\r\033[K%s[%d/%d]%s %-52s %s\n", colorDim, r.filesDone, r.total, colorReset, rel, statusStr)
  247. return nil
  248. }
  249. func (r *runner) printProgress(currentFile string, start time.Time) {
  250. rel, _ := filepath.Rel("testdata/sqllogictest", currentFile)
  251. if rel == "" {
  252. rel = filepath.Base(currentFile)
  253. }
  254. elapsed := time.Since(start)
  255. elapsedStr := elapsed.Round(time.Second).String()
  256. var etaStr string
  257. if r.filesDone > 0 {
  258. rate := float64(r.filesDone) / elapsed.Seconds()
  259. eta := time.Duration(float64(r.total-r.filesDone) / rate * float64(time.Second)).Round(time.Second)
  260. etaStr = "eta " + eta.String()
  261. } else {
  262. etaStr = "eta --"
  263. }
  264. checked := r.passed + r.failed
  265. var rateStr string
  266. if checked > 0 {
  267. pct := 100.0 * float64(r.passed) / float64(checked)
  268. color := colorRed
  269. if r.failed == 0 {
  270. color = colorGreen
  271. } else if pct >= 90 {
  272. color = colorYellow
  273. }
  274. rateStr = fmt.Sprintf("%s%.1f%%%s", color, pct, colorReset)
  275. } else {
  276. rateStr = " --.--%"
  277. }
  278. var throughputStr string
  279. if secs := elapsed.Seconds(); secs > 0 && checked > 0 {
  280. qps := float64(checked) / secs
  281. switch {
  282. case qps >= 1_000_000:
  283. throughputStr = fmt.Sprintf("%.1fM q/s", qps/1_000_000)
  284. case qps >= 1_000:
  285. throughputStr = fmt.Sprintf("%.1fk q/s", qps/1_000)
  286. default:
  287. throughputStr = fmt.Sprintf("%.0f q/s", qps)
  288. }
  289. } else {
  290. throughputStr = "-- q/s"
  291. }
  292. fmt.Printf("\r\033[K%s[%d/%d]%s %-40s %s pass=%-6d %sfail=%-5d%s skip=%-5d %s / %s %s%s%s",
  293. colorDim, r.filesDone+1, r.total, colorReset,
  294. rel, rateStr,
  295. r.passed,
  296. colorRed, r.failed, colorReset,
  297. r.skipped,
  298. elapsedStr, etaStr,
  299. colorCyan, throughputStr, colorReset,
  300. )
  301. }
  302. // collectCreatedTables scans records for CREATE TABLE statements and returns
  303. // the table names so they can be pre-dropped before each test file runs.
  304. func collectCreatedViews(records []*record) []string {
  305. seen := map[string]bool{}
  306. var views []string
  307. for _, rec := range records {
  308. if !rec.isStatement {
  309. continue
  310. }
  311. fields := strings.Fields(rec.sql)
  312. if len(fields) < 3 {
  313. continue
  314. }
  315. if !strings.EqualFold(fields[0], "CREATE") || !strings.EqualFold(fields[1], "VIEW") {
  316. continue
  317. }
  318. idx := 2
  319. if strings.EqualFold(fields[idx], "IF") && len(fields) > idx+2 {
  320. idx = 5
  321. }
  322. if idx < len(fields) {
  323. name := strings.TrimSuffix(fields[idx], ";")
  324. if name != "" && !seen[name] {
  325. seen[name] = true
  326. views = append(views, name)
  327. }
  328. }
  329. }
  330. return views
  331. }
  332. func collectCreatedTables(records []*record) []string {
  333. seen := map[string]bool{}
  334. var tables []string
  335. for _, rec := range records {
  336. if !rec.isStatement {
  337. continue
  338. }
  339. fields := strings.Fields(rec.sql)
  340. if len(fields) < 3 {
  341. continue
  342. }
  343. if !strings.EqualFold(fields[0], "CREATE") || !strings.EqualFold(fields[1], "TABLE") {
  344. continue
  345. }
  346. idx := 2
  347. if strings.EqualFold(fields[idx], "IF") && len(fields) > idx+2 {
  348. idx = 5 // CREATE TABLE IF NOT EXISTS <name>
  349. }
  350. if idx < len(fields) {
  351. name := strings.TrimSuffix(strings.TrimSuffix(fields[idx], "("), ";")
  352. if name != "" && !seen[name] {
  353. seen[name] = true
  354. tables = append(tables, name)
  355. }
  356. }
  357. }
  358. return tables
  359. }
  360. func (r *runner) runRecord(rec *record, labelCache map[string][]string) {
  361. resp, err := r.execQuery(rec.sql)
  362. if err != nil {
  363. r.fail(rec, "http error: %v", err)
  364. return
  365. }
  366. if rec.isStatement {
  367. if rec.expectOK {
  368. if resp.Error != nil {
  369. r.fail(rec, "expected ok, got error: %s", resp.Error.Message)
  370. } else {
  371. r.pass(rec)
  372. }
  373. } else {
  374. if resp.Error == nil {
  375. r.fail(rec, "expected error, got ok")
  376. } else {
  377. r.pass(rec)
  378. }
  379. }
  380. return
  381. }
  382. // query record
  383. if resp.Error != nil {
  384. r.fail(rec, "unexpected error: %s", resp.Error.Message)
  385. return
  386. }
  387. got := r.formatResults(resp, rec.typeStr)
  388. ncols := len(rec.typeStr)
  389. if ncols == 0 {
  390. ncols = 1
  391. }
  392. // Apply sort before any comparison.
  393. switch rec.sortMode {
  394. case "rowsort":
  395. got = sortRows(got, ncols)
  396. case "valuesort":
  397. g := append([]string(nil), got...)
  398. sort.Strings(g)
  399. got = g
  400. }
  401. // Label caching: if this query has a label, compare against first occurrence.
  402. if rec.label != "" {
  403. if cached, seen := labelCache[rec.label]; seen {
  404. if !equalSlices(got, cached) {
  405. r.fail(rec, "label %q result mismatch\n want: %v\n got: %v", rec.label, cached, got)
  406. } else {
  407. r.pass(rec)
  408. }
  409. return
  410. }
  411. // First occurrence: store and fall through to normal expected-value check.
  412. labelCache[rec.label] = got
  413. }
  414. // hash format: "N values hashing to <md5>"
  415. if len(rec.expected) == 1 {
  416. parts := strings.Fields(rec.expected[0])
  417. if len(parts) == 5 && parts[1] == "values" && parts[2] == "hashing" && parts[3] == "to" {
  418. wantCount, _ := strconv.Atoi(parts[0])
  419. wantHash := parts[4]
  420. if len(got) != wantCount {
  421. r.fail(rec, "hash record: want %d values got %d", wantCount, len(got))
  422. return
  423. }
  424. h := md5.Sum([]byte(strings.Join(got, "\n") + "\n"))
  425. gotHash := fmt.Sprintf("%x", h)
  426. if gotHash != wantHash {
  427. r.fail(rec, "hash mismatch: want %s got %s", wantHash, gotHash)
  428. return
  429. }
  430. r.pass(rec)
  431. return
  432. }
  433. }
  434. exp := rec.expected
  435. switch rec.sortMode {
  436. case "rowsort":
  437. exp = sortRows(exp, ncols)
  438. case "valuesort":
  439. e := append([]string(nil), exp...)
  440. sort.Strings(e)
  441. exp = e
  442. }
  443. if !equalSlices(got, exp) {
  444. r.fail(rec, "result mismatch\n want: %v\n got: %v", exp, got)
  445. } else {
  446. r.pass(rec)
  447. }
  448. }
  449. // ── formatting ───────────────────────────────────────────────────────────────
  450. func (r *runner) formatResults(resp *queryResponse, typeStr string) []string {
  451. var vals []string
  452. for _, row := range resp.Rows {
  453. for i, v := range row {
  454. ct := byte('T')
  455. if i < len(typeStr) {
  456. ct = typeStr[i]
  457. }
  458. vals = append(vals, formatValue(v, ct))
  459. }
  460. }
  461. return vals
  462. }
  463. // formatValue converts a JSON value to the string representation expected by
  464. // the sqllogictest format. Type chars: I=integer, R=real (%.3g), T=text.
  465. func formatValue(v interface{}, colType byte) string {
  466. if v == nil {
  467. return "NULL"
  468. }
  469. switch colType {
  470. case 'I':
  471. switch n := v.(type) {
  472. case float64:
  473. return strconv.FormatInt(int64(n), 10)
  474. case int64:
  475. return strconv.FormatInt(n, 10)
  476. case int:
  477. return strconv.Itoa(n)
  478. case bool:
  479. if n {
  480. return "1"
  481. }
  482. return "0"
  483. case string:
  484. if i, err := strconv.ParseInt(n, 10, 64); err == nil {
  485. return strconv.FormatInt(i, 10)
  486. }
  487. return "0"
  488. default:
  489. return fmt.Sprintf("%v", v)
  490. }
  491. case 'R':
  492. switch n := v.(type) {
  493. case float64:
  494. return strconv.FormatFloat(n, 'g', 3, 64)
  495. case int64:
  496. return strconv.FormatFloat(float64(n), 'g', 3, 64)
  497. case int:
  498. return strconv.FormatFloat(float64(n), 'g', 3, 64)
  499. case string:
  500. if f, err := strconv.ParseFloat(n, 64); err == nil {
  501. return strconv.FormatFloat(f, 'g', 3, 64)
  502. }
  503. return "0"
  504. default:
  505. return fmt.Sprintf("%v", v)
  506. }
  507. default: // T
  508. switch s := v.(type) {
  509. case string:
  510. return s
  511. case bool:
  512. if s {
  513. return "1"
  514. }
  515. return "0"
  516. case float64:
  517. if s == math.Trunc(s) && !math.IsInf(s, 0) {
  518. return strconv.FormatInt(int64(s), 10)
  519. }
  520. return fmt.Sprintf("%g", s)
  521. default:
  522. return fmt.Sprintf("%v", v)
  523. }
  524. }
  525. }
  526. // ── helpers ───────────────────────────────────────────────────────────────────
  527. func sortRows(vals []string, ncols int) []string {
  528. if ncols <= 0 || len(vals) == 0 {
  529. return vals
  530. }
  531. nrows := len(vals) / ncols
  532. rows := make([][]string, nrows)
  533. for i := range rows {
  534. s, e := i*ncols, i*ncols+ncols
  535. if e > len(vals) {
  536. e = len(vals)
  537. }
  538. rows[i] = vals[s:e]
  539. }
  540. sort.Slice(rows, func(i, j int) bool {
  541. for k := 0; k < len(rows[i]) && k < len(rows[j]); k++ {
  542. if rows[i][k] != rows[j][k] {
  543. return rows[i][k] < rows[j][k]
  544. }
  545. }
  546. return len(rows[i]) < len(rows[j])
  547. })
  548. out := make([]string, 0, len(vals))
  549. for _, row := range rows {
  550. out = append(out, row...)
  551. }
  552. return out
  553. }
  554. func equalSlices(a, b []string) bool {
  555. if len(a) != len(b) {
  556. return false
  557. }
  558. for i := range a {
  559. if a[i] != b[i] {
  560. return false
  561. }
  562. }
  563. return true
  564. }
  565. func (r *runner) execQuery(sql string) (*queryResponse, error) {
  566. if r.usePG {
  567. return r.execQueryPG(sql)
  568. }
  569. return r.execQueryHTTP(sql)
  570. }
  571. func (r *runner) execQueryHTTP(sql string) (*queryResponse, error) {
  572. body, _ := json.Marshal(queryRequest{SQL: sql})
  573. resp, err := r.client.Post(r.baseURL+"/query", "application/json", bytes.NewReader(body))
  574. if err != nil {
  575. return nil, err
  576. }
  577. defer resp.Body.Close()
  578. var qr queryResponse
  579. if err := json.NewDecoder(resp.Body).Decode(&qr); err != nil {
  580. return nil, fmt.Errorf("decode response: %w", err)
  581. }
  582. return &qr, nil
  583. }
  584. func (r *runner) execQueryPG(sql string) (*queryResponse, error) {
  585. ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second)
  586. defer cancel()
  587. // Check if it's a query or statement
  588. sqlUpper := strings.TrimSpace(strings.ToUpper(sql))
  589. isSelect := strings.HasPrefix(sqlUpper, "SELECT") ||
  590. strings.HasPrefix(sqlUpper, "PRAGMA") ||
  591. strings.HasPrefix(sqlUpper, "EXPLAIN")
  592. var qr queryResponse
  593. if isSelect {
  594. // Execute query and get results
  595. rows, err := r.pgDB.QueryContext(ctx, sql)
  596. if err != nil {
  597. qr.Error = &struct {
  598. Code string `json:"code"`
  599. Message string `json:"message"`
  600. }{
  601. Code: "QUERY_ERROR",
  602. Message: err.Error(),
  603. }
  604. return &qr, nil
  605. }
  606. defer rows.Close()
  607. // Get column information
  608. colTypes, err := rows.ColumnTypes()
  609. if err != nil {
  610. return nil, fmt.Errorf("get column types: %w", err)
  611. }
  612. for _, ct := range colTypes {
  613. qr.Columns = append(qr.Columns, struct {
  614. Name string `json:"name"`
  615. Type string `json:"type"`
  616. }{
  617. Name: ct.Name(),
  618. Type: ct.DatabaseTypeName(),
  619. })
  620. }
  621. // Read all rows
  622. for rows.Next() {
  623. values := make([]interface{}, len(colTypes))
  624. valuePtrs := make([]interface{}, len(colTypes))
  625. for i := range values {
  626. valuePtrs[i] = &values[i]
  627. }
  628. if err := rows.Scan(valuePtrs...); err != nil {
  629. return nil, fmt.Errorf("scan row: %w", err)
  630. }
  631. // Convert byte arrays to strings (PostgreSQL returns some types as []byte)
  632. for i, v := range values {
  633. if b, ok := v.([]byte); ok {
  634. values[i] = string(b)
  635. }
  636. }
  637. qr.Rows = append(qr.Rows, values)
  638. }
  639. if err := rows.Err(); err != nil {
  640. return nil, fmt.Errorf("rows error: %w", err)
  641. }
  642. } else {
  643. // Execute statement (INSERT, UPDATE, DELETE, CREATE, etc.)
  644. _, err := r.pgDB.ExecContext(ctx, sql)
  645. if err != nil {
  646. qr.Error = &struct {
  647. Code string `json:"code"`
  648. Message string `json:"message"`
  649. }{
  650. Code: "EXEC_ERROR",
  651. Message: err.Error(),
  652. }
  653. return &qr, nil
  654. }
  655. }
  656. return &qr, nil
  657. }
  658. func (r *runner) pass(rec *record) {
  659. r.passed++
  660. if r.verbose && r.logW != nil {
  661. fmt.Fprintf(r.logW, " ok %s:%d\n", rec.file, rec.line)
  662. }
  663. }
  664. func (r *runner) fail(rec *record, format string, args ...interface{}) {
  665. r.failed++
  666. msg := fmt.Sprintf(format, args...)
  667. sql := strings.ReplaceAll(strings.TrimSpace(rec.sql), "\n", " ")
  668. if len(sql) > 120 {
  669. sql = sql[:117] + "..."
  670. }
  671. line := fmt.Sprintf("FAIL %s:%d: %s\n SQL: %s\n", rec.file, rec.line, msg, sql)
  672. if r.logW != nil {
  673. fmt.Fprint(r.logW, line)
  674. r.logW.Flush()
  675. } else {
  676. fmt.Print(line)
  677. }
  678. }
  679. // ── parser ────────────────────────────────────────────────────────────────────
  680. // parseFile reads a sqllogictest file and returns all records.
  681. func parseFile(path string, f *os.File) ([]*record, error) {
  682. scanner := bufio.NewScanner(f)
  683. var lines []lineInfo
  684. n := 0
  685. for scanner.Scan() {
  686. n++
  687. text := scanner.Text()
  688. if !strings.HasPrefix(strings.TrimSpace(text), "#") {
  689. lines = append(lines, lineInfo{text: text, num: n})
  690. }
  691. }
  692. if err := scanner.Err(); err != nil {
  693. return nil, err
  694. }
  695. // split into blocks separated by blank lines
  696. var blocks [][]lineInfo
  697. var cur []lineInfo
  698. for _, li := range lines {
  699. if strings.TrimSpace(li.text) == "" {
  700. if len(cur) > 0 {
  701. blocks = append(blocks, cur)
  702. cur = nil
  703. }
  704. } else {
  705. cur = append(cur, li)
  706. }
  707. }
  708. if len(cur) > 0 {
  709. blocks = append(blocks, cur)
  710. }
  711. var records []*record
  712. haltSeen := false
  713. skipNext := false
  714. for _, block := range blocks {
  715. if haltSeen {
  716. break
  717. }
  718. // consume skipif / onlyif lines at the top of the block
  719. i := 0
  720. for i < len(block) {
  721. lower := strings.ToLower(strings.TrimSpace(block[i].text))
  722. if strings.HasPrefix(lower, "skipif ") {
  723. engine := strings.TrimSpace(block[i].text[7:])
  724. if strings.EqualFold(engine, engineName) {
  725. skipNext = true
  726. }
  727. i++
  728. } else if strings.HasPrefix(lower, "onlyif ") {
  729. engine := strings.TrimSpace(block[i].text[7:])
  730. if !strings.EqualFold(engine, engineName) {
  731. skipNext = true
  732. }
  733. i++
  734. } else {
  735. break
  736. }
  737. }
  738. if i >= len(block) {
  739. continue
  740. }
  741. directiveLine := block[i]
  742. parts := strings.Fields(directiveLine.text)
  743. if len(parts) == 0 {
  744. continue
  745. }
  746. rec := &record{file: path, line: directiveLine.num, skip: skipNext}
  747. skipNext = false
  748. body := block[i+1:]
  749. switch parts[0] {
  750. case "halt":
  751. haltSeen = true
  752. continue
  753. case "statement":
  754. rec.isStatement = true
  755. rec.expectOK = len(parts) > 1 && parts[1] == "ok"
  756. var sqlLines []string
  757. for _, li := range body {
  758. sqlLines = append(sqlLines, li.text)
  759. }
  760. rec.sql = strings.Join(sqlLines, "\n")
  761. case "query":
  762. rec.isQuery = true
  763. if len(parts) > 1 {
  764. rec.typeStr = strings.ToUpper(parts[1])
  765. }
  766. if len(parts) > 2 {
  767. rec.sortMode = parts[2]
  768. } else {
  769. rec.sortMode = "nosort"
  770. }
  771. if len(parts) > 3 {
  772. rec.label = parts[3]
  773. }
  774. inResults := false
  775. var sqlLines []string
  776. for _, li := range body {
  777. if strings.TrimSpace(li.text) == "----" {
  778. inResults = true
  779. continue
  780. }
  781. if inResults {
  782. rec.expected = append(rec.expected, strings.TrimSpace(li.text))
  783. } else {
  784. sqlLines = append(sqlLines, li.text)
  785. }
  786. }
  787. rec.sql = strings.Join(sqlLines, "\n")
  788. default:
  789. continue
  790. }
  791. if strings.TrimSpace(rec.sql) == "" {
  792. continue
  793. }
  794. records = append(records, rec)
  795. }
  796. return records, nil
  797. }