diff --git a/.github/workflows/TestingCI.yml b/.github/workflows/TestingCI.yml
index aba81b4..66d42c1 100644
--- a/.github/workflows/TestingCI.yml
+++ b/.github/workflows/TestingCI.yml
@@ -74,5 +74,58 @@ jobs:
restore-keys: ${{ runner.os }}-cargo-
- name: Build
run: cargo build --verbose
+ # Serialized like the Linux jobs. The REPL tests spawn the binary as a
+ # subprocess and several tests share fixtures under tests/data/, so
+ # running them in parallel here made macOS and Windows flaky in a way
+ # Linux never showed.
- name: Run tests
- run: cargo test --verbose
\ No newline at end of file
+ run: cargo test --verbose
+ env:
+ RUST_TEST_THREADS: 1
+
+ lint:
+ name: Format and Clippy
+ runs-on: ubuntu-latest
+ steps:
+ - uses: actions/checkout@v4
+ - name: Install Rust toolchain
+ uses: dtolnay/rust-toolchain@stable
+ with:
+ components: rustfmt, clippy
+ - name: Cache dependencies
+ uses: actions/cache@v4
+ with:
+ path: |
+ ~/.cargo/registry
+ ~/.cargo/git
+ target
+ key: ${{ runner.os }}-cargo-lint-${{ hashFiles('**/Cargo.lock') }}
+ restore-keys: ${{ runner.os }}-cargo-lint-
+ - name: Check formatting
+ run: cargo fmt --all -- --check
+ - name: Clippy
+ run: cargo clippy --all-targets -- -D warnings
+
+ msrv:
+ name: Minimum Supported Rust Version
+ runs-on: ubuntu-latest
+ steps:
+ - uses: actions/checkout@v4
+ # Pinned to the rust-version declared in Cargo.toml. That field claimed
+ # 1.70 for a long time while the dependency tree had moved well past it,
+ # because nothing ever checked.
+ - name: Install Rust toolchain
+ uses: dtolnay/rust-toolchain@1.88
+ - name: Cache dependencies
+ uses: actions/cache@v4
+ with:
+ path: |
+ ~/.cargo/registry
+ ~/.cargo/git
+ target
+ key: ${{ runner.os }}-cargo-msrv-${{ hashFiles('**/Cargo.lock') }}
+ restore-keys: ${{ runner.os }}-cargo-msrv-
+ # Library and binaries only: rust-version is a promise to consumers, and
+ # dev-dependencies are not part of it.
+ - name: Check build at MSRV
+ run: cargo check --verbose
diff --git a/CLAUDE.md b/CLAUDE.md
index 70966aa..1053081 100644
--- a/CLAUDE.md
+++ b/CLAUDE.md
@@ -50,10 +50,22 @@ SQL execution uses a bytecode VM inspired by SQLite's architecture:
### VM Module (`src/vm/`)
-- `bytecode.rs`: Defines bytecode instruction set
-- `compiler.rs`: Compiles SQL AST to bytecode
+- `bytecode.rs`: Bytecode instruction set and register/program types
+- `compiler.rs`: SQL AST to bytecode. Holds `code_expr`, the single expression
+ compiler (the analogue of SQLite's `sqlite3ExprCode`) used from the SELECT
+ list, WHERE, HAVING, ORDER BY and SET clauses alike, plus `NameCtx`, which
+ carries the FROM sources a name can resolve against
+- `compiler_aggregate.rs`: aggregates, GROUP BY, HAVING
+- `compiler_join.rs`: joins
+- `compiler_dml.rs`: INSERT / UPDATE / DELETE
+- `compiler_ddl.rs`: CREATE / DROP / ALTER / TRUNCATE
+- `compiler_window.rs`: window functions
+- `ast_compat.rs`: thin accessors normalizing sqlparser AST shapes
- `engine.rs`: Executes bytecode instructions
+Each statement is compiled and executed separately, so a statement observes the
+effects of the ones before it.
+
### File Handling
- **`file_handler.rs`**: Manages loading files into database tables
@@ -62,9 +74,10 @@ SQL execution uses a bytecode VM inspired by SQLite's architecture:
### SQL Features
-- `aggregate.rs`: COUNT, SUM, AVG, MIN, MAX functions
-- `join.rs`: Cross join and INNER JOIN implementations
-- `string_functions.rs`: UPPER, LOWER, TRIM, SUBSTR, REPLACE
+Expression compilation, including all scalar and string functions, lives in
+`src/vm/compiler.rs` behind `code_expr`. Aggregates are compiled in
+`src/vm/compiler_aggregate.rs` and executed by the `AggStep`/`AggFinal`
+opcodes; `src/aggregate.rs` retains only the function-name lookup.
### Safe Writeback Model
@@ -79,4 +92,15 @@ Tests are in `tests/` organized by functionality:
- `helpers/`: Test utilities
- `common.rs`: Shared test helpers including REPL script runner
+Two suites carry the correctness contract:
+
+- `tests/golden/`: exact-output characterization tests. They assert the
+ COMPLETE stdout -- every row, in order -- via `assert_query`. The rest of the
+ suite uses `predicate::str::contains`, which is substring matching: a test
+ expecting `"Engineering,3"` passes even when extra wrong rows are emitted and
+ regardless of row order. Prefer `assert_query` for new tests.
+- `tests/defects/`: `#[ignore]`d tests encoding CORRECT behaviour for known
+ defects. Run with `cargo test --test mod defects -- --ignored`. Remove the
+ `#[ignore]` when a defect is fixed; the list only shrinks.
+
Tests use `assert_cmd` for CLI testing and `tempfile` for temporary test data.
diff --git a/Cargo.lock b/Cargo.lock
index 4e5e4b2..7d926f7 100644
--- a/Cargo.lock
+++ b/Cargo.lock
@@ -67,6 +67,15 @@ version = "1.0.100"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a23eb6b1614318a8071c9b2521f36b424b2c83db5eb3a0fead4a6c0809af6e61"
+[[package]]
+name = "ar_archive_writer"
+version = "0.5.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "73cd58deff2140a0a8eae87e417bd01db68a33e148aa93d1e8cd837e55e312b6"
+dependencies = [
+ "object",
+]
+
[[package]]
name = "assert_cmd"
version = "2.1.2"
@@ -105,6 +114,16 @@ dependencies = [
"serde",
]
+[[package]]
+name = "cc"
+version = "1.4.2"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "5d262e149917187838d5b42777c8253bcb64500067342904e7d429499a6f277e"
+dependencies = [
+ "find-msvc-tools",
+ "shlex",
+]
+
[[package]]
name = "cfg-if"
version = "1.0.4"
@@ -238,6 +257,12 @@ dependencies = [
"windows-sys 0.59.0",
]
+[[package]]
+name = "find-msvc-tools"
+version = "0.1.10"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "26b73573e6edcd2af0cdf47bd6cb58f0b3839491263c314eaad1ccf24430e1de"
+
[[package]]
name = "float-cmp"
version = "0.10.0"
@@ -355,6 +380,15 @@ dependencies = [
"autocfg",
]
+[[package]]
+name = "object"
+version = "0.39.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "2e5a6c098c7a3b6547378093f5cc30bc54fd361ce711e05293a5cc589562739b"
+dependencies = [
+ "memchr",
+]
+
[[package]]
name = "once_cell"
version = "1.21.3"
@@ -406,6 +440,16 @@ dependencies = [
"unicode-ident",
]
+[[package]]
+name = "psm"
+version = "0.1.32"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "4dcd034599e63b970727f70d79e02d62390a4a84f7c6b827c27c46d5ac3fa622"
+dependencies = [
+ "ar_archive_writer",
+ "cc",
+]
+
[[package]]
name = "quote"
version = "1.0.43"
@@ -431,6 +475,26 @@ dependencies = [
"nibble_vec",
]
+[[package]]
+name = "recursive"
+version = "0.1.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0786a43debb760f491b1bc0269fe5e84155353c67482b9e60d0cfb596054b43e"
+dependencies = [
+ "recursive-proc-macro-impl",
+ "stacker",
+]
+
+[[package]]
+name = "recursive-proc-macro-impl"
+version = "0.1.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "76009fbe0614077fc1a2ce255e3a1881a2e3a3527097d5dc6d8212c585e7e38b"
+dependencies = [
+ "quote",
+ "syn",
+]
+
[[package]]
name = "regex"
version = "1.12.2"
@@ -530,6 +594,12 @@ dependencies = [
"syn",
]
+[[package]]
+name = "shlex"
+version = "2.0.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "f8fadd59c855ef2080decdef8ff161eb6661b86933c9d82e5ba29dc602a55aba"
+
[[package]]
name = "smallvec"
version = "1.15.1"
@@ -555,11 +625,25 @@ dependencies = [
[[package]]
name = "sqlparser"
-version = "0.36.1"
+version = "0.62.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
-checksum = "2eaa1e88e78d2c2460d78b7dc3f0c08dbb606ab4222f9aff36f420d36e307d87"
+checksum = "13c6d1b651dc4edf07eead2a0c6c78016ce971bc2c10da5266861b13f25e7cec"
dependencies = [
"log",
+ "recursive",
+]
+
+[[package]]
+name = "stacker"
+version = "0.1.25"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "707f49d46706bacf8a2b00d51dace3f9de527c13eec3778f570c411f89e69967"
+dependencies = [
+ "cc",
+ "cfg-if",
+ "libc",
+ "psm",
+ "windows-sys 0.61.2",
]
[[package]]
diff --git a/Cargo.toml b/Cargo.toml
index d6c0cc1..ac98a9c 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -10,13 +10,24 @@ homepage = "https://github.com/jgarzik/sqawk"
readme = "README.md"
keywords = ["csv", "tsv", "sql", "delimited", "awk"]
categories = ["command-line-utilities", "text-processing", "parser-implementations"]
-# Minimum version of Rust required
-rust-version = "1.70.0"
-exclude = ["tests/", "doc/", ".github/", "CLAUDE.md"]
+# Minimum supported Rust version.
+#
+# Verified by building with that toolchain, not inferred. The binding
+# constraints are transitive: `home` (via rustyline) and `psm` (via
+# sqlparser -> recursive -> stacker) both require 1.88.
+rust-version = "1.88"
+exclude = [
+ "tests/",
+ "doc/",
+ ".github/",
+ "CLAUDE.md",
+ "generated-icon.png",
+ ".replit",
+]
[dependencies]
clap = { version = "4", features = ["derive"] }
-sqlparser = "0.36"
+sqlparser = "0.62"
csv = "1"
anyhow = "1"
regex = "1"
diff --git a/README.md b/README.md
index 9ecb126..b8b3a53 100644
--- a/README.md
+++ b/README.md
@@ -12,6 +12,11 @@ Sqawk is an SQL-based command-line tool for processing delimiter-separated files
- **Joins** - INNER, LEFT, RIGHT, FULL OUTER, and CROSS joins with ON conditions
- **Aggregates** - COUNT, SUM, AVG, MIN, MAX with GROUP BY support
- **Functions** - String (UPPER, LOWER, SUBSTR, REPLACE, etc.), math (ABS, ROUND, etc.), date/time
+- **Subqueries** - Scalar, `IN (SELECT ...)`, and `EXISTS`, including correlated
+- **Set Operations** - UNION, UNION ALL, INTERSECT, EXCEPT
+- **Window Functions** - ROW_NUMBER, RANK, DENSE_RANK, LAG, LEAD, and aggregates with `OVER (PARTITION BY ... ORDER BY ...)`
+- **DDL** - CREATE TABLE, CREATE TABLE AS SELECT, DROP, ALTER TABLE ADD COLUMN, TRUNCATE
+- **Expressions** - CASE, CAST, COALESCE, NULLIF, BETWEEN, IN, LIKE/ILIKE, `||`, arithmetic
- **File Formats** - CSV, TSV, and custom delimiters; headerless files via `--tabledef`
- **Safe by Default** - Files unchanged unless `--write` flag is specified
- **Interactive REPL** - Explore data interactively with `-i` flag
@@ -38,6 +43,29 @@ sqawk -s "SELECT department, AVG(salary) FROM employees GROUP BY department" emp
sqawk -s "UPDATE data SET status = 'archived' WHERE year < 2020" data.csv --write
```
+## `tsq` - test data generator
+
+`cargo install sqawk` also installs `tsq`, which generates deterministic
+multi-table CSV data plus a corpus of SQL queries for exercising sqawk.
+
+```sh
+tsq --seed 42 --rows 1000 --output-dir /tmp/sqawk-test
+sqawk -s "SELECT * FROM customers LIMIT 10" /tmp/sqawk-test/data/customers.csv
+```
+
+It writes `data/` (customers, products, orders, order_items, reviews, with
+realistic foreign-key relationships), `queries/` (numbered `.sql` files
+covering selects, joins, aggregates, subqueries and window functions),
+`verify/run_verification.sh`, and a `metadata.json` recording the seed and row
+counts. The same seed always produces the same data.
+
+| Option | Meaning |
+| --- | --- |
+| `-s`, `--seed` | Seed for reproducible generation; a random one is printed if omitted |
+| `-r`, `--rows` | Base customer row count; other tables scale proportionally (default 100000) |
+| `-o`, `--output-dir` | Where to write the generated tree (required) |
+| `-v`, `--verbose` | Show generation progress |
+
## Documentation
- [User Guide](doc/user_guide.md) - Installation, CLI options, and examples
diff --git a/c.txt b/c.txt
new file mode 100644
index 0000000..224ef20
--- /dev/null
+++ b/c.txt
@@ -0,0 +1,68 @@
+src/vm/mod.rs
+Comment on lines +147 to +151
+ let applied = apply_modifications(database, modifications, verbose)?;
+ modified_tables.extend(applied.modified_tables);
+ if applied.affected_rows > 0 {
+ affected_rows = applied.affected_rows;
+ }
+Copilot commented now
+Copilot
+Copilot
+AI
+1m ago
+Medium
+affected_rows is only updated when applied.affected_rows > 0. This means a DML statement that affects 0 rows (e.g., UPDATE ... WHERE matches nothing) will incorrectly keep the previous non-zero count, so .changes/row-change reporting can be wrong.
+
+Suggested changeset 1 (1)
+src/vm/mod.rs
+Copilot uses AI. Check for mistakes.
+@jgarzik Reply...
+tests/tsq/mod.rs
+Comment on lines +113 to +117
+ let stdout = String::from_utf8_lossy(&output.stdout);
+ assert!(
+ stdout.contains("40"),
+ "expected 40 customers, got: {stdout}"
+ );
+Copilot commented now
+Copilot
+Copilot
+AI
+1m ago
+Low
+This test uses stdout.contains("40"), which can produce false positives (e.g. output 140 would still pass). Given the repository guidance to prefer exact-output assertions for new coverage, this should assert the full result more precisely.
+
+Suggested changeset 1 (1)
+tests/tsq/mod.rs
+Original file line number Diff line number Diff line change
+ String::from_utf8_lossy(&output.stderr)
+ );
+ let stdout = String::from_utf8_lossy(&output.stdout);
+ assert!(
+ stdout.contains("40"),
+ let lines: Vec<&str> = stdout.lines().collect();
+ assert_eq!(
+ lines,
+ ["COUNT", "40"],
+ "expected 40 customers, got: {stdout}"
+ );
+ Ok(())
+Copilot generated this review using guidance from repository custom instructions.
+@jgarzik Reply...
+src/vm/bytecode.rs
+Comment on lines 234 to +237
+ // Window function operations
+ WindowAggStep, // P1 = func type, P2 = value reg, P3 = accum reg, P4 = window spec. Step window aggregate.
+ WindowValue, // P1 = accum reg, P2 = dest reg, P3 = func type. Get current window value.
+ WindowValue,
+ WindowFinalize, // Back-fill result column P1 with each partition's final value // P1 = accum reg, P2 = dest reg, P3 = func type. Get current window value.
+Copilot commented now
+Copilot
+Copilot
+AI
+1m ago
+Low
+The inline comment for WindowFinalize appears to be copy/pasted from WindowValue and describes the wrong operands, which makes the opcode contract unclear for future changes.
+
+Copilot uses AI. Check for mistakes.
+
diff --git a/doc/database.md b/doc/database.md
index e8c9ebd..462abc5 100644
--- a/doc/database.md
+++ b/doc/database.md
@@ -30,7 +30,15 @@ Tables maintain an internal mapping of column names to their indices for efficie
- **Rows**: Represented as vectors of values with one element per column
- **Columns**: Identified by name, with automatic type inference based on content
- **Schema**: Dynamically determined from the file's header row
-- **Column Types**: Column types are not explicitly declared but inferred at runtime
+- **Column Types**: Not declared; inferred at runtime **per value**, so a single
+ column may hold integers, floats and text
+- **NULL**: An empty field reads as NULL. Comparisons follow SQL three-valued
+ logic, so a comparison involving NULL is UNKNOWN and its row is filtered out
+ rather than raising an error. Sorting and grouping treat NULL as equal to
+ NULL, which SQL comparison must not — the two use separate paths deliberately
+- **Coercion**: A comparison mixing a number and a string converts the string
+ when it parses as a number, and compares textually otherwise. See
+ [SQL Reference](sql_reference.md) for the exact rules
### Table Lifecycle
@@ -59,14 +67,43 @@ The `Table` struct represents an in-memory table with:
- Methods for accessing and manipulating rows
- Projection capabilities (selecting subsets of columns)
-### SQL Executor
+### SQL Execution: a bytecode VM
-The `SqlExecutor` implements SQL parsing and execution:
-- Uses `sqlparser` crate to parse SQL statements
-- Converts parsed AST to operations on in-memory tables
-- Handles WHERE clause evaluation
-- Maintains a set of modified table names (`modified_tables`) to track changes
-- Provides `save_modified_tables()` method that only writes back tables that were actually modified
+SQL is not interpreted against the tables directly. It is compiled to bytecode
+and run on a register virtual machine, following SQLite's architecture:
+
+1. **Parse** — the `sqlparser` crate produces an AST
+2. **Compile** — `vm::compiler` lowers the AST to a bytecode program
+3. **Execute** — `vm::engine` runs that program against table cursors
+
+`SqlExecutor` is a thin wrapper over this pipeline: it forwards to the VM,
+tracks which tables were modified, and provides `save_modified_tables()`, which
+writes back only tables that actually changed.
+
+The compiler has one recursive expression compiler, `code_expr` (the analogue
+of SQLite's `sqlite3ExprCode`), used from every position an expression can
+appear: the SELECT list, WHERE, HAVING, ORDER BY, SET clauses, join conditions,
+function arguments and correlated subqueries. Name resolution is carried
+separately in a `NameCtx`, which models the visible FROM sources, whether a
+column lives behind a cursor or in an already-loaded register, and the
+inner/outer scopes a correlated subquery needs. A node supported in one
+position is therefore supported in all of them.
+
+### Statement Execution Model
+
+A script is parsed once, then each statement is compiled and executed on its
+own, with its modifications applied before the next statement compiles. Two
+consequences follow:
+
+- Each statement produces its own result set, with its own column schema. They
+ are not merged.
+- A statement observes the database as the previous one left it, which is what
+ `CREATE TABLE t; INSERT INTO t ...; SELECT * FROM t` requires.
+
+A subquery in `FROM` (a derived table) is materialized before its enclosing
+statement compiles: the subquery runs, its result is registered under the
+alias, and the statement is rewritten to reference that name. The registration
+lasts only for that statement.
### File Handlers
@@ -152,6 +189,11 @@ The storage backend is selected automatically based on the data source:
- **Stdin/pipes**: Uses traditional in-memory storage
- **Write operations**: Mmap tables convert to memory on first modification
+### Row Identity Across Updates
+
+UPDATE replaces a row in place rather than deleting and re-appending it, so a
+table keeps its row order and `--write` does not reorder the user's file.
+
### Write Operation Behavior
When a table with mmap storage is modified (INSERT, UPDATE, DELETE):
@@ -197,9 +239,15 @@ The database engine supports combining data from multiple tables through SQL-sta
### Join Syntax
-- **ON constraints**: Supported for specifying join conditions (`JOIN t2 ON t1.id = t2.id`)
-- **USING constraints**: Not supported
+- **ON constraints**: Supported, and the condition is a full expression, not
+ just a single comparison: `ON a.id = b.id AND UPPER(a.name) = 'X'`
+- **Table aliases**: Supported (`FROM users u JOIN orders o ON u.id = o.user_id`)
+- **USING / NATURAL**: Not supported
- **Multi-table joins**: Supports chaining 3+ tables in a single query
+- **Aggregates over joins**: Supported for inner joins, which are rewritten to
+ the equivalent comma join. Deliberately NOT applied to outer joins, since
+ moving the ON condition into WHERE would discard the NULL-extended rows an
+ outer join exists to produce
### Technical Implementation
@@ -213,9 +261,16 @@ The database engine supports combining data from multiple tables through SQL-sta
The database engine has several architectural limitations:
- **No Index Structure**: All operations perform full table scans
-- **No Transaction Support**: Operations are applied immediately with no rollback capability
-- **Schema Flexibility**: Types are inferred rather than enforced
+- **No Transaction Support**: A statement's modifications are applied once it
+ completes; there is no BEGIN/COMMIT and no rollback across statements
+- **Schema Flexibility**: Types are inferred per value rather than enforced
- **No Constraints System**: Referential integrity not enforced
+- **Whole-dataset residency**: A table must fit in memory, or be addressable
+ via mmap
+
+Unsupported SQL of note: CTEs (`WITH`), explicit window frames
+(`ROWS BETWEEN ...`), `USING`/`NATURAL` joins, and aggregates over outer
+joins.
---
diff --git a/doc/sql_reference.md b/doc/sql_reference.md
index 47ef66f..0c68d6f 100644
--- a/doc/sql_reference.md
+++ b/doc/sql_reference.md
@@ -125,6 +125,63 @@ SELECT * FROM t1
JOIN t3 ON t2.id = t3.t2_id
```
+## Subqueries
+
+```sql
+-- Scalar subquery
+SELECT name FROM employees WHERE salary = (SELECT MAX(salary) FROM employees)
+
+-- IN / NOT IN
+SELECT name FROM users WHERE id IN (SELECT user_id FROM orders)
+SELECT name FROM users WHERE id NOT IN (SELECT user_id FROM orders)
+
+-- EXISTS / NOT EXISTS
+SELECT name FROM users WHERE EXISTS (SELECT 1 FROM orders WHERE orders.user_id = users.id)
+
+-- Correlated: the inner query references the outer row
+SELECT name FROM employees e
+ WHERE salary > (SELECT AVG(salary) FROM employees WHERE department = e.department)
+```
+
+```sql
+-- Derived table: a subquery in FROM, which must be given an alias
+SELECT name FROM (SELECT name, salary FROM employees WHERE salary > 70000) t
+SELECT COUNT(*) FROM (SELECT department FROM employees GROUP BY department) t
+```
+
+A derived table is materialized before the outer query runs and exists only for
+the statement that declares it. Its alias may not shadow a real table.
+
+## Window Functions
+
+```sql
+SELECT name, ROW_NUMBER() OVER (ORDER BY salary DESC) FROM employees
+SELECT name, RANK() OVER (ORDER BY department) FROM employees
+SELECT name, DENSE_RANK() OVER (ORDER BY department) FROM employees
+SELECT name, LAG(salary) OVER (ORDER BY salary) FROM employees
+SELECT name, LEAD(salary) OVER (ORDER BY salary) FROM employees
+
+-- Aggregates over a window
+SELECT name, SUM(salary) OVER (PARTITION BY department) FROM employees
+```
+
+Supported functions: `ROW_NUMBER`, `RANK`, `DENSE_RANK`, `LAG`, `LEAD`, and the
+aggregates `COUNT`, `SUM`, `AVG`, `MIN`, `MAX`.
+
+The frame follows the standard: **without** `ORDER BY` in the `OVER` clause the
+frame is the whole partition, so every row sees the partition total; **with**
+`ORDER BY` the frame grows row by row, giving a running value.
+
+```sql
+-- 210000 on every Engineering row
+SELECT name, SUM(salary) OVER (PARTITION BY department) FROM employees
+
+-- 65000, 135000, 210000 across the Engineering rows
+SELECT name, SUM(salary) OVER (PARTITION BY department ORDER BY salary) FROM employees
+```
+
+Explicit frame clauses (`ROWS BETWEEN ...`) are not supported.
+
## Set Operations
```sql
@@ -191,6 +248,28 @@ ALTER TABLE name ADD COLUMN col_name TYPE
TRUNCATE TABLE name
```
+## NULL handling
+
+An empty field reads as NULL.
+
+Comparison follows SQL three-valued logic: any comparison involving NULL is
+UNKNOWN, and a row whose `WHERE` evaluates to UNKNOWN is not returned. This
+means a NULL row satisfies neither `x > 5` nor `x <= 5`, and `NULL = NULL` is
+UNKNOWN rather than true. Use `IS NULL` / `IS NOT NULL` to test for NULL.
+
+Aggregates skip NULLs: `COUNT(*)` counts rows while `COUNT(col)` counts
+non-NULL values. `ORDER BY` sorts NULLs first.
+
+## Type coercion
+
+Values are typed per cell as they are read, so one column may hold integers,
+floats and text.
+
+When a comparison mixes a number and a string, the string is converted to a
+number if it parses as one, and the two are compared textually otherwise. This
+matches arithmetic, so `x + '1'` and `x > '1'` agree — awk-like rather than
+strict SQL, which suits untyped delimited input.
+
## Writeback
Modifications (INSERT, UPDATE, DELETE) remain in-memory unless `--write` flag is specified.
diff --git a/doc/user_guide.md b/doc/user_guide.md
index 2cd0f94..0ff9633 100644
--- a/doc/user_guide.md
+++ b/doc/user_guide.md
@@ -52,6 +52,8 @@ Sqawk is an SQL-based command-line tool for processing delimiter-separated files
- No database setup or schema definition required
- Automatic type inference and cross-file operations
- Powerful SQL dialect including joins, sorting, filtering, and aggregations
+- Subqueries (scalar, `IN`, `EXISTS`, correlated) and derived tables
+- Set operations (UNION, INTERSECT, EXCEPT) and window functions
- Interactive REPL mode for SQL exploration and execution
- Safe operation with explicit write-back control
@@ -75,7 +77,7 @@ To build from source:
1. Clone the repository:
```sh
- git clone https://github.com/username/sqawk.git
+ git clone https://github.com/jgarzik/sqawk.git
cd sqawk
```
@@ -178,13 +180,22 @@ This opens an interactive SQL prompt where you can:
| Command | Description |
|---------|-------------|
-| `.exit` or `.quit` | Exit the REPL |
-| `.help` | Display help information about available commands |
-| `.save [table]` | Immediately save changes to all modified tables or a specific table |
-| `.schema [table]` | Show schema for a specific table or all tables |
-| `.tables` | List all available tables |
-| `.verbose [on/off]` | Toggle verbose mode on/off |
-| `.write [on/off]` | Toggle write mode on/off (default is off) |
+| `.cd DIRECTORY` | Change the working directory |
+| `.changes [on\|off]` | Show the number of rows changed by each statement |
+| `.exit [CODE]` | Exit the REPL, optionally with an exit code |
+| `.help` | List the available commands |
+| `.load [TABLE=]FILE` | Load `FILE`, optionally under the name `TABLE` |
+| `.print STRING...` | Print a literal string |
+| `.quit` | Exit the REPL |
+| `.save [TABLE]` | Save changes to all modified tables, or just `TABLE` |
+| `.schema [TABLE]` | Show the schema for `TABLE`, or for every table |
+| `.show [WHAT]` | Show current settings and status |
+| `.stats [on\|off]` | Toggle statistics display |
+| `.tables [PATTERN]` | List tables, optionally matching a LIKE pattern |
+| `.version` | Show source, library and compiler versions |
+| `.write [on\|off]` | Toggle writing changes to files (default off) |
+
+`.help` inside the REPL always lists the authoritative set.
**Example REPL Session:**
@@ -663,7 +674,10 @@ sqawk> .exit
## Working with Large Files
-Sqawk loads all data into memory, which provides excellent performance but requires consideration when working with large files:
+Sqawk holds a whole table at once. On-disk files are memory-mapped, so loading
+is zero-copy and read-only queries allocate almost nothing; a table is copied to
+the heap the first time it is modified. Either way the dataset must fit, which
+is worth planning for with large files:
**Tips for handling large files:**
@@ -687,7 +701,7 @@ Sqawk loads all data into memory, which provides excellent performance but requi
4. **Monitor memory usage**: Particularly when joining large tables, be aware of memory constraints
```sh
# Using a more targeted join condition reduces memory requirements
- sqawk -s "SELECT a.id, b.name FROM large_a INNER JOIN large_b ON a.id = b.id WHERE a.region = 'West'" large_a.csv large_b.csv
+ sqawk -s "SELECT a.id, b.name FROM large_a a INNER JOIN large_b b ON a.id = b.id WHERE a.region = 'West'" large_a.csv large_b.csv
```
## Troubleshooting
@@ -704,12 +718,22 @@ Sqawk loads all data into memory, which provides excellent performance but requi
- For tab-delimited files, use `-F '\t'`
- Ensure consistent delimiters throughout your files
-3. **Type conversion errors**:
- - Sqawk automatically infers types but sometimes needs hints
- - Use explicit casts in SQL when needed: `CAST(value AS INT)`
- - Check that numeric columns don't contain non-numeric characters
-
-4. **CSV parsing errors with malformed rows**:
+3. **Unexpected results from comparisons**:
+ - Types are inferred per value, so one column can hold numbers and text
+ - Comparing a number against a numeric string coerces and compares
+ numerically, so `WHERE salary > '60000'` behaves as you would expect
+ - Use explicit casts when you want to force one interpretation:
+ `CAST(value AS INT)`
+ - See the SQL reference for the full coercion and NULL rules
+
+4. **Rows unexpectedly missing from results**:
+ - An empty field reads as NULL, and any comparison involving NULL is
+ UNKNOWN, so such rows satisfy neither `x > 5` nor `x <= 5`
+ - This follows SQL, and is not an error -- use `IS NULL` / `IS NOT NULL`
+ to test for NULL explicitly
+ - `COUNT(*)` counts rows while `COUNT(col)` counts non-NULL values
+
+5. **CSV parsing errors with malformed rows**:
- Error messages about "field count mismatch" indicate rows with inconsistent numbers of fields
- Error messages include line numbers to help locate problematic rows
- Common causes include:
@@ -718,27 +742,27 @@ Sqawk loads all data into memory, which provides excellent performance but requi
- Newlines within quoted fields
- Use the error recovery options described in the File Format Support section to handle malformed rows
-5. **Issues with comment lines**:
+6. **Issues with comment lines**:
- Comments must start at the beginning of a line with the comment character
- Comment characters appearing within data (not at the start of a line) are treated as regular data
- If you're seeing unexpected parsing errors, check if comment lines are properly formatted
-6. **Memory limitations**:
+7. **Memory limitations**:
- If processing very large files, filter data early in your queries
- Consider processing in batches or using more targeted queries
- Select only the columns you need rather than using SELECT *
-7. **Changes not saved**:
+8. **Changes not saved**:
- Remember to use the `--write` flag to save changes
- Only modified tables are written back
- Check verbose output (`-v`) to confirm which tables were modified
-8. **SQL syntax errors**:
+9. **SQL syntax errors**:
- Try running your query in interactive mode to get immediate feedback
- Use the `-v` verbose flag to see the exact SQL being executed
- Verify SQL statement syntax, particularly quotes, parentheses, and required clauses
-9. **Special characters in files**:
+10. **Special characters in files**:
- For files with quotes or special characters, Sqawk follows CSV escaping rules
- If encountering parsing issues, check for malformed CSV data
@@ -766,14 +790,15 @@ For more help, use the verbose mode (`-v`) to see detailed information about pro
2. **Use version control** or backups before modifying important data files
3. **Qualify column names** with table names in multi-table queries
4. **Use verbose mode** (`-v`) when learning or debugging
-5. **Chain SQL statements** for complex operations rather than using complex subqueries
+5. **Chain SQL statements** with `-s` or `;` when it reads more clearly; each
+ statement produces its own result and sees the effects of the ones before it
6. **Test on sample data** before processing large files
### Additional Resources
- [SQL Language Reference](sql_reference.md) - Complete guide to Sqawk's SQL dialect
-- [GitHub Repository](https://github.com/username/sqawk) - Source code and issue tracking
-- [Release Notes](https://github.com/username/sqawk/releases) - Latest features and bug fixes
+- [GitHub Repository](https://github.com/jgarzik/sqawk) - Source code and issue tracking
+- [Release Notes](https://github.com/jgarzik/sqawk/releases) - Latest features and bug fixes
---
diff --git a/src/join.rs b/src/join.rs
deleted file mode 100644
index 54a122c..0000000
--- a/src/join.rs
+++ /dev/null
@@ -1,172 +0,0 @@
-//! Join module for the sqawk query engine
-//!
-//! This module implements table join operations for SQL queries in the sqawk utility.
-//! It provides functionality for:
-//!
-//! - Cross joins (Cartesian product of rows from both tables)
-//! - Inner joins with filtering based on specified conditions
-//! - A flexible join type system mapping SQL join types to internal representations
-//! - Column naming with qualification to prevent ambiguity in join results
-//!
-//! Currently, the module fully implements cross joins and has a foundation for more
-//! complex join types (inner, left, right, full) with placeholder implementations
-//! that will be expanded in future versions.
-
-use sqlparser::ast::{Expr, JoinOperator};
-
-use crate::error::{SqawkError, SqawkResult};
-use crate::table::Table;
-
-/// Join types supported by sqawk
-#[derive(Debug, PartialEq, Clone, Copy)]
-pub enum JoinType {
- /// Inner join - returns rows when there is a match in both tables
- Inner,
- /// Left join - returns all rows from the left table and matching rows from the right
- Left,
- /// Right join - returns all rows from the right table and matching rows from the left
- Right,
- /// Full join - returns rows when there is a match in one of the tables
- Full,
- /// Cross join - returns the Cartesian product of rows from both tables
- Cross,
-}
-
-impl From<&JoinOperator> for JoinType {
- fn from(op: &JoinOperator) -> Self {
- match op {
- JoinOperator::Inner(_) => JoinType::Inner,
- JoinOperator::LeftOuter(_) => JoinType::Left,
- JoinOperator::RightOuter(_) => JoinType::Right,
- JoinOperator::FullOuter(_) => JoinType::Full,
- JoinOperator::CrossJoin => JoinType::Cross,
- // Default to inner join for other types
- _ => JoinType::Inner,
- }
- }
-}
-
-/// Executor for join operations
-///
-/// This struct handles the execution of different types of joins between tables.
-pub struct JoinExecutor {
- // Placeholder for future join state
-}
-
-impl Default for JoinExecutor {
- fn default() -> Self {
- Self::new()
- }
-}
-
-impl JoinExecutor {
- /// Create a new join executor
- pub fn new() -> Self {
- JoinExecutor {}
- }
-
- /// Execute a join operation between two tables
- ///
- /// # Arguments
- /// * `left` - The left table
- /// * `right` - The right table
- /// * `join_type` - The type of join to perform
- /// * `condition` - Optional join condition (for non-cross joins)
- ///
- /// # Returns
- /// * The resulting joined table
- pub fn execute_join(
- &mut self,
- left: &Table,
- right: &Table,
- join_type: JoinType,
- condition: Option<&Expr>,
- ) -> SqawkResult
{
- match join_type {
- JoinType::Cross => self.execute_cross_join(left, right),
- JoinType::Inner => {
- if let Some(on_expr) = condition {
- self.execute_inner_join(left, right, on_expr)
- } else {
- // Inner join without condition is equivalent to cross join
- self.execute_cross_join(left, right)
- }
- }
- _ => Err(SqawkError::UnsupportedSqlFeature(format!(
- "Join type {:?} is not implemented yet",
- join_type
- ))),
- }
- }
-
- /// Execute a cross join (Cartesian product)
- ///
- /// # Arguments
- /// * `left` - The left table
- /// * `right` - The right table
- ///
- /// # Returns
- /// * The resulting cross-joined table
- fn execute_cross_join(&self, left: &Table, right: &Table) -> SqawkResult {
- // Use the Table's cross_join method
- left.cross_join(right)
- }
-
- /// Execute an inner join with an ON condition
- ///
- /// # Arguments
- /// * `left` - The left table
- /// * `right` - The right table
- /// * `on_expr` - The ON condition expression
- ///
- /// # Returns
- /// * The resulting inner-joined table
- fn execute_inner_join(
- &self,
- left: &Table,
- right: &Table,
- _on_expr: &Expr,
- ) -> SqawkResult {
- // For now, we'll just do a cross join
- // In a real implementation, we would evaluate the ON condition for each row combination
-
- // Create a new table with combined columns from both tables
- let mut result_columns = Vec::with_capacity(left.column_count() + right.column_count());
-
- // Add qualified columns from left table
- for col in left.columns() {
- result_columns.push(format!("left.{}", col));
- }
-
- // Add qualified columns from right table
- for col in right.columns() {
- result_columns.push(format!("right.{}", col));
- }
-
- let mut result = Table::new("joined", result_columns, None);
-
- // The naive nested loop join approach
- for left_row in left.rows().iter() {
- for right_row in right.rows().iter() {
- // Join the rows
- let mut new_row = Vec::with_capacity(left_row.len() + right_row.len());
-
- // Add values from left row
- for value in left_row {
- new_row.push(value.clone());
- }
-
- // Add values from right row
- for value in right_row {
- new_row.push(value.clone());
- }
-
- // Add the combined row to the result table
- // TODO: Evaluate ON condition here
- let _ = result.add_row(new_row);
- }
- }
-
- Ok(result)
- }
-}
diff --git a/src/main.rs b/src/main.rs
index f29eaf2..0aaeab5 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -140,23 +140,21 @@ fn main() -> Result<()> {
.execute(sql)
.with_context(|| format!("Failed to execute SQL: {sql}"))?;
- // Step 4: Output results to stdout (for SELECT queries)
- match result {
- // For SELECT queries that return data
- Some(table) => {
- if config.verbose() {
- let row_count = table.row_count();
- println!("Query returned {row_count} rows");
- }
- // Print the result table in delimiter-separated format
- table.print_to_stdout()?;
- }
- // For statements that don't return data (UPDATE, DELETE, INSERT)
- None => {
- if config.verbose() {
- println!("Query executed successfully (no results to display)");
- }
+ // Step 4: Output results to stdout (for SELECT queries).
+ //
+ // One result set per statement, each with its own header. They used to
+ // share a single buffer and schema, so `SELECT a; SELECT b` printed
+ // every row under the last statement's header.
+ for table in &result {
+ if config.verbose() {
+ let row_count = table.row_count();
+ println!("Query returned {row_count} rows");
}
+ // Print the result table in delimiter-separated format
+ table.print_to_stdout()?;
+ }
+ if result.is_empty() && config.verbose() {
+ println!("Query executed successfully (no results to display)");
}
}
diff --git a/src/repl.rs b/src/repl.rs
index 26b22d6..0eb31db 100644
--- a/src/repl.rs
+++ b/src/repl.rs
@@ -452,10 +452,11 @@ impl<'a> Repl<'a> {
} else if self.show_changes {
// For non-SELECT statements that don't return rows (INSERT, UPDATE, DELETE)
// Try to display the number of affected rows if show_changes is enabled
+ // Reported even when zero. With `.changes on` a statement that
+ // matched nothing otherwise produced no output at all, leaving it
+ // ambiguous whether it had run.
if let Ok(affected_rows) = self.executor.get_affected_row_count() {
- if affected_rows > 0 {
- println!("{} rows affected", affected_rows);
- }
+ println!("{} rows affected", affected_rows);
}
}
@@ -692,8 +693,12 @@ impl<'a> Repl<'a> {
}
/// Show version information
+ ///
+ /// Read from the manifest rather than written out, so it cannot drift.
+ /// This reported 0.1.1 for a long time while the crate was at 0.8.0.
fn show_version(&self) -> Result<()> {
- println!("Sqawk version 0.1.1");
+ println!("Sqawk version {}", env!("CARGO_PKG_VERSION"));
+ println!("sqlparser {}", crate::vm::SQLPARSER_VERSION);
Ok(())
}
diff --git a/src/sql_executor.rs b/src/sql_executor.rs
index e78542a..0dc8c69 100644
--- a/src/sql_executor.rs
+++ b/src/sql_executor.rs
@@ -69,8 +69,10 @@ impl<'a> SqlExecutor<'a> {
/// * `sql` - SQL statement to execute
///
/// # Returns
- /// * `SqawkResult>` - Result of the operation, possibly containing a table
- pub fn execute(&mut self, sql: &str) -> SqawkResult > {
+ /// * One result table per statement that produced rows, in order. A
+ /// multi-statement script yields one result set per statement; they are
+ /// deliberately NOT merged, since each has its own column schema.
+ pub fn execute(&mut self, sql: &str) -> SqawkResult> {
if self.config.verbose() {
println!("Executing SQL: {}", sql);
}
@@ -86,18 +88,15 @@ impl<'a> SqlExecutor<'a> {
// Track affected rows from the last statement
self.affected_rows = result.affected_rows;
- // Set delimiter on result table to match config (for consistent output format)
- let table = match result.table {
- Some(mut t) => {
- if let Some(delim) = self.config.field_separator() {
- t.set_delimiter(delim);
- }
- Some(t)
+ // Set delimiter on each result table to match config
+ let mut tables = result.tables;
+ if let Some(delim) = self.config.field_separator() {
+ for t in &mut tables {
+ t.set_delimiter(delim.clone());
}
- None => None,
- };
+ }
- Ok(table)
+ Ok(tables)
}
/// Save all modified tables back to their source files
@@ -198,7 +197,9 @@ impl<'a> SqlExecutor<'a> {
pub fn execute_sql(&mut self, sql: &str) -> Result> {
let result = self.execute(sql)?;
- Ok(result.map(|table| ResultSet {
+ // The REPL shows one result set at a time; a multi-statement line
+ // reports the last statement's.
+ Ok(result.last().map(|table| ResultSet {
columns: table.columns().to_vec(),
rows: table.rows_as_strings(),
}))
diff --git a/src/table.rs b/src/table.rs
index 880b733..64a2def 100644
--- a/src/table.rs
+++ b/src/table.rs
@@ -854,6 +854,14 @@ impl Table {
/// Returns the name of the table as a string slice.
/// This is useful for operations that need to access the table's name
/// such as logging, error messages, or generating SQL output.
+ /// Rename the table.
+ ///
+ /// Used when a derived table is materialized and registered under its
+ /// alias, so the table reports the name the outer query refers to.
+ pub fn set_name(&mut self, name: String) {
+ self.name = name;
+ }
+
pub fn name(&self) -> &str {
&self.name
}
@@ -1022,6 +1030,26 @@ impl Table {
self.modified = true;
}
+ /// Replace a single row in place, preserving its position.
+ ///
+ /// Used by UPDATE. Expressing an update as delete-plus-insert appends the
+ /// new row, which reorders the table -- and with --write, the user's file.
+ pub fn replace_row(&mut self, row_index: usize, row: Row) -> SqawkResult<()> {
+ self.storage.ensure_mutable();
+ let rows = self.storage.rows_mut().ok_or_else(|| {
+ SqawkError::InvalidSqlQuery("Table storage is not mutable".to_string())
+ })?;
+ if row_index >= rows.len() {
+ return Err(SqawkError::InvalidSqlQuery(format!(
+ "Row index {} out of range for table '{}'",
+ row_index, self.name
+ )));
+ }
+ rows[row_index] = row;
+ self.modified = true;
+ Ok(())
+ }
+
/// Add a column to the table with a specified data type
#[cfg(test)]
pub fn add_column(&mut self, name: String, data_type_str: String) {
diff --git a/src/vm/ast_compat.rs b/src/vm/ast_compat.rs
new file mode 100644
index 0000000..eaa569c
--- /dev/null
+++ b/src/vm/ast_compat.rs
@@ -0,0 +1,185 @@
+//! Small accessors that normalize sqlparser AST shapes.
+//!
+//! sqlparser treats every AST change as breaking, and a few of its types have
+//! grown structure that the compiler does not care about. Rather than spread
+//! that structure across ~150 call sites, the compiler goes through the
+//! accessors here.
+//!
+//! This is deliberately thin: it adapts *shape*, never *meaning*. Anything that
+//! requires a semantic decision (which join kinds are supported, how DISTINCT
+//! affects an aggregate) belongs in the compiler proper.
+
+use sqlparser::ast::{
+ Assignment, AssignmentTarget, CreateTableOptions, DuplicateTreatment, Expr, FromTable,
+ Function, FunctionArg, FunctionArguments, GroupByExpr, LimitClause, ObjectName, OrderBy,
+ OrderByExpr, OrderByKind, Query, SqlOption, TableObject, TableWithJoins,
+};
+
+use crate::error::{SqawkError, SqawkResult};
+
+/// The positional argument list of a function call.
+///
+/// `FunctionArguments` distinguishes three cases that the compiler mostly does
+/// not need to: no parentheses at all (`CURRENT_TIMESTAMP`), a bare subquery
+/// argument, and an ordinary parenthesized list. Only the last carries
+/// positional arguments; the other two yield an empty slice.
+pub(crate) fn func_args(func: &Function) -> &[FunctionArg] {
+ match &func.args {
+ FunctionArguments::List(list) => &list.args,
+ FunctionArguments::None | FunctionArguments::Subquery(_) => &[],
+ }
+}
+
+/// Whether a function call was written with `DISTINCT`, as in
+/// `COUNT(DISTINCT department)`.
+///
+/// Older sqlparser exposed this as a plain `Function::distinct` bool; it now
+/// lives in the argument list as a `DuplicateTreatment`. `ALL` is the default
+/// and means the same thing as omitting it.
+pub(crate) fn func_is_distinct(func: &Function) -> bool {
+ match &func.args {
+ FunctionArguments::List(list) => {
+ matches!(list.duplicate_treatment, Some(DuplicateTreatment::Distinct))
+ }
+ FunctionArguments::None | FunctionArguments::Subquery(_) => false,
+ }
+}
+
+/// The GROUP BY key expressions.
+///
+/// `GROUP BY ALL` has no explicit key list; it yields an empty slice here and
+/// is rejected by the caller that cares.
+pub(crate) fn group_by_exprs(group_by: &GroupByExpr) -> &[sqlparser::ast::Expr] {
+ match group_by {
+ GroupByExpr::Expressions(exprs, _) => exprs,
+ GroupByExpr::All(_) => &[],
+ }
+}
+
+/// Whether this GROUP BY is the unsupported `GROUP BY ALL` form.
+pub(crate) fn is_group_by_all(group_by: &GroupByExpr) -> bool {
+ matches!(group_by, GroupByExpr::All(_))
+}
+
+/// The ORDER BY key expressions of a query.
+///
+/// `Query::order_by` is now an `Option` whose kind distinguishes an
+/// explicit key list from `ORDER BY ALL`. Absent ORDER BY and `ORDER BY ALL`
+/// both yield an empty slice; the caller that cares rejects the latter.
+pub(crate) fn query_order_by(query: &Query) -> &[OrderByExpr] {
+ match &query.order_by {
+ Some(OrderBy {
+ kind: OrderByKind::Expressions(exprs),
+ ..
+ }) => exprs,
+ _ => &[],
+ }
+}
+
+/// Whether this query uses the unsupported `ORDER BY ALL` form.
+pub(crate) fn is_order_by_all(query: &Query) -> bool {
+ matches!(
+ &query.order_by,
+ Some(OrderBy {
+ kind: OrderByKind::All(_),
+ ..
+ })
+ )
+}
+
+/// Sort direction for one ORDER BY key: `true` ascending, `false` descending.
+///
+/// Ascending is the SQL default when the direction is omitted. `asc` moved
+/// from `OrderByExpr` into a nested `OrderByOptions`.
+pub(crate) fn order_by_is_asc(expr: &OrderByExpr) -> bool {
+ expr.options.asc.unwrap_or(true)
+}
+
+/// The LIMIT expression, if any.
+///
+/// `Query::limit` and `Query::offset` were merged into a single
+/// `limit_clause` that also models MySQL's reversed `LIMIT , `.
+pub(crate) fn query_limit(query: &Query) -> Option<&Expr> {
+ match &query.limit_clause {
+ Some(LimitClause::LimitOffset { limit, .. }) => limit.as_ref(),
+ Some(LimitClause::OffsetCommaLimit { limit, .. }) => Some(limit),
+ None => None,
+ }
+}
+
+/// The OFFSET expression, if any. See [`query_limit`].
+pub(crate) fn query_offset(query: &Query) -> Option<&Expr> {
+ match &query.limit_clause {
+ Some(LimitClause::LimitOffset { offset, .. }) => offset.as_ref().map(|o| &o.value),
+ Some(LimitClause::OffsetCommaLimit { offset, .. }) => Some(offset),
+ None => None,
+ }
+}
+
+/// The target table of an INSERT.
+///
+/// `Insert::table` is now a `TableObject`, which can also be a table function.
+pub(crate) fn insert_target(table: &TableObject) -> SqawkResult<&ObjectName> {
+ match table {
+ TableObject::TableName(name) => Ok(name),
+ TableObject::TableFunction(_) | TableObject::TableQuery(_) => Err(
+ SqawkError::UnsupportedSqlFeature("INSERT target must be a plain table name".into()),
+ ),
+ }
+}
+
+/// The source tables of a DELETE.
+///
+/// `Delete::from` distinguishes `DELETE FROM t` from `DELETE t`; sqawk treats
+/// them the same.
+pub(crate) fn delete_from(from: &FromTable) -> &[TableWithJoins] {
+ match from {
+ FromTable::WithFromKeyword(t) | FromTable::WithoutKeyword(t) => t,
+ }
+}
+
+/// The assigned column of an UPDATE `SET` clause.
+///
+/// `Assignment::id` (a `Vec`) became `target: AssignmentTarget`.
+pub(crate) fn assignment_column(assignment: &Assignment) -> SqawkResult {
+ match &assignment.target {
+ AssignmentTarget::ColumnName(name) => Ok(object_name_last(name)),
+ AssignmentTarget::Tuple(_) => Err(SqawkError::UnsupportedSqlFeature(
+ "Tuple assignment in UPDATE is not supported".into(),
+ )),
+ }
+}
+
+/// The `WITH (...)` options attached to a CREATE TABLE.
+///
+/// Previously a plain `Vec`; the dialect-specific spellings are now
+/// distinguished by `CreateTableOptions`, which sqawk does not care about.
+pub(crate) fn create_table_options(options: &CreateTableOptions) -> &[SqlOption] {
+ match options {
+ CreateTableOptions::With(o)
+ | CreateTableOptions::Options(o)
+ | CreateTableOptions::Plain(o)
+ | CreateTableOptions::TableProperties(o) => o,
+ _ => &[],
+ }
+}
+
+/// A CREATE TABLE option as a `(name, value)` pair, for the options sqawk
+/// understands (`delimiter`, `header`, ...). Non key-value forms yield `None`.
+pub(crate) fn sql_option_key_value(option: &SqlOption) -> Option<(String, String)> {
+ match option {
+ SqlOption::KeyValue { key, value } => Some((key.value.clone(), value.to_string())),
+ _ => None,
+ }
+}
+
+/// The final segment of an `ObjectName` -- the bare table or function name.
+pub(crate) fn object_name_last(name: &ObjectName) -> String {
+ match name.0.last() {
+ Some(part) => match part.as_ident() {
+ Some(ident) => ident.value.clone(),
+ None => part.to_string(),
+ },
+ None => String::new(),
+ }
+}
diff --git a/src/vm/bytecode.rs b/src/vm/bytecode.rs
index e0cc97b..e7f9dbe 100644
--- a/src/vm/bytecode.rs
+++ b/src/vm/bytecode.rs
@@ -24,6 +24,17 @@ pub const AGG_MIN: i64 = 3;
/// MAX aggregate function type
pub const AGG_MAX: i64 = 4;
+/// Flag bit OR-ed into an AggStep/AggFinal function type to request DISTINCT.
+///
+/// Carried as a flag rather than as five extra constants so that DISTINCT
+/// composes with every aggregate: `COUNT(DISTINCT x)`, `SUM(DISTINCT x)` and
+/// the rest all use the same mechanism. Mask with `AGG_TYPE_MASK` to recover
+/// the base function.
+pub const AGG_DISTINCT: i64 = 0x100;
+
+/// Mask selecting the base aggregate function out of a function-type word.
+pub const AGG_TYPE_MASK: i64 = 0xFF;
+
/// A column in a result schema
#[derive(Debug, Clone)]
pub struct ResultColumn {
@@ -90,6 +101,7 @@ pub enum OpCode {
Column, // Read column value into register
InsertRow, // Insert row from registers P2..P2+P3 into cursor P1
DeleteRow, // Delete current row at cursor P1
+ UpdateRow, // Replace current row at cursor P1 with registers P2..P2+P3 IN PLACE
// Data manipulation
Integer, // Load integer constant
@@ -123,6 +135,29 @@ pub enum OpCode {
// Null operations
IsNull, // Set P2 to 1 if P1 is NULL, 0 otherwise
+ // Logical negation
+ //
+ // Distinct from the comparison opcodes because it must propagate NULL:
+ // NOT NULL is NULL (UNKNOWN), not true. Replaces the several hand-rolled
+ // inversions used for NOT LIKE / NOT IN / IS NOT NULL and makes a bare
+ // `NOT ` expressible at all.
+ Not, // P2 = NULL if P1 is NULL, else 1 if P1 is falsy, else 0
+
+ // Vector comparison, modelled on SQLite's OP_Compare / OP_Jump.
+ //
+ // Compare sets an internal flag from a pairwise comparison of two register
+ // vectors; Jump then branches three ways on that flag. Together they make
+ // multi-column key comparison two instructions regardless of key width,
+ // which is what GROUP BY needs to detect a group change across ALL key
+ // columns rather than just the first.
+ //
+ // These use INTERNAL ordering, not SQL comparison: NULL sorts equal to
+ // NULL, so a NULL group key groups with other NULL keys. SQL `=` must not
+ // behave that way, which is exactly why this is a separate opcode rather
+ // than a reuse of Eq.
+ Compare, // Compare P3 registers starting at P1 against P3 starting at P2
+ Jump, // Branch on the last Compare: P1 if Less, P2 if Equal, P3 if Greater
+
// Conditional jumps
IfZ, // Jump to P2 if register P1 contains 0
IfPos, // Jump to P2 if register P1 is positive (> 0)
@@ -132,7 +167,8 @@ pub enum OpCode {
Noop, // No operation
// Set operations
- Distinct, // Remove duplicate rows from results
+ SortResults, // Sort accumulated results. P4 = "col:asc,col:desc,..." over RESULT columns.
+ Distinct, // Remove duplicate rows from results
Limit, // Limit results to P1 rows. P2 = offset (rows to skip, already applied). Post-processing opcode.
Intersect, // Keep only rows that exist in both left and right result sets
Except, // Keep only rows from left that don't exist in right result set
@@ -198,6 +234,8 @@ pub enum OpCode {
// Window function operations
WindowAggStep, // P1 = func type, P2 = value reg, P3 = accum reg, P4 = window spec. Step window aggregate.
WindowValue, // P1 = accum reg, P2 = dest reg, P3 = func type. Get current window value.
+ WindowFinalize, // P1 = result column index. Give every row of a partition that partition's
+ // final window value, which for an unordered frame is the partition total.
}
/// A SQL VM instruction with opcode and parameters
@@ -276,6 +314,13 @@ pub struct Program {
pub instructions: Vec,
/// Schema for the result set (built at compile time)
pub result_schema: ResultSchema,
+ /// Number of registers the compiler allocated for this program.
+ ///
+ /// The engine sizes its register file from this. It must not be inferred
+ /// from the instruction stream: p1/p2/p3 hold literals as often as they
+ /// hold register numbers (`Integer 999999 -> r` being the obvious case),
+ /// so scanning for the maximum both over- and under-estimates.
+ pub register_count: i64,
}
impl Default for Program {
@@ -289,6 +334,7 @@ impl Program {
Self {
instructions: Vec::with_capacity(DEFAULT_INSTRUCTION_CAPACITY),
result_schema: ResultSchema::new(),
+ register_count: 0,
}
}
diff --git a/src/vm/compiler.rs b/src/vm/compiler.rs
index 3db2cc1..3490705 100644
--- a/src/vm/compiler.rs
+++ b/src/vm/compiler.rs
@@ -4,10 +4,14 @@
//! into bytecode instructions that can be executed by the SQL VM.
//! It implements a visitor pattern to walk the AST generated by sqlparser.
+use super::ast_compat::{
+ create_table_options, delete_from, func_args, group_by_exprs, insert_target, is_group_by_all,
+ is_order_by_all, object_name_last, order_by_is_asc, query_limit, query_offset, query_order_by,
+};
use sqlparser::ast::{
- BinaryOperator, DataType as SqlDataType, Expr, Function, FunctionArg, FunctionArgExpr,
- ObjectName, Query, Select, SelectItem, SetExpr, SetOperator, SetQuantifier, Statement,
- TableWithJoins, UnaryOperator, Value,
+ BinaryOperator, CaseWhen, DataType as SqlDataType, Expr, Function, FunctionArg,
+ FunctionArgExpr, Ident, JoinConstraint, ObjectName, Query, Select, SelectItem, SetExpr,
+ SetOperator, SetQuantifier, Statement, TableWithJoins, UnaryOperator, Value, ValueWithSpan,
};
use sqlparser::dialect::HiveDialect;
use sqlparser::parser::Parser;
@@ -28,10 +32,237 @@ pub(crate) type MultiTableProjection = (Vec>, Vec<(usize, usize)>, Re
/// Outer column reference found in a correlated subquery (Phase 4B)
#[derive(Debug, Clone)]
struct OuterColumnRef {
- /// The table qualifier (alias or table name)
+ /// The table qualifier (alias or table name) the outer query used.
+ ///
+ /// Only the qualifier is needed: which column it names is resolved by
+ /// NameCtx against the outer table, so recording it here would be a second
+ /// copy of that lookup.
qualifier: String,
- /// The column name
- column: String,
+}
+
+/// A symbolic jump target.
+///
+/// Jump destinations used to be computed arithmetically, as offsets from
+/// `program.len()` counting instructions that had not been emitted yet -- e.g.
+/// `let skip_offset = if limit_reg.is_some() { 3 } else { 2 }`. That couples
+/// every jump to the exact shape of the code after it, so inserting,
+/// reordering, or conditionally omitting an instruction silently retargets the
+/// jump. The failure mode is a wrong answer, not a crash.
+///
+/// A `Label` is created before it is known where it points, used as a jump
+/// destination any number of times, and bound to an address later with
+/// [`SqlCompiler::resolve`]. Unresolved labels are a compile-time bug and are
+/// caught by an assertion when the program is finished.
+#[derive(Clone, Copy, Debug, PartialEq, Eq)]
+pub(crate) struct Label(usize);
+
+/// One FROM source visible to expression compilation.
+pub(crate) struct Src<'t> {
+ /// VM cursor id for this source.
+ pub cursor: i64,
+ /// Base register holding this source's columns, when the loop has already
+ /// materialized them.
+ ///
+ /// Multi-table joins load every column into registers before evaluating
+ /// join conditions, so a column reference there is a register offset
+ /// rather than a `Column` read from a cursor. Modelling both here means
+ /// one expression compiler serves both loop shapes.
+ pub reg_base: Option,
+ /// The table behind the cursor.
+ pub table: &'t Table,
+ /// How the query refers to it: the alias if there is one, else the table
+ /// name. Stored ASCII-lowercased, since SQL identifiers are matched
+ /// case-insensitively.
+ pub refname: String,
+}
+
+/// A resolved column: where its value lives at run time.
+#[derive(Clone, Copy, Debug)]
+pub(crate) enum ColumnRef {
+ /// Read column `col` from cursor `cursor`.
+ Cursor { cursor: i64, col: usize },
+ /// The value is already in this register.
+ Register(i64),
+}
+
+impl ColumnRef {
+ fn new(src: &Src<'_>, col: usize) -> Self {
+ match src.reg_base {
+ Some(base) => ColumnRef::Register(base + col as i64),
+ None => ColumnRef::Cursor {
+ cursor: src.cursor,
+ col,
+ },
+ }
+ }
+}
+
+/// Name resolution scope for [`SqlCompiler::code_expr`].
+///
+/// The reason two expression compilers drifted apart is that both threaded a
+/// bare `(table, cursor_idx)` pair, a shape that cannot express a join, an
+/// alias, or a post-aggregation scope. Everything a name can resolve against
+/// lives here instead, so one compiler serves the SELECT list, WHERE, HAVING,
+/// ORDER BY and SET clauses.
+pub(crate) struct NameCtx<'t> {
+ /// FROM sources in cursor order. Empty for `SELECT 1` (no FROM).
+ pub srcs: Vec>,
+
+ /// Enclosing query's sources, for a correlated subquery.
+ ///
+ /// SQL resolves a name in the innermost scope first and only then looks
+ /// outward, so these are consulted only when `srcs` has no match. Keeping
+ /// them in a separate list rather than appending to `srcs` preserves that
+ /// precedence, and keeps the ambiguity check meaningful: two tables in one
+ /// FROM are ambiguous, but an inner name shadowing an outer one is not.
+ pub outer: Vec>,
+
+ /// Aggregate results available to this expression, keyed by the rendered
+ /// call text (`"sum(salary)"`).
+ ///
+ /// After aggregation an aggregate is a value like any other, so an
+ /// expression may contain one -- `SUM(salary) + 1`. Binding them here lets
+ /// the ordinary expression compiler handle the surrounding arithmetic
+ /// instead of needing a separate post-aggregation compiler.
+ pub agg_bindings: Vec<(String, i64)>,
+}
+
+impl<'t> NameCtx<'t> {
+ /// The common single-table scope, equivalent to the old
+ /// `(table, cursor_idx)` pair.
+ pub(crate) fn single(table: &'t Table, cursor: i64) -> Self {
+ NameCtx {
+ srcs: vec![Src {
+ cursor,
+ reg_base: None,
+ table,
+ refname: table.name().to_ascii_lowercase(),
+ }],
+ outer: Vec::new(),
+ agg_bindings: Vec::new(),
+ }
+ }
+
+ /// A scope over several FROM sources, for joins.
+ pub(crate) fn multi(srcs: Vec>) -> Self {
+ NameCtx {
+ srcs,
+ outer: Vec::new(),
+ agg_bindings: Vec::new(),
+ }
+ }
+
+ /// A scope with no FROM sources, for constant expressions.
+ #[allow(dead_code)] // used by constant-expression contexts
+ pub(crate) fn empty() -> Self {
+ NameCtx {
+ srcs: Vec::new(),
+ outer: Vec::new(),
+ agg_bindings: Vec::new(),
+ }
+ }
+
+ /// Normalized key for an aggregate call, used to bind and to look up.
+ pub(crate) fn agg_key(expr: &Expr) -> String {
+ expr.to_string().to_ascii_lowercase()
+ }
+
+ /// Register holding the result of an already-computed aggregate.
+ fn lookup_agg(&self, expr: &Expr) -> Option {
+ let key = Self::agg_key(expr);
+ self.agg_bindings
+ .iter()
+ .find(|(k, _)| *k == key)
+ .map(|(_, r)| *r)
+ }
+
+ /// Resolve an unqualified column name across every source.
+ ///
+ /// Case-insensitive: `SELECT NAME` already worked while
+ /// `WHERE NAME = 'x'` did not, because the projection path lowercased and
+ /// `find_column_index` compared exactly.
+ ///
+ /// Returns `Err` if no source has the column, or if more than one does --
+ /// silently picking the first would make the result depend on FROM order.
+ fn resolve_unqualified(&self, name: &str) -> SqawkResult {
+ let want = name.to_ascii_lowercase();
+ let mut found: Option = None;
+ for src in &self.srcs {
+ for (i, col) in src.table.column_metadata().iter().enumerate() {
+ if col.name.to_ascii_lowercase() == want {
+ if found.is_some() {
+ return Err(SqawkError::InvalidSqlQuery(format!(
+ "Column '{}' is ambiguous across the tables in FROM",
+ name
+ )));
+ }
+ found = Some(ColumnRef::new(src, i));
+ break;
+ }
+ }
+ }
+ if let Some(col) = found {
+ return Ok(col);
+ }
+ // Not in this scope: look outward, nearest first.
+ for src in &self.outer {
+ for (i, col) in src.table.column_metadata().iter().enumerate() {
+ if col.name.to_ascii_lowercase() == want {
+ return Ok(ColumnRef::new(src, i));
+ }
+ }
+ }
+ Err(SqawkError::ColumnNotFound(name.to_string()))
+ }
+
+ /// Resolve a `table.column` (or `alias.column`) reference.
+ fn resolve_qualified(&self, qualifier: &str, name: &str) -> SqawkResult {
+ let qual = qualifier.to_ascii_lowercase();
+ let want = name.to_ascii_lowercase();
+ for src in &self.srcs {
+ if src.refname != qual {
+ continue;
+ }
+ for (i, col) in src.table.column_metadata().iter().enumerate() {
+ if col.name.to_ascii_lowercase() == want {
+ return Ok(ColumnRef::new(src, i));
+ }
+ }
+ return Err(SqawkError::ColumnNotFound(format!(
+ "{}.{}",
+ qualifier, name
+ )));
+ }
+ // Qualifier not in this scope: try the enclosing one, which is how a
+ // correlated subquery refers to its outer row.
+ for src in &self.outer {
+ if src.refname != qual {
+ continue;
+ }
+ for (i, col) in src.table.column_metadata().iter().enumerate() {
+ if col.name.to_ascii_lowercase() == want {
+ return Ok(ColumnRef::new(src, i));
+ }
+ }
+ }
+ // Unknown qualifier entirely: fall back to an unqualified match so a
+ // single-table query can still say `employees.age` when the cursor was
+ // opened under a different refname.
+ self.resolve_unqualified(name)
+ }
+}
+
+/// Classify a join operator for the inner-join rewrite.
+///
+/// `Some(Some(c))` = inner join with constraint `c`; `Some(None)` = cross join;
+/// `None` = not rewritable (outer joins and the exotic forms).
+fn classify_join_for_rewrite(op: &sqlparser::ast::JoinOperator) -> Option> {
+ use sqlparser::ast::JoinOperator as J;
+ match op {
+ J::Join(c) | J::Inner(c) => Some(Some(c)),
+ J::CrossJoin(_) => Some(None),
+ _ => None,
+ }
}
/// SQL statement compiler that generates bytecode for the VM engine
@@ -53,6 +284,17 @@ pub struct SqlCompiler<'a> {
/// Current outer table alias for correlated subquery detection
pub(crate) current_outer_alias: Option,
+
+ /// Resolved address of each label, indexed by label id. `None` until
+ /// [`SqlCompiler::resolve`] binds it.
+ label_addrs: Vec>,
+
+ /// Jumps awaiting resolution: (instruction address, target label).
+ pending_jumps: Vec<(usize, Label)>,
+
+ /// Three-way `Jump`s awaiting resolution:
+ /// (instruction address, less, equal, greater).
+ pending_jumps_three: Vec<(usize, Label, Label, Label)>,
}
impl<'a> SqlCompiler<'a> {
@@ -65,7 +307,102 @@ impl<'a> SqlCompiler<'a> {
register_counter: 0,
verbose,
current_outer_alias: None,
+ label_addrs: Vec::new(),
+ pending_jumps: Vec::new(),
+ pending_jumps_three: Vec::new(),
+ }
+ }
+
+ /// Create a new, unresolved jump target.
+ pub(crate) fn label(&mut self) -> Label {
+ self.label_addrs.push(None);
+ Label(self.label_addrs.len() - 1)
+ }
+
+ /// Bind `label` to the next instruction to be emitted.
+ ///
+ /// Panics if the label was already resolved -- a label denotes exactly one
+ /// address, and rebinding one silently redirects every jump to it.
+ pub(crate) fn resolve(&mut self, label: Label) {
+ let addr = self.program.len();
+ debug_assert!(
+ self.label_addrs[label.0].is_none(),
+ "label {:?} resolved twice (already at {:?}, now {})",
+ label,
+ self.label_addrs[label.0],
+ addr
+ );
+ self.label_addrs[label.0] = Some(addr);
+ }
+
+ /// Emit a jump instruction whose destination (p2) is `target`.
+ ///
+ /// The destination is patched in by [`SqlCompiler::finish_labels`], so the
+ /// target may be resolved before or after this call.
+ pub(crate) fn emit_jump_to(
+ &mut self,
+ opcode: OpCode,
+ p1: i64,
+ target: Label,
+ p3: i64,
+ p4: Option,
+ comment: &str,
+ ) {
+ let addr = self.program.len();
+ // p2 is a placeholder; finish_labels overwrites it.
+ self.emit(opcode, p1, -1, p3, p4, comment);
+ self.pending_jumps.push((addr, target));
+ }
+
+ /// Emit a `Jump`: a three-way branch on the most recent `Compare`.
+ ///
+ /// All three destinations are labels, patched by `finish_labels`.
+ pub(crate) fn emit_jump_to_three(
+ &mut self,
+ less: Label,
+ equal: Label,
+ greater: Label,
+ comment: &str,
+ ) {
+ let addr = self.program.len();
+ self.emit(OpCode::Jump, -1, -1, -1, None, comment);
+ self.pending_jumps_three.push((addr, less, equal, greater));
+ }
+
+ /// Patch every pending jump with its label's address.
+ ///
+ /// Called once when a program is finished. An unresolved label means a
+ /// code path forgot to mark its destination, which would leave a jump
+ /// pointing at the placeholder -- so it fails loudly rather than
+ /// generating a program that runs and returns nonsense.
+ pub(crate) fn finish_labels(&mut self) -> SqawkResult<()> {
+ for (addr, label) in std::mem::take(&mut self.pending_jumps) {
+ let target = self.label_addrs[label.0].ok_or_else(|| {
+ SqawkError::InvalidSqlQuery(format!(
+ "internal error: jump at {} targets unresolved label {:?}",
+ addr, label
+ ))
+ })?;
+ self.patch_jump(addr, target);
+ }
+ for (addr, less, equal, greater) in std::mem::take(&mut self.pending_jumps_three) {
+ let resolve = |l: Label| -> SqawkResult {
+ self.label_addrs[l.0].map(|a| a as i64).ok_or_else(|| {
+ SqawkError::InvalidSqlQuery(format!(
+ "internal error: Jump at {} targets unresolved label {:?}",
+ addr, l
+ ))
+ })
+ };
+ let (p1, p2, p3) = (resolve(less)?, resolve(equal)?, resolve(greater)?);
+ if let Some(inst) = self.program.instructions.get_mut(addr) {
+ inst.p1 = p1;
+ inst.p2 = p2;
+ inst.p3 = p3;
+ }
}
+ self.label_addrs.clear();
+ Ok(())
}
/// Allocate a new register and return its index
@@ -149,95 +486,6 @@ impl<'a> SqlCompiler<'a> {
}
}
- /// Emit AND logic: result = left_reg && right_reg
- /// Returns the register containing the result (1 if both true, 0 otherwise)
- pub(crate) fn emit_and(&mut self, left_reg: i64, right_reg: i64) -> i64 {
- let result_reg = self.allocate_register();
-
- // Start with 0 (false)
- self.emit(
- OpCode::Integer,
- 0,
- result_reg,
- 0,
- None,
- &format!("r[{}] = 0 (AND init)", result_reg),
- );
-
- // If left is 0, skip to end (result stays 0)
- let skip_addr1 = self.program.len();
- self.emit(OpCode::IfZ, left_reg, 0, 0, None, "Skip if left is false");
-
- // If right is 0, skip to end (result stays 0)
- let skip_addr2 = self.program.len();
- self.emit(OpCode::IfZ, right_reg, 0, 0, None, "Skip if right is false");
-
- // Both true, set result to 1
- self.emit(
- OpCode::Integer,
- 1,
- result_reg,
- 0,
- None,
- &format!("r[{}] = 1 (AND true)", result_reg),
- );
-
- // Patch skip addresses to end
- let end_addr = self.program.len();
- self.patch_jump(skip_addr1, end_addr);
- self.patch_jump(skip_addr2, end_addr);
-
- result_reg
- }
-
- /// Emit OR logic: result = left_reg || right_reg
- /// Returns the register containing the result (1 if either true, 0 otherwise)
- pub(crate) fn emit_or(&mut self, left_reg: i64, right_reg: i64) -> i64 {
- let result_reg = self.allocate_register();
-
- // Start with 1 (optimistic - assume true)
- self.emit(
- OpCode::Integer,
- 1,
- result_reg,
- 0,
- None,
- &format!("r[{}] = 1 (OR init)", result_reg),
- );
-
- // If left is non-zero, skip to end (result stays 1)
- let skip_addr1 = self.program.len();
- self.emit(OpCode::IfPos, left_reg, 0, 0, None, "Skip if left is true");
-
- // If right is non-zero, skip to end (result stays 1)
- let skip_addr2 = self.program.len();
- self.emit(
- OpCode::IfPos,
- right_reg,
- 0,
- 0,
- None,
- "Skip if right is true",
- );
-
- // Both false, set result to 0
- self.emit(
- OpCode::Integer,
- 0,
- result_reg,
- 0,
- None,
- &format!("r[{}] = 0 (OR false)", result_reg),
- );
-
- // Patch skip addresses to end
- let end_addr = self.program.len();
- self.patch_jump(skip_addr1, end_addr);
- self.patch_jump(skip_addr2, end_addr);
-
- result_reg
- }
-
/// Emit a comparison operation and return the result register
pub(crate) fn emit_comparison(
&mut self,
@@ -346,11 +594,52 @@ impl<'a> SqlCompiler<'a> {
// TODO: Add direct query compilation support
/// Compile an SQL string into bytecode
+ /// Compile a single already-parsed statement into its own program.
+ ///
+ /// Each statement gets its own program so it can be executed, and its
+ /// modifications applied, before the next one is compiled -- which is what
+ /// lets a later statement see an earlier one's effects.
+ pub fn compile_statement_program(&mut self, statement: &Statement) -> SqawkResult {
+ self.program = Program::new();
+ self.reset_registers();
+ self.table_map.clear();
+ self.label_addrs.clear();
+ self.pending_jumps.clear();
+ self.pending_jumps_three.clear();
+
+ let init_addr = self.program.len();
+ self.emit(OpCode::Init, 0, 0, 0, None, "Start address filled in later");
+
+ self.compile_statement(statement)?;
+
+ self.emit(OpCode::Halt, 0, 0, 0, None, "End execution");
+
+ let start = init_addr + 1;
+ if let Some(instruction) = self.program.instructions.get_mut(init_addr) {
+ instruction.p2 = start as i64;
+ instruction.comment = Some(format!("Start at {}", start).into());
+ }
+
+ self.finish_labels()?;
+ self.program.register_count = self.register_counter;
+ Ok(self.program.clone())
+ }
+
+ /// Parse and compile `sql` into a single program.
+ ///
+ /// Convenience wrapper used by the compiler unit tests, which examine
+ /// generated bytecode for one statement. Execution goes through
+ /// `compile_statement_program` per statement instead, so that each
+ /// statement can observe the previous one's effects.
+ #[allow(dead_code)]
pub fn compile(&mut self, sql: &str) -> SqawkResult {
// Reset state for new compilation
self.program = Program::new();
self.reset_registers();
self.table_map.clear();
+ self.label_addrs.clear();
+ self.pending_jumps.clear();
+ self.pending_jumps_three.clear();
// Parse the SQL statement
let dialect = HiveDialect {};
@@ -390,6 +679,13 @@ impl<'a> SqlCompiler<'a> {
instruction.comment = Some(format!("Start at {}", transaction_addr).into());
}
+ // Patch every symbolic jump now that all addresses are known. This
+ // errors rather than emitting a program with a placeholder target.
+ self.finish_labels()?;
+
+ // Hand the engine the real register high-water mark.
+ self.program.register_count = self.register_counter;
+
Ok(self.program.clone())
}
@@ -399,35 +695,39 @@ impl<'a> SqlCompiler<'a> {
fn compile_statement(&mut self, statement: &Statement) -> SqawkResult<()> {
match statement {
Statement::Query(query) => self.compile_query(query),
- Statement::Insert {
- table_name,
- columns,
- source,
- ..
- } => self.compile_insert(table_name, columns, source),
- Statement::Delete {
- from, selection, ..
- } => self.compile_delete(from, selection.as_ref()),
- Statement::Update {
- table,
- assignments,
- selection,
- ..
- } => self.compile_update(table, assignments, selection.as_ref()),
- Statement::CreateTable {
- name,
- columns,
- hive_formats,
- location,
- with_options,
- query,
- ..
- } => {
- if let Some(q) = query {
+ Statement::Insert(insert) => {
+ let source = insert.source.as_ref().ok_or_else(|| {
+ SqawkError::UnsupportedSqlFeature("INSERT without a source".into())
+ })?;
+ // `Insert::columns` became Vec; sqawk only accepts
+ // bare column names, so flatten each to its last segment.
+ let columns: Vec = insert
+ .columns
+ .iter()
+ .map(|c| Ident::new(object_name_last(c)))
+ .collect();
+ self.compile_insert(insert_target(&insert.table)?, &columns, source)
+ }
+ Statement::Delete(delete) => {
+ self.compile_delete(delete_from(&delete.from), delete.selection.as_ref())
+ }
+ Statement::Update(update) => self.compile_update(
+ &update.table,
+ &update.assignments,
+ update.selection.as_ref(),
+ ),
+ Statement::CreateTable(create) => {
+ if let Some(q) = &create.query {
// CREATE TABLE ... AS SELECT
- self.compile_create_table_as_select(name, q)
+ self.compile_create_table_as_select(&create.name, q)
} else {
- self.compile_create_table(name, columns, hive_formats, location, with_options)
+ self.compile_create_table(
+ &create.name,
+ &create.columns,
+ &create.hive_formats,
+ &create.location,
+ create_table_options(&create.table_options),
+ )
}
}
Statement::Drop {
@@ -436,10 +736,21 @@ impl<'a> SqlCompiler<'a> {
if_exists,
..
} => self.compile_drop(object_type, names, *if_exists),
- Statement::AlterTable {
- name, operation, ..
- } => self.compile_alter_table(name, operation),
- Statement::Truncate { table_name, .. } => self.compile_truncate(table_name),
+ Statement::AlterTable(alter) => {
+ // `operation` became `operations`; sqawk handles one at a time.
+ match alter.operations.as_slice() {
+ [operation] => self.compile_alter_table(&alter.name, operation),
+ _ => Err(SqawkError::UnsupportedSqlFeature(
+ "ALTER TABLE with multiple operations is not supported".into(),
+ )),
+ }
+ }
+ Statement::Truncate(truncate) => match truncate.table_names.as_slice() {
+ [target] => self.compile_truncate(&target.name),
+ _ => Err(SqawkError::UnsupportedSqlFeature(
+ "TRUNCATE of multiple tables is not supported".into(),
+ )),
+ },
_ => Err(SqawkError::UnsupportedSqlFeature(format!(
"Unsupported SQL statement type: {:?}",
statement
@@ -449,9 +760,17 @@ impl<'a> SqlCompiler<'a> {
/// Compile a SQL query
fn compile_query(&mut self, query: &Query) -> SqawkResult<()> {
+ // `ORDER BY ALL` carries no key list, so it would otherwise look
+ // identical to "no ORDER BY" and be silently ignored. Reject it.
+ if is_order_by_all(query) {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "ORDER BY ALL is not supported".into(),
+ ));
+ }
+
// Check if we have ORDER BY, LIMIT, or OFFSET
- let has_order_by = !query.order_by.is_empty();
- let has_limit = query.limit.is_some() || query.offset.is_some();
+ let has_order_by = !query_order_by(query).is_empty();
+ let has_limit = query_limit(query).is_some() || query_offset(query).is_some();
// Compile the body (SELECT or set operation)
match &*query.body {
@@ -554,7 +873,7 @@ impl<'a> SqlCompiler<'a> {
"Keep only rows in both result sets",
);
}
- SetOperator::Except => {
+ SetOperator::Except | SetOperator::Minus => {
self.emit(
OpCode::Except,
left_results_marker.unwrap_or(0),
@@ -589,8 +908,79 @@ impl<'a> SqlCompiler<'a> {
}
}
+ /// Rewrite `A INNER JOIN B ON cond` into `FROM A, B WHERE cond`.
+ ///
+ /// The implicit (comma) join path already supports aggregates and GROUP BY;
+ /// the explicit-JOIN path does not, and rejects `COUNT(*)` in its
+ /// projection resolver. For an INNER join the two spellings mean the same
+ /// thing, so an aggregate query over an explicit join is answered by
+ /// handing it to the machinery that can already do it, rather than by
+ /// duplicating that machinery.
+ ///
+ /// Returns `None` when the rewrite does not apply: outer joins are NOT
+ /// equivalent to a comma join plus WHERE, since the WHERE would discard
+ /// exactly the NULL-extended rows the outer join exists to produce.
+ fn rewrite_inner_join_as_implicit(select: &Select) -> Option {
+ if select.from.len() != 1 {
+ return None;
+ }
+ let twj = &select.from[0];
+ if twj.joins.is_empty() {
+ return None;
+ }
+
+ let mut froms = vec![TableWithJoins {
+ relation: twj.relation.clone(),
+ joins: Vec::new(),
+ }];
+ let mut condition = select.selection.clone();
+
+ for join in &twj.joins {
+ let constraint = match classify_join_for_rewrite(&join.join_operator)? {
+ Some(c) => c,
+ // CROSS JOIN contributes a source but no condition.
+ None => {
+ froms.push(TableWithJoins {
+ relation: join.relation.clone(),
+ joins: Vec::new(),
+ });
+ continue;
+ }
+ };
+ let on_expr = match constraint {
+ JoinConstraint::On(expr) => expr.clone(),
+ _ => return None,
+ };
+ froms.push(TableWithJoins {
+ relation: join.relation.clone(),
+ joins: Vec::new(),
+ });
+ condition = Some(match condition {
+ Some(existing) => Expr::BinaryOp {
+ left: Box::new(existing),
+ op: BinaryOperator::And,
+ right: Box::new(on_expr),
+ },
+ None => on_expr,
+ });
+ }
+
+ let mut rewritten = select.clone();
+ rewritten.from = froms;
+ rewritten.selection = condition;
+ Some(rewritten)
+ }
+
/// Compile a SELECT statement
fn compile_select(&mut self, select: &Select) -> SqawkResult<()> {
+ // `GROUP BY ALL` carries no key list, so it would otherwise look
+ // identical to "no GROUP BY" and silently collapse the query. Reject it.
+ if is_group_by_all(&select.group_by) {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "GROUP BY ALL is not supported".into(),
+ ));
+ }
+
if select.from.is_empty() {
self.compile_select_literal(&select.projection)?;
} else {
@@ -602,7 +992,7 @@ impl<'a> SqlCompiler<'a> {
}
// Check for GROUP BY or aggregates - requires specialized handling
- let has_group_by = !select.group_by.is_empty();
+ let has_group_by = !group_by_exprs(&select.group_by).is_empty();
let has_aggregates = self.has_aggregates(&select.projection);
if has_group_by || has_aggregates {
@@ -610,11 +1000,14 @@ impl<'a> SqlCompiler<'a> {
let query = Query {
with: None,
body: Box::new(SetExpr::Select(Box::new(select.clone()))),
- order_by: vec![],
- limit: None,
- offset: None,
+ order_by: None,
+ limit_clause: None,
fetch: None,
locks: vec![],
+ for_clause: None,
+ settings: None,
+ format_clause: None,
+ pipe_operators: vec![],
};
return self.compile_implicit_join_with_group_by(select, &query);
}
@@ -624,10 +1017,20 @@ impl<'a> SqlCompiler<'a> {
let table_with_joins = &select.from[0];
if !table_with_joins.joins.is_empty() {
+ // An aggregate or GROUP BY over an explicit INNER join is
+ // answered by rewriting it to the equivalent comma join, whose
+ // path already supports both.
+ let needs_aggregation = !group_by_exprs(&select.group_by).is_empty()
+ || self.has_aggregates(&select.projection);
+ if needs_aggregation {
+ if let Some(rewritten) = Self::rewrite_inner_join_as_implicit(select) {
+ return self.compile_select(&rewritten);
+ }
+ }
self.compile_join(table_with_joins, &select.projection, &select.selection)?;
} else {
// Check for GROUP BY, window functions, or aggregates
- let has_group_by = !select.group_by.is_empty();
+ let has_group_by = !group_by_exprs(&select.group_by).is_empty();
let has_window_functions = self.has_window_functions(&select.projection);
let has_aggregates = self.has_aggregates(&select.projection);
@@ -697,7 +1100,7 @@ impl<'a> SqlCompiler<'a> {
select: &Select,
query: &Query,
) -> SqawkResult<()> {
- let has_order_by = !query.order_by.is_empty();
+ let has_order_by = !query_order_by(query).is_empty();
if select.from.is_empty() {
self.compile_select_literal(&select.projection)?;
@@ -712,12 +1115,25 @@ impl<'a> SqlCompiler<'a> {
}
// Check for GROUP BY or aggregates - requires specialized handling
- let has_group_by = !select.group_by.is_empty();
+ let has_group_by = !group_by_exprs(&select.group_by).is_empty();
let has_aggregates = self.has_aggregates(&select.projection);
if has_group_by || has_aggregates {
- // Multi-table query with GROUP BY or aggregates
- return self.compile_implicit_join_with_group_by(select, query);
+ // Multi-table query with GROUP BY or aggregates. Like the
+ // single-table grouped path, this emits rows directly, so
+ // ORDER BY and LIMIT are applied to the finished result set --
+ // both were silently dropped here.
+ self.compile_implicit_join_with_group_by(select, query)?;
+ return self.emit_result_post_processing(select, query);
+ }
+
+ if !query_order_by(query).is_empty() {
+ // ORDER BY needs the whole result before it can sort, so the
+ // scan must NOT apply LIMIT as it goes -- doing so would keep
+ // the first N rows and then sort those, rather than the top N.
+ // Sorting and limiting both happen afterwards.
+ self.compile_implicit_join(select)?;
+ return self.emit_result_post_processing(select, query);
}
return self.compile_implicit_join_with_query(select, query);
@@ -727,13 +1143,22 @@ impl<'a> SqlCompiler<'a> {
// Check for explicit JOINs (FROM table1 JOIN table2 ON ...)
if !table_with_joins.joins.is_empty() {
- // For multi-table JOINs with ORDER BY, compile the JOIN normally
- // ORDER BY support for JOINs requires post-processing which isn't fully implemented
- // For now, compile the JOIN and skip ORDER BY
- if self.verbose && has_order_by {
- eprintln!("Note: ORDER BY with explicit JOINs - compiling join without sort");
+ // See compile_select: aggregates over an explicit INNER join are
+ // handled by rewriting to the comma-join form.
+ let needs_aggregation = !group_by_exprs(&select.group_by).is_empty()
+ || self.has_aggregates(&select.projection);
+ if needs_aggregation {
+ if let Some(rewritten) = Self::rewrite_inner_join_as_implicit(select) {
+ return self.compile_select_with_post_processing(&rewritten, query);
+ }
}
- return self.compile_join(table_with_joins, &select.projection, &select.selection);
+ // A join emits rows directly rather than through the cursor-based
+ // sorter, which is why ORDER BY used to be dropped here outright.
+ // Sorting the finished result set works regardless of how many
+ // cursors produced it.
+ self.compile_join(table_with_joins, &select.projection, &select.selection)?;
+ self.emit_result_post_processing(select, query)?;
+ return Ok(());
}
// Get table and columns for ORDER BY resolution
@@ -754,39 +1179,25 @@ impl<'a> SqlCompiler<'a> {
// Check for GROUP BY - delegate to specialized handler
// GROUP BY with ORDER BY/LIMIT handles these internally
- let has_group_by = !select.group_by.is_empty();
+ let has_group_by = !group_by_exprs(&select.group_by).is_empty();
let has_aggregates = self.has_aggregates(&select.projection);
if has_group_by || has_aggregates {
if has_group_by {
- return self.compile_select_with_group_by(select, table, &table_name);
+ self.compile_select_with_group_by(select, table, &table_name)?;
} else {
- return self.compile_select_with_aggregate(select, table, &table_name);
+ self.compile_select_with_aggregate(select, table, &table_name)?;
}
+ // These paths emit rows directly rather than through the sorter,
+ // so ORDER BY and LIMIT are applied to the finished result set.
+ // `query` used to be in scope here and simply never passed on,
+ // which is why both were silently dropped on every grouped query.
+ self.emit_result_post_processing(select, query)?;
+ return Ok(());
}
// Check if projection contains function expressions or computed expressions
- let has_computed_expr = select.projection.iter().any(|item| {
- let expr = match item {
- SelectItem::UnnamedExpr(e) => Some(e),
- SelectItem::ExprWithAlias { expr: e, .. } => Some(e),
- _ => None,
- };
- if let Some(e) = expr {
- matches!(
- e,
- Expr::Function(_)
- | Expr::Trim { .. }
- | Expr::Substring { .. }
- | Expr::BinaryOp { .. }
- | Expr::UnaryOp { .. }
- | Expr::Ceil { .. }
- | Expr::Floor { .. }
- )
- } else {
- false
- }
- });
+ let has_computed_expr = Self::projection_needs_expr(&select.projection);
// If we have functions/computed expressions and only LIMIT (no ORDER BY), use compile_table_scan
// which handles function expressions properly
@@ -822,14 +1233,30 @@ impl<'a> SqlCompiler<'a> {
let sorter_id = 0i64;
let col_count = columns.len();
- // Build sort spec from ORDER BY
- let sort_spec = self.build_sort_spec(&query.order_by, table, columns)?;
+ // Sorter payload is [sort keys.., output columns..].
+ //
+ // It used to hold only the output columns, and the sort spec mapped
+ // each ORDER BY key to a position among them -- so ordering by a
+ // column you did not select was rejected with "ORDER BY column must be
+ // in SELECT list", and ordering by an expression was rejected outright.
+ // Carrying the keys separately removes both restrictions: a key is
+ // just another value computed per row, whether or not it is projected.
+ let order_by = query_order_by(query);
+ let key_count = order_by.len();
+ let total_cols = key_count + col_count;
+
+ let sort_spec: String = order_by
+ .iter()
+ .enumerate()
+ .map(|(i, ob)| format!("{}:{}", i, if order_by_is_asc(ob) { "asc" } else { "desc" }))
+ .collect::>()
+ .join(",");
// Open sorter
self.emit(
OpCode::SorterOpen,
sorter_id,
- col_count as i64,
+ total_cols as i64,
0,
Some(sort_spec),
"",
@@ -858,29 +1285,34 @@ impl<'a> SqlCompiler<'a> {
let loop_start = self.program.len();
- // Load columns into registers
- let start_reg = self.allocate_registers(col_count);
+ // WHERE first, so filtered rows never reach the sorter.
+ let next_label = self.label();
+ if let Some(where_expr) = &select.selection {
+ let cond_reg = self.compile_where_condition(where_expr, table, cursor_idx as usize)?;
+ self.emit_jump_to(
+ OpCode::IfZ,
+ cond_reg,
+ next_label,
+ 0,
+ None,
+ "Skip row if WHERE is false",
+ );
+ }
+
+ // One contiguous block: sort keys first, then the output columns.
+ let start_reg = self.allocate_registers(total_cols);
+
+ let ctx = NameCtx::single(table, cursor_idx);
+ for (i, ob) in order_by.iter().enumerate() {
+ self.code_expr(&ob.expr, &ctx, Some(start_reg + i as i64))?;
+ }
for (i, col_idx) in columns.iter().enumerate() {
self.emit(
OpCode::Column,
cursor_idx,
*col_idx as i64,
- start_reg + i as i64,
- None,
- "",
- );
- }
-
- // Compile WHERE clause if present
- if let Some(where_expr) = &select.selection {
- let cond_reg = self.compile_where_condition(where_expr, table, cursor_idx as usize)?;
- // Skip SorterInsert, go to Next
- self.emit(
- OpCode::IfZ,
- cond_reg,
- (self.program.len() + 2) as i64,
- 0,
+ start_reg + (key_count + i) as i64,
None,
"",
);
@@ -891,12 +1323,13 @@ impl<'a> SqlCompiler<'a> {
OpCode::SorterInsert,
sorter_id,
start_reg,
- col_count as i64,
+ total_cols as i64,
None,
"",
);
// Next row
+ self.resolve(next_label);
self.emit(OpCode::Next, cursor_idx, loop_start as i64, 0, None, "");
let after_scan = self.program.len();
@@ -944,75 +1377,78 @@ impl<'a> SqlCompiler<'a> {
OpCode::SorterData,
sorter_id,
start_reg,
- col_count as i64,
+ total_cols as i64,
None,
"",
);
+ // Symbolic targets for the drain loop. Using labels here matters
+ // because whether a LIMIT instruction exists changes the distance
+ // between the OFFSET logic and SorterNext -- the exact coupling that
+ // the old `skip_offset = if limit_reg.is_some() { 3 } else { 2 }`
+ // encoded by hand.
+ let result_row_label = self.label();
+ let sorter_next_label = self.label();
+ let after_loop_label = self.label();
+
// Handle OFFSET - skip first N rows
if let Some(off_reg) = offset_reg {
+ let decrement_label = self.label();
+
// If offset counter > 0, decrement and skip this row
- self.emit(
+ self.emit_jump_to(
OpCode::IfPos,
off_reg,
- (self.program.len() + 2) as i64, // Jump to decrement and next (DecrJumpZero)
+ decrement_label,
0,
None,
- "",
+ "offset remaining -> skip this row",
);
// Offset exhausted, continue to output
- let output_addr = self.program.len();
- self.emit(
+ self.emit_jump_to(
OpCode::Goto,
0,
- 0, // Will be patched to ResultRow
- 0,
- None,
- "",
- );
- // Decrement offset and go to next
- self.emit(
- OpCode::DecrJumpZero,
- off_reg,
- (self.program.len() + 1) as i64, // Continue to SorterNext
+ result_row_label,
0,
None,
- "",
+ "offset exhausted -> output row",
);
- // Skip offset depends on whether there's a LIMIT instruction after ResultRow
- let skip_offset = if limit_reg.is_some() { 3 } else { 2 };
- self.emit(
- OpCode::Goto,
- 0,
- (self.program.len() + skip_offset) as i64, // Skip to SorterNext
- 0,
- None,
- "",
- );
- // Patch the Goto to ResultRow
- let result_row_addr = self.program.len() as i64;
- if let Some(inst) = self.program.instructions.get_mut(output_addr) {
- inst.p2 = result_row_addr;
- }
+
+ self.resolve(decrement_label);
+ // Decrement offset, then fall through to the skip.
+ let continue_label = self.label();
+ self.emit_jump_to(OpCode::DecrJumpZero, off_reg, continue_label, 0, None, "");
+ self.resolve(continue_label);
+ self.emit_jump_to(OpCode::Goto, 0, sorter_next_label, 0, None, "skip row");
}
// Output result row
- self.emit(OpCode::ResultRow, start_reg, col_count as i64, 0, None, "");
+ self.resolve(result_row_label);
+ // Emit only the output half: the sort keys are internal.
+ self.emit(
+ OpCode::ResultRow,
+ start_reg + key_count as i64,
+ col_count as i64,
+ 0,
+ None,
+ "",
+ );
// Handle LIMIT
if let Some(lim_reg) = limit_reg {
// DecrJumpZero exits when limit reached
- self.emit(
+ self.emit_jump_to(
OpCode::DecrJumpZero,
lim_reg,
- (self.program.len() + 2) as i64, // Jump past SorterNext to end
+ after_loop_label,
0,
None,
- "",
+ "limit reached -> exit loop",
);
}
// Next sorted row
+ self.resolve(sorter_next_label);
self.emit(
OpCode::SorterNext,
sorter_id,
@@ -1023,6 +1459,7 @@ impl<'a> SqlCompiler<'a> {
);
// Close cursor
+ self.resolve(after_loop_label);
self.emit(OpCode::Close, cursor_idx, 0, 0, None, "");
// Handle DISTINCT - emit Distinct opcode to deduplicate results
@@ -1135,16 +1572,23 @@ impl<'a> SqlCompiler<'a> {
);
}
+ // A row that fails WHERE must resume at Next -- NOT at a fixed offset.
+ // The distance from here to Next depends on whether OFFSET and LIMIT
+ // emit instructions, so the old `program.len() + 2` landed on the
+ // LIMIT counter's DecrJumpZero (making LIMIT count rows *scanned*
+ // rather than rows *returned*) or, with OFFSET, on ResultRow itself.
+ let next_label = self.label();
+
// WHERE clause
if let Some(where_expr) = &select.selection {
let cond_reg = self.compile_where_condition(where_expr, table, cursor_idx as usize)?;
- self.emit(
+ self.emit_jump_to(
OpCode::IfZ,
cond_reg,
- (self.program.len() + 2) as i64,
+ next_label,
0,
None,
- "",
+ "Skip row if WHERE is false",
);
}
@@ -1217,6 +1661,7 @@ impl<'a> SqlCompiler<'a> {
// Next row
let next_addr = self.program.len();
+ self.resolve(next_label);
self.emit(OpCode::Next, cursor_idx, loop_start as i64, 0, None, "");
let after_loop = self.program.len();
@@ -1257,58 +1702,14 @@ impl<'a> SqlCompiler<'a> {
Ok(())
}
- /// Build sort spec string from ORDER BY clause
- fn build_sort_spec(
- &self,
- order_by: &[sqlparser::ast::OrderByExpr],
- table: &Table,
- projected_columns: &[usize],
- ) -> SqawkResult {
- let mut specs = Vec::new();
- for expr in order_by {
- // Find column index in original table
- let table_col_idx = match &expr.expr {
- Expr::Identifier(ident) => {
- let col_name = ident.value.to_lowercase();
- table
- .column_index(&col_name)
- .ok_or(SqawkError::ColumnNotFound(col_name))?
- }
- Expr::CompoundIdentifier(parts) => {
- let col_name = parts
- .last()
- .map(|p| p.value.to_lowercase())
- .unwrap_or_default();
- table
- .column_index(&col_name)
- .ok_or(SqawkError::ColumnNotFound(col_name))?
- }
- _ => {
- return Err(SqawkError::UnsupportedSqlFeature(
- "Complex ORDER BY expressions not supported".into(),
- ))
- }
- };
- // Map to position in projected columns
- let col_idx = projected_columns
- .iter()
- .position(|&c| c == table_col_idx)
- .ok_or_else(|| {
- SqawkError::InvalidSqlQuery(
- "ORDER BY column must be in SELECT list".to_string(),
- )
- })?;
- let asc = expr.asc.unwrap_or(true);
- specs.push(format!("{}:{}", col_idx, if asc { "asc" } else { "desc" }));
- }
- Ok(specs.join(","))
- }
-
/// Extract LIMIT and OFFSET values from query
pub(crate) fn extract_limit_offset(&self, query: &Query) -> SqawkResult<(Option, i64)> {
- let limit = if let Some(limit_expr) = &query.limit {
+ let limit = if let Some(limit_expr) = query_limit(query) {
match limit_expr {
- Expr::Value(Value::Number(n, _)) => Some(n.parse::().map_err(|_| {
+ Expr::Value(ValueWithSpan {
+ value: Value::Number(n, _),
+ ..
+ }) => Some(n.parse::().map_err(|_| {
SqawkError::InvalidSqlQuery(format!("Invalid LIMIT value: {}", n))
})?),
_ => {
@@ -1321,9 +1722,12 @@ impl<'a> SqlCompiler<'a> {
None
};
- let offset = if let Some(offset_clause) = &query.offset {
- match &offset_clause.value {
- Expr::Value(Value::Number(n, _)) => n.parse::().map_err(|_| {
+ let offset = if let Some(offset_clause) = query_offset(query) {
+ match offset_clause {
+ Expr::Value(ValueWithSpan {
+ value: Value::Number(n, _),
+ ..
+ }) => n.parse::().map_err(|_| {
SqawkError::InvalidSqlQuery(format!("Invalid OFFSET value: {}", n))
})?,
_ => {
@@ -1436,7 +1840,7 @@ impl<'a> SqlCompiler<'a> {
let result_reg = self.allocate_register();
match expr {
- sqlparser::ast::Expr::Value(value) => {
+ sqlparser::ast::Expr::Value(ValueWithSpan { value, .. }) => {
// Load the appropriate value based on type
match value {
sqlparser::ast::Value::Number(num, _) => {
@@ -1577,20 +1981,20 @@ impl<'a> SqlCompiler<'a> {
.name
.0
.iter()
- .map(|id| id.value.as_str())
+ .map(|id| id.as_ident().map(|i| i.value.as_str()).unwrap_or_default())
.collect::>()
.join(".")
.to_uppercase();
match func_name.as_str() {
"ABS" | "ROUND" | "CEIL" | "CEILING" | "FLOOR" => {
- if func.args.is_empty() {
+ if func_args(func).is_empty() {
return Err(SqawkError::InvalidSqlQuery(format!(
"{} requires one argument",
func_name
)));
}
- let arg_expr = self.extract_function_arg_expr(&func.args[0])?;
+ let arg_expr = self.extract_function_arg_expr(&func_args(func)[0])?;
let src_reg = self.compile_expr(&arg_expr)?;
self.emit(
OpCode::MathFunc,
@@ -1614,13 +2018,13 @@ impl<'a> SqlCompiler<'a> {
}
"DATE" | "TIME" => {
// Single-argument date/time functions
- if func.args.is_empty() {
+ if func_args(func).is_empty() {
return Err(SqawkError::InvalidSqlQuery(format!(
"{} requires one argument",
func_name
)));
}
- let arg_expr = self.extract_function_arg_expr(&func.args[0])?;
+ let arg_expr = self.extract_function_arg_expr(&func_args(func)[0])?;
let src_reg = self.compile_expr(&arg_expr)?;
self.emit(
OpCode::DateFunc,
@@ -1706,29 +2110,7 @@ impl<'a> SqlCompiler<'a> {
let table = self.database.get_table(&table_name).unwrap();
// Check if projection contains function expressions or computed expressions
- let has_function_expr = projection.iter().any(|item| {
- let expr = match item {
- SelectItem::UnnamedExpr(e) => Some(e),
- SelectItem::ExprWithAlias { expr: e, .. } => Some(e),
- _ => None,
- };
- if let Some(e) = expr {
- matches!(
- e,
- Expr::Function(_)
- | Expr::Trim { .. }
- | Expr::Substring { .. }
- | Expr::Position { .. }
- | Expr::Overlay { .. }
- | Expr::BinaryOp { .. }
- | Expr::UnaryOp { .. }
- | Expr::Ceil { .. }
- | Expr::Floor { .. }
- )
- } else {
- false
- }
- });
+ let has_function_expr = Self::projection_needs_expr(projection);
// Build result schema with column names and types
let schema = self.build_result_schema(projection, table);
@@ -1780,6 +2162,13 @@ impl<'a> SqlCompiler<'a> {
let value_reg = result_regs[idx];
match item {
+ // Multiple aliases for one item (`expr AS (a, b)`) is a
+ // non-standard extension sqawk does not implement.
+ SelectItem::ExprWithAliases { .. } => {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "Multiple column aliases are not supported".into(),
+ ))
+ }
SelectItem::Wildcard(_) => {
// For wildcard, load all columns
for col_idx in 0..table.column_count() {
@@ -1807,208 +2196,15 @@ impl<'a> SqlCompiler<'a> {
}
}
SelectItem::UnnamedExpr(expr) | SelectItem::ExprWithAlias { expr, .. } => {
- match expr {
- Expr::Function(func) => {
- self.compile_function(func, table, cursor_idx, value_reg)?;
- }
- Expr::Trim {
- expr: trim_expr, ..
- } => {
- // Compile TRIM expression
- let src_reg = self.allocate_register();
- self.compile_where_operand(trim_expr, table, cursor_idx, src_reg)?;
- self.emit(
- OpCode::StringFunc,
- src_reg,
- value_reg,
- 0,
- Some("TRIM".to_string()),
- &format!("r[{}] = TRIM(r[{}])", value_reg, src_reg),
- );
- }
- Expr::Substring {
- expr: sub_expr,
- substring_from,
- substring_for,
- ..
- } => {
- // Compile SUBSTRING expression
- let src_reg = self.allocate_register();
- self.compile_where_operand(sub_expr, table, cursor_idx, src_reg)?;
-
- // Compile start position
- let start_reg = self.allocate_register();
- if let Some(from_expr) = substring_from {
- self.compile_where_operand(
- from_expr, table, cursor_idx, start_reg,
- )?;
- } else {
- self.emit(
- OpCode::Integer,
- 1,
- start_reg,
- 0,
- None,
- &format!("r[{}] = 1 (default start)", start_reg),
- );
- }
-
- // Check if we have a length
- let func_spec = if let Some(for_expr) = substring_for {
- let len_reg = self.allocate_register();
- self.compile_where_operand(
- for_expr, table, cursor_idx, len_reg,
- )?;
- format!("SUBSTR:{}", len_reg)
- } else {
- "SUBSTR".to_string()
- };
-
- self.emit(
- OpCode::StringFunc,
- src_reg,
- value_reg,
- start_reg,
- Some(func_spec),
- &format!("r[{}] = SUBSTR(...)", value_reg),
- );
- }
- Expr::Identifier(ident) => {
- let col_name = ident.value.to_lowercase();
- let col_idx = table
- .column_index(&col_name)
- .ok_or_else(|| SqawkError::ColumnNotFound(col_name.clone()))?;
- self.emit(
- OpCode::Column,
- cursor_idx as i64,
- col_idx as i64,
- value_reg,
- None,
- &format!("r[{}] = column {}", value_reg, col_idx),
- );
- }
- Expr::CompoundIdentifier(parts) => {
- let col_name = parts
- .last()
- .map(|p| p.value.to_lowercase())
- .unwrap_or_default();
- let col_idx = table
- .column_index(&col_name)
- .ok_or_else(|| SqawkError::ColumnNotFound(col_name.clone()))?;
- self.emit(
- OpCode::Column,
- cursor_idx as i64,
- col_idx as i64,
- value_reg,
- None,
- &format!("r[{}] = column {}", value_reg, col_idx),
- );
- }
- Expr::BinaryOp { left, op, right } => {
- // Handle arithmetic binary operators
- let left_reg = self.allocate_register();
- let right_reg = self.allocate_register();
-
- self.compile_where_operand(left, table, cursor_idx, left_reg)?;
- self.compile_where_operand(right, table, cursor_idx, right_reg)?;
-
- let opcode = match op {
- BinaryOperator::Plus => OpCode::Add,
- BinaryOperator::Minus => OpCode::Subtract,
- BinaryOperator::Multiply => OpCode::Multiply,
- BinaryOperator::Divide => OpCode::Divide,
- BinaryOperator::Modulo => OpCode::Remainder,
- _ => {
- return Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported binary operator in SELECT: {:?}",
- op
- )));
- }
- };
-
- self.emit(
- opcode,
- left_reg,
- right_reg,
- value_reg,
- None,
- &format!(
- "r[{}] = r[{}] {:?} r[{}]",
- value_reg, left_reg, op, right_reg
- ),
- );
- }
- Expr::UnaryOp { op, expr: inner } => {
- // Handle unary operators (e.g., -value)
- match op {
- sqlparser::ast::UnaryOperator::Minus => {
- let inner_reg = self.allocate_register();
- self.compile_where_operand(
- inner, table, cursor_idx, inner_reg,
- )?;
- // Negate using 0 - value (avoids Integer -1 special case)
- let zero_reg = self.allocate_register();
- self.emit(
- OpCode::Integer,
- 0,
- zero_reg,
- 0,
- None,
- &format!("r[{}] = 0", zero_reg),
- );
- self.emit(
- OpCode::Subtract,
- zero_reg,
- inner_reg,
- value_reg,
- None,
- &format!("r[{}] = -r[{}]", value_reg, inner_reg),
- );
- }
- sqlparser::ast::UnaryOperator::Plus => {
- self.compile_where_operand(
- inner, table, cursor_idx, value_reg,
- )?;
- }
- _ => {
- return Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported unary operator: {:?}",
- op
- )));
- }
- }
- }
- Expr::Ceil { expr: inner, .. } => {
- let src_reg = self.allocate_register();
- self.compile_where_operand(inner, table, cursor_idx, src_reg)?;
- self.emit(
- OpCode::MathFunc,
- src_reg,
- value_reg,
- 0,
- Some("CEIL".to_string()),
- &format!("r[{}] = CEIL(r[{}])", value_reg, src_reg),
- );
- }
- Expr::Floor { expr: inner, .. } => {
- let src_reg = self.allocate_register();
- self.compile_where_operand(inner, table, cursor_idx, src_reg)?;
- self.emit(
- OpCode::MathFunc,
- src_reg,
- value_reg,
- 0,
- Some("FLOOR".to_string()),
- &format!("r[{}] = FLOOR(r[{}])", value_reg, src_reg),
- );
- }
- _ => {
- return Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported expression in SELECT: {:?}",
- expr
- )));
- }
- }
+ // One expression compiler, as everywhere else. This was
+ // a parallel implementation understanding Function,
+ // Trim, Substring, Position, Overlay, BinaryOp, UnaryOp,
+ // Ceil and Floor, and rejecting everything else -- so
+ // IS NULL, LIKE, BETWEEN, IN and subqueries were
+ // unusable in a projection despite being supported in
+ // WHERE.
+ let ctx = NameCtx::single(table, cursor_idx as i64);
+ self.code_expr(expr, &ctx, Some(value_reg))?;
}
SelectItem::QualifiedWildcard(_, _) => {
for col_idx in 0..table.column_count() {
@@ -2185,6 +2381,13 @@ impl<'a> SqlCompiler<'a> {
let value_reg = result_regs[idx];
match item {
+ // Multiple aliases for one item (`expr AS (a, b)`) is a
+ // non-standard extension sqawk does not implement.
+ SelectItem::ExprWithAliases { .. } => {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "Multiple column aliases are not supported".into(),
+ ))
+ }
SelectItem::Wildcard(_) | SelectItem::QualifiedWildcard(_, _) => {
for col_idx in 0..table.column_count() {
if col_idx > 0 {
@@ -2314,122 +2517,17 @@ impl<'a> SqlCompiler<'a> {
cursor_idx: usize,
dest_reg: i64,
) -> SqawkResult<()> {
- match expr {
- Expr::Identifier(ident) => {
- let col_name = ident.value.to_lowercase();
- if let Some(col_idx) = table.column_index(&col_name) {
- self.emit(
- OpCode::Column,
- cursor_idx as i64,
- col_idx as i64,
- dest_reg,
- None,
- &format!("r[{}] = {}", dest_reg, col_name),
- );
- }
- }
- Expr::Function(func) => {
- let func_name = func.name.to_string().to_uppercase();
- if let Some(sqlparser::ast::FunctionArg::Unnamed(
- sqlparser::ast::FunctionArgExpr::Expr(Expr::Identifier(ident)),
- )) = func.args.first()
- {
- let col_name = ident.value.to_lowercase();
- if let Some(col_idx) = table.column_index(&col_name) {
- let src_reg = self.allocate_register();
- self.emit(
- OpCode::Column,
- cursor_idx as i64,
- col_idx as i64,
- src_reg,
- None,
- &format!("Load {} for {}", col_name, func_name),
- );
-
- // Handle different string functions
- match func_name.as_str() {
- "UPPER" | "LOWER" | "TRIM" | "LTRIM" | "RTRIM" | "LENGTH" => {
- self.emit(
- OpCode::StringFunc,
- src_reg,
- dest_reg,
- 0,
- Some(func_name.clone()),
- &format!("{}()", func_name),
- );
- }
- "SUBSTR" | "SUBSTRING" => {
- // Handle SUBSTR with start and length
- let (start, len) = self.extract_substr_args(func)?;
- let start_reg = self.allocate_register();
- self.emit(
- OpCode::Integer,
- start,
- start_reg,
- 0,
- None,
- "SUBSTR start",
- );
- let len_reg = self.allocate_register();
- self.emit(OpCode::Integer, len, len_reg, 0, None, "SUBSTR length");
- // P3 = start register, P4 = "SUBSTR:len_reg"
- self.emit(
- OpCode::StringFunc,
- src_reg,
- dest_reg,
- start_reg,
- Some(format!("SUBSTR:{}", len_reg)),
- "SUBSTR()",
- );
- }
- _ => {
- // Copy source to dest for unknown functions
- self.emit(OpCode::Copy, src_reg, dest_reg, 0, None, "Copy value");
- }
- };
- }
- }
- }
- _ => {
- // For other expressions, try to resolve as column
- if let Ok(col_idx) = self.resolve_column_expr(expr, table) {
- self.emit(
- OpCode::Column,
- cursor_idx as i64,
- col_idx as i64,
- dest_reg,
- None,
- &format!("r[{}] = column {}", dest_reg, col_idx),
- );
- }
- }
- }
- Ok(())
- }
-
- /// Extract SUBSTR arguments (start, length) from function
- fn extract_substr_args(&self, func: &sqlparser::ast::Function) -> SqawkResult<(i64, i64)> {
- let mut start = 1i64;
- let mut length = 100i64; // Default length
-
- let args: Vec<_> = func.args.iter().collect();
- if args.len() >= 2 {
- if let sqlparser::ast::FunctionArg::Unnamed(sqlparser::ast::FunctionArgExpr::Expr(
- Expr::Value(sqlparser::ast::Value::Number(n, _)),
- )) = &args[1]
- {
- start = n.parse().unwrap_or(1);
- }
- }
- if args.len() >= 3 {
- if let sqlparser::ast::FunctionArg::Unnamed(sqlparser::ast::FunctionArgExpr::Expr(
- Expr::Value(sqlparser::ast::Value::Number(n, _)),
- )) = &args[2]
- {
- length = n.parse().unwrap_or(100);
- }
- }
- Ok((start, length))
+ // Delegate to the single expression compiler.
+ //
+ // This function used to be a parallel implementation that understood
+ // only bare columns and a fixed list of single-argument functions.
+ // Its `_` arm tried resolve_column_expr and, on failure, emitted NO
+ // INSTRUCTION AT ALL -- leaving the destination register at its
+ // default NULL. That is why `SELECT name, salary*2 FROM employees
+ // LIMIT 3` produced NULLs while the same query without LIMIT (which
+ // took a different path) produced correct values.
+ let ctx = NameCtx::single(table, cursor_idx as i64);
+ self.code_expr(expr, &ctx, Some(dest_reg)).map(|_| ())
}
/// Compile a WHERE condition and return the register containing the result
@@ -2469,12 +2567,11 @@ impl<'a> SqlCompiler<'a> {
Expr::Case {
operand,
conditions,
- results,
else_result,
+ ..
} => self.compile_case(
operand.as_deref(),
conditions,
- results,
else_result.as_deref(),
table,
cursor_idx,
@@ -2603,6 +2700,27 @@ impl<'a> SqlCompiler<'a> {
// EXISTS (SELECT ...) - check if subquery returns any rows
self.compile_exists_subquery(subquery, *negated, table, cursor_idx)
}
+ // `NOT `. Previously unsupported in any position: the
+ // predicate compiler had no arm, and the several NOT-forms that
+ // did work (NOT LIKE, NOT IN, IS NOT NULL) each hand-rolled their
+ // own inversion instead of sharing one.
+ Expr::UnaryOp {
+ op: UnaryOperator::Not,
+ expr: inner,
+ } => {
+ let inner_reg = self.compile_where_condition(inner, table, cursor_idx)?;
+ let dest = self.allocate_register();
+ self.emit(
+ OpCode::Not,
+ inner_reg,
+ dest,
+ 0,
+ None,
+ &format!("r[{}] = NOT r[{}]", dest, inner_reg),
+ );
+ Ok(dest)
+ }
+ Expr::Nested(inner) => self.compile_where_condition(inner, table, cursor_idx),
_ => Err(SqawkError::UnsupportedSqlFeature(format!(
"Unsupported WHERE clause expression: {:?}",
where_expr
@@ -2649,8 +2767,8 @@ impl<'a> SqlCompiler<'a> {
/// Extract the column index from an aggregate function argument
fn extract_agg_column_index(&self, func: &Function, table: &Table) -> SqawkResult {
// Handle COUNT(*) specially
- if !func.args.is_empty() {
- match &func.args[0] {
+ if !func_args(func).is_empty() {
+ match &func_args(func)[0] {
FunctionArg::Unnamed(FunctionArgExpr::Wildcard) => {
// COUNT(*) - use first column (doesn't matter which for count)
return Ok(0);
@@ -2797,7 +2915,11 @@ impl<'a> SqlCompiler<'a> {
.name
.0
.first()
- .map(|id| id.value.to_uppercase())
+ .map(|id| {
+ id.as_ident()
+ .map(|i| i.value.to_uppercase())
+ .unwrap_or_default()
+ })
.unwrap_or_default();
// Get the aggregate function
@@ -3040,7 +3162,7 @@ impl<'a> SqlCompiler<'a> {
TableValue::Null
}
}
- Expr::Value(val) => match val {
+ Expr::Value(ValueWithSpan { value: val, .. }) => match val {
Value::Number(n, _) => {
if let Ok(i) = n.parse::() {
TableValue::Integer(i)
@@ -3488,7 +3610,8 @@ impl<'a> SqlCompiler<'a> {
let table_name = name
.0
.last()
- .map(|p| p.value.to_lowercase())
+ .and_then(|p| p.as_ident())
+ .map(|i| i.value.to_lowercase())
.unwrap_or_default();
let alias_name = alias.as_ref().map(|a| a.name.value.to_lowercase());
Some((table_name, alias_name))
@@ -3525,7 +3648,6 @@ impl<'a> SqlCompiler<'a> {
match expr {
Expr::CompoundIdentifier(parts) if parts.len() == 2 => {
let qualifier = parts[0].value.to_lowercase();
- let column = parts[1].value.to_lowercase();
// Check if this references the outer table (not the subquery's own table)
let is_outer_ref = {
@@ -3546,7 +3668,7 @@ impl<'a> SqlCompiler<'a> {
};
if is_outer_ref {
- refs.push(OuterColumnRef { qualifier, column });
+ refs.push(OuterColumnRef { qualifier });
}
}
Expr::BinaryOp { left, right, .. } => {
@@ -3641,8 +3763,8 @@ impl<'a> SqlCompiler<'a> {
Expr::Case {
operand,
conditions,
- results,
else_result,
+ ..
} => {
if let Some(op) = operand {
self.collect_outer_refs_from_expr(
@@ -3653,18 +3775,16 @@ impl<'a> SqlCompiler<'a> {
refs,
);
}
- for cond in conditions {
+ for branch in conditions {
self.collect_outer_refs_from_expr(
- cond,
+ &branch.condition,
outer_table_name,
outer_alias,
subquery_table,
refs,
);
- }
- for result in results {
self.collect_outer_refs_from_expr(
- result,
+ &branch.result,
outer_table_name,
outer_alias,
subquery_table,
@@ -3682,7 +3802,7 @@ impl<'a> SqlCompiler<'a> {
}
}
Expr::Function(func) => {
- for arg in &func.args {
+ for arg in func_args(func) {
if let FunctionArg::Unnamed(FunctionArgExpr::Expr(e)) = arg {
self.collect_outer_refs_from_expr(
e,
@@ -4090,6 +4210,17 @@ impl<'a> SqlCompiler<'a> {
/// Compile an operand in a correlated WHERE clause
#[allow(clippy::too_many_arguments)]
+ /// Compile an operand of a correlated subquery's condition.
+ ///
+ /// The subquery's own tables form the inner scope and the enclosing
+ /// query's the outer one, which is exactly `NameCtx`'s two-level
+ /// resolution: an unqualified name binds to the subquery, and a name it
+ /// does not have is looked up outward.
+ ///
+ /// This was a private operand compiler handling identifiers, qualified
+ /// identifiers and literals, so a correlated condition could not contain
+ /// arithmetic or a function call even though a plain WHERE could.
+ #[allow(clippy::too_many_arguments)]
fn compile_correlated_operand(
&mut self,
expr: &Expr,
@@ -4100,123 +4231,31 @@ impl<'a> SqlCompiler<'a> {
outer_refs: &[OuterColumnRef],
dest_reg: i64,
) -> SqawkResult<()> {
- match expr {
- Expr::Identifier(ident) => {
- // Unqualified column - must be from subquery table
- let col_name = ident.value.to_lowercase();
- let col_idx = subquery_table
- .column_index(&col_name)
- .ok_or(SqawkError::ColumnNotFound(col_name.clone()))?;
- self.emit(
- OpCode::Column,
- sub_cursor_idx as i64,
- col_idx as i64,
- dest_reg,
- None,
- &format!("r[{}] = subquery.{}", dest_reg, col_name),
- );
- Ok(())
- }
- Expr::CompoundIdentifier(parts) if parts.len() == 2 => {
- let qualifier = parts[0].value.to_lowercase();
- let col_name = parts[1].value.to_lowercase();
-
- // Check if this is an outer reference
- let is_outer = outer_refs
- .iter()
- .any(|r| r.qualifier == qualifier && r.column == col_name)
- || qualifier == outer_table.name().to_lowercase();
-
- if is_outer {
- // Load from outer table
- let col_idx = outer_table
- .column_index(&col_name)
- .ok_or(SqawkError::ColumnNotFound(col_name.clone()))?;
- self.emit(
- OpCode::Column,
- outer_cursor_idx as i64,
- col_idx as i64,
- dest_reg,
- None,
- &format!("r[{}] = outer.{}", dest_reg, col_name),
- );
- } else {
- // Load from subquery table
- let col_idx = subquery_table
- .column_index(&col_name)
- .ok_or(SqawkError::ColumnNotFound(col_name.clone()))?;
- self.emit(
- OpCode::Column,
- sub_cursor_idx as i64,
- col_idx as i64,
- dest_reg,
- None,
- &format!("r[{}] = subquery.{}", dest_reg, col_name),
- );
- }
- Ok(())
- }
- Expr::Value(v) => {
- match v {
- Value::Number(n, _) => {
- if let Ok(i) = n.parse::() {
- self.emit(
- OpCode::Integer,
- i,
- dest_reg,
- 0,
- None,
- &format!("r[{}] = {}", dest_reg, i),
- );
- } else if let Ok(f) = n.parse::() {
- self.emit(
- OpCode::String,
- 0,
- dest_reg,
- 0,
- Some(f.to_string()),
- &format!("r[{}] = {}", dest_reg, f),
- );
- self.emit(
- OpCode::Cast,
- dest_reg,
- dest_reg,
- 0,
- Some("REAL".to_string()),
- "",
- );
- }
- }
- Value::SingleQuotedString(s) => {
- self.emit(
- OpCode::String,
- 0,
- dest_reg,
- 0,
- Some(s.clone()),
- &format!("r[{}] = '{}'", dest_reg, s),
- );
- }
- _ => {
- return Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported value type in correlated subquery: {:?}",
- v
- )))
- }
- }
- Ok(())
- }
- _ => Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported operand in correlated subquery: {:?}",
- expr
- ))),
- }
+ // An alias the outer query used for its table must resolve too, not
+ // just the table's own name.
+ let outer_refname = outer_refs
+ .first()
+ .map(|r| r.qualifier.clone())
+ .unwrap_or_else(|| outer_table.name().to_ascii_lowercase());
+
+ let mut ctx = NameCtx::single(subquery_table, sub_cursor_idx as i64);
+ ctx.outer = vec![
+ Src {
+ cursor: outer_cursor_idx as i64,
+ reg_base: None,
+ table: outer_table,
+ refname: outer_refname,
+ },
+ Src {
+ cursor: outer_cursor_idx as i64,
+ reg_base: None,
+ table: outer_table,
+ refname: outer_table.name().to_ascii_lowercase(),
+ },
+ ];
+ self.code_expr(expr, &ctx, Some(dest_reg)).map(|_| ())
}
- /// Compile a correlated scalar subquery (Phase 4B)
- ///
- /// For each outer row, iterates through the subquery table and computes
- /// an aggregate over matching rows.
fn compile_correlated_scalar_subquery(
&mut self,
subquery: &Query,
@@ -4452,7 +4491,10 @@ impl<'a> SqlCompiler<'a> {
// Get the pattern string
let pattern_str = match pattern {
- Expr::Value(Value::SingleQuotedString(s)) => s.clone(),
+ Expr::Value(ValueWithSpan {
+ value: Value::SingleQuotedString(s),
+ ..
+ }) => s.clone(),
_ => {
return Err(SqawkError::UnsupportedSqlFeature(
"LIKE pattern must be a string literal".to_string(),
@@ -4745,45 +4787,94 @@ impl<'a> SqlCompiler<'a> {
Ok(result_reg)
}
- /// Compile a CASE expression
+ /// Compile a SUBSTRING/SUBSTR expression into `target_reg`.
///
- /// Supports both:
- /// - Simple CASE: CASE expr WHEN val1 THEN res1 ... ELSE default END
- /// - Searched CASE: CASE WHEN cond1 THEN res1 ... ELSE default END
+ /// Shared by the projection and WHERE paths: `SUBSTR(x, 1, 3)` and
+ /// `SUBSTRING(x FROM 1 FOR 3)` are the same `Expr::Substring` node, and
+ /// both positions need to compile it identically.
#[allow(clippy::too_many_arguments)]
- fn compile_case(
+ fn compile_substring(
&mut self,
- operand: Option<&Expr>,
- conditions: &[Expr],
- results: &[Expr],
- else_result: Option<&Expr>,
+ sub_expr: &Expr,
+ substring_from: Option<&Expr>,
+ substring_for: Option<&Expr>,
table: &Table,
cursor_idx: usize,
- target_reg: Option,
- ) -> SqawkResult {
- if conditions.len() != results.len() {
- return Err(SqawkError::InvalidSqlQuery(
- "CASE: conditions and results must have same length".to_string(),
- ));
- }
-
- let result_reg = target_reg.unwrap_or_else(|| self.allocate_register());
- let temp_reg = self.allocate_register();
+ target_reg: i64,
+ ) -> SqawkResult<()> {
+ let src_reg = self.allocate_register();
+ self.compile_where_operand(sub_expr, table, cursor_idx, src_reg)?;
- // For simple CASE, compile the operand once
- let operand_reg = if let Some(op_expr) = operand {
- let reg = self.allocate_register();
- self.compile_where_operand(op_expr, table, cursor_idx, reg)?;
- Some(reg)
+ // Start position defaults to 1 when omitted.
+ let start_reg = self.allocate_register();
+ if let Some(from_expr) = substring_from {
+ self.compile_where_operand(from_expr, table, cursor_idx, start_reg)?;
} else {
- None
- };
-
- // Track jump addresses that need to go to the end
- let mut jump_to_end_addrs = Vec::new();
+ self.emit(
+ OpCode::Integer,
+ 1,
+ start_reg,
+ 0,
+ None,
+ &format!("r[{}] = 1 (default start)", start_reg),
+ );
+ }
+
+ // A length, if given, is passed to the engine via the P4 spec.
+ let func_spec = if let Some(for_expr) = substring_for {
+ let len_reg = self.allocate_register();
+ self.compile_where_operand(for_expr, table, cursor_idx, len_reg)?;
+ format!("SUBSTR:{}", len_reg)
+ } else {
+ "SUBSTR".to_string()
+ };
+
+ self.emit(
+ OpCode::StringFunc,
+ src_reg,
+ target_reg,
+ start_reg,
+ Some(func_spec),
+ &format!("r[{}] = SUBSTR(...)", target_reg),
+ );
+ Ok(())
+ }
+
+ /// Compile a CASE expression
+ ///
+ /// Supports both:
+ /// - Simple CASE: CASE expr WHEN val1 THEN res1 ... ELSE default END
+ /// - Searched CASE: CASE WHEN cond1 THEN res1 ... ELSE default END
+ fn compile_case(
+ &mut self,
+ operand: Option<&Expr>,
+ conditions: &[CaseWhen],
+ else_result: Option<&Expr>,
+ table: &Table,
+ cursor_idx: usize,
+ target_reg: Option,
+ ) -> SqawkResult {
+ let result_reg = target_reg.unwrap_or_else(|| self.allocate_register());
+ let temp_reg = self.allocate_register();
+
+ // For simple CASE, compile the operand once
+ let operand_reg = if let Some(op_expr) = operand {
+ let reg = self.allocate_register();
+ self.compile_where_operand(op_expr, table, cursor_idx, reg)?;
+ Some(reg)
+ } else {
+ None
+ };
+
+ // Track jump addresses that need to go to the end
+ let mut jump_to_end_addrs = Vec::new();
// Process each WHEN clause
- for (condition, result_expr) in conditions.iter().zip(results.iter()) {
+ for CaseWhen {
+ condition,
+ result: result_expr,
+ } in conditions.iter()
+ {
let cond_reg = self.allocate_register();
if let Some(op_reg) = operand_reg {
@@ -4872,7 +4963,7 @@ impl<'a> SqlCompiler<'a> {
func.name
.0
.iter()
- .map(|id| id.value.as_str())
+ .map(|id| id.as_ident().map(|i| i.value.as_str()).unwrap_or_default())
.collect::>()
.join(".")
.to_uppercase()
@@ -4882,11 +4973,10 @@ impl<'a> SqlCompiler<'a> {
fn compile_func_coalesce(
&mut self,
func: &Function,
- table: &Table,
- cursor_idx: usize,
+ ctx: &NameCtx<'_>,
target_reg: i64,
) -> SqawkResult<()> {
- if func.args.is_empty() {
+ if func_args(func).is_empty() {
return Err(SqawkError::InvalidSqlQuery(
"COALESCE requires at least one argument".to_string(),
));
@@ -4894,13 +4984,13 @@ impl<'a> SqlCompiler<'a> {
let mut jump_to_end_addrs = Vec::new();
- for arg in &func.args {
+ for arg in func_args(func) {
let expr = self.extract_function_arg_expr(arg)?;
let arg_reg = self.allocate_register();
let null_check_reg = self.allocate_register();
// Compile the argument
- self.compile_where_operand(&expr, table, cursor_idx, arg_reg)?;
+ self.code_expr_into(&expr, ctx, arg_reg)?;
// Check if it's NULL
self.emit(
@@ -4996,26 +5086,25 @@ impl<'a> SqlCompiler<'a> {
fn compile_func_nullif(
&mut self,
func: &Function,
- table: &Table,
- cursor_idx: usize,
+ ctx: &NameCtx<'_>,
target_reg: i64,
) -> SqawkResult<()> {
- if func.args.len() != 2 {
+ if func_args(func).len() != 2 {
return Err(SqawkError::InvalidSqlQuery(
"NULLIF requires exactly two arguments".to_string(),
));
}
- let expr1 = self.extract_function_arg_expr(&func.args[0])?;
- let expr2 = self.extract_function_arg_expr(&func.args[1])?;
+ let expr1 = self.extract_function_arg_expr(&func_args(func)[0])?;
+ let expr2 = self.extract_function_arg_expr(&func_args(func)[1])?;
let reg1 = self.allocate_register();
let reg2 = self.allocate_register();
let cmp_reg = self.allocate_register();
// Compile both arguments
- self.compile_where_operand(&expr1, table, cursor_idx, reg1)?;
- self.compile_where_operand(&expr2, table, cursor_idx, reg2)?;
+ self.code_expr_into(&expr1, ctx, reg1)?;
+ self.code_expr_into(&expr2, ctx, reg2)?;
// Compare them
self.emit(
@@ -5088,22 +5177,21 @@ impl<'a> SqlCompiler<'a> {
&mut self,
func_name: &str,
func: &Function,
- table: &Table,
- cursor_idx: usize,
+ ctx: &NameCtx<'_>,
target_reg: i64,
) -> SqawkResult<()> {
- if func.args.is_empty() {
+ if func_args(func).is_empty() {
return Err(SqawkError::InvalidSqlQuery(format!(
"{} requires one argument",
func_name
)));
}
- let expr = self.extract_function_arg_expr(&func.args[0])?;
+ let expr = self.extract_function_arg_expr(&func_args(func)[0])?;
let src_reg = self.allocate_register();
// Compile the argument
- self.compile_where_operand(&expr, table, cursor_idx, src_reg)?;
+ self.code_expr_into(&expr, ctx, src_reg)?;
// Emit StringFunc opcode
self.emit(
@@ -5122,31 +5210,30 @@ impl<'a> SqlCompiler<'a> {
fn compile_func_substr(
&mut self,
func: &Function,
- table: &Table,
- cursor_idx: usize,
+ ctx: &NameCtx<'_>,
target_reg: i64,
) -> SqawkResult<()> {
- if func.args.len() < 2 {
+ if func_args(func).len() < 2 {
return Err(SqawkError::InvalidSqlQuery(
"SUBSTR requires at least two arguments".to_string(),
));
}
- let str_expr = self.extract_function_arg_expr(&func.args[0])?;
- let start_expr = self.extract_function_arg_expr(&func.args[1])?;
+ let str_expr = self.extract_function_arg_expr(&func_args(func)[0])?;
+ let start_expr = self.extract_function_arg_expr(&func_args(func)[1])?;
let str_reg = self.allocate_register();
let start_reg = self.allocate_register();
// Compile string and start arguments
- self.compile_where_operand(&str_expr, table, cursor_idx, str_reg)?;
- self.compile_where_operand(&start_expr, table, cursor_idx, start_reg)?;
+ self.code_expr_into(&str_expr, ctx, str_reg)?;
+ self.code_expr_into(&start_expr, ctx, start_reg)?;
// Check if we have a length argument
- let func_spec = if func.args.len() >= 3 {
- let len_expr = self.extract_function_arg_expr(&func.args[2])?;
+ let func_spec = if func_args(func).len() >= 3 {
+ let len_expr = self.extract_function_arg_expr(&func_args(func)[2])?;
let len_reg = self.allocate_register();
- self.compile_where_operand(&len_expr, table, cursor_idx, len_reg)?;
+ self.code_expr_into(&len_expr, ctx, len_reg)?;
format!("SUBSTR:{}", len_reg)
} else {
"SUBSTR".to_string()
@@ -5169,28 +5256,27 @@ impl<'a> SqlCompiler<'a> {
fn compile_func_replace(
&mut self,
func: &Function,
- table: &Table,
- cursor_idx: usize,
+ ctx: &NameCtx<'_>,
target_reg: i64,
) -> SqawkResult<()> {
- if func.args.len() != 3 {
+ if func_args(func).len() != 3 {
return Err(SqawkError::InvalidSqlQuery(
"REPLACE requires three arguments".to_string(),
));
}
- let str_expr = self.extract_function_arg_expr(&func.args[0])?;
- let from_expr = self.extract_function_arg_expr(&func.args[1])?;
- let to_expr = self.extract_function_arg_expr(&func.args[2])?;
+ let str_expr = self.extract_function_arg_expr(&func_args(func)[0])?;
+ let from_expr = self.extract_function_arg_expr(&func_args(func)[1])?;
+ let to_expr = self.extract_function_arg_expr(&func_args(func)[2])?;
let str_reg = self.allocate_register();
let from_reg = self.allocate_register();
let to_reg = self.allocate_register();
// Compile all three arguments
- self.compile_where_operand(&str_expr, table, cursor_idx, str_reg)?;
- self.compile_where_operand(&from_expr, table, cursor_idx, from_reg)?;
- self.compile_where_operand(&to_expr, table, cursor_idx, to_reg)?;
+ self.code_expr_into(&str_expr, ctx, str_reg)?;
+ self.code_expr_into(&from_expr, ctx, from_reg)?;
+ self.code_expr_into(&to_expr, ctx, to_reg)?;
// Emit StringFunc opcode with from_reg and to_reg in P4
self.emit(
@@ -5209,27 +5295,26 @@ impl<'a> SqlCompiler<'a> {
fn compile_func_concat(
&mut self,
func: &Function,
- table: &Table,
- cursor_idx: usize,
+ ctx: &NameCtx<'_>,
target_reg: i64,
) -> SqawkResult<()> {
- if func.args.len() < 2 {
+ if func_args(func).len() < 2 {
return Err(SqawkError::InvalidSqlQuery(
"CONCAT requires at least two arguments".to_string(),
));
}
// Compile first argument to src_reg
- let first_expr = self.extract_function_arg_expr(&func.args[0])?;
+ let first_expr = self.extract_function_arg_expr(&func_args(func)[0])?;
let src_reg = self.allocate_register();
- self.compile_where_operand(&first_expr, table, cursor_idx, src_reg)?;
+ self.code_expr_into(&first_expr, ctx, src_reg)?;
// Compile remaining arguments and build P4 spec
- let mut arg_regs = Vec::with_capacity(func.args.len().saturating_sub(1));
- for arg in func.args.iter().skip(1) {
+ let mut arg_regs = Vec::with_capacity(func_args(func).len().saturating_sub(1));
+ for arg in func_args(func).iter().skip(1) {
let expr = self.extract_function_arg_expr(arg)?;
let reg = self.allocate_register();
- self.compile_where_operand(&expr, table, cursor_idx, reg)?;
+ self.code_expr_into(&expr, ctx, reg)?;
arg_regs.push(reg.to_string());
}
@@ -5253,25 +5338,24 @@ impl<'a> SqlCompiler<'a> {
&mut self,
func_name: &str,
func: &Function,
- table: &Table,
- cursor_idx: usize,
+ ctx: &NameCtx<'_>,
target_reg: i64,
) -> SqawkResult<()> {
- if func.args.len() != 2 {
+ if func_args(func).len() != 2 {
return Err(SqawkError::InvalidSqlQuery(format!(
"{} requires exactly two arguments",
func_name
)));
}
- let str_expr = self.extract_function_arg_expr(&func.args[0])?;
- let len_expr = self.extract_function_arg_expr(&func.args[1])?;
+ let str_expr = self.extract_function_arg_expr(&func_args(func)[0])?;
+ let len_expr = self.extract_function_arg_expr(&func_args(func)[1])?;
let str_reg = self.allocate_register();
let len_reg = self.allocate_register();
- self.compile_where_operand(&str_expr, table, cursor_idx, str_reg)?;
- self.compile_where_operand(&len_expr, table, cursor_idx, len_reg)?;
+ self.code_expr_into(&str_expr, ctx, str_reg)?;
+ self.code_expr_into(&len_expr, ctx, len_reg)?;
self.emit(
OpCode::StringFunc,
@@ -5293,22 +5377,21 @@ impl<'a> SqlCompiler<'a> {
&mut self,
func_name: &str,
func: &Function,
- table: &Table,
- cursor_idx: usize,
+ ctx: &NameCtx<'_>,
target_reg: i64,
) -> SqawkResult<()> {
- if func.args.is_empty() {
+ if func_args(func).is_empty() {
return Err(SqawkError::InvalidSqlQuery(format!(
"{} requires one argument",
func_name
)));
}
- let expr = self.extract_function_arg_expr(&func.args[0])?;
+ let expr = self.extract_function_arg_expr(&func_args(func)[0])?;
let src_reg = self.allocate_register();
// Compile the argument
- self.compile_where_operand(&expr, table, cursor_idx, src_reg)?;
+ self.code_expr_into(&expr, ctx, src_reg)?;
// Emit MathFunc opcode
self.emit(
@@ -5342,22 +5425,21 @@ impl<'a> SqlCompiler<'a> {
&mut self,
func_name: &str,
func: &Function,
- table: &Table,
- cursor_idx: usize,
+ ctx: &NameCtx<'_>,
target_reg: i64,
) -> SqawkResult<()> {
- if func.args.is_empty() {
+ if func_args(func).is_empty() {
return Err(SqawkError::InvalidSqlQuery(format!(
"{} requires one argument",
func_name
)));
}
- let expr = self.extract_function_arg_expr(&func.args[0])?;
+ let expr = self.extract_function_arg_expr(&func_args(func)[0])?;
let src_reg = self.allocate_register();
// Compile the argument
- self.compile_where_operand(&expr, table, cursor_idx, src_reg)?;
+ self.code_expr_into(&expr, ctx, src_reg)?;
// Emit DateFunc opcode
self.emit(
@@ -5376,33 +5458,28 @@ impl<'a> SqlCompiler<'a> {
fn compile_function(
&mut self,
func: &Function,
- table: &Table,
- cursor_idx: usize,
+ ctx: &NameCtx<'_>,
target_reg: i64,
) -> SqawkResult<()> {
let func_name = Self::get_function_name(func);
match func_name.as_str() {
- "COALESCE" => self.compile_func_coalesce(func, table, cursor_idx, target_reg),
- "NULLIF" => self.compile_func_nullif(func, table, cursor_idx, target_reg),
+ "COALESCE" => self.compile_func_coalesce(func, ctx, target_reg),
+ "NULLIF" => self.compile_func_nullif(func, ctx, target_reg),
"UPPER" | "LOWER" | "TRIM" | "LTRIM" | "RTRIM" | "LENGTH" => {
- self.compile_func_string_simple(&func_name, func, table, cursor_idx, target_reg)
- }
- "SUBSTR" | "SUBSTRING" => self.compile_func_substr(func, table, cursor_idx, target_reg),
- "REPLACE" => self.compile_func_replace(func, table, cursor_idx, target_reg),
- "CONCAT" => self.compile_func_concat(func, table, cursor_idx, target_reg),
- "LEFT" | "RIGHT" => {
- self.compile_func_left_right(&func_name, func, table, cursor_idx, target_reg)
+ self.compile_func_string_simple(&func_name, func, ctx, target_reg)
}
+ "SUBSTR" | "SUBSTRING" => self.compile_func_substr(func, ctx, target_reg),
+ "REPLACE" => self.compile_func_replace(func, ctx, target_reg),
+ "CONCAT" => self.compile_func_concat(func, ctx, target_reg),
+ "LEFT" | "RIGHT" => self.compile_func_left_right(&func_name, func, ctx, target_reg),
"ABS" | "ROUND" | "CEIL" | "CEILING" | "FLOOR" => {
- self.compile_func_math(&func_name, func, table, cursor_idx, target_reg)
+ self.compile_func_math(&func_name, func, ctx, target_reg)
}
"NOW" | "CURRENT_TIMESTAMP" | "CURRENT_DATE" | "CURRENT_TIME" => {
self.compile_func_datetime_noarg(&func_name, target_reg)
}
- "DATE" | "TIME" => {
- self.compile_func_datetime(&func_name, func, table, cursor_idx, target_reg)
- }
+ "DATE" | "TIME" => self.compile_func_datetime(&func_name, func, ctx, target_reg),
_ => Err(SqawkError::UnsupportedSqlFeature(format!(
"Unsupported function: {}",
func_name
@@ -5436,44 +5513,50 @@ impl<'a> SqlCompiler<'a> {
// Handle logical operators (AND, OR) specially
match op {
BinaryOperator::And => {
- // For AND, both conditions must be true
+ // For AND, both conditions must be true.
+ //
+ // The result register MUST be initialized before the
+ // short-circuit test on the left operand. Registers persist
+ // across loop iterations, so if the left-false jump skips the
+ // initialization, the register still holds this expression's
+ // result from a PREVIOUS row -- and a row failing the left
+ // condition inherits the previous row's verdict.
+ let result_reg = self.allocate_register();
+ self.emit(
+ OpCode::Integer,
+ 0,
+ result_reg,
+ 0,
+ None,
+ &format!("r[{}] = 0 (default AND result)", result_reg),
+ );
+
+ let end_label = self.label();
+
// Compile left condition
let left_result = self.compile_where_condition(left, table, cursor_idx)?;
- // If left is false, skip right evaluation and return 0
- let skip_to_false = self.program.len();
- self.emit(
+ // If left is false, the result stays 0; skip right entirely.
+ self.emit_jump_to(
OpCode::IfZ,
left_result,
- 0, // Will be patched
+ end_label,
0,
None,
- "if left is false, skip to false",
+ "if left is false, AND is false",
);
// Compile right condition
let right_result = self.compile_where_condition(right, table, cursor_idx)?;
- // Result is right_result (if we got here, left was true)
- let result_reg = self.allocate_register();
- self.emit(
- OpCode::Integer,
- 0,
- result_reg,
- 0,
- None,
- &format!("r[{}] = 0 (default AND result)", result_reg),
- );
-
- // If right is false, skip setting result to 1
- let skip_to_end = self.program.len();
- self.emit(
+ // If right is false, the result stays 0.
+ self.emit_jump_to(
OpCode::IfZ,
right_result,
- 0, // Will be patched to skip to end
+ end_label,
0,
None,
- "if right is false, skip to end",
+ "if right is false, AND is false",
);
// Both are true, set result to 1
@@ -5486,34 +5569,7 @@ impl<'a> SqlCompiler<'a> {
&format!("r[{}] = 1 (AND true)", result_reg),
);
- // Jump to end
- let goto_end = self.program.len();
- self.emit(
- OpCode::Goto,
- 0,
- 0, // Will be patched
- 0,
- None,
- "goto end",
- );
-
- // False label (left was false)
- let false_label = self.program.len();
- if let Some(inst) = self.program.instructions.get_mut(skip_to_false) {
- inst.p2 = false_label as i64;
- }
-
- // Result is already 0 (set above), so we just fall through
-
- // End label
- let end_label = self.program.len();
- if let Some(inst) = self.program.instructions.get_mut(skip_to_end) {
- inst.p2 = end_label as i64;
- }
- if let Some(inst) = self.program.instructions.get_mut(goto_end) {
- inst.p2 = end_label as i64;
- }
-
+ self.resolve(end_label);
return Ok(result_reg);
}
BinaryOperator::Or => {
@@ -5640,12 +5696,476 @@ impl<'a> SqlCompiler<'a> {
}
/// Compile a WHERE clause operand (column reference or literal)
+ /// Compile `expr` into a register, returning that register.
+ ///
+ /// This is the single expression compiler -- the analogue of SQLite's
+ /// `sqlite3ExprCode()`. It is used from every position an expression can
+ /// appear in, so a node supported in one position is supported in all of
+ /// them. sqawk previously had two: `compile_expr` (no `Identifier` arm at
+ /// all) and `compile_projection_expr`/`resolve_column_expr` ("Only column
+ /// references supported in SELECT"). That split is why `CASE` and `CAST`
+ /// worked in WHERE but not in a SELECT list, and why `SUBSTR` in WHERE
+ /// broke on the sqlparser upgrade while `SUBSTR` in a projection did not.
+ ///
+ /// `into` names the destination register when the caller needs the value
+ /// at a specific place -- which is how contiguous blocks for `ResultRow`
+ /// and `SorterInsert` are built without the copy-shuffling the old
+ /// projection paths did.
+ ///
+ /// Invariant: on `Ok(reg)`, code has been emitted that defines `reg` on
+ /// every path. There is deliberately no arm that returns `Ok` without
+ /// emitting -- that bug (an `_` arm that silently emitted nothing, leaving
+ /// the register at its default NULL) is why `SELECT name, salary*2 ...
+ /// LIMIT 3` returned NULL.
+ pub(crate) fn code_expr(
+ &mut self,
+ expr: &Expr,
+ ctx: &NameCtx<'_>,
+ into: Option,
+ ) -> SqawkResult {
+ let dest = into.unwrap_or_else(|| self.allocate_register());
+
+ match expr {
+ // Column references resolve through the context, so they work
+ // across joins and aliases and are matched case-insensitively.
+ Expr::Identifier(ident) => {
+ let col = ctx.resolve_unqualified(&ident.value)?;
+ self.emit_column_ref(col, dest, &ident.value);
+ }
+
+ // `table.column` / `alias.column`. The old WHERE compiler had no
+ // arm for this at all, so `FROM employees e WHERE e.age > 30`
+ // failed on a single-table query.
+ Expr::CompoundIdentifier(parts) if parts.len() >= 2 => {
+ let qualifier = &parts[parts.len() - 2].value;
+ let name = &parts[parts.len() - 1].value;
+ let col = ctx.resolve_qualified(qualifier, name)?;
+ self.emit_column_ref(col, dest, &format!("{}.{}", qualifier, name));
+ }
+
+ Expr::Nested(inner) => {
+ return self.code_expr(inner, ctx, Some(dest));
+ }
+
+ // Literals resolve against nothing, so they are compiled here
+ // rather than delegated. The fallback needs a single FROM source
+ // and would otherwise reject a constant in a join condition --
+ // `ON a.x = b.y AND b.id > 102` failing on the `102`.
+ Expr::Value(ValueWithSpan { value, .. }) => {
+ let literal = match value {
+ Value::Number(n, _) => match n.parse::() {
+ Ok(i) => TableValue::Integer(i),
+ Err(_) => match n.parse::() {
+ Ok(f) => TableValue::Float(f),
+ Err(_) => TableValue::String(n.clone().into()),
+ },
+ },
+ Value::SingleQuotedString(v) | Value::DoubleQuotedString(v) => {
+ TableValue::String(v.clone().into())
+ }
+ Value::Boolean(b) => TableValue::Boolean(*b),
+ Value::Null => TableValue::Null,
+ other => TableValue::String(other.to_string().into()),
+ };
+ self.emit_value_literal(dest, &literal)?;
+ }
+
+ // An aggregate call in a position where aggregates have already
+ // been computed: read the bound register rather than trying to
+ // compile SUM as a scalar function.
+ Expr::Function(_) if ctx.lookup_agg(expr).is_some() => {
+ let src = ctx.lookup_agg(expr).unwrap();
+ if src != dest {
+ self.emit(
+ OpCode::Copy,
+ src,
+ dest,
+ 0,
+ None,
+ &format!("r[{}] = aggregate r[{}]", dest, src),
+ );
+ }
+ }
+
+ // String concatenation. sqlparser has always produced this node;
+ // there was simply no arm for it, so `a || b` was reported as an
+ // unsupported operator. It lowers onto the CONCAT string function
+ // the engine already implements.
+ Expr::BinaryOp {
+ left,
+ op: BinaryOperator::StringConcat,
+ right,
+ } => {
+ let left_reg = self.code_expr(left, ctx, None)?;
+ let right_reg = self.code_expr(right, ctx, None)?;
+ self.emit(
+ OpCode::StringFunc,
+ left_reg,
+ dest,
+ 0,
+ Some(format!("CONCAT:{}", right_reg)),
+ &format!("r[{}] = r[{}] || r[{}]", dest, left_reg, right_reg),
+ );
+ }
+
+ // Logical AND / OR, short-circuiting.
+ //
+ // The result register is initialized BEFORE the short-circuit
+ // test. Registers persist across loop iterations, so a jump that
+ // skips the initialization leaves this expression holding its
+ // verdict from the PREVIOUS row -- which is exactly how AND came
+ // to report true for a row failing both operands.
+ Expr::BinaryOp {
+ left,
+ op: op @ (BinaryOperator::And | BinaryOperator::Or),
+ right,
+ } => {
+ let is_and = matches!(op, BinaryOperator::And);
+ // AND starts false and is set true only if both hold; OR
+ // starts false and is set true as soon as either does.
+ self.emit(
+ OpCode::Integer,
+ 0,
+ dest,
+ 0,
+ None,
+ &format!("r[{}] = 0 ({:?} init)", dest, op),
+ );
+ let done = self.label();
+
+ let left_reg = self.code_expr(left, ctx, None)?;
+ if is_and {
+ self.emit_jump_to(OpCode::IfZ, left_reg, done, 0, None, "AND: left false");
+ } else {
+ let check_right = self.label();
+ self.emit_jump_to(
+ OpCode::IfZ,
+ left_reg,
+ check_right,
+ 0,
+ None,
+ "OR: left false, try right",
+ );
+ self.emit(OpCode::Integer, 1, dest, 0, None, "OR: left true");
+ self.emit_jump_to(OpCode::Goto, 0, done, 0, None, "");
+ self.resolve(check_right);
+ }
+
+ let right_reg = self.code_expr(right, ctx, None)?;
+ self.emit_jump_to(OpCode::IfZ, right_reg, done, 0, None, "right false");
+ self.emit(OpCode::Integer, 1, dest, 0, None, "both/either true");
+
+ self.resolve(done);
+ }
+
+ // Arithmetic and comparison operators, compiled here rather than
+ // delegated, so the name-resolution context propagates into the
+ // operands. Delegating rebuilt a bare single-table context and
+ // dropped any aggregate bindings, which is what made
+ // `SUM(salary) + 1` report SUM as an unknown scalar function.
+ Expr::BinaryOp { left, op, right }
+ if matches!(
+ op,
+ BinaryOperator::Plus
+ | BinaryOperator::Minus
+ | BinaryOperator::Multiply
+ | BinaryOperator::Divide
+ | BinaryOperator::Modulo
+ | BinaryOperator::Gt
+ | BinaryOperator::Lt
+ | BinaryOperator::GtEq
+ | BinaryOperator::LtEq
+ | BinaryOperator::Eq
+ | BinaryOperator::NotEq
+ ) =>
+ {
+ let left_reg = self.code_expr(left, ctx, None)?;
+ let right_reg = self.code_expr(right, ctx, None)?;
+ let opcode = match op {
+ BinaryOperator::Plus => OpCode::Add,
+ BinaryOperator::Minus => OpCode::Subtract,
+ BinaryOperator::Multiply => OpCode::Multiply,
+ BinaryOperator::Divide => OpCode::Divide,
+ BinaryOperator::Modulo => OpCode::Remainder,
+ BinaryOperator::Gt => OpCode::Gt,
+ BinaryOperator::Lt => OpCode::Lt,
+ BinaryOperator::GtEq => OpCode::Ge,
+ BinaryOperator::LtEq => OpCode::Le,
+ BinaryOperator::Eq => OpCode::Eq,
+ BinaryOperator::NotEq => OpCode::Ne,
+ _ => unreachable!("guarded above"),
+ };
+ self.emit(
+ opcode,
+ left_reg,
+ right_reg,
+ dest,
+ None,
+ &format!("r[{}] = r[{}] {:?} r[{}]", dest, left_reg, op, right_reg),
+ );
+ }
+
+ // Logical negation, including NULL propagation. Previously
+ // unsupported in any position.
+ Expr::UnaryOp {
+ op: UnaryOperator::Not,
+ expr: inner,
+ } => {
+ let inner_reg = self.code_expr(inner, ctx, None)?;
+ self.emit(
+ OpCode::Not,
+ inner_reg,
+ dest,
+ 0,
+ None,
+ &format!("r[{}] = NOT r[{}]", dest, inner_reg),
+ );
+ }
+
+ // Function calls resolve their arguments in the CALLER's scope.
+ //
+ // Reaching them through the single-source fallback rebuilt a
+ // context containing only the innermost table, so inside a
+ // correlated subquery `UPPER(e.department)` silently resolved
+ // `e.department` against the SUBQUERY table -- making both sides of
+ // a comparison identical and matching every row.
+ Expr::Function(func) => {
+ self.compile_function(func, ctx, dest)?;
+ }
+
+ // TRIM / CEIL / FLOOR are their own Expr variants rather than
+ // Function calls. They were previously handled only inside the
+ // projection path, so they worked in a SELECT list and failed in
+ // WHERE -- the same asymmetry SUBSTR had.
+ Expr::Trim {
+ expr: inner,
+ trim_where: None,
+ trim_what: None,
+ ..
+ } => {
+ let src = self.code_expr(inner, ctx, None)?;
+ self.emit(
+ OpCode::StringFunc,
+ src,
+ dest,
+ 0,
+ Some("TRIM".to_string()),
+ &format!("r[{}] = TRIM(r[{}])", dest, src),
+ );
+ }
+ Expr::Ceil { expr: inner, .. } => {
+ let src = self.code_expr(inner, ctx, None)?;
+ self.emit(
+ OpCode::MathFunc,
+ src,
+ dest,
+ 0,
+ Some("CEIL".to_string()),
+ &format!("r[{}] = CEIL(r[{}])", dest, src),
+ );
+ }
+ Expr::Floor { expr: inner, .. } => {
+ let src = self.code_expr(inner, ctx, None)?;
+ self.emit(
+ OpCode::MathFunc,
+ src,
+ dest,
+ 0,
+ Some("FLOOR".to_string()),
+ &format!("r[{}] = FLOOR(r[{}])", dest, src),
+ );
+ }
+
+ // Predicate-shaped expressions used as VALUES.
+ //
+ // These are compiled by the condition compiler, which is a third
+ // parallel path. Routing them here means `SELECT x IS NULL`,
+ // `SELECT name LIKE 'A%'` and `SELECT age BETWEEN 1 AND 2` work as
+ // projected values, not just as filters -- SQL makes no
+ // distinction between the two positions.
+ Expr::IsNull(_)
+ | Expr::IsNotNull(_)
+ | Expr::Like { .. }
+ | Expr::ILike { .. }
+ | Expr::Between { .. }
+ | Expr::InList { .. } => {
+ let (table, cursor) = self.legacy_single_source(ctx, expr)?;
+ let src = self.compile_where_condition(expr, table, cursor)?;
+ if src != dest {
+ self.emit(
+ OpCode::Copy,
+ src,
+ dest,
+ 0,
+ None,
+ &format!("r[{}] = r[{}]", dest, src),
+ );
+ }
+ }
+
+ // Everything else still goes through the original operand
+ // compiler. Arms migrate out of it into the match above as the
+ // pipeline work proceeds; routing through here first means every
+ // call site picks up the new capabilities immediately.
+ other => {
+ let (table, cursor) = self.legacy_single_source(ctx, other)?;
+ self.compile_operand_legacy(other, table, cursor, dest)?;
+ }
+ }
+
+ Ok(dest)
+ }
+
+ /// Emit whatever loads a resolved column into `dest`.
+ fn emit_column_ref(&mut self, col: ColumnRef, dest: i64, label: &str) {
+ match col {
+ ColumnRef::Cursor { cursor, col } => self.emit(
+ OpCode::Column,
+ cursor,
+ col as i64,
+ dest,
+ None,
+ &format!("r[{}] = {} (column {})", dest, label, col),
+ ),
+ ColumnRef::Register(src) => {
+ if src != dest {
+ self.emit(
+ OpCode::Copy,
+ src,
+ dest,
+ 0,
+ None,
+ &format!("r[{}] = {} (r[{}])", dest, label, src),
+ );
+ }
+ }
+ }
+ }
+
+ /// `code_expr` with the argument order the function compilers use.
+ fn code_expr_into(
+ &mut self,
+ expr: &Expr,
+ ctx: &NameCtx<'_>,
+ target_reg: i64,
+ ) -> SqawkResult<()> {
+ self.code_expr(expr, ctx, Some(target_reg)).map(|_| ())
+ }
+
+ /// Compile `expr` as a boolean predicate, returning its register.
+ #[allow(dead_code)] // wired up by the staged SELECT pipeline
+ pub(crate) fn code_predicate(&mut self, expr: &Expr, ctx: &NameCtx<'_>) -> SqawkResult {
+ let (table, cursor) = self.legacy_single_source(ctx, expr)?;
+ self.compile_where_condition(expr, table, cursor)
+ }
+
+ /// Extract the single `(table, cursor)` pair the not-yet-migrated code
+ /// still expects.
+ ///
+ /// With one FROM source this is trivial. With several, the arms that have
+ /// not yet been migrated -- function calls above all -- still take a single
+ /// `(table, cursor)`, so the source is inferred from the column references
+ /// inside `expr`. A function almost always applies to columns of one
+ /// table, and when it genuinely spans two this errors rather than
+ /// resolving against an arbitrary one.
+ fn legacy_single_source<'t>(
+ &self,
+ ctx: &NameCtx<'t>,
+ expr: &Expr,
+ ) -> SqawkResult<(&'t Table, usize)> {
+ match ctx.srcs.len() {
+ 1 => return Ok((ctx.srcs[0].table, ctx.srcs[0].cursor as usize)),
+ 0 => {
+ return Err(SqawkError::UnsupportedSqlFeature(format!(
+ "Expression requires a FROM clause: {:?}",
+ expr
+ )))
+ }
+ _ => {}
+ }
+
+ let mut found: Option = None;
+ let mut ambiguous = false;
+ Self::for_each_column_ref(expr, &mut |qualifier, name| {
+ let idx = match qualifier {
+ Some(q) => ctx
+ .srcs
+ .iter()
+ .position(|s| s.refname == q.to_ascii_lowercase()),
+ None => ctx.srcs.iter().position(|s| {
+ s.table
+ .column_metadata()
+ .iter()
+ .any(|c| c.name.eq_ignore_ascii_case(name))
+ }),
+ };
+ if let Some(i) = idx {
+ match found {
+ Some(prev) if prev != i => ambiguous = true,
+ _ => found = Some(i),
+ }
+ }
+ });
+
+ match (found, ambiguous) {
+ (Some(i), false) => Ok((ctx.srcs[i].table, ctx.srcs[i].cursor as usize)),
+ (_, true) => Err(SqawkError::UnsupportedSqlFeature(format!(
+ "Expression spans multiple tables: {:?}",
+ expr
+ ))),
+ (None, _) => Ok((ctx.srcs[0].table, ctx.srcs[0].cursor as usize)),
+ }
+ }
+
+ /// Visit every column reference in `expr` as `(qualifier, name)`.
+ fn for_each_column_ref(expr: &Expr, f: &mut impl FnMut(Option<&str>, &str)) {
+ match expr {
+ Expr::Identifier(ident) => f(None, &ident.value),
+ Expr::CompoundIdentifier(parts) if parts.len() >= 2 => f(
+ Some(&parts[parts.len() - 2].value),
+ &parts[parts.len() - 1].value,
+ ),
+ Expr::BinaryOp { left, right, .. } => {
+ Self::for_each_column_ref(left, f);
+ Self::for_each_column_ref(right, f);
+ }
+ Expr::UnaryOp { expr: inner, .. }
+ | Expr::Nested(inner)
+ | Expr::Cast { expr: inner, .. }
+ | Expr::IsNull(inner)
+ | Expr::IsNotNull(inner) => Self::for_each_column_ref(inner, f),
+ Expr::Function(func) => {
+ for arg in func_args(func) {
+ if let FunctionArg::Unnamed(FunctionArgExpr::Expr(e)) = arg {
+ Self::for_each_column_ref(e, f);
+ }
+ }
+ }
+ _ => {}
+ }
+ }
+
+ /// Shim preserving the old operand entry point.
+ ///
+ /// Every existing caller keeps working and gains the new capabilities
+ /// (qualified names, case-insensitive resolution, `||`, `NOT`) for free.
pub(crate) fn compile_where_operand(
&mut self,
expr: &Expr,
table: &Table,
cursor_idx: usize,
target_reg: i64,
+ ) -> SqawkResult<()> {
+ let ctx = NameCtx::single(table, cursor_idx as i64);
+ self.code_expr(expr, &ctx, Some(target_reg)).map(|_| ())
+ }
+
+ fn compile_operand_legacy(
+ &mut self,
+ expr: &Expr,
+ table: &Table,
+ cursor_idx: usize,
+ target_reg: i64,
) -> SqawkResult<()> {
match expr {
Expr::Identifier(ident) => {
@@ -5665,7 +6185,7 @@ impl<'a> SqlCompiler<'a> {
),
);
}
- Expr::Value(value) => {
+ Expr::Value(ValueWithSpan { value, .. }) => {
// This is a literal value
match value {
Value::Number(num, _) => {
@@ -5726,14 +6246,13 @@ impl<'a> SqlCompiler<'a> {
Expr::Case {
operand,
conditions,
- results,
else_result,
+ ..
} => {
// Compile CASE expression, storing result directly in target_reg
self.compile_case(
operand.as_deref(),
conditions,
- results,
else_result.as_deref(),
table,
cursor_idx,
@@ -5804,7 +6323,11 @@ impl<'a> SqlCompiler<'a> {
}
Expr::Function(func) => {
// Handle COALESCE and NULLIF functions
- self.compile_function(func, table, cursor_idx, target_reg)?;
+ self.compile_function(
+ func,
+ &NameCtx::single(table, cursor_idx as i64),
+ target_reg,
+ )?;
}
Expr::Subquery(subquery) => {
// Handle scalar subquery as an operand
@@ -5995,6 +6518,25 @@ impl<'a> SqlCompiler<'a> {
}
}
}
+ // `SUBSTR(x, 1, 3)` parses as `Expr::Substring`, not as a
+ // `Expr::Function`, so it needs an arm of its own here. (Older
+ // sqlparser reserved this variant for the `SUBSTRING(x FROM 1 FOR
+ // 3)` spelling and routed the comma form through Function.)
+ Expr::Substring {
+ expr: sub_expr,
+ substring_from,
+ substring_for,
+ ..
+ } => {
+ self.compile_substring(
+ sub_expr,
+ substring_from.as_deref(),
+ substring_for.as_deref(),
+ table,
+ cursor_idx,
+ target_reg,
+ )?;
+ }
_ => {
return Err(SqawkError::UnsupportedSqlFeature(format!(
"Unsupported WHERE operand: {:?}",
@@ -6023,11 +6565,13 @@ impl<'a> SqlCompiler<'a> {
| SqlDataType::Integer(_)
| SqlDataType::BigInt(_)
| SqlDataType::SmallInt(_) => "INTEGER".to_string(),
- SqlDataType::Real | SqlDataType::Float(_) | SqlDataType::Double => "REAL".to_string(),
+ SqlDataType::Real | SqlDataType::Float(_) | SqlDataType::Double(_) => {
+ "REAL".to_string()
+ }
SqlDataType::Text
| SqlDataType::Varchar(_)
| SqlDataType::Char(_)
- | SqlDataType::String => "TEXT".to_string(),
+ | SqlDataType::String(_) => "TEXT".to_string(),
SqlDataType::Boolean => "BOOLEAN".to_string(),
_ => format!("{:?}", data_type),
}
@@ -6041,7 +6585,129 @@ impl<'a> SqlCompiler<'a> {
// Use the last part of the object name as the table name
// (schemas and catalogs are ignored for now)
- Ok(name.0.last().unwrap().value.clone())
+ Ok(object_name_last(name))
+ }
+
+ /// Whether a projection item needs expression compilation rather than a
+ /// bare column load.
+ ///
+ /// Phrased as "not a plain column reference" rather than as an allowlist
+ /// of computed forms. The allowlists this replaces named nine expression
+ /// kinds and omitted CASE, CAST, IS NULL, LIKE, BETWEEN, IN and
+ /// subqueries, so those fell through to the column-index path and were
+ /// rejected with "Only column references supported in SELECT" -- even
+ /// though the expression compiler handles them. An allowlist has to be
+ /// updated every time an expression kind is added; this cannot drift.
+ fn projection_needs_expr(projection: &[SelectItem]) -> bool {
+ projection.iter().any(|item| {
+ let expr = match item {
+ SelectItem::UnnamedExpr(e) | SelectItem::ExprWithAlias { expr: e, .. } => e,
+ _ => return false,
+ };
+ !matches!(expr, Expr::Identifier(_) | Expr::CompoundIdentifier(_))
+ })
+ }
+
+ /// Emit post-processing ORDER BY / LIMIT / OFFSET over the result set.
+ ///
+ /// Used by the paths that produce rows without going through the
+ /// cursor-based sorter -- GROUP BY, plain aggregates and joins. Those
+ /// paths previously dropped ORDER BY and LIMIT on the floor: `query` was
+ /// simply never passed to them, and for joins the code said so outright
+ /// ("ORDER BY support for JOINs [...] isn't fully implemented. For now,
+ /// compile the JOIN and skip ORDER BY").
+ ///
+ /// ORDER BY keys are matched against the SELECT list, since after
+ /// aggregation the only columns that still exist are the projected ones.
+ fn emit_result_post_processing(&mut self, select: &Select, query: &Query) -> SqawkResult<()> {
+ // ORDER BY
+ let order_by = query_order_by(query);
+ if !order_by.is_empty() {
+ let mut keys: Vec = Vec::with_capacity(order_by.len());
+ for ob in order_by {
+ let pos =
+ Self::projection_position(&select.projection, &ob.expr).ok_or_else(|| {
+ SqawkError::InvalidSqlQuery(format!(
+ "ORDER BY expression must appear in the SELECT list: {}",
+ ob.expr
+ ))
+ })?;
+ keys.push(format!(
+ "{}:{}",
+ pos,
+ if order_by_is_asc(ob) { "asc" } else { "desc" }
+ ));
+ }
+ let spec = keys.join(",");
+ self.emit(
+ OpCode::SortResults,
+ 0,
+ 0,
+ 0,
+ Some(spec.clone()),
+ &format!("ORDER BY {}", spec),
+ );
+ }
+
+ // LIMIT / OFFSET
+ let limit = match query_limit(query) {
+ Some(e) => Some(Self::const_i64(e, "LIMIT")?),
+ None => None,
+ };
+ let offset = match query_offset(query) {
+ Some(e) => Self::const_i64(e, "OFFSET")?,
+ None => 0,
+ };
+ if limit.is_some() || offset > 0 {
+ self.emit(
+ OpCode::Limit,
+ limit.unwrap_or(i64::MAX),
+ offset,
+ 0,
+ None,
+ &format!("Limit {:?} offset {}", limit, offset),
+ );
+ }
+ Ok(())
+ }
+
+ /// Position of `needle` within the SELECT list, if it appears there.
+ ///
+ /// Matches an alias by name, and otherwise compares the rendered
+ /// expression, so `ORDER BY COUNT(*)` finds `COUNT(*)` in the projection
+ /// and `ORDER BY department` finds a bare column.
+ fn projection_position(projection: &[SelectItem], needle: &Expr) -> Option {
+ let want = needle.to_string().to_ascii_lowercase();
+ for (i, item) in projection.iter().enumerate() {
+ let matches = match item {
+ SelectItem::UnnamedExpr(e) => e.to_string().to_ascii_lowercase() == want,
+ SelectItem::ExprWithAlias { expr, alias } => {
+ alias.value.to_ascii_lowercase() == want
+ || expr.to_string().to_ascii_lowercase() == want
+ }
+ _ => false,
+ };
+ if matches {
+ return Some(i);
+ }
+ }
+ None
+ }
+
+ /// Evaluate a LIMIT/OFFSET operand, which must be a constant.
+ fn const_i64(expr: &Expr, what: &str) -> SqawkResult {
+ match expr {
+ Expr::Value(ValueWithSpan {
+ value: Value::Number(n, _),
+ ..
+ }) => n
+ .parse::()
+ .map_err(|_| SqawkError::InvalidSqlQuery(format!("Invalid {} value: {}", what, n))),
+ _ => Err(SqawkError::UnsupportedSqlFeature(format!(
+ "Only constant {} supported",
+ what
+ ))),
+ }
}
/// Resolve projection items to column indices
@@ -6054,6 +6720,13 @@ impl<'a> SqlCompiler<'a> {
for item in projection {
match item {
+ // Multiple aliases for one item (`expr AS (a, b)`) is a
+ // non-standard extension sqawk does not implement.
+ SelectItem::ExprWithAliases { .. } => {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "Multiple column aliases are not supported".into(),
+ ))
+ }
SelectItem::Wildcard(_) => {
// SELECT * - include all columns
for i in 0..table.column_count() {
@@ -6122,11 +6795,11 @@ impl<'a> SqlCompiler<'a> {
if matches!(func_name.as_str(), "COUNT" | "SUM" | "AVG" | "MIN" | "MAX") {
func_name
} else if let Some(FunctionArg::Unnamed(FunctionArgExpr::Wildcard)) =
- func.args.first()
+ func_args(func).first()
{
format!("{}(*)", func_name)
} else if let Some(FunctionArg::Unnamed(FunctionArgExpr::Expr(arg_expr))) =
- func.args.first()
+ func_args(func).first()
{
format!(
"{}({})",
@@ -6137,6 +6810,12 @@ impl<'a> SqlCompiler<'a> {
format!("{}()", func_name)
}
}
+ // SUBSTR/SUBSTRING is its own Expr variant rather than a Function,
+ // so it needs an arm here to keep the descriptive column name
+ // (`SUBSTR(name)`) instead of falling through to `expr`.
+ Expr::Substring { expr, .. } => {
+ format!("SUBSTR({})", self.get_column_name_from_expr(expr))
+ }
_ => "expr".to_string(),
}
}
@@ -6154,6 +6833,10 @@ impl<'a> SqlCompiler<'a> {
for item in projection {
match item {
+ // Multiple aliases for one item (`expr AS (a, b)`) is a
+ // non-standard extension sqawk does not implement. This
+ // function cannot fail; the compile path rejects the item.
+ SelectItem::ExprWithAliases { .. } => {}
SelectItem::Wildcard(_) | SelectItem::QualifiedWildcard(_, _) => {
// Add all columns from the table with their types
for meta in col_metadata {
@@ -6210,7 +6893,7 @@ impl<'a> SqlCompiler<'a> {
} else {
// Try to infer from argument
if let Some(FunctionArg::Unnamed(FunctionArgExpr::Expr(arg))) =
- func.args.first()
+ func_args(func).first()
{
let arg_type = self.infer_expr_type(arg, table);
if matches!(arg_type, DataType::Integer) {
@@ -6223,7 +6906,7 @@ impl<'a> SqlCompiler<'a> {
"MIN" | "MAX" => {
// MIN/MAX preserves the type of the argument
if let Some(FunctionArg::Unnamed(FunctionArgExpr::Expr(arg))) =
- func.args.first()
+ func_args(func).first()
{
return self.infer_expr_type(arg, table);
}
@@ -6235,7 +6918,7 @@ impl<'a> SqlCompiler<'a> {
"COALESCE" => {
// COALESCE returns the type of its first non-null argument
if let Some(FunctionArg::Unnamed(FunctionArgExpr::Expr(arg))) =
- func.args.first()
+ func_args(func).first()
{
return self.infer_expr_type(arg, table);
}
@@ -6244,7 +6927,7 @@ impl<'a> SqlCompiler<'a> {
_ => DataType::Text,
}
}
- Expr::Value(val) => {
+ Expr::Value(ValueWithSpan { value: val, .. }) => {
match val {
Value::Number(_, _) => DataType::Integer, // Could be Float, but Integer is common
Value::SingleQuotedString(_) | Value::DoubleQuotedString(_) => DataType::Text,
@@ -6274,13 +6957,13 @@ impl<'a> SqlCompiler<'a> {
}
}
Expr::Case {
- results,
+ conditions,
else_result,
..
} => {
// CASE returns the type of its first result
- if let Some(first_result) = results.first() {
- return self.infer_expr_type(first_result, table);
+ if let Some(first) = conditions.first() {
+ return self.infer_expr_type(&first.result, table);
}
if let Some(else_expr) = else_result {
return self.infer_expr_type(else_expr, table);
@@ -6294,7 +6977,7 @@ impl<'a> SqlCompiler<'a> {
| SqlDataType::Integer(_)
| SqlDataType::BigInt(_)
| SqlDataType::SmallInt(_) => DataType::Integer,
- SqlDataType::Real | SqlDataType::Float(_) | SqlDataType::Double => {
+ SqlDataType::Real | SqlDataType::Float(_) | SqlDataType::Double(_) => {
DataType::Float
}
SqlDataType::Boolean => DataType::Boolean,
@@ -6309,10 +6992,10 @@ impl<'a> SqlCompiler<'a> {
fn has_aggregates(&self, projection: &[SelectItem]) -> bool {
for item in projection {
match item {
- SelectItem::UnnamedExpr(expr) | SelectItem::ExprWithAlias { expr, .. } => {
- if self.is_aggregate_expr(expr) {
- return true;
- }
+ SelectItem::UnnamedExpr(expr) | SelectItem::ExprWithAlias { expr, .. }
+ if self.is_aggregate_expr(expr) =>
+ {
+ return true;
}
_ => {}
}
@@ -6320,10 +7003,13 @@ impl<'a> SqlCompiler<'a> {
false
}
- /// Check if an expression is or contains an aggregate function
- fn is_aggregate_expr(&self, expr: &Expr) -> bool {
+ /// Whether `expr` is directly an aggregate call.
+ ///
+ /// A free function so the grouped-projection planner can use it without a
+ /// compiler borrow.
+ pub(crate) fn is_aggregate_call(expr: &Expr) -> bool {
match expr {
- Expr::Function(func) => {
+ Expr::Function(func) if func.over.is_none() => {
let name = func.name.to_string().to_uppercase();
matches!(name.as_str(), "COUNT" | "SUM" | "AVG" | "MIN" | "MAX")
}
@@ -6331,14 +7017,35 @@ impl<'a> SqlCompiler<'a> {
}
}
+ /// Check if an expression is or CONTAINS an aggregate function.
+ ///
+ /// Recursive: this used to match only a top-level call, so
+ /// `SELECT SUM(salary) + 1` was not recognised as an aggregate query and
+ /// took the plain table-scan path, where SUM is not a known scalar
+ /// function.
+ fn is_aggregate_expr(&self, expr: &Expr) -> bool {
+ if Self::is_aggregate_call(expr) {
+ return true;
+ }
+ match expr {
+ Expr::BinaryOp { left, right, .. } => {
+ self.is_aggregate_expr(left) || self.is_aggregate_expr(right)
+ }
+ Expr::UnaryOp { expr: inner, .. }
+ | Expr::Nested(inner)
+ | Expr::Cast { expr: inner, .. } => self.is_aggregate_expr(inner),
+ _ => false,
+ }
+ }
+
/// Check if the projection contains any window functions
fn has_window_functions(&self, projection: &[SelectItem]) -> bool {
for item in projection {
match item {
- SelectItem::UnnamedExpr(expr) | SelectItem::ExprWithAlias { expr, .. } => {
- if self.is_window_expr(expr) {
- return true;
- }
+ SelectItem::UnnamedExpr(expr) | SelectItem::ExprWithAlias { expr, .. }
+ if self.is_window_expr(expr) =>
+ {
+ return true;
}
_ => {}
}
diff --git a/src/vm/compiler_aggregate.rs b/src/vm/compiler_aggregate.rs
index b70a990..446f281 100644
--- a/src/vm/compiler_aggregate.rs
+++ b/src/vm/compiler_aggregate.rs
@@ -4,12 +4,44 @@
use sqlparser::ast::{Expr, FunctionArg, FunctionArgExpr, Select, SelectItem};
-use super::bytecode::{OpCode, ResultSchema};
-use super::compiler::SqlCompiler;
+use super::ast_compat::{func_args, func_is_distinct, group_by_exprs};
+use super::bytecode::{OpCode, ResultSchema, AGG_DISTINCT};
+use super::compiler::{NameCtx, SqlCompiler};
use crate::error::{SqawkError, SqawkResult};
use crate::table::Table;
+/// One column of a grouped result, in SELECT-list order.
+#[derive(Clone, Copy, Debug)]
+enum GroupOut {
+ /// The n-th GROUP BY key.
+ Key(usize),
+ /// The n-th aggregate.
+ Agg(usize),
+}
+
impl<'a> SqlCompiler<'a> {
+ /// Collect every aggregate call within `expr`, in order of appearance.
+ ///
+ /// Detection used to look only at the top level, so `SUM(salary) + 1` was
+ /// not recognised as an aggregate query at all and fell through to the
+ /// plain table-scan path, where SUM is not a known scalar function.
+ fn collect_aggregates<'e>(expr: &'e Expr, out: &mut Vec<&'e Expr>) {
+ if Self::is_aggregate_call(expr) {
+ out.push(expr);
+ return;
+ }
+ match expr {
+ Expr::BinaryOp { left, right, .. } => {
+ Self::collect_aggregates(left, out);
+ Self::collect_aggregates(right, out);
+ }
+ Expr::UnaryOp { expr: inner, .. }
+ | Expr::Nested(inner)
+ | Expr::Cast { expr: inner, .. } => Self::collect_aggregates(inner, out),
+ _ => {}
+ }
+ }
+
pub(crate) fn compile_select_with_aggregate(
&mut self,
select: &Select,
@@ -59,20 +91,23 @@ impl<'a> SqlCompiler<'a> {
);
}
- // For each aggregate in the projection, emit AggStep
- let mut acc_regs = Vec::new();
+ // Gather every aggregate appearing anywhere in the projection, then
+ // step each one per row.
+ let mut agg_calls: Vec<&Expr> = Vec::new();
for item in &select.projection {
- match item {
- SelectItem::UnnamedExpr(expr) | SelectItem::ExprWithAlias { expr, .. } => {
- if let Some((func_type, col_reg)) =
- self.compile_aggregate_step(expr, table, cursor_idx)?
- {
- let acc_reg = self.allocate_register();
- acc_regs.push((acc_reg, func_type));
- self.emit(OpCode::AggStep, func_type, col_reg, acc_reg, None, "");
- }
- }
- _ => {}
+ if let SelectItem::UnnamedExpr(expr) | SelectItem::ExprWithAlias { expr, .. } = item {
+ Self::collect_aggregates(expr, &mut agg_calls);
+ }
+ }
+
+ let mut acc_regs = Vec::new();
+ for call in &agg_calls {
+ if let Some((func_type, col_reg)) =
+ self.compile_aggregate_step(call, table, cursor_idx)?
+ {
+ let acc_reg = self.allocate_register();
+ acc_regs.push(acc_reg);
+ self.emit(OpCode::AggStep, func_type, col_reg, acc_reg, None, "");
}
}
@@ -94,25 +129,48 @@ impl<'a> SqlCompiler<'a> {
inst.p2 = after_loop as i64;
}
- // Finalize aggregates and output result
- let result_start_reg = self.allocate_registers(acc_regs.len());
-
- for (i, (acc_reg, _)) in acc_regs.iter().enumerate() {
+ // Finalize each aggregate into its own register.
+ let agg_start_reg = self.allocate_registers(acc_regs.len().max(1));
+ for (i, acc_reg) in acc_regs.iter().enumerate() {
self.emit(
OpCode::AggFinal,
*acc_reg,
- result_start_reg + i as i64,
+ agg_start_reg + i as i64,
0,
None,
"",
);
}
+ // Bind them by call text, then compile each projection item as an
+ // ordinary expression. That is what makes `SUM(salary) + 1` work: the
+ // arithmetic is compiled by the shared expression compiler, which
+ // resolves the inner SUM to its finalized register.
+ let mut ctx = NameCtx::single(table, cursor_idx);
+ for (i, call) in agg_calls.iter().enumerate() {
+ ctx.agg_bindings
+ .push((NameCtx::agg_key(call), agg_start_reg + i as i64));
+ }
+
+ let out_items: Vec<&Expr> = select
+ .projection
+ .iter()
+ .filter_map(|item| match item {
+ SelectItem::UnnamedExpr(e) | SelectItem::ExprWithAlias { expr: e, .. } => Some(e),
+ _ => None,
+ })
+ .collect();
+
+ let result_start_reg = self.allocate_registers(out_items.len().max(1));
+ for (i, expr) in out_items.iter().enumerate() {
+ self.code_expr(expr, &ctx, Some(result_start_reg + i as i64))?;
+ }
+
// Output result row
self.emit(
OpCode::ResultRow,
result_start_reg,
- acc_regs.len() as i64,
+ out_items.len() as i64,
0,
None,
"",
@@ -134,28 +192,36 @@ impl<'a> SqlCompiler<'a> {
match expr {
Expr::Function(func) => {
let name = func.name.to_string().to_uppercase();
- let func_type = Self::agg_func_type(&name);
+ // DISTINCT rides along as a flag bit on the function type.
+ // It used to be dropped entirely -- func.distinct was never
+ // read -- so COUNT(DISTINCT department) counted rows.
+ let func_type = Self::agg_func_type(&name)
+ | if func_is_distinct(func) {
+ AGG_DISTINCT
+ } else {
+ 0
+ };
// Check for COUNT(*)
- if let Some(FunctionArg::Unnamed(FunctionArgExpr::Wildcard)) = func.args.first() {
+ if let Some(FunctionArg::Unnamed(FunctionArgExpr::Wildcard)) =
+ func_args(func).first()
+ {
// COUNT(*) - use -1 to indicate no column
return Ok(Some((func_type, -1)));
}
- // Get the column argument
+ // Compile the argument as a full expression.
+ //
+ // This used to call resolve_column_expr, which yields a column
+ // INDEX and so only accepts a bare column reference --
+ // `SUM(salary + age)` was rejected with "Only column
+ // references supported in SELECT". An aggregate argument is
+ // an ordinary expression evaluated once per row.
if let Some(FunctionArg::Unnamed(FunctionArgExpr::Expr(arg_expr))) =
- func.args.first()
+ func_args(func).first()
{
- let col_idx = self.resolve_column_expr(arg_expr, table)?;
- let col_reg = self.allocate_register();
- self.emit(
- OpCode::Column,
- cursor_idx,
- col_idx as i64,
- col_reg,
- None,
- "",
- );
+ let ctx = NameCtx::single(table, cursor_idx);
+ let col_reg = self.code_expr(arg_expr, &ctx, None)?;
return Ok(Some((func_type, col_reg)));
}
@@ -189,6 +255,91 @@ impl<'a> SqlCompiler<'a> {
schema
}
+ /// Emit one grouped output row: keys, finalized aggregates, HAVING, ResultRow.
+ ///
+ /// Written once and called from both flush points -- the group-change flush
+ /// inside the drain loop and the final flush after it. Those were
+ /// copy-pasted duplicates, which is how HAVING came to be applied slightly
+ /// differently in each.
+ ///
+ /// Values are emitted in PROJECTION order. The result block used to be laid
+ /// out positionally as [group keys.., aggregates..] while the schema was
+ /// built from the SELECT list, so `SELECT COUNT(*), department ... GROUP BY
+ /// department` printed the header `COUNT,department` above the values
+ /// `Engineering,3`. Emitting through out_plan makes header and value order
+ /// the same by construction, and means a group key that is not projected
+ /// contributes no column instead of a phantom `col1`.
+ #[allow(clippy::too_many_arguments)]
+ fn emit_group_flush(
+ &mut self,
+ select: &Select,
+ table: &Table,
+ out_plan: &[GroupOut],
+ group_col_indices: &[usize],
+ agg_info: &[(i64, Option)],
+ group_key_reg: i64,
+ acc_base_reg: i64,
+ result_start_reg: i64,
+ ) -> SqawkResult<()> {
+ // Finalize every accumulator into its slot in the internal
+ // [keys.., aggs..] block, which HAVING resolves against.
+ for i in 0..agg_info.len() {
+ self.emit(
+ OpCode::AggFinal,
+ acc_base_reg + i as i64,
+ result_start_reg + group_col_indices.len() as i64 + i as i64,
+ 0,
+ None,
+ "",
+ );
+ }
+ for i in 0..group_col_indices.len() {
+ self.emit(
+ OpCode::Copy,
+ group_key_reg + i as i64,
+ result_start_reg + i as i64,
+ 0,
+ None,
+ "",
+ );
+ }
+
+ // HAVING gates the row, and is evaluated against the internal block.
+ let after_row = self.label();
+ if let Some(having_expr) = &select.having {
+ let having_reg = self.compile_having_condition(
+ having_expr,
+ table,
+ group_col_indices,
+ agg_info,
+ result_start_reg,
+ group_key_reg,
+ )?;
+ self.emit_jump_to(OpCode::IfZ, having_reg, after_row, 0, None, "HAVING failed");
+ }
+
+ // Gather into a contiguous block in projection order and emit.
+ let out_start = self.allocate_registers(out_plan.len().max(1));
+ for (i, slot) in out_plan.iter().enumerate() {
+ let src = match slot {
+ GroupOut::Key(k) => result_start_reg + *k as i64,
+ GroupOut::Agg(a) => result_start_reg + group_col_indices.len() as i64 + *a as i64,
+ };
+ self.emit(OpCode::Copy, src, out_start + i as i64, 0, None, "");
+ }
+ self.emit(
+ OpCode::ResultRow,
+ out_start,
+ out_plan.len() as i64,
+ 0,
+ None,
+ "",
+ );
+
+ self.resolve(after_row);
+ Ok(())
+ }
+
/// Compile a SELECT with GROUP BY
pub(crate) fn compile_select_with_group_by(
&mut self,
@@ -201,7 +352,7 @@ impl<'a> SqlCompiler<'a> {
// Determine GROUP BY column indices
let mut group_col_indices: Vec = Vec::new();
- for expr in select.group_by.iter() {
+ for expr in group_by_exprs(&select.group_by).iter() {
match expr {
Expr::Identifier(ident) => {
let col_name = ident.value.to_lowercase();
@@ -238,12 +389,12 @@ impl<'a> SqlCompiler<'a> {
if matches!(name.as_str(), "COUNT" | "SUM" | "AVG" | "MIN" | "MAX") {
let func_type = Self::agg_func_type(&name);
if let Some(FunctionArg::Unnamed(FunctionArgExpr::Wildcard)) =
- func.args.first()
+ func_args(func).first()
{
agg_info.push((func_type, None));
} else if let Some(FunctionArg::Unnamed(FunctionArgExpr::Expr(
arg_expr,
- ))) = func.args.first()
+ ))) = func_args(func).first()
{
let col_idx = self.resolve_column_expr(arg_expr, table)?;
agg_info.push((func_type, Some(col_idx)));
@@ -255,6 +406,45 @@ impl<'a> SqlCompiler<'a> {
}
}
+ // Map each projected item to the group key or aggregate it names, in
+ // SELECT-list order. Emitting through this is what keeps the values
+ // aligned with the header the schema builder produces.
+ let mut out_plan: Vec = Vec::with_capacity(select.projection.len());
+ let mut agg_seen = 0usize;
+ for item in &select.projection {
+ let expr = match item {
+ SelectItem::UnnamedExpr(e) | SelectItem::ExprWithAlias { expr: e, .. } => e,
+ _ => {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "Only expressions are supported in a GROUP BY projection".into(),
+ ))
+ }
+ };
+ if Self::is_aggregate_call(expr) {
+ out_plan.push(GroupOut::Agg(agg_seen));
+ agg_seen += 1;
+ continue;
+ }
+ // Otherwise it must name a GROUP BY key: in a grouped query no
+ // other column has a single well-defined value per row.
+ let col_idx = self.resolve_column_expr(expr, table).map_err(|_| {
+ SqawkError::InvalidSqlQuery(format!(
+ "'{}' must appear in GROUP BY or be used in an aggregate",
+ expr
+ ))
+ })?;
+ let key_pos = group_col_indices
+ .iter()
+ .position(|c| *c == col_idx)
+ .ok_or_else(|| {
+ SqawkError::InvalidSqlQuery(format!(
+ "'{}' must appear in GROUP BY or be used in an aggregate",
+ expr
+ ))
+ })?;
+ out_plan.push(GroupOut::Key(key_pos));
+ }
+
// Build result schema with column names and types
let schema = self.build_aggregate_result_schema(&select.projection, table);
self.program.set_result_schema(schema);
@@ -297,6 +487,26 @@ impl<'a> SqlCompiler<'a> {
let loop_start = self.program.len();
+ // WHERE filtering.
+ //
+ // This function never read select.selection, so WHERE was silently
+ // ignored on EVERY grouped query: `... WHERE salary > 60000 GROUP BY
+ // department` returned all departments with unfiltered counts. Filter
+ // before feeding the sorter so excluded rows never reach an
+ // accumulator.
+ let next_label = self.label();
+ if let Some(where_expr) = &select.selection {
+ let cond_reg = self.compile_where_condition(where_expr, table, cursor_idx as usize)?;
+ self.emit_jump_to(
+ OpCode::IfZ,
+ cond_reg,
+ next_label,
+ 0,
+ None,
+ "Skip row if WHERE is false",
+ );
+ }
+
// Allocate registers for row data
let row_start_reg = self.allocate_registers(total_cols);
@@ -341,6 +551,7 @@ impl<'a> SqlCompiler<'a> {
);
// Next row
+ self.resolve(next_label);
self.emit(OpCode::Next, cursor_idx, loop_start as i64, 0, None, "");
let after_scan = self.program.len();
@@ -391,72 +602,44 @@ impl<'a> SqlCompiler<'a> {
let first_row_jump_addr = self.program.len();
self.emit(OpCode::IfPos, first_row_reg, 0, 0, None, "");
- // Compare first group column with saved group key
- let cmp_reg = self.allocate_register();
- self.emit(OpCode::Ne, row_start_reg, group_key_reg, cmp_reg, None, "");
-
- // If group NOT changed (cmp_reg is 0), skip output and go to step aggregates
- let skip_output_addr = self.program.len();
- self.emit(OpCode::IfZ, cmp_reg, 0, 0, None, "");
-
- // Group changed - output previous group
- // Copy group key to result
- for i in 0..group_col_indices.len() {
- self.emit(
- OpCode::Copy,
- group_key_reg + i as i64,
- result_start_reg + i as i64,
- 0,
- None,
- "",
- );
- }
-
- // Finalize aggregates
- for (i, _) in agg_info.iter().enumerate() {
- let dest_reg = result_start_reg + group_col_indices.len() as i64 + i as i64;
- self.emit(
- OpCode::AggFinal,
- acc_base_reg + i as i64,
- dest_reg,
- 0,
- None,
- "",
- );
- }
-
- // Apply HAVING filter if present
- let mut having_skip_addr: Option = None;
- if let Some(having_expr) = &select.having {
- let having_reg = self.compile_having_condition(
- having_expr,
- &group_col_indices,
- &agg_info,
- result_start_reg,
- group_key_reg,
- )?;
- // If HAVING condition is false (0), skip ResultRow
- having_skip_addr = Some(self.program.len());
- self.emit(OpCode::IfZ, having_reg, 0, 0, None, "");
- }
-
- // Output result row
+ // Compare ALL group key columns against the saved key.
+ //
+ // This was a single `Ne` against column 0, so `GROUP BY department,
+ // role` grouped on department alone and reported whichever role
+ // happened to arrive first. Compare/Jump handles the whole key vector
+ // in two instructions regardless of width, and treats NULL as equal to
+ // NULL so NULL keys group together rather than starting a new group on
+ // every row.
self.emit(
- OpCode::ResultRow,
- result_start_reg,
- (group_col_indices.len() + agg_info.len()) as i64,
- 0,
+ OpCode::Compare,
+ row_start_reg,
+ group_key_reg,
+ group_col_indices.len() as i64,
None,
- "",
+ "Compare group key vector",
);
-
- // Patch HAVING skip jump if present
- if let Some(addr) = having_skip_addr {
- let after_result = self.program.len();
- if let Some(inst) = self.program.instructions.get_mut(addr) {
- inst.p2 = after_result as i64;
- }
- }
+ // Equal -> same group, skip the flush. Less/Greater -> group changed.
+ let flush_label = self.label();
+ let skip_output_label = self.label();
+ self.emit_jump_to_three(
+ flush_label,
+ skip_output_label,
+ flush_label,
+ "Group changed?",
+ );
+ self.resolve(flush_label);
+
+ // Group changed - output the group just completed.
+ self.emit_group_flush(
+ select,
+ table,
+ &out_plan,
+ &group_col_indices,
+ &agg_info,
+ group_key_reg,
+ acc_base_reg,
+ result_start_reg,
+ )?;
// Reset accumulators for new group
for (i, _) in agg_info.iter().enumerate() {
@@ -486,13 +669,8 @@ impl<'a> SqlCompiler<'a> {
// Clear first row flag
self.emit(OpCode::Integer, 0, first_row_reg, 0, None, "");
- // === Step aggregates section (skip output jumps here) ===
- let step_agg_addr = self.program.len();
-
- // Patch skip output jump
- if let Some(inst) = self.program.instructions.get_mut(skip_output_addr) {
- inst.p2 = step_agg_addr as i64;
- }
+ // === Step aggregates section (same-group rows jump here) ===
+ self.resolve(skip_output_label);
// Step aggregates for current row
for (i, (func_type, _)) in agg_info.iter().enumerate() {
@@ -517,64 +695,17 @@ impl<'a> SqlCompiler<'a> {
"",
);
- // Output final group
- // Copy group key to result
- for i in 0..group_col_indices.len() {
- self.emit(
- OpCode::Copy,
- group_key_reg + i as i64,
- result_start_reg + i as i64,
- 0,
- None,
- "",
- );
- }
-
- // Finalize aggregates
- for (i, _) in agg_info.iter().enumerate() {
- let dest_reg = result_start_reg + group_col_indices.len() as i64 + i as i64;
- self.emit(
- OpCode::AggFinal,
- acc_base_reg + i as i64,
- dest_reg,
- 0,
- None,
- "",
- );
- }
-
- // Apply HAVING filter for final group if present
- let mut final_having_skip_addr: Option = None;
- if let Some(having_expr) = &select.having {
- let having_reg = self.compile_having_condition(
- having_expr,
- &group_col_indices,
- &agg_info,
- result_start_reg,
- group_key_reg,
- )?;
- // If HAVING condition is false (0), skip ResultRow
- final_having_skip_addr = Some(self.program.len());
- self.emit(OpCode::IfZ, having_reg, 0, 0, None, "");
- }
-
- // Output final result row
- self.emit(
- OpCode::ResultRow,
+ // Output the final group, which no group-change ever flushed.
+ self.emit_group_flush(
+ select,
+ table,
+ &out_plan,
+ &group_col_indices,
+ &agg_info,
+ group_key_reg,
+ acc_base_reg,
result_start_reg,
- (group_col_indices.len() + agg_info.len()) as i64,
- 0,
- None,
- "",
- );
-
- // Patch final HAVING skip jump if present
- if let Some(addr) = final_having_skip_addr {
- let after_result = self.program.len();
- if let Some(inst) = self.program.instructions.get_mut(addr) {
- inst.p2 = after_result as i64;
- }
- }
+ )?;
// Close cursor
self.emit(OpCode::Close, cursor_idx, 0, 0, None, "");
@@ -584,9 +715,11 @@ impl<'a> SqlCompiler<'a> {
/// Compile a HAVING condition and return the register containing the result (1 or 0)
/// agg_result_regs maps aggregate function signatures to their result registers
+ #[allow(clippy::too_many_arguments)]
fn compile_having_condition(
&mut self,
having_expr: &Expr,
+ table: &Table,
group_col_indices: &[usize],
agg_info: &[(i64, Option)],
result_start_reg: i64,
@@ -596,6 +729,7 @@ impl<'a> SqlCompiler<'a> {
Expr::BinaryOp { left, op, right } => {
let left_reg = self.compile_having_operand(
left,
+ table,
group_col_indices,
agg_info,
result_start_reg,
@@ -603,6 +737,7 @@ impl<'a> SqlCompiler<'a> {
)?;
let right_reg = self.compile_having_operand(
right,
+ table,
group_col_indices,
agg_info,
result_start_reg,
@@ -618,9 +753,11 @@ impl<'a> SqlCompiler<'a> {
}
/// Compile an operand in a HAVING condition
+ #[allow(clippy::too_many_arguments)]
fn compile_having_operand(
&mut self,
expr: &Expr,
+ table: &Table,
group_col_indices: &[usize],
agg_info: &[(i64, Option)],
result_start_reg: i64,
@@ -628,14 +765,33 @@ impl<'a> SqlCompiler<'a> {
) -> SqawkResult {
match expr {
Expr::Function(func) => {
- // Find the aggregate function in agg_info and return its result register
+ // Match the aggregate on BOTH its function type and its
+ // argument column.
+ //
+ // This used to compare func_type alone, so with two aggregates
+ // of the same kind -- `SELECT department, SUM(salary),
+ // SUM(age) ... HAVING SUM(age) > 60` -- HAVING bound to
+ // SUM(salary), whichever appeared first, and filtered on
+ // entirely the wrong column. MIN vs MAX bound correctly, which
+ // is why it stayed hidden.
let name = func.name.to_string().to_uppercase();
- let func_type = Self::agg_func_type(&name);
+ let func_type = Self::agg_func_type(&name)
+ | if func_is_distinct(func) {
+ AGG_DISTINCT
+ } else {
+ 0
+ };
+
+ let wanted_col: Option = match func_args(func).first() {
+ Some(FunctionArg::Unnamed(FunctionArgExpr::Wildcard)) => None,
+ Some(FunctionArg::Unnamed(FunctionArgExpr::Expr(arg_expr))) => {
+ Some(self.resolve_column_expr(arg_expr, table)?)
+ }
+ _ => None,
+ };
- // Find matching aggregate in agg_info by function type
- for (i, (agg_func_type, _)) in agg_info.iter().enumerate() {
- if *agg_func_type == func_type {
- // Found matching aggregate - copy its result to a new register
+ for (i, (agg_func_type, agg_col)) in agg_info.iter().enumerate() {
+ if *agg_func_type == func_type && *agg_col == wanted_col {
let agg_result_reg =
result_start_reg + group_col_indices.len() as i64 + i as i64;
let reg = self.allocate_register();
@@ -645,24 +801,46 @@ impl<'a> SqlCompiler<'a> {
}
Err(SqawkError::UnsupportedSqlFeature(format!(
- "Aggregate function in HAVING not found in SELECT: {}",
- name
+ "Aggregate in HAVING must also appear in SELECT: {}",
+ expr
)))
}
- Expr::Value(value) => {
- let reg = self.allocate_register();
- if let sqlparser::ast::Value::Number(n, _) = value {
- if let Ok(i) = n.parse::() {
- self.emit(OpCode::Integer, i, reg, 0, None, "");
- }
- }
- Ok(reg)
+ Expr::Value(_) => {
+ // Compile the literal through the shared expression compiler.
+ //
+ // This arm used to handle Value::Number only, and silently
+ // emitted NOTHING for anything else -- so `HAVING department =
+ // 'Sales'` compared the group key against an uninitialized
+ // register and matched no group.
+ self.code_expr(expr, &NameCtx::single(table, 0), None)
}
- Expr::Identifier(_) => {
- // GROUP BY column reference - copy from first group key register
- // (simplified - for multi-column GROUP BY would need column name matching)
+ Expr::Identifier(ident) => {
+ // A GROUP BY key referenced in HAVING. This used to copy
+ // group_key_reg unconditionally -- the FIRST key column --
+ // whatever name was written, so any reference to a later key
+ // silently filtered on the first one.
+ let col_name = ident.value.to_ascii_lowercase();
+ let col_idx = table
+ .column_index(&col_name)
+ .ok_or_else(|| SqawkError::ColumnNotFound(col_name.clone()))?;
+ let key_pos = group_col_indices
+ .iter()
+ .position(|c| *c == col_idx)
+ .ok_or_else(|| {
+ SqawkError::InvalidSqlQuery(format!(
+ "Column '{}' in HAVING must appear in GROUP BY",
+ ident.value
+ ))
+ })?;
let reg = self.allocate_register();
- self.emit(OpCode::Copy, group_key_reg, reg, 0, None, "");
+ self.emit(
+ OpCode::Copy,
+ group_key_reg + key_pos as i64,
+ reg,
+ 0,
+ None,
+ "",
+ );
Ok(reg)
}
_ => Err(SqawkError::UnsupportedSqlFeature(format!(
diff --git a/src/vm/compiler_ddl.rs b/src/vm/compiler_ddl.rs
index 0c2e590..0679e6e 100644
--- a/src/vm/compiler_ddl.rs
+++ b/src/vm/compiler_ddl.rs
@@ -3,9 +3,10 @@
//! This module extends SqlCompiler with DDL statement compilation methods.
use sqlparser::ast::{
- AlterTableOperation, ObjectName, ObjectType, SelectItem, SetExpr, TableFactor, Value,
+ AlterTableOperation, ObjectName, ObjectType, SelectItem, SetExpr, TableFactor,
};
+use super::ast_compat::sql_option_key_value;
use super::bytecode::OpCode;
use super::compiler::SqlCompiler;
use crate::error::{SqawkError, SqawkResult};
@@ -429,10 +430,10 @@ impl<'a> SqlCompiler<'a> {
// Extract delimiter from WITH options
let delimiter = with_options.iter().find_map(|opt| {
- if opt.name.value.to_lowercase() == "delimiter" {
- if let Value::SingleQuotedString(s) = &opt.value {
- return Some(s.clone());
- }
+ let (key, value) = sql_option_key_value(opt)?;
+ if key.to_lowercase() == "delimiter" {
+ // `value` is rendered SQL, so a string literal arrives quoted.
+ return Some(value.trim_matches('\'').to_string());
}
None
});
diff --git a/src/vm/compiler_dml.rs b/src/vm/compiler_dml.rs
index 508861f..f6adb79 100644
--- a/src/vm/compiler_dml.rs
+++ b/src/vm/compiler_dml.rs
@@ -4,8 +4,9 @@
use sqlparser::ast::{Expr, Select, SelectItem, SetExpr, TableFactor};
+use super::ast_compat::assignment_column;
use super::bytecode::OpCode;
-use super::compiler::SqlCompiler;
+use super::compiler::{NameCtx, SqlCompiler};
use crate::error::{SqawkError, SqawkResult};
impl<'a> SqlCompiler<'a> {
@@ -464,13 +465,7 @@ impl<'a> SqlCompiler<'a> {
let mut assignment_map: Vec<(usize, &Expr)> = Vec::new();
for assignment in assignments {
// Handle different column identifier formats
- let col_name = if !assignment.id.is_empty() {
- assignment.id[0].value.clone()
- } else {
- return Err(SqawkError::InvalidSqlQuery(
- "Invalid assignment target".to_string(),
- ));
- };
+ let col_name = assignment_column(assignment)?;
let col_idx = table_ref
.column_index(&col_name)
@@ -536,27 +531,34 @@ impl<'a> SqlCompiler<'a> {
None
};
- // Apply assignments
+ // Apply assignments.
+ //
+ // The right-hand side is compiled in the row's scope, so a column may
+ // appear there: `SET salary = salary + 1`. compile_expr, which this
+ // used to call, has no Identifier arm at all and rejected it.
+ //
+ // Reading through a Column op against the cursor yields the ORIGINAL
+ // row value, which is the required semantics -- every RHS is evaluated
+ // against the row as it was before any assignment in this statement.
+ let update_ctx = NameCtx::single(table_ref, 0);
for (col_idx, expr) in &assignment_map {
- self.compile_expr_into_register(expr, base_reg + *col_idx)?;
+ self.code_expr(expr, &update_ctx, Some((base_reg + *col_idx) as i64))?;
}
- // Delete old row and insert new row
- self.emit(
- OpCode::DeleteRow,
- 0,
- 0,
- 0,
- Some(table_name.clone()),
- "Delete old row",
- );
+ // Replace the row in place.
+ //
+ // This was DeleteRow followed by InsertRow, and the insert appends, so
+ // every updated row jumped to the end of the table -- persisted to the
+ // user's file under --write. UpdateRow keeps the row where it is, and
+ // it also makes the affected-row count exact instead of inferred from
+ // matching insert and delete tallies.
self.emit(
- OpCode::InsertRow,
+ OpCode::UpdateRow,
0,
base_reg as i64,
column_count as i64,
Some(table_name.clone()),
- "Insert updated row",
+ "Replace row in place",
);
// Patch skip if present
diff --git a/src/vm/compiler_join.rs b/src/vm/compiler_join.rs
index f6753af..04243b4 100644
--- a/src/vm/compiler_join.rs
+++ b/src/vm/compiler_join.rs
@@ -3,8 +3,8 @@
//! This module extends SqlCompiler with JOIN compilation methods.
use sqlparser::ast::{
- BinaryOperator, Expr, FunctionArg, FunctionArgExpr, JoinConstraint, JoinOperator, Query,
- Select, SelectItem, TableWithJoins, Value,
+ Expr, FunctionArg, FunctionArgExpr, JoinConstraint, JoinOperator, Query, Select, SelectItem,
+ TableWithJoins,
};
/// Represents an item in a multi-table aggregate projection
@@ -16,12 +16,132 @@ enum MultiTableProjectionItem {
Aggregate(i64, Option<(usize, usize)>),
}
+use super::ast_compat::{func_args, group_by_exprs};
use super::bytecode::{OpCode, ResultSchema};
-use super::compiler::{MultiTableProjection, SqlCompiler};
+use super::compiler::{MultiTableProjection, NameCtx, SqlCompiler, Src};
use crate::error::{SqawkError, SqawkResult};
use crate::table::Table;
+/// Classify a `JoinOperator` into sqawk's join kind plus its constraint.
+///
+/// This is the single place where `JoinOperator` is destructured, and the match
+/// is deliberately EXHAUSTIVE -- no `_` arm. sqlparser treats any AST change as
+/// breaking, and `JoinOperator` gains variants across versions: newer releases
+/// parse a plain `LEFT JOIN` as `Left` rather than `LeftOuter`, reserving
+/// `LeftOuter` for the explicit `LEFT OUTER JOIN` spelling. Behind a catch-all
+/// that rename turns every `LEFT JOIN` in the suite into a runtime error (or,
+/// worse, silently into an INNER JOIN); with an exhaustive match it is a
+/// compile error that names the site.
+///
+/// `CrossJoin` returns `None` for its constraint -- callers handle it before
+/// asking for an ON expression.
+fn classify_join(op: &JoinOperator) -> SqawkResult<(&'static str, Option<&JoinConstraint>)> {
+ Ok(match op {
+ // A bare `JOIN ... ON` is an INNER JOIN. sqlparser distinguishes the
+ // spelling (`Join` vs `Inner`); SQL does not.
+ JoinOperator::Join(c) | JoinOperator::Inner(c) => ("INNER", Some(c)),
+ // Likewise `LEFT JOIN` and `LEFT OUTER JOIN` are the same operation,
+ // as are the RIGHT pair. sqlparser reserves the `*Outer` variants for
+ // the explicit OUTER spelling.
+ JoinOperator::Left(c) | JoinOperator::LeftOuter(c) => ("LEFT", Some(c)),
+ JoinOperator::Right(c) | JoinOperator::RightOuter(c) => ("RIGHT", Some(c)),
+ JoinOperator::FullOuter(c) => ("FULL", Some(c)),
+ JoinOperator::CrossJoin(_) => ("CROSS", None),
+ JoinOperator::Semi(_) | JoinOperator::LeftSemi(_) | JoinOperator::RightSemi(_) => {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "SEMI JOIN not supported".into(),
+ ))
+ }
+ JoinOperator::Anti(_) | JoinOperator::LeftAnti(_) | JoinOperator::RightAnti(_) => {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "ANTI JOIN not supported".into(),
+ ))
+ }
+ JoinOperator::CrossApply => {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "CROSS APPLY not supported".into(),
+ ))
+ }
+ JoinOperator::OuterApply => {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "OUTER APPLY not supported".into(),
+ ))
+ }
+ JoinOperator::AsOf { .. } => {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "ASOF JOIN not supported".into(),
+ ))
+ }
+ // MySQL optimizer hint: semantically a plain INNER JOIN, but it also
+ // forces a join order. Rejected rather than silently ignoring the hint.
+ JoinOperator::StraightJoin(_) => {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "STRAIGHT_JOIN not supported".into(),
+ ))
+ }
+ JoinOperator::ArrayJoin | JoinOperator::LeftArrayJoin | JoinOperator::InnerArrayJoin => {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "ARRAY JOIN not supported".into(),
+ ))
+ }
+ })
+}
+
+/// The name a query uses to refer to a FROM item: its alias if it has one,
+/// otherwise the table name.
+///
+/// SQL lets a join alias its inputs, and `FROM users u JOIN orders o ON
+/// u.id = o.user_id` must resolve `u` and `o`. Aliases already worked for a
+/// single table and for comma joins; explicit joins resolved against the table
+/// name alone, so any aliased explicit join failed with "Table 'u' not found".
+fn table_refname(factor: &sqlparser::ast::TableFactor, fallback: &str) -> String {
+ match factor {
+ sqlparser::ast::TableFactor::Table {
+ alias: Some(alias), ..
+ } => alias.name.value.to_ascii_lowercase(),
+ _ => fallback.to_ascii_lowercase(),
+ }
+}
+
+/// Resolved join projection: the left table's column indices, the right
+/// table's column indices, the output schema, and where each projected column
+/// lives as (side, index within that side) in SELECT-list order.
+type JoinProjection = (Vec, Vec, ResultSchema, Vec<(usize, usize)>);
+
impl<'a> SqlCompiler<'a> {
+ /// Emit a join result row in SELECT-list order.
+ ///
+ /// Join registers are laid out left-table columns then right-table
+ /// columns, but the schema is built in projection order, so emitting the
+ /// raw block put values under the wrong headers whenever the SELECT list
+ /// interleaved the tables: `SELECT orders.id, users.name` printed the
+ /// header `orders.id,users.name` above the values `John,101`.
+ fn emit_join_result_row(
+ &mut self,
+ out_plan: &[(usize, usize)],
+ left_start_reg: i64,
+ left_len: usize,
+ comment: &str,
+ ) {
+ let out = self.allocate_registers(out_plan.len().max(1));
+ for (i, (side, idx)) in out_plan.iter().enumerate() {
+ let src = if *side == 0 {
+ left_start_reg + *idx as i64
+ } else {
+ left_start_reg + left_len as i64 + *idx as i64
+ };
+ self.emit(OpCode::Copy, src, out + i as i64, 0, None, "");
+ }
+ self.emit(
+ OpCode::ResultRow,
+ out,
+ out_plan.len() as i64,
+ 0,
+ None,
+ comment,
+ );
+ }
+
pub(crate) fn compile_join(
&mut self,
table_with_joins: &TableWithJoins,
@@ -63,22 +183,23 @@ impl<'a> SqlCompiler<'a> {
return Err(SqawkError::TableNotFound(right_table_name));
}
+ // How the query refers to each side: alias if given, else table name.
+ let left_ref = table_refname(&table_with_joins.relation, &left_table_name);
+ let right_ref = table_refname(&join.relation, &right_table_name);
+
// Check for CROSS JOIN - handle separately
- if matches!(&join.join_operator, JoinOperator::CrossJoin) {
+ if matches!(&join.join_operator, JoinOperator::CrossJoin(_)) {
return self.compile_cross_join(&left_table_name, &right_table_name, projection);
}
- // Determine join type and constraint
- let (join_type, join_condition) = match &join.join_operator {
- JoinOperator::LeftOuter(constraint) => ("LEFT", constraint),
- JoinOperator::RightOuter(constraint) => ("RIGHT", constraint),
- JoinOperator::FullOuter(constraint) => ("FULL", constraint),
- JoinOperator::Inner(constraint) => ("INNER", constraint),
- JoinOperator::CrossJoin => unreachable!("Handled above"),
- _ => {
- return Err(SqawkError::UnsupportedSqlFeature(
- "Unsupported join type".into(),
- ))
+ // Determine join type and constraint. See `classify_join`.
+ let (join_type, join_condition) = match classify_join(&join.join_operator)? {
+ ("CROSS", _) => unreachable!("Handled above"),
+ (kind, Some(constraint)) => (kind, constraint),
+ (kind, None) => {
+ return Err(SqawkError::UnsupportedSqlFeature(format!(
+ "{kind} JOIN without a constraint is not supported"
+ )))
}
};
@@ -114,12 +235,12 @@ impl<'a> SqlCompiler<'a> {
let right_table = self.database.get_table(&right_table_name)?;
// Determine which columns to output based on projection
- let (left_cols, right_cols, schema) = self.resolve_join_projection(
+ let (left_cols, right_cols, schema, out_plan) = self.resolve_join_projection(
projection,
left_table,
right_table,
- &left_table_name,
- &right_table_name,
+ &left_ref,
+ &right_ref,
)?;
// Set result schema
@@ -238,12 +359,14 @@ impl<'a> SqlCompiler<'a> {
self.emit_column_loads(inner_cursor as i64, &inner_cols, inner_start_reg, "inner");
// Compile join condition
- let condition_reg = self.compile_join_condition(
+ let condition_reg = self.compile_join_condition_with_refs(
on_expr,
left_table,
right_table,
left_cursor,
right_cursor,
+ &left_ref,
+ &right_ref,
)?;
// If ON condition fails, skip to next inner iteration (placeholder - will be patched)
@@ -260,12 +383,14 @@ impl<'a> SqlCompiler<'a> {
// Compile WHERE clause filtering if present
let skip_where_addr = if let Some(where_expr) = where_clause {
// Use compile_join_condition to handle table.column references in WHERE
- let where_cond_reg = self.compile_join_condition(
+ let where_cond_reg = self.compile_join_condition_with_refs(
where_expr,
left_table,
right_table,
left_cursor,
right_cursor,
+ &left_ref,
+ &right_ref,
)?;
// If WHERE condition fails, skip to next inner iteration
@@ -294,12 +419,10 @@ impl<'a> SqlCompiler<'a> {
);
// Output combined row (matched)
- self.emit(
- OpCode::ResultRow,
+ self.emit_join_result_row(
+ &out_plan,
left_start_reg,
- (left_cols.len() + right_cols.len()) as i64,
- 0,
- None,
+ left_cols.len(),
"Output matched row",
);
@@ -331,14 +454,18 @@ impl<'a> SqlCompiler<'a> {
// LEFT JOIN: outer=left, inner=right, fill right with NULL
// RIGHT JOIN: outer=right, inner=left, fill left with NULL
if join_type == "LEFT" || join_type == "RIGHT" || join_type == "FULL" {
- // Check match flag - if match found (non-zero), skip NULL row output
- // If no match (zero), continue to NullRow
- // Jump calculation: IfPos (current), NullRow (+1), ResultRow (+2), Next outer (+3)
- // So if match found, skip to +3 (Next outer)
- self.emit(
+ // If a match was found, skip the NULL-filled row entirely.
+ //
+ // This was a hand-counted `program.len() + 3` justified by a
+ // comment enumerating the three instructions that followed. The
+ // count silently stopped being 3 the moment the result row needed
+ // more than one instruction to emit, so the jump landed inside the
+ // block it was supposed to skip and produced spurious NULL rows.
+ let skip_null_row = self.label();
+ self.emit_jump_to(
OpCode::IfPos,
match_reg,
- (self.program.len() + 3) as i64,
+ skip_null_row,
0,
None,
"Skip NULL row if match found",
@@ -355,14 +482,13 @@ impl<'a> SqlCompiler<'a> {
);
// Output row with NULLs
- self.emit(
- OpCode::ResultRow,
+ self.emit_join_result_row(
+ &out_plan,
left_start_reg,
- (left_cols.len() + right_cols.len()) as i64,
- 0,
- None,
+ left_cols.len(),
"Output outer row with NULLs for inner",
);
+ self.resolve(skip_null_row);
}
// Next outer
@@ -432,20 +558,23 @@ impl<'a> SqlCompiler<'a> {
self.emit_column_loads(left_cursor as i64, &left_cols, left_start_reg, "left");
// Check join condition
- let condition_reg2 = self.compile_join_condition(
+ let condition_reg2 = self.compile_join_condition_with_refs(
on_expr,
left_table,
right_table,
left_cursor,
right_cursor,
+ &left_ref,
+ &right_ref,
)?;
// If condition matches, mark it and skip to next right row
// IfZ skips over MarkMatch if condition is false
- self.emit(
+ let skip_mark = self.label();
+ self.emit_jump_to(
OpCode::IfZ,
condition_reg2,
- (self.program.len() + 2) as i64,
+ skip_mark,
0,
None,
"Skip if condition false",
@@ -460,6 +589,7 @@ impl<'a> SqlCompiler<'a> {
None,
"Mark that right row has a match",
);
+ self.resolve(skip_mark);
// Next left (continue checking)
self.emit(
@@ -475,10 +605,11 @@ impl<'a> SqlCompiler<'a> {
// After checking all left rows: if no match, output right row with NULLs for left
// If match found (IfPos), skip to next right row
- self.emit(
+ let skip_null_left = self.label();
+ self.emit_jump_to(
OpCode::IfPos,
match_reg,
- (self.program.len() + 3) as i64,
+ skip_null_left,
0,
None,
"Skip if right row had a match",
@@ -495,14 +626,13 @@ impl<'a> SqlCompiler<'a> {
);
// Output row with NULLs for left
- self.emit(
- OpCode::ResultRow,
+ self.emit_join_result_row(
+ &out_plan,
left_start_reg,
- (left_cols.len() + right_cols.len()) as i64,
- 0,
- None,
+ left_cols.len(),
"Output unmatched right row with NULLs",
);
+ self.resolve(skip_null_left);
// Next right row
self.emit(
@@ -598,7 +728,7 @@ impl<'a> SqlCompiler<'a> {
let right_table = self.database.get_table(right_table_name)?;
// Determine which columns to output based on projection
- let (left_cols, right_cols, schema) = self.resolve_join_projection(
+ let (left_cols, right_cols, schema, out_plan) = self.resolve_join_projection(
projection,
left_table,
right_table,
@@ -670,12 +800,10 @@ impl<'a> SqlCompiler<'a> {
self.emit_column_loads(right_cursor as i64, &right_cols, right_start_reg, "right");
// Output combined row
- self.emit(
- OpCode::ResultRow,
+ self.emit_join_result_row(
+ &out_plan,
left_start_reg,
- (left_cols.len() + right_cols.len()) as i64,
- 0,
- None,
+ left_cols.len(),
"Output cross join row",
);
@@ -966,7 +1094,7 @@ impl<'a> SqlCompiler<'a> {
// Determine GROUP BY column references
let mut group_col_refs: Vec<(usize, usize)> = Vec::new(); // (table_idx, col_idx)
- for expr in select.group_by.iter() {
+ for expr in group_by_exprs(&select.group_by).iter() {
let (tbl_idx, col_idx) = self.resolve_multi_table_column(expr, &tables, &table_refs)?;
group_col_refs.push((tbl_idx, col_idx));
}
@@ -991,11 +1119,11 @@ impl<'a> SqlCompiler<'a> {
if matches!(name.as_str(), "COUNT" | "SUM" | "AVG" | "MIN" | "MAX") {
let func_type = Self::agg_func_type(&name);
let col_ref = if let Some(FunctionArg::Unnamed(FunctionArgExpr::Wildcard)) =
- func.args.first()
+ func_args(func).first()
{
None // COUNT(*)
} else if let Some(FunctionArg::Unnamed(FunctionArgExpr::Expr(arg_expr))) =
- func.args.first()
+ func_args(func).first()
{
let (tbl_idx, col_idx) =
self.resolve_multi_table_column(arg_expr, &tables, &table_refs)?;
@@ -1935,144 +2063,37 @@ impl<'a> SqlCompiler<'a> {
}
/// Compile a WHERE condition for multi-table joins
+ /// Compile a WHERE condition spanning several tables.
+ ///
+ /// Delegates to the shared expression compiler with one name-resolution
+ /// source per table. Cursor `i` reads `tables[i]`, which is how the
+ /// implicit-join loop opens them.
+ ///
+ /// This replaces a private condition/operand pair that understood only
+ /// AND, OR and the six comparisons over compound identifiers and literals.
+ /// Anything else -- a function call, arithmetic, IS NULL -- was rejected in
+ /// a multi-table WHERE while working in a single-table one.
fn compile_multi_table_condition(
&mut self,
expr: &Expr,
tables: &[&Table],
table_names: &[String],
) -> SqawkResult {
- match expr {
- Expr::BinaryOp { left, op, right } => {
- match op {
- BinaryOperator::And => {
- let left_reg =
- self.compile_multi_table_condition(left, tables, table_names)?;
- let right_reg =
- self.compile_multi_table_condition(right, tables, table_names)?;
- Ok(self.emit_and(left_reg, right_reg))
- }
- BinaryOperator::Or => {
- let left_reg =
- self.compile_multi_table_condition(left, tables, table_names)?;
- let right_reg =
- self.compile_multi_table_condition(right, tables, table_names)?;
- Ok(self.emit_or(left_reg, right_reg))
- }
- // Comparison operators
- BinaryOperator::Eq
- | BinaryOperator::NotEq
- | BinaryOperator::Lt
- | BinaryOperator::LtEq
- | BinaryOperator::Gt
- | BinaryOperator::GtEq => {
- let left_reg = self.allocate_register();
- let right_reg = self.allocate_register();
- self.compile_multi_table_operand(left, tables, table_names, left_reg)?;
- self.compile_multi_table_operand(right, tables, table_names, right_reg)?;
- self.emit_comparison(op, left_reg, right_reg)
- }
- _ => Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported binary operator in multi-table WHERE: {:?}",
- op
- ))),
- }
- }
- _ => Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported expression type in multi-table WHERE: {:?}",
- expr
- ))),
- }
+ let srcs = tables
+ .iter()
+ .zip(table_names.iter())
+ .enumerate()
+ .map(|(i, (table, name))| Src {
+ cursor: i as i64,
+ reg_base: None,
+ table,
+ refname: name.to_ascii_lowercase(),
+ })
+ .collect();
+ let ctx = NameCtx::multi(srcs);
+ self.code_expr(expr, &ctx, None)
}
- /// Compile an operand for a multi-table condition (column ref or literal)
- fn compile_multi_table_operand(
- &mut self,
- expr: &Expr,
- tables: &[&Table],
- table_names: &[String],
- target_reg: i64,
- ) -> SqawkResult<()> {
- match expr {
- Expr::CompoundIdentifier(parts) => {
- if parts.len() == 2 {
- let tbl_name = &parts[0].value;
- let col_name = &parts[1].value;
-
- // Find which table this column belongs to
- for (i, table_name) in table_names.iter().enumerate() {
- if tbl_name.eq_ignore_ascii_case(table_name) {
- let columns = tables[i].columns();
- if let Some(col_idx) = columns
- .iter()
- .position(|c| c.eq_ignore_ascii_case(col_name))
- {
- self.emit(
- OpCode::Column,
- i as i64,
- col_idx as i64,
- target_reg,
- None,
- &format!("r[{}] = {}.{}", target_reg, table_name, col_name),
- );
- return Ok(());
- } else {
- return Err(SqawkError::ColumnNotFound(col_name.clone()));
- }
- }
- }
- Err(SqawkError::TableNotFound(tbl_name.clone()))
- } else {
- Err(SqawkError::UnsupportedSqlFeature(format!(
- "Compound identifier with {} parts not supported",
- parts.len()
- )))
- }
- }
- Expr::Value(value) => {
- match value {
- Value::Number(num, _) => {
- if let Ok(int_val) = num.parse::() {
- self.emit(
- OpCode::Integer,
- int_val,
- target_reg,
- 0,
- None,
- &format!("r[{}] = {}", target_reg, int_val),
- );
- } else {
- return Err(SqawkError::UnsupportedSqlFeature(format!(
- "Non-integer literal: {}",
- num
- )));
- }
- }
- Value::SingleQuotedString(s) => {
- self.emit(
- OpCode::String,
- 0,
- target_reg,
- 0,
- Some(s.clone()),
- &format!("r[{}] = '{}'", target_reg, s),
- );
- }
- _ => {
- return Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported literal type: {:?}",
- value
- )));
- }
- }
- Ok(())
- }
- _ => Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported operand type in multi-table condition: {:?}",
- expr
- ))),
- }
- }
- /// Compile a multi-table JOIN (3+ tables)
pub(crate) fn compile_multi_join(
&mut self,
table_with_joins: &TableWithJoins,
@@ -2117,9 +2138,12 @@ impl<'a> SqlCompiler<'a> {
}
};
- // Only support INNER JOIN for multi-table joins
- match &join.join_operator {
- JoinOperator::Inner(JoinConstraint::On(expr)) => {
+ // Only support INNER JOIN for multi-table joins. Routed through
+ // `classify_join` so a new/renamed JoinOperator variant is a
+ // compile error here too, rather than silently falling into the
+ // "Only INNER JOIN" error arm.
+ match classify_join(&join.join_operator)? {
+ ("INNER", Some(JoinConstraint::On(expr))) => {
join_conditions.push(expr);
}
_ => {
@@ -2232,7 +2256,7 @@ impl<'a> SqlCompiler<'a> {
for (i, cond) in join_conditions.iter().enumerate() {
// Use table_refs (aliases) for column resolution
let cond_reg =
- self.compile_multi_join_condition(cond, &table_refs, &table_start_regs)?;
+ self.compile_multi_join_condition(cond, &tables, &table_refs, &table_start_regs)?;
// If condition fails, skip to next innermost iteration (placeholder)
let skip_addr = self.program.len();
@@ -2367,112 +2391,37 @@ impl<'a> SqlCompiler<'a> {
Ok(())
}
- /// Compile a join condition for multi-table joins
+ /// Compile a join condition for multi-table joins.
+ ///
+ /// Here every table's columns have already been loaded into registers, so
+ /// a column reference is a register offset rather than a cursor read. That
+ /// is what `Src::reg_base` models, which lets this share the expression
+ /// compiler with the cursor-based paths instead of carrying its own
+ /// condition/operand pair -- one that understood a single top-level
+ /// comparison and nothing else.
fn compile_multi_join_condition(
&mut self,
expr: &Expr,
- table_names: &[String],
- table_start_regs: &[i64],
- ) -> SqawkResult {
- match expr {
- Expr::BinaryOp { left, op, right } => {
- let left_reg =
- self.compile_multi_join_operand(left, table_names, table_start_regs)?;
- let right_reg =
- self.compile_multi_join_operand(right, table_names, table_start_regs)?;
- self.emit_comparison(op, left_reg, right_reg)
- }
- _ => Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported multi-join condition: {:?}",
- expr
- ))),
- }
- }
-
- /// Compile an operand for multi-table join conditions
- fn compile_multi_join_operand(
- &mut self,
- expr: &Expr,
- table_names: &[String],
+ tables: &[&Table],
+ table_refs: &[String],
table_start_regs: &[i64],
) -> SqawkResult {
- match expr {
- Expr::CompoundIdentifier(parts) => {
- if parts.len() == 2 {
- let table_alias = parts[0].value.to_lowercase();
- let col_name = parts[1].value.to_lowercase();
-
- // Find which table this refers to
- for (i, name) in table_names.iter().enumerate() {
- if name.to_lowercase() == table_alias {
- let table = self.database.get_table(name)?;
- if let Some(col_idx) = table.column_index(&col_name) {
- // The column value should already be loaded in register
- let reg = table_start_regs[i] + col_idx as i64;
- return Ok(reg);
- }
- }
- }
- }
- Err(SqawkError::UnsupportedSqlFeature(format!(
- "Could not resolve column reference: {:?}",
- expr
- )))
- }
- Expr::Identifier(ident) => {
- // Unqualified column name - search all tables
- let col_name = ident.value.to_lowercase();
- for (i, name) in table_names.iter().enumerate() {
- let table = self.database.get_table(name)?;
- if let Some(col_idx) = table.column_index(&col_name) {
- let reg = table_start_regs[i] + col_idx as i64;
- return Ok(reg);
- }
- }
- Err(SqawkError::UnsupportedSqlFeature(format!(
- "Could not find column: {}",
- col_name
- )))
- }
- Expr::Value(val) => {
- let reg = self.allocate_register();
- match val {
- sqlparser::ast::Value::Number(n, _) => {
- if let Ok(i) = n.parse::() {
- self.emit(OpCode::Integer, i, reg, 0, None, "Load constant");
- } else {
- return Err(SqawkError::UnsupportedSqlFeature(
- "Only integer constants supported".into(),
- ));
- }
- }
- sqlparser::ast::Value::SingleQuotedString(s) => {
- self.emit(
- OpCode::String,
- 0,
- reg,
- 0,
- Some(s.clone()),
- "Load string constant",
- );
- }
- _ => {
- return Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported value type: {:?}",
- val
- )))
- }
- }
- Ok(reg)
- }
- _ => Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported operand in multi-join: {:?}",
- expr
- ))),
- }
+ let srcs = tables
+ .iter()
+ .zip(table_refs.iter())
+ .zip(table_start_regs.iter())
+ .enumerate()
+ .map(|(i, ((table, name), base))| Src {
+ cursor: i as i64,
+ reg_base: Some(*base),
+ table,
+ refname: name.to_ascii_lowercase(),
+ })
+ .collect();
+ let ctx = NameCtx::multi(srcs);
+ self.code_expr(expr, &ctx, None)
}
- /// Resolve projection for join - returns (left_col_indices, right_col_indices, schema)
fn resolve_join_projection(
&self,
projection: &[SelectItem],
@@ -2480,7 +2429,7 @@ impl<'a> SqlCompiler<'a> {
right_table: &Table,
left_table_name: &str,
right_table_name: &str,
- ) -> SqawkResult<(Vec, Vec, ResultSchema)> {
+ ) -> SqawkResult {
let tables: Vec<&Table> = vec![left_table, right_table];
let table_names: Vec<&str> = vec![left_table_name, right_table_name];
@@ -2488,7 +2437,15 @@ impl<'a> SqlCompiler<'a> {
for item in projection {
if matches!(item, SelectItem::Wildcard(_)) {
let (table_cols, schema) = Self::build_wildcard_schema(&tables, &table_names);
- return Ok((table_cols[0].clone(), table_cols[1].clone(), schema));
+ // Wildcard output IS left-then-right, so the plan is the
+ // identity over that layout.
+ let plan = table_cols[0]
+ .iter()
+ .enumerate()
+ .map(|(i, _)| (0usize, i))
+ .chain(table_cols[1].iter().enumerate().map(|(i, _)| (1usize, i)))
+ .collect();
+ return Ok((table_cols[0].clone(), table_cols[1].clone(), schema, plan));
}
}
@@ -2496,6 +2453,11 @@ impl<'a> SqlCompiler<'a> {
let mut left_cols = Vec::new();
let mut right_cols = Vec::new();
let mut schema = ResultSchema::new();
+ // Where each projected column lives: (side, index within that side).
+ // Registers are laid out left-then-right, but the SELECT list may name
+ // them in any order, so the emission order has to be recorded rather
+ // than assumed.
+ let mut out_plan: Vec<(usize, usize)> = Vec::with_capacity(projection.len());
for item in projection {
let (tbl_name, col_name, output_name) = match item {
@@ -2519,8 +2481,10 @@ impl<'a> SqlCompiler<'a> {
let (tbl_idx, col_idx) =
Self::find_column_in_tables(&tbl_name, &col_name, &tables, &table_names)?;
if tbl_idx == 0 {
+ out_plan.push((0, left_cols.len()));
left_cols.push(col_idx);
} else {
+ out_plan.push((1, right_cols.len()));
right_cols.push(col_idx);
}
schema.add_column(
@@ -2529,7 +2493,7 @@ impl<'a> SqlCompiler<'a> {
);
}
- Ok((left_cols, right_cols, schema))
+ Ok((left_cols, right_cols, schema, out_plan))
}
/// Extract table and column name from a compound identifier (e.g., users.name)
@@ -2560,29 +2524,16 @@ impl<'a> SqlCompiler<'a> {
}
}
- /// Compile a join condition expression
- fn compile_join_condition(
- &mut self,
- expr: &Expr,
- left_table: &Table,
- right_table: &Table,
- left_cursor: usize,
- right_cursor: usize,
- ) -> SqawkResult {
- // Use table names as references (no aliases)
- let left_ref = left_table.name().to_string();
- let right_ref = right_table.name().to_string();
- self.compile_join_condition_with_refs(
- expr,
- left_table,
- right_table,
- left_cursor,
- right_cursor,
- &left_ref,
- &right_ref,
- )
- }
-
+ /// Compile a two-table join condition through the shared expression
+ /// compiler.
+ ///
+ /// Qualifiers resolve against `left_ref`/`right_ref` -- the alias each side
+ /// was given, or its table name if it had none.
+ ///
+ /// This replaces a private family of three functions: a condition compiler
+ /// that understood only a single top-level BinaryOp, and an operand
+ /// compiler beneath it. So `ON a.x = b.y AND a.z = 1` was rejected, as was
+ /// any function call in an ON clause, even though both work in a WHERE.
#[allow(clippy::too_many_arguments)]
fn compile_join_condition_with_refs(
&mut self,
@@ -2594,163 +2545,21 @@ impl<'a> SqlCompiler<'a> {
left_ref: &str,
right_ref: &str,
) -> SqawkResult {
- // Handle a.col = b.col style conditions
- match expr {
- Expr::BinaryOp { left, op, right } => {
- let left_reg = self.compile_join_operand_with_refs(
- left,
- left_table,
- right_table,
- left_cursor,
- right_cursor,
- left_ref,
- right_ref,
- )?;
- let right_reg = self.compile_join_operand_with_refs(
- right,
- left_table,
- right_table,
- left_cursor,
- right_cursor,
- left_ref,
- right_ref,
- )?;
- self.emit_comparison(op, left_reg, right_reg)
- }
- _ => Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported join condition expression: {:?}",
- expr
- ))),
- }
- }
-
- /// Compile a join operand with explicit table references (aliases or names)
- #[allow(clippy::too_many_arguments)]
- fn compile_join_operand_with_refs(
- &mut self,
- expr: &Expr,
- left_table: &Table,
- right_table: &Table,
- left_cursor: usize,
- right_cursor: usize,
- left_ref: &str,
- right_ref: &str,
- ) -> SqawkResult {
- match expr {
- Expr::CompoundIdentifier(parts) => {
- // Handle table.column syntax
- if parts.len() == 2 {
- let table_alias = parts[0].value.to_lowercase();
- let col_name = parts[1].value.to_lowercase();
-
- // Try to find column in left or right table
- // Check against table references (which may be aliases)
- let left_ref_lower = left_ref.to_lowercase();
- let right_ref_lower = right_ref.to_lowercase();
-
- if table_alias == left_ref_lower {
- if let Some(col_idx) = left_table.column_index(&col_name) {
- let reg = self.allocate_register();
- self.emit(
- OpCode::Column,
- left_cursor as i64,
- col_idx as i64,
- reg,
- None,
- &format!("Load {}.{} for join", table_alias, col_name),
- );
- return Ok(reg);
- }
- } else if table_alias == right_ref_lower {
- if let Some(col_idx) = right_table.column_index(&col_name) {
- let reg = self.allocate_register();
- self.emit(
- OpCode::Column,
- right_cursor as i64,
- col_idx as i64,
- reg,
- None,
- &format!("Load {}.{} for join", table_alias, col_name),
- );
- return Ok(reg);
- }
- }
-
- Err(SqawkError::UnsupportedSqlFeature(format!(
- "Column {}.{} not found in join tables",
- table_alias, col_name
- )))
- } else {
- Err(SqawkError::UnsupportedSqlFeature(
- "Expected table.column in join condition".into(),
- ))
- }
- }
- Expr::Identifier(ident) => {
- // Unqualified column name - try to find in either table
- let col_name = ident.value.to_lowercase();
-
- // Try left table first
- if let Some(col_idx) = left_table.column_index(&col_name) {
- let reg = self.allocate_register();
- self.emit(
- OpCode::Column,
- left_cursor as i64,
- col_idx as i64,
- reg,
- None,
- &format!("Load {} from left for join", col_name),
- );
- return Ok(reg);
- }
-
- // Try right table
- if let Some(col_idx) = right_table.column_index(&col_name) {
- let reg = self.allocate_register();
- self.emit(
- OpCode::Column,
- right_cursor as i64,
- col_idx as i64,
- reg,
- None,
- &format!("Load {} from right for join", col_name),
- );
- return Ok(reg);
- }
-
- Err(SqawkError::UnsupportedSqlFeature(format!(
- "Column {} not found in either join table",
- col_name
- )))
- }
- Expr::Value(val) => {
- // Literal value in join condition
- let reg = self.allocate_register();
- match val {
- Value::Number(n, _) => {
- if let Ok(i) = n.parse::() {
- self.emit(OpCode::Integer, i, reg, 0, None, "Load literal for join");
- }
- }
- Value::SingleQuotedString(s) => {
- self.emit(
- OpCode::String,
- 0,
- reg,
- 0,
- Some(s.clone()),
- "Load string literal for join",
- );
- }
- _ => {}
- }
- Ok(reg)
- }
- _ => Err(SqawkError::UnsupportedSqlFeature(format!(
- "Unsupported join operand: {:?}",
- expr
- ))),
- }
+ let ctx = NameCtx::multi(vec![
+ Src {
+ cursor: left_cursor as i64,
+ reg_base: None,
+ table: left_table,
+ refname: left_ref.to_ascii_lowercase(),
+ },
+ Src {
+ cursor: right_cursor as i64,
+ reg_base: None,
+ table: right_table,
+ refname: right_ref.to_ascii_lowercase(),
+ },
+ ]);
+ self.code_expr(expr, &ctx, None)
}
/// Resolve a column expression to (table_idx, col_idx) for multi-table queries
diff --git a/src/vm/compiler_window.rs b/src/vm/compiler_window.rs
index dbcf100..b4335a0 100644
--- a/src/vm/compiler_window.rs
+++ b/src/vm/compiler_window.rs
@@ -2,8 +2,9 @@
//!
//! This module extends SqlCompiler with window function compilation methods.
-use sqlparser::ast::{Expr, Select, SelectItem, Value, WindowType};
+use sqlparser::ast::{Expr, Select, SelectItem, Value, ValueWithSpan, WindowType};
+use super::ast_compat::{func_args, order_by_is_asc};
use super::bytecode::{OpCode, ResultSchema};
use super::compiler::SqlCompiler;
use crate::error::{SqawkError, SqawkResult};
@@ -280,9 +281,9 @@ impl<'a> SqlCompiler<'a> {
{
if func.name.to_string().to_uppercase() == *func_name && func.over.is_some()
{
- if !func.args.is_empty() {
+ if !func_args(func).is_empty() {
if let Ok(Expr::Identifier(ident)) =
- self.extract_function_arg_expr(&func.args[0])
+ self.extract_function_arg_expr(&func_args(func)[0])
{
if let Some(src_idx) = table.column_index(&ident.value) {
agg_col_pos =
@@ -355,9 +356,9 @@ impl<'a> SqlCompiler<'a> {
let p4 = if func_name == "LAG" || func_name == "LEAD" {
// LAG/LEAD(column, offset, default)
// Get column position in ephemeral cursor
- let col_pos = if !func.args.is_empty() {
+ let col_pos = if !func_args(func).is_empty() {
if let Ok(Expr::Identifier(ident)) =
- self.extract_function_arg_expr(&func.args[0])
+ self.extract_function_arg_expr(&func_args(func)[0])
{
if let Some(src_idx) = table.column_index(&ident.value) {
base_columns
@@ -375,9 +376,11 @@ impl<'a> SqlCompiler<'a> {
};
// Get offset (default 1)
- let offset = if func.args.len() >= 2 {
- if let Ok(Expr::Value(Value::Number(n, _))) =
- self.extract_function_arg_expr(&func.args[1])
+ let offset = if func_args(func).len() >= 2 {
+ if let Ok(Expr::Value(ValueWithSpan {
+ value: Value::Number(n, _),
+ ..
+ })) = self.extract_function_arg_expr(&func_args(func)[1])
{
n.parse::().unwrap_or(1)
} else {
@@ -489,6 +492,49 @@ impl<'a> SqlCompiler<'a> {
inst.p2 = after_output as i64;
}
+ // An aggregate window with PARTITION BY and no ORDER BY has the whole
+ // partition as its frame, so every row must see the partition total.
+ // The streaming pass produces a RUNNING total, whose value at the last
+ // row of each partition IS the total -- back-fill from there.
+ //
+ // With an ORDER BY the running total is the correct answer, so this is
+ // emitted only when the window is unordered.
+ if order_cols.is_empty() && !partition_cols.is_empty() {
+ let agg_output_positions: Vec = select
+ .projection
+ .iter()
+ .enumerate()
+ .filter(|(_, item)| {
+ matches!(
+ item,
+ SelectItem::UnnamedExpr(Expr::Function(f))
+ | SelectItem::ExprWithAlias {
+ expr: Expr::Function(f),
+ ..
+ } if f.over.is_some()
+ && matches!(
+ f.name.to_string().to_uppercase().as_str(),
+ "SUM" | "COUNT" | "AVG" | "MIN" | "MAX"
+ )
+ )
+ })
+ .map(|(i, _)| i)
+ .collect();
+
+ // One per aggregate window column: finalizing only the first left
+ // any others showing a running value.
+ for pos in agg_output_positions {
+ self.emit(
+ OpCode::WindowFinalize,
+ pos as i64,
+ 0,
+ 0,
+ None,
+ "Give every row its partition's final window value",
+ );
+ }
+ }
+
// Halt
self.emit(OpCode::Halt, 0, 0, 0, None, "");
@@ -514,9 +560,9 @@ impl<'a> SqlCompiler<'a> {
window_funcs.push((name, func.over.clone()));
// If it's an aggregate window function, add argument column
- if !func.args.is_empty() {
+ if !func_args(func).is_empty() {
if let Ok(Expr::Identifier(ident)) =
- self.extract_function_arg_expr(&func.args[0])
+ self.extract_function_arg_expr(&func_args(func)[0])
{
if let Some(col_idx) = table.column_index(&ident.value) {
if !base_columns.contains(&col_idx) {
@@ -562,6 +608,17 @@ impl<'a> SqlCompiler<'a> {
let mut order_asc = Vec::new();
if let Some(WindowType::WindowSpec(spec)) = window_type {
+ // An explicit frame is parsed but not implemented. Accepting it
+ // silently returns the default frame's answer -- `ROWS BETWEEN 1
+ // PRECEDING AND CURRENT ROW` produced a running total over the
+ // whole partition rather than a two-row sliding window -- so it is
+ // rejected instead. A clear error beats a plausible wrong number.
+ if spec.window_frame.is_some() {
+ return Err(SqawkError::UnsupportedSqlFeature(
+ "Explicit window frames (ROWS/RANGE BETWEEN) are not supported".into(),
+ ));
+ }
+
// Extract PARTITION BY columns
for expr in &spec.partition_by {
if let Expr::Identifier(ident) = expr {
@@ -576,7 +633,7 @@ impl<'a> SqlCompiler<'a> {
if let Expr::Identifier(ident) = &order_expr.expr {
if let Some(col_idx) = table.column_index(&ident.value) {
order_cols.push(col_idx);
- order_asc.push(order_expr.asc.unwrap_or(true));
+ order_asc.push(order_by_is_asc(order_expr));
}
}
}
diff --git a/src/vm/engine.rs b/src/vm/engine.rs
index 18d2ac4..ea47271 100644
--- a/src/vm/engine.rs
+++ b/src/vm/engine.rs
@@ -5,7 +5,8 @@
//! and the program counter during execution.
use super::bytecode::{
- Instruction, OpCode, Program, Register, AGG_AVG, AGG_COUNT, AGG_MAX, AGG_MIN, AGG_SUM,
+ Instruction, OpCode, Program, Register, AGG_AVG, AGG_COUNT, AGG_DISTINCT, AGG_MAX, AGG_MIN,
+ AGG_SUM, AGG_TYPE_MASK,
};
use crate::capacity::{
DEFAULT_ACCUMULATOR_CAPACITY, DEFAULT_CURSOR_CAPACITY, DEFAULT_MODIFICATIONS_CAPACITY,
@@ -330,6 +331,9 @@ struct AggAccumulator {
count: i64,
/// Accumulated value (sum for SUM/AVG, min/max for MIN/MAX)
value: Option,
+ /// Values already accumulated, for DISTINCT aggregates. `None` when the
+ /// aggregate is not DISTINCT, so the ordinary path allocates nothing.
+ seen: Option>,
}
impl Sorter {
@@ -412,6 +416,16 @@ pub enum TableModification {
table_name: String,
row_index: usize,
},
+ /// Replace a row in place, preserving its position.
+ ///
+ /// UPDATE used to be expressed as Delete plus Insert, and the insert
+ /// appends, so every updated row jumped to the end of the table -- and
+ /// with --write that reordering was written to the user's file.
+ Replace {
+ table_name: String,
+ row_index: usize,
+ row: Vec,
+ },
CreateTable {
table_name: String,
columns: Vec,
@@ -452,6 +466,19 @@ pub struct VmEngine<'a> {
#[allow(dead_code)]
cursor_source_order: Vec,
+ /// Index into `results` at which each window partition begins.
+ ///
+ /// Recorded during the streaming window pass so a following WindowFinalize
+ /// can identify each partition's extent without needing the partition key
+ /// to be a projected column.
+ window_partition_starts: Vec,
+
+ /// Result of the most recent Compare, consumed by Jump.
+ ///
+ /// Mirrors SQLite, where OP_Compare leaves its verdict for a following
+ /// OP_Jump rather than materializing it into a register.
+ compare_flag: std::cmp::Ordering,
+
/// Current transaction state (None, Active, Committed, or RolledBack)
/// Tracks the lifecycle of a transaction and enforces proper operation sequencing
transaction_state: TransactionState,
@@ -474,6 +501,8 @@ impl<'a> VmEngine<'a> {
database,
program: Program::new(),
pc: 0,
+ compare_flag: std::cmp::Ordering::Equal,
+ window_partition_starts: Vec::new(),
registers: Vec::with_capacity(DEFAULT_REGISTER_CAPACITY),
cursors: HashMap::with_capacity(DEFAULT_CURSOR_CAPACITY),
sorters: HashMap::with_capacity(DEFAULT_SORTER_CAPACITY),
@@ -507,20 +536,31 @@ impl<'a> VmEngine<'a> {
self.cursor_source_order.clear();
self.pending_modifications.clear();
self.once_flags.clear();
-
- // Allocate enough registers for the program
- let max_reg = self
- .program
- .instructions
- .iter()
- .map(|i| i.p1.max(i.p2).max(i.p3))
- .max()
- .unwrap_or(10);
-
- // Allocate a few extra registers just in case
- // Cap max_reg to a reasonable value to avoid overflow (p1/p2/p3 sometimes contain values, not register numbers)
- let max_reg = max_reg.min(10000);
- self.registers = vec![Register::Null; (max_reg + 5) as usize];
+ self.window_partition_starts.clear();
+
+ // Size the register file from the compiler's allocation count.
+ //
+ // This used to be inferred as max(p1, p2, p3) over the instruction
+ // stream, capped at 10000. That was wrong in both directions: those
+ // operands hold literals as often as register numbers, so `Integer
+ // 999999 -> r` inflated the file to the cap, while a program that
+ // genuinely needed more than 10000 registers got a short one -- and
+ // Column errors on a short file rather than growing it.
+ //
+ // Hand-assembled test programs carry no count, so fall back to the old
+ // scan when register_count is 0 but instructions exist.
+ let reg_count = if self.program.register_count > 0 {
+ self.program.register_count
+ } else {
+ self.program
+ .instructions
+ .iter()
+ .map(|i| i.p1.max(i.p2).max(i.p3))
+ .max()
+ .unwrap_or(10)
+ .min(10000)
+ };
+ self.registers = vec![Register::Null; (reg_count + 5) as usize];
// Reset transaction state
self.transaction_state = TransactionState::None;
@@ -779,6 +819,43 @@ impl<'a> VmEngine<'a> {
Ok(ExecuteResult::Continue)
}
+ OpCode::UpdateRow => {
+ // Replace the current row at cursor P1 with registers
+ // P2..P2+P3, keeping its position in the table.
+ let cursor_idx = inst.p1 as usize;
+ let base_reg = inst.p2 as usize;
+ let col_count = inst.p3 as usize;
+
+ let (table_name, row_idx) = if let Some(cursor) = self.cursors.get(&cursor_idx) {
+ (cursor.table_name().to_string(), cursor.current_row_index())
+ } else {
+ return Err(SqawkError::VmError(format!(
+ "Invalid cursor for UpdateRow: {}",
+ cursor_idx
+ )));
+ };
+
+ let table_name = inst
+ .p4
+ .as_deref()
+ .map(|s| s.to_string())
+ .unwrap_or(table_name);
+
+ if let Some(idx) = row_idx {
+ let mut row = Vec::with_capacity(col_count);
+ for i in 0..col_count {
+ row.push(Value::from(self.get_register(base_reg + i)?));
+ }
+ self.pending_modifications.push(TableModification::Replace {
+ table_name,
+ row_index: idx,
+ row,
+ });
+ }
+
+ Ok(ExecuteResult::Continue)
+ }
+
OpCode::DeleteRow => {
// Delete current row at cursor P1
let cursor_idx = inst.p1 as usize;
@@ -1486,6 +1563,11 @@ impl<'a> VmEngine<'a> {
}
};
+ if partition_changed {
+ // The row about to be emitted starts a new partition.
+ self.window_partition_starts.push(self.results.len());
+ }
+
// Check if order value changed (for RANK)
let prev_ord_val = self.get_register(prev_ord_reg)?;
let order_changed = if is_first_row {
@@ -1662,6 +1744,34 @@ impl<'a> VmEngine<'a> {
Ok(ExecuteResult::Continue)
}
+ OpCode::WindowFinalize => {
+ // Give every row of a partition that partition's FINAL window
+ // value, which for an unordered frame is the partition total.
+ //
+ // The streaming pass accumulates as it goes, so without this a
+ // window aggregate with PARTITION BY and no ORDER BY returns a
+ // RUNNING total. SQL defines the frame in that case as the
+ // whole partition, so every row should see the same value.
+ let col = inst.p1 as usize;
+ let starts = self.window_partition_starts.clone();
+ for (i, &start) in starts.iter().enumerate() {
+ let end = starts.get(i + 1).copied().unwrap_or(self.results.len());
+ if end == 0 || start >= end {
+ continue;
+ }
+ let final_value = match self.results[end - 1].get(col) {
+ Some(v) => v.clone(),
+ None => continue,
+ };
+ for row in &mut self.results[start..end] {
+ if col < row.len() {
+ row[col] = final_value.clone();
+ }
+ }
+ }
+ Ok(ExecuteResult::Continue)
+ }
+
OpCode::WindowValue => {
// Get current window function value
// P1 = state base reg, P2 = dest reg, P3 = func type
@@ -2292,6 +2402,59 @@ impl<'a> VmEngine<'a> {
Ok(ExecuteResult::Continue)
}
+ OpCode::Not => {
+ // Logical negation with NULL propagation: NOT UNKNOWN is
+ // UNKNOWN, not true.
+ let src = self.get_register(inst.p1 as usize)?;
+ let out = match src {
+ Register::Null => Register::Null,
+ other => Register::Integer(if Self::register_is_truthy(&other) {
+ 0
+ } else {
+ 1
+ }),
+ };
+ self.set_register(inst.p2 as usize, out)?;
+ Ok(ExecuteResult::Continue)
+ }
+
+ OpCode::Compare => {
+ // Pairwise-compare two register vectors of length P3, starting
+ // at P1 and P2, and stash the Ordering for the next Jump.
+ //
+ // Uses internal ordering (compare_values), where NULL == NULL.
+ // That is deliberate: this drives grouping and sorting, where
+ // NULL keys must group together, whereas SQL `=` must yield
+ // UNKNOWN for NULL. Keeping them separate is what lets the
+ // three-valued-logic change apply to Eq without breaking
+ // GROUP BY on a nullable key.
+ let lhs_start = inst.p1 as usize;
+ let rhs_start = inst.p2 as usize;
+ let count = inst.p3 as usize;
+
+ let mut ordering = std::cmp::Ordering::Equal;
+ for i in 0..count {
+ let a: Value = self.get_register(lhs_start + i)?.into();
+ let b: Value = self.get_register(rhs_start + i)?.into();
+ ordering = compare_values(&a, &b);
+ if ordering != std::cmp::Ordering::Equal {
+ break;
+ }
+ }
+ self.compare_flag = ordering;
+ Ok(ExecuteResult::Continue)
+ }
+
+ OpCode::Jump => {
+ // Three-way branch on the last Compare.
+ let target = match self.compare_flag {
+ std::cmp::Ordering::Less => inst.p1,
+ std::cmp::Ordering::Equal => inst.p2,
+ std::cmp::Ordering::Greater => inst.p3,
+ };
+ Ok(ExecuteResult::Jump(target as usize))
+ }
+
// JumpIfTrue and JumpIfFalse have been replaced by SQLite-style opcodes:
// - IfZ (jump if zero/false)
// - IfPos (jump if positive)
@@ -2301,6 +2464,56 @@ impl<'a> VmEngine<'a> {
Ok(ExecuteResult::Continue)
}
+ OpCode::SortResults => {
+ // Sort the accumulated result rows.
+ //
+ // A post-processing opcode in the same family as Distinct and
+ // Limit: it runs over self.results after the producing loop has
+ // finished, so it is independent of how those rows were
+ // produced. That is what lets ORDER BY work over GROUP BY
+ // output and over joins, neither of which can feed the
+ // cursor-based sorter used by a plain table scan.
+ let spec = inst.p4.as_deref().unwrap_or("");
+ let keys: Vec<(usize, bool)> = spec
+ .split(',')
+ .filter(|s| !s.is_empty())
+ .filter_map(|part| {
+ let (idx, dir) = part.split_once(':')?;
+ Some((
+ idx.trim().parse().ok()?,
+ !dir.trim().eq_ignore_ascii_case("desc"),
+ ))
+ })
+ .collect();
+
+ if !keys.is_empty() {
+ // Sort results and their index tracking together so the
+ // two do not drift apart.
+ let track = self.result_row_indices.len() == self.results.len();
+ let mut order: Vec = (0..self.results.len()).collect();
+ order.sort_by(|&a, &b| {
+ for &(col, asc) in &keys {
+ let va = self.results[a].get(col).unwrap_or(&Value::Null);
+ let vb = self.results[b].get(col).unwrap_or(&Value::Null);
+ let ord = compare_values(va, vb);
+ let ord = if asc { ord } else { ord.reverse() };
+ if ord != std::cmp::Ordering::Equal {
+ return ord;
+ }
+ }
+ std::cmp::Ordering::Equal
+ });
+ self.results = order.iter().map(|&i| self.results[i].clone()).collect();
+ if track {
+ self.result_row_indices = order
+ .iter()
+ .map(|&i| self.result_row_indices[i].clone())
+ .collect();
+ }
+ }
+ Ok(ExecuteResult::Continue)
+ }
+
OpCode::Distinct => {
// Remove duplicate rows from results
// This is used by UNION (without ALL) to deduplicate combined results
@@ -2335,18 +2548,34 @@ impl<'a> VmEngine<'a> {
let limit = inst.p1 as usize;
let offset = inst.p2 as usize;
+ // Index tracking is trimmed alongside the rows. It used to be
+ // left untouched, so after a post-processing LIMIT the two
+ // vectors disagreed and the debug validator reported a
+ // mismatch on every limited query.
+ let track = self.result_row_indices.len() == self.results.len();
+
// First apply OFFSET - skip first N rows
if offset > 0 && !self.results.is_empty() {
if offset >= self.results.len() {
self.results.clear();
+ if track {
+ self.result_row_indices.clear();
+ }
} else {
self.results = self.results.drain(offset..).collect();
+ if track {
+ self.result_row_indices =
+ self.result_row_indices.drain(offset..).collect();
+ }
}
}
// Then apply LIMIT - truncate to N rows
if self.results.len() > limit {
self.results.truncate(limit);
+ if track && self.result_row_indices.len() > limit {
+ self.result_row_indices.truncate(limit);
+ }
}
if self.verbose {
@@ -2860,7 +3089,8 @@ impl<'a> VmEngine<'a> {
// P1 = function type (0=COUNT, 1=SUM, 2=AVG, 3=MIN, 4=MAX)
// P2 = value register
// P3 = accumulator register (used as key for accumulator storage)
- let func_type = inst.p1;
+ let func_type = inst.p1 & AGG_TYPE_MASK;
+ let is_distinct = inst.p1 & AGG_DISTINCT != 0;
let value_reg = inst.p2 as usize;
let acc_reg = inst.p3 as usize;
@@ -2882,6 +3112,11 @@ impl<'a> VmEngine<'a> {
func_type,
count: 0,
value: None,
+ seen: if is_distinct {
+ Some(std::collections::HashSet::new())
+ } else {
+ None
+ },
},
);
// Mark the register as non-Null to indicate accumulator is active
@@ -2893,8 +3128,27 @@ impl<'a> VmEngine<'a> {
func_type,
count: 0,
value: None,
+ seen: if is_distinct {
+ Some(std::collections::HashSet::new())
+ } else {
+ None
+ },
});
+ // DISTINCT: ignore a value already accumulated for this group.
+ // NULLs are not tracked -- every aggregate already skips them.
+ if is_distinct {
+ if let Some(val) = value.as_ref() {
+ if !matches!(val, Value::Null) {
+ let key = format!("{:?}", val);
+ let seen = acc.seen.get_or_insert_with(Default::default);
+ if !seen.insert(key) {
+ return Ok(ExecuteResult::Continue);
+ }
+ }
+ }
+ }
+
match func_type {
AGG_COUNT => {
// COUNT - count non-NULL values (or all rows if value is None)
@@ -3410,6 +3664,20 @@ impl<'a> VmEngine<'a> {
///
/// # Returns
/// The register value or an error if the index is out of bounds
+ /// Whether a register counts as true.
+ ///
+ /// Matches IfZ's notion of "zero" exactly, so `Not` and the conditional
+ /// jumps cannot disagree about what a value means.
+ fn register_is_truthy(reg: &Register) -> bool {
+ match reg {
+ Register::Integer(v) => *v != 0,
+ Register::Float(v) => *v != 0.0,
+ Register::String(v) => !v.is_empty(),
+ Register::Boolean(v) => *v,
+ Register::Null => false,
+ }
+ }
+
fn get_register(&self, idx: usize) -> SqawkResult {
if idx < self.registers.len() {
Ok(self.registers[idx].clone())
@@ -3450,76 +3718,86 @@ impl<'a> VmEngine<'a> {
/// - p3: destination register for result (1 for true, 0 for false)
/// - opcode: the comparison type (Lt, Le, Gt, Ge, Eq, Ne)
fn execute_comparison(&mut self, inst: &Instruction) -> SqawkResult {
+ use std::cmp::Ordering;
+
let reg1 = self.get_register(inst.p1 as usize)?;
let reg2 = self.get_register(inst.p2 as usize)?;
- let result = match inst.opcode {
- OpCode::Lt => {
- self.compare_registers(®1, ®2, |a, b| a < b, |a, b| a < b, |a, b| a < b)?
- }
- OpCode::Le => {
- self.compare_registers(®1, ®2, |a, b| a <= b, |a, b| a <= b, |a, b| a <= b)?
- }
- OpCode::Gt => {
- self.compare_registers(®1, ®2, |a, b| a > b, |a, b| a > b, |a, b| a > b)?
- }
- OpCode::Ge => {
- self.compare_registers(®1, ®2, |a, b| a >= b, |a, b| a >= b, |a, b| a >= b)?
- }
- OpCode::Eq => self.compare_registers_eq(®1, ®2)?,
- OpCode::Ne => !self.compare_registers_eq(®1, ®2)?,
- _ => {
- return Err(SqawkError::VmError(format!(
- "Invalid comparison opcode: {:?}",
- inst.opcode
- )))
+ // SQL three-valued logic: a comparison involving NULL is UNKNOWN, not
+ // an error and not false. The result register holds NULL for UNKNOWN.
+ //
+ // IfZ already treats a NULL register as zero, so every WHERE, HAVING
+ // and ON site rejects an UNKNOWN row without any control-flow change.
+ let result = match Self::sql_compare(®1, ®2) {
+ None => Register::Null,
+ Some(ord) => {
+ let truth = match inst.opcode {
+ OpCode::Lt => ord == Ordering::Less,
+ OpCode::Le => ord != Ordering::Greater,
+ OpCode::Gt => ord == Ordering::Greater,
+ OpCode::Ge => ord != Ordering::Less,
+ OpCode::Eq => ord == Ordering::Equal,
+ OpCode::Ne => ord != Ordering::Equal,
+ _ => {
+ return Err(SqawkError::VmError(format!(
+ "Invalid comparison opcode: {:?}",
+ inst.opcode
+ )))
+ }
+ };
+ Register::Integer(if truth { 1 } else { 0 })
}
};
- let result_value = Register::Integer(if result { 1 } else { 0 });
- self.set_register(inst.p3 as usize, result_value)?;
+ self.set_register(inst.p3 as usize, result)?;
Ok(ExecuteResult::Continue)
}
- /// Compare two registers for ordering operations (Lt, Le, Gt, Ge)
- fn compare_registers(
- &self,
- reg1: &Register,
- reg2: &Register,
- cmp_int: Fi,
- cmp_float: Ff,
- cmp_str: Fs,
- ) -> SqawkResult
- where
- Fi: Fn(i64, i64) -> bool,
- Ff: Fn(f64, f64) -> bool,
- Fs: Fn(&str, &str) -> bool,
- {
- match (reg1, reg2) {
- (Register::Integer(a), Register::Integer(b)) => Ok(cmp_int(*a, *b)),
- (Register::Float(a), Register::Float(b)) => Ok(cmp_float(*a, *b)),
- (Register::String(a), Register::String(b)) => Ok(cmp_str(a, b)),
- (Register::Integer(a), Register::Float(b)) => Ok(cmp_float(*a as f64, *b)),
- (Register::Float(a), Register::Integer(b)) => Ok(cmp_float(*a, *b as f64)),
- _ => Err(SqawkError::VmError(format!(
- "Cannot compare incompatible types: {:?} and {:?}",
- reg1, reg2
- ))),
+ /// SQL comparison of two registers.
+ ///
+ /// `None` means UNKNOWN, which is the result whenever either operand is
+ /// NULL -- including NULL against NULL. (Grouping and sorting need the
+ /// opposite, NULL equal to NULL, which is why the Compare opcode has its
+ /// own internal ordering rather than sharing this.)
+ ///
+ /// Mixed string/number operands are coerced when the string parses as a
+ /// number, and compared as text otherwise. This is the awk-shaped choice
+ /// rather than the SQLite one, for two reasons: a CSV column is untyped
+ /// text that sqawk guesses a type for PER CELL, so a single column can
+ /// hold both; and arithmetic already coerces this way, so `x + '1'` and
+ /// `x > '1'` agree instead of disagreeing.
+ fn sql_compare(a: &Register, b: &Register) -> Option {
+ // Any NULL operand yields UNKNOWN.
+ if matches!(a, Register::Null) || matches!(b, Register::Null) {
+ return None;
}
- }
- /// Compare two registers for equality (Eq, Ne)
- fn compare_registers_eq(&self, reg1: &Register, reg2: &Register) -> SqawkResult {
- match (reg1, reg2) {
- (Register::Integer(a), Register::Integer(b)) => Ok(a == b),
- (Register::Float(a), Register::Float(b)) => Ok(a == b),
- (Register::String(a), Register::String(b)) => Ok(a == b),
- (Register::Boolean(a), Register::Boolean(b)) => Ok(a == b),
- (Register::Null, Register::Null) => Ok(true),
- (Register::Integer(a), Register::Float(b)) => Ok((*a as f64) == *b),
- (Register::Float(a), Register::Integer(b)) => Ok(*a == (*b as f64)),
- // Different types are not equal
- _ => Ok(false),
+ fn as_number(reg: &Register) -> Option {
+ match reg {
+ Register::Integer(i) => Some(*i as f64),
+ Register::Float(f) => Some(*f),
+ Register::Boolean(v) => Some(if *v { 1.0 } else { 0.0 }),
+ Register::String(s) => s.trim().parse::().ok(),
+ Register::Null => None,
+ }
+ }
+
+ match (a, b) {
+ // Integer against integer compares as i64, not through f64, so
+ // values above 2^53 stay exact.
+ (Register::Integer(x), Register::Integer(y)) => Some(x.cmp(y)),
+ (Register::String(x), Register::String(y)) => Some(x.cmp(y)),
+ (Register::Boolean(x), Register::Boolean(y)) => Some(x.cmp(y)),
+ _ => match (as_number(a), as_number(b)) {
+ (Some(x), Some(y)) => x.partial_cmp(&y),
+ // A string that does not parse as a number falls back to a
+ // textual comparison against the other side's rendering.
+ _ => {
+ let x: Value = a.clone().into();
+ let y: Value = b.clone().into();
+ Some(x.to_string().cmp(&y.to_string()))
+ }
+ },
}
}
diff --git a/src/vm/mod.rs b/src/vm/mod.rs
index 57a189f..75966e7 100644
--- a/src/vm/mod.rs
+++ b/src/vm/mod.rs
@@ -7,6 +7,13 @@
//!
//! For more information on this approach, see: https://www.sqlite.org/opcode.html
+/// The sqlparser version sqawk is built against.
+///
+/// Reported by the REPL's `.version`, since which SQL dialect version is in
+/// use is the single most useful thing to know when a statement is rejected.
+pub const SQLPARSER_VERSION: &str = "0.62";
+
+pub mod ast_compat;
pub mod bytecode;
pub mod compiler;
mod compiler_aggregate;
@@ -25,13 +32,17 @@ use std::collections::HashSet;
use crate::capacity::DEFAULT_TABLE_CAPACITY;
use crate::database::Database;
-use crate::error::SqawkResult;
+use crate::error::{SqawkError, SqawkResult};
use crate::table::Table;
/// Result of VM execution including both the result table and modified table names
pub struct VmExecutionResult {
- /// The result table (for SELECT queries)
- pub table: Option,
+ /// One result table per statement that produced rows, in order.
+ ///
+ /// A multi-statement script produces one result set per statement; folding
+ /// them into a single table would print every statement's rows under one
+ /// header.
+ pub tables: Vec,
/// Names of tables that were modified (INSERT, UPDATE, DELETE, CREATE TABLE)
pub modified_tables: HashSet,
/// Number of rows affected by the last DML statement (INSERT, UPDATE, DELETE)
@@ -53,10 +64,7 @@ pub fn execute_vm(
) -> SqawkResult {
if verbose {
println!("VM Engine: Executing SQL via bytecode: {}", sql);
- }
- // In verbose mode, print messages about SQL features being applied
- if verbose {
let sql_upper = sql.to_uppercase();
if sql_upper.contains("DISTINCT") {
eprintln!("Applying DISTINCT");
@@ -75,41 +83,218 @@ pub fn execute_vm(
}
}
- // PHASE 1: SQL → BYTECODE COMPILATION
+ // Parse once, then compile and execute ONE STATEMENT AT A TIME.
+ //
+ // Every statement used to be folded into a single bytecode program sharing
+ // one result buffer and one schema, which had two consequences:
+ //
+ // - `SELECT a; SELECT b` printed the rows of both statements under the
+ // LAST statement's header.
+ // - A statement could not see the effects of the one before it, because
+ // modifications are applied only after the whole program finishes. So
+ // `UPDATE ...; SELECT ...` ran the SELECT against the pre-update table
+ // and, sharing the buffer, emitted nothing useful.
+ //
+ // Executing per statement gives each its own result set and makes each
+ // statement observe the database as the previous one left it, which is
+ // also what `CREATE TABLE t; INSERT INTO t ...` requires.
+ let dialect = sqlparser::dialect::HiveDialect {};
+ let statements =
+ sqlparser::parser::Parser::parse_sql(&dialect, sql).map_err(SqawkError::SqlParseError)?;
+
+ if statements.is_empty() {
+ return Err(SqawkError::InvalidSqlQuery(
+ "No SQL statements found".to_string(),
+ ));
+ }
- // Use the SqlCompiler to convert the SQL statement to bytecode
- let mut compiler = compiler::SqlCompiler::new(database, verbose);
- let program = compiler.compile(sql)?;
+ let mut tables: Vec = Vec::new();
+ let mut modified_tables = HashSet::with_capacity(DEFAULT_TABLE_CAPACITY);
+ let mut affected_rows: usize = 0;
- if verbose {
- println!("Phase 1 complete - Generated bytecode:");
- println!("{}", program);
- }
+ for statement in &statements {
+ // PHASE 0: MATERIALIZE DERIVED TABLES.
+ //
+ // `FROM (SELECT ...) t` has no representation in the compiler, which
+ // only knows how to open a cursor on a named table. Rather than teach
+ // every FROM site about subqueries, each derived table is executed
+ // first and registered under its alias, and the statement is rewritten
+ // to reference that name. The compiler then sees an ordinary table.
+ let mut statement = statement.clone();
+ let derived = materialize_derived_tables(&mut statement, database, verbose)?;
+
+ // PHASE 1: SQL -> BYTECODE, against the CURRENT database state.
+ let program = {
+ let mut compiler = compiler::SqlCompiler::new(database, verbose);
+ compiler.compile_statement_program(&statement)?
+ };
+
+ if verbose {
+ println!("Generated bytecode:");
+ println!("{}", program);
+ }
- // PHASE 2: BYTECODE → EXECUTION & RESULTS
+ // PHASE 2: EXECUTE
+ let mut vm = engine::VmEngine::new_mut(database, verbose);
+ vm.init(program);
+ vm.execute()?;
+
+ let result_table = vm.create_result_table()?;
+ let modifications = vm.take_modifications();
+ drop(vm);
+
+ // PHASE 3: APPLY MODIFICATIONS before the next statement compiles.
+ let applied = apply_modifications(database, modifications, verbose)?;
+ modified_tables.extend(applied.modified_tables);
+
+ // A DML statement reports its own count, INCLUDING zero.
+ //
+ // Guarding on `applied.affected_rows > 0` conflated two different
+ // things: a statement that was not DML at all (a SELECT contributes no
+ // modifications), and a DML statement that matched no rows. The latter
+ // then inherited the previous statement's count, so
+ // `UPDATE ... WHERE ; UPDATE ... WHERE `
+ // reported 3 rows affected instead of 0.
+ //
+ // Which statements count is the same set SQL's `changes()` uses.
+ if is_row_counting_dml(&statement) {
+ affected_rows = applied.affected_rows;
+ }
- // Initialize the VM engine with the bytecode program
- let mut vm = engine::VmEngine::new_mut(database, verbose);
- vm.init(program);
+ // Derived tables live only for the statement that declared them.
+ for name in derived {
+ database.remove_table(&name);
+ }
- if verbose {
- println!("Phase 2 starting - Executing bytecode in VM");
+ if let Some(t) = result_table {
+ tables.push(t);
+ }
}
- // Execute all instructions in the program
- vm.execute()?;
+ Ok(VmExecutionResult {
+ tables,
+ modified_tables,
+ affected_rows,
+ })
+}
- if verbose {
- println!("Phase 2 complete - Execution finished");
+/// Execute every derived table in `statement`, register each under its alias,
+/// and rewrite the statement to reference the registered name.
+///
+/// Returns the names registered, so the caller can drop them once the
+/// statement finishes -- a derived table is scoped to the statement that
+/// declares it.
+fn materialize_derived_tables(
+ statement: &mut sqlparser::ast::Statement,
+ database: &mut Database,
+ verbose: bool,
+) -> SqawkResult> {
+ let mut registered = Vec::new();
+ if let sqlparser::ast::Statement::Query(query) = statement {
+ materialize_in_query(query, database, verbose, &mut registered)?;
+ }
+ Ok(registered)
+}
+
+fn materialize_in_query(
+ query: &mut sqlparser::ast::Query,
+ database: &mut Database,
+ verbose: bool,
+ registered: &mut Vec,
+) -> SqawkResult<()> {
+ if let sqlparser::ast::SetExpr::Select(select) = &mut *query.body {
+ for twj in &mut select.from {
+ materialize_in_factor(&mut twj.relation, database, verbose, registered)?;
+ for join in &mut twj.joins {
+ materialize_in_factor(&mut join.relation, database, verbose, registered)?;
+ }
+ }
+ }
+ Ok(())
+}
+
+fn materialize_in_factor(
+ factor: &mut sqlparser::ast::TableFactor,
+ database: &mut Database,
+ verbose: bool,
+ registered: &mut Vec,
+) -> SqawkResult<()> {
+ let (subquery, alias) = match factor {
+ sqlparser::ast::TableFactor::Derived {
+ subquery, alias, ..
+ } => (subquery.clone(), alias.clone()),
+ _ => return Ok(()),
+ };
+
+ // SQL requires a derived table to be named; without one there is nothing
+ // for the outer query to refer to.
+ let alias = alias.ok_or_else(|| {
+ SqawkError::InvalidSqlQuery("Subquery in FROM must have an alias".to_string())
+ })?;
+ let name = alias.name.value.to_ascii_lowercase();
+
+ if database.has_table(&name) {
+ return Err(SqawkError::InvalidSqlQuery(format!(
+ "Derived table alias '{}' shadows an existing table",
+ name
+ )));
}
- // Get results and modifications before dropping VM (to release database borrow)
- let result_table = vm.create_result_table();
- let modifications = vm.take_modifications();
+ // Nested derived tables are materialized innermost-first.
+ let mut inner = *subquery;
+ materialize_in_query(&mut inner, database, verbose, registered)?;
+
+ let inner_sql = inner.to_string();
+ let result = execute_vm(&inner_sql, database, verbose)?;
+ let mut table = result.tables.into_iter().next_back().ok_or_else(|| {
+ SqawkError::InvalidSqlQuery("Subquery in FROM produced no result".to_string())
+ })?;
+ table.set_name(name.clone());
+ database.add_table(name.clone(), table)?;
+ registered.push(name.clone());
+
+ *factor = sqlparser::ast::TableFactor::Table {
+ name: sqlparser::ast::ObjectName(vec![sqlparser::ast::ObjectNamePart::Identifier(
+ sqlparser::ast::Ident::new(name),
+ )]),
+ alias: None,
+ args: None,
+ with_hints: Vec::new(),
+ version: None,
+ partitions: Vec::new(),
+ with_ordinality: false,
+ json_path: None,
+ sample: None,
+ index_hints: Vec::new(),
+ };
+ Ok(())
+}
- // Drop VM to release database borrow
- drop(vm);
+/// Whether a statement's affected-row count is meaningful to report.
+///
+/// The INSERT/UPDATE/DELETE family, matching what SQL's `changes()` covers. A
+/// SELECT or DDL statement leaves the previous count alone rather than
+/// resetting it to zero.
+fn is_row_counting_dml(statement: &sqlparser::ast::Statement) -> bool {
+ use sqlparser::ast::Statement;
+ matches!(
+ statement,
+ Statement::Insert(_) | Statement::Update(_) | Statement::Delete(_)
+ )
+}
+/// Tables touched and rows affected by one statement's modifications.
+struct AppliedModifications {
+ modified_tables: HashSet,
+ affected_rows: usize,
+}
+
+/// Apply a statement's pending modifications to the database.
+fn apply_modifications(
+ database: &mut Database,
+ modifications: Vec,
+ verbose: bool,
+) -> SqawkResult {
// PHASE 3: APPLY MODIFICATIONS TO DATABASE
// Track which tables were modified
@@ -126,6 +311,10 @@ pub fn execute_vm(
let mut insert_counts_by_table: std::collections::HashMap =
std::collections::HashMap::new();
+ // Rows replaced in place by UPDATE, per table.
+ let mut replace_counts_by_table: std::collections::HashMap =
+ std::collections::HashMap::new();
+
// First pass: collect deletions and apply non-delete modifications
for modification in modifications {
match modification {
@@ -137,6 +326,19 @@ pub fn execute_vm(
.or_insert(0) += 1;
modified_tables.insert(table_name);
}
+ engine::TableModification::Replace {
+ table_name,
+ row_index,
+ row,
+ } => {
+ let table = database.get_table_mut(&table_name)?;
+ table.replace_row(row_index, row)?;
+ // Tallied per table and reported once, not once per row.
+ *replace_counts_by_table
+ .entry(table_name.clone())
+ .or_insert(0) += 1;
+ modified_tables.insert(table_name);
+ }
engine::TableModification::Delete {
table_name,
row_index,
@@ -273,23 +475,14 @@ pub fn execute_vm(
}
}
- let result = result_table?;
-
- if verbose {
- if let Some(ref table) = result {
- println!(
- "Result table created with {} rows and {} columns",
- table.row_count(),
- table.column_count()
- );
- } else {
- println!("Query executed successfully with no result table");
+ for (_table, count) in replace_counts_by_table {
+ affected_rows += count;
+ if verbose {
+ eprintln!("Updated {} rows", count);
}
}
- // Return both the result table, modified table names, and affected row count
- Ok(VmExecutionResult {
- table: result,
+ Ok(AppliedModifications {
modified_tables,
affected_rows,
})
diff --git a/src/vm/tests.rs b/src/vm/tests.rs
index 326cbde..9b45049 100644
--- a/src/vm/tests.rs
+++ b/src/vm/tests.rs
@@ -1559,3 +1559,149 @@ mod comparison_tests {
}
}
}
+
+/// Tests for the Not / Compare / Jump opcodes.
+///
+/// These exist to support the compiler unification: `Not` makes `NOT `
+/// expressible and replaces several hand-rolled inversions, while
+/// `Compare`/`Jump` make multi-column key comparison two instructions
+/// regardless of key width -- which is what GROUP BY needs in order to detect
+/// a group change across ALL key columns instead of only the first.
+mod compare_jump_not_tests {
+ use super::bytecode_tests::{create_instruction, execute_bytecode_program};
+ use super::*;
+
+ /// Run a program and return the single result row.
+ fn run_row(instructions: Vec) -> Vec {
+ let database = Database::new();
+ let table = execute_bytecode_program(instructions, &database)
+ .expect("program failed")
+ .expect("expected a result table");
+ assert_eq!(table.row_count(), 1, "expected exactly one result row");
+ table.rows()[0].clone()
+ }
+
+ fn init() -> Instruction {
+ create_instruction(OpCode::Init, 0, 1, 0, None, None)
+ }
+
+ fn halt() -> Instruction {
+ create_instruction(OpCode::Halt, 0, 0, 0, None, None)
+ }
+
+ #[test]
+ fn not_inverts_truthy_and_falsy() {
+ // r1 = 5 -> NOT r1 = 0 ; r3 = 0 -> NOT r3 = 1
+ let row = run_row(vec![
+ init(),
+ create_instruction(OpCode::Integer, 5, 1, 0, None, None),
+ create_instruction(OpCode::Not, 1, 2, 0, None, None),
+ create_instruction(OpCode::Integer, 0, 3, 0, None, None),
+ create_instruction(OpCode::Not, 3, 4, 0, None, None),
+ // Emit r2, r4 as a two-column row (registers must be contiguous).
+ create_instruction(OpCode::Copy, 2, 10, 0, None, None),
+ create_instruction(OpCode::Copy, 4, 11, 0, None, None),
+ create_instruction(OpCode::ResultRow, 10, 2, 0, None, None),
+ halt(),
+ ]);
+ assert_eq!(row[0], Value::Integer(0), "NOT 5 should be 0");
+ assert_eq!(row[1], Value::Integer(1), "NOT 0 should be 1");
+ }
+
+ #[test]
+ fn not_propagates_null() {
+ // NOT NULL must be NULL (SQL UNKNOWN), not true.
+ let row = run_row(vec![
+ init(),
+ create_instruction(OpCode::Null, 0, 1, 0, None, None),
+ create_instruction(OpCode::Not, 1, 2, 0, None, None),
+ create_instruction(OpCode::ResultRow, 2, 1, 0, None, None),
+ halt(),
+ ]);
+ assert_eq!(
+ row[0],
+ Value::Null,
+ "NOT NULL must stay NULL, not become true"
+ );
+ }
+
+ /// Build a program that compares two 2-register vectors and reports which
+ /// branch Jump takes: 1 = Less, 2 = Equal, 3 = Greater.
+ fn compare_branch(lhs: [i64; 2], rhs: [i64; 2]) -> Value {
+ let row = run_row(vec![
+ init(),
+ // lhs in r1,r2 ; rhs in r3,r4
+ create_instruction(OpCode::Integer, lhs[0], 1, 0, None, None),
+ create_instruction(OpCode::Integer, lhs[1], 2, 0, None, None),
+ create_instruction(OpCode::Integer, rhs[0], 3, 0, None, None),
+ create_instruction(OpCode::Integer, rhs[1], 4, 0, None, None),
+ // Compare r1..r2 against r3..r4 (2 columns)
+ create_instruction(OpCode::Compare, 1, 3, 2, None, None),
+ // Jump: less -> 7, equal -> 9, greater -> 11
+ create_instruction(OpCode::Jump, 7, 9, 11, None, None),
+ // 7: less
+ create_instruction(OpCode::Integer, 1, 5, 0, None, None),
+ create_instruction(OpCode::Goto, 0, 13, 0, None, None),
+ // 9: equal
+ create_instruction(OpCode::Integer, 2, 5, 0, None, None),
+ create_instruction(OpCode::Goto, 0, 13, 0, None, None),
+ // 11: greater
+ create_instruction(OpCode::Integer, 3, 5, 0, None, None),
+ create_instruction(OpCode::Goto, 0, 13, 0, None, None),
+ // 13: emit
+ create_instruction(OpCode::ResultRow, 5, 1, 0, None, None),
+ halt(),
+ ]);
+ row[0].clone()
+ }
+
+ #[test]
+ fn compare_jump_branches_three_ways() {
+ assert_eq!(compare_branch([1, 1], [1, 1]), Value::Integer(2), "equal");
+ assert_eq!(compare_branch([1, 1], [1, 2]), Value::Integer(1), "less");
+ assert_eq!(compare_branch([1, 2], [1, 1]), Value::Integer(3), "greater");
+ }
+
+ #[test]
+ fn compare_uses_all_key_columns() {
+ // The whole point: a difference in the SECOND column must be detected.
+ // The old single-Ne group-change check compared only the first, which
+ // is why multi-column GROUP BY collapsed to one key.
+ assert_ne!(
+ compare_branch([7, 1], [7, 2]),
+ Value::Integer(2),
+ "vectors differing only in the second column must not compare equal"
+ );
+ }
+
+ #[test]
+ fn compare_treats_nulls_as_equal() {
+ // Internal ordering, not SQL `=`: two NULL group keys must group
+ // together. SQL equality must yield UNKNOWN for NULL, which is exactly
+ // why Compare is a separate opcode from Eq.
+ let row = run_row(vec![
+ init(),
+ create_instruction(OpCode::Null, 0, 1, 0, None, None),
+ create_instruction(OpCode::Null, 0, 2, 0, None, None),
+ create_instruction(OpCode::Compare, 1, 2, 1, None, None),
+ create_instruction(OpCode::Jump, 5, 7, 9, None, None),
+ // 5: less
+ create_instruction(OpCode::Integer, 1, 3, 0, None, None),
+ create_instruction(OpCode::Goto, 0, 11, 0, None, None),
+ // 7: equal
+ create_instruction(OpCode::Integer, 2, 3, 0, None, None),
+ create_instruction(OpCode::Goto, 0, 11, 0, None, None),
+ // 9: greater
+ create_instruction(OpCode::Integer, 3, 3, 0, None, None),
+ create_instruction(OpCode::Goto, 0, 11, 0, None, None),
+ // 11: emit
+ create_instruction(OpCode::ResultRow, 3, 1, 0, None, None),
+ halt(),
+ ]);
+ assert_eq!(
+ row[0],
+ Value::Integer(2),
+ "NULL should compare equal to NULL"
+ );
+ }
+}
diff --git a/tests/data/customers.csv b/tests/data/customers.csv
new file mode 100644
index 0000000..2babfbf
--- /dev/null
+++ b/tests/data/customers.csv
@@ -0,0 +1,4 @@
+id,name
+1,Ann
+2,Ben
+3,Cara
diff --git a/tests/data/nullable.csv b/tests/data/nullable.csv
new file mode 100644
index 0000000..6bb9f7d
--- /dev/null
+++ b/tests/data/nullable.csv
@@ -0,0 +1,6 @@
+id,name,score,grade
+1,Alice,90,A
+2,Bob,,B
+3,Charlie,70,
+4,,50,D
+5,Eve,90,A
diff --git a/tests/data/purchases.csv b/tests/data/purchases.csv
new file mode 100644
index 0000000..5f8c5a5
--- /dev/null
+++ b/tests/data/purchases.csv
@@ -0,0 +1,5 @@
+id,customer_id,item
+10,1,Book
+11,1,Pen
+12,2,Desk
+13,99,Orphan
diff --git a/tests/data/repeats.csv b/tests/data/repeats.csv
new file mode 100644
index 0000000..cffb03f
--- /dev/null
+++ b/tests/data/repeats.csv
@@ -0,0 +1,6 @@
+id,v
+1,10
+2,10
+3,20
+4,30
+5,30
diff --git a/tests/defects/mod.rs b/tests/defects/mod.rs
new file mode 100644
index 0000000..0345a97
--- /dev/null
+++ b/tests/defects/mod.rs
@@ -0,0 +1,538 @@
+//! Red-list: the correctness defects found in the 2026-08 audit.
+//!
+//! Every test in this module encodes **correct** SQL behaviour for a defect
+//! found in that audit. Each was `#[ignore]`d when written, because sqawk got
+//! it wrong, and un-ignored as it was fixed.
+//!
+//! As of this commit the red list is EMPTY: every test here passes, and they
+//! now serve as regression cover for the defects rather than as a to-do list.
+//! Add a new `#[ignore]`d test here when a defect is found, and remove the
+//! attribute when it is fixed.
+//!
+//! Run the red list with:
+//! cargo test --test mod defects -- --ignored
+//!
+//! These use the exact-output helpers (`assert_query`), not the substring
+//! `contains` assertions used elsewhere in the suite. Substring matching is
+//! precisely why every defect below survived a 302-test green suite: a test
+//! asserting `"Engineering,3"` passes even when three extra wrong rows are
+//! emitted, and passes regardless of row order.
+
+use crate::helpers::{
+ assert_after_write, assert_query, departments_csv, employees_csv, nullable_csv, orders_csv,
+ users_csv,
+};
+
+// ---------------------------------------------------------------------------
+// Group 1: WHERE / GROUP BY / HAVING -- silent wrong answers
+// ---------------------------------------------------------------------------
+
+/// Defect 1: `compile_select_with_group_by` never reads `select.selection`,
+/// so WHERE is silently ignored on *every* GROUP BY query.
+/// Currently returns all four departments with unfiltered counts.
+#[test]
+fn where_is_applied_to_group_by() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT department, COUNT(*) FROM employees WHERE salary > 60000 GROUP BY department",
+ "department,COUNT\n\
+ Engineering,3\n\
+ Sales,2",
+ )
+}
+
+/// Defect 2: group-change detection compares only the first key column
+/// (single `Ne` at compiler_aggregate.rs:396), so multi-column GROUP BY
+/// collapses to the first column. Currently returns 4 rows instead of 6,
+/// with an arbitrary `role` value.
+#[test]
+fn group_by_uses_all_key_columns() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT department, role, COUNT(*) FROM employees GROUP BY department, role",
+ "department,role,COUNT\n\
+ Engineering,Developer,2\n\
+ Engineering,Manager,1\n\
+ HR,Intern,1\n\
+ Marketing,Analyst,1\n\
+ Marketing,Specialist,1\n\
+ Sales,Director,2",
+ )
+}
+
+/// Defect 3a: `query` is never passed into the GROUP BY compiler
+/// (compiler.rs:760-766), so ORDER BY is silently dropped.
+#[test]
+fn order_by_applies_to_group_by() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT department, COUNT(*) FROM employees GROUP BY department ORDER BY department DESC",
+ "department,COUNT\n\
+ Sales,2\n\
+ Marketing,2\n\
+ HR,1\n\
+ Engineering,3",
+ )
+}
+
+/// Defect 3b: same root cause -- LIMIT is silently dropped on GROUP BY.
+#[test]
+fn limit_applies_to_group_by() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT department, COUNT(*) FROM employees GROUP BY department LIMIT 1",
+ "department,COUNT\n\
+ Engineering,3",
+ )
+}
+
+/// Defect 7: `compile_having_operand` matches aggregates by function *type*,
+/// ignoring the argument column, so `HAVING SUM(age)` binds to `SUM(salary)`.
+/// Needs two same-type aggregates to surface -- MIN vs MAX bind correctly.
+#[test]
+fn having_binds_to_the_named_aggregate() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT department, SUM(salary), SUM(age) FROM employees \
+ GROUP BY department HAVING SUM(age) > 60",
+ "department,SUM,SUM\n\
+ Engineering,210000,97\n\
+ Sales,170000,85",
+ )
+}
+
+/// Defect 6: `func.distinct` is never read, so COUNT(DISTINCT x) counts rows.
+#[test]
+fn count_distinct_deduplicates() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT COUNT(DISTINCT department) FROM employees",
+ "COUNT\n4",
+ )
+}
+
+/// Defect A: the GROUP BY result block is laid out positionally as
+/// [group keys.., aggregates..] while the schema is built in projection order,
+/// so headers and values disagree. Currently prints `COUNT,department` over
+/// `Engineering,3`.
+#[test]
+fn group_by_respects_projection_order() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT COUNT(*), department FROM employees GROUP BY department",
+ "COUNT,department\n\
+ 3,Engineering\n\
+ 1,HR\n\
+ 2,Marketing\n\
+ 2,Sales",
+ )
+}
+
+/// Defect A (second symptom): a GROUP BY key that is not in the projection
+/// still emits a phantom output column, currently named `col1`.
+#[test]
+fn group_by_emits_only_projected_columns() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT COUNT(*) FROM employees GROUP BY department",
+ "COUNT\n3\n1\n2\n2",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Group 2: JOIN projection and ordering
+// ---------------------------------------------------------------------------
+
+/// Defect B: JOIN always emits left-table columns then right-table columns,
+/// while the schema is built in SELECT order. Currently prints header
+/// `orders.id,users.name` over values `John,101` -- swapped.
+#[test]
+fn join_respects_projection_order() -> Result<(), Box> {
+ assert_query(
+ &[users_csv(), orders_csv()],
+ "SELECT orders.id, users.name FROM users JOIN orders ON users.id = orders.user_id",
+ "orders.id,users.name\n\
+ 101,John\n\
+ 103,John\n\
+ 105,John\n\
+ 102,Jane\n\
+ 104,Jane",
+ )
+}
+
+/// Defect 4: ORDER BY is deliberately skipped on explicit JOINs
+/// (compiler.rs:729-737 -- "isn't fully implemented. For now [...] skip").
+#[test]
+fn order_by_applies_to_join() -> Result<(), Box> {
+ assert_query(
+ &[users_csv(), orders_csv()],
+ "SELECT users.name, orders.id FROM users JOIN orders ON users.id = orders.user_id \
+ ORDER BY orders.id DESC",
+ "users.name,orders.id\n\
+ John,105\n\
+ Jane,104\n\
+ John,103\n\
+ Jane,102\n\
+ John,101",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Group 3: projection expressions
+// ---------------------------------------------------------------------------
+
+/// Defect 5: the `_` arm of `compile_projection_expr` (compiler.rs:2393-2405)
+/// emits **no instruction** on failure, leaving the register at its default
+/// NULL. The same query without LIMIT returns correct values, because it takes
+/// a different compilation path.
+#[test]
+fn arithmetic_projection_survives_limit() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name, salary*2 FROM employees LIMIT 3",
+ "name,expr\n\
+ Alice,140000\n\
+ Bob,110000\n\
+ Charlie,130000",
+ )
+}
+
+/// `CASE` is routed in WHERE but hits `resolve_column_expr` in a projection,
+/// which errors with "Only column references supported in SELECT".
+#[test]
+fn case_works_in_projection() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name, CASE WHEN age > 30 THEN 'old' ELSE 'young' END FROM employees LIMIT 3",
+ "name,expr\n\
+ Alice,young\n\
+ Bob,young\n\
+ Charlie,old",
+ )
+}
+
+/// `||` parses fine in sqlparser 0.36 -- the compiler simply has no arm for
+/// `BinaryOperator::StringConcat` (compiler.rs:1922).
+#[test]
+fn string_concat_operator_works() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name || '!' FROM employees LIMIT 2",
+ "expr\n\
+ Alice!\n\
+ Bob!",
+ )
+}
+
+/// A comparison used as a projected *value* rather than a filter.
+#[test]
+fn comparison_works_in_projection() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name, salary > 60000 FROM employees LIMIT 3",
+ "name,expr\n\
+ Alice,1\n\
+ Bob,0\n\
+ Charlie,1",
+ )
+}
+
+/// Aggregate over an expression -- `resolve_column_expr` rejects anything that
+/// is not a bare column reference.
+#[test]
+fn aggregate_accepts_expression_argument() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT SUM(salary + age) FROM employees",
+ "SUM\n540257",
+ )
+}
+
+/// An aggregate nested inside an expression is not even *detected* as an
+/// aggregate query (`is_aggregate_expr` only matches a top-level Function),
+/// so it falls into the plain table-scan path.
+#[test]
+fn aggregate_nested_in_expression_works() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT SUM(salary) + 1 FROM employees",
+ "expr\n540001",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Group 4: WHERE clause expression gaps
+// ---------------------------------------------------------------------------
+
+/// `NOT` has no arm in the WHERE compiler at all.
+#[test]
+fn not_works_in_where() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name FROM employees WHERE NOT (age > 40)",
+ "name\n\
+ Alice\nBob\nCharlie\nDavid\nEve\nFrank\nHenry",
+ )
+}
+
+/// Table-alias-qualified columns are rejected in a single-table WHERE:
+/// `compile_where_operand` has no `CompoundIdentifier` arm.
+#[test]
+fn alias_qualified_column_works_in_where() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name FROM employees e WHERE e.age > 30",
+ "name\n\
+ Charlie\nDavid\nFrank\nGrace",
+ )
+}
+
+/// Column resolution is case-insensitive in a projection but case-*sensitive*
+/// in WHERE: `SELECT NAME` works while `WHERE NAME = 'Alice'` errors with
+/// "Column 'NAME' not found" (find_column_index, compiler.rs:6010).
+#[test]
+fn column_resolution_is_case_insensitive() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name FROM employees WHERE NAME = 'Alice'",
+ "name\nAlice",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Group 5: NULL semantics and type coercion
+// ---------------------------------------------------------------------------
+
+/// The worst defect for a CSV tool: comparing against a NULL raises a runtime
+/// error instead of evaluating to UNKNOWN and filtering the row out.
+/// `nullable.csv` row 2 has an empty `score`.
+#[test]
+fn null_comparison_yields_unknown() -> Result<(), Box> {
+ assert_query(
+ &[nullable_csv()],
+ "SELECT name FROM nullable WHERE score > 60",
+ "name\n\
+ Alice\nCharlie\nEve",
+ )
+}
+
+/// A NULL must not satisfy an equality test either -- but it must not error.
+///
+/// NOT ignored: this already passes, because equality is routed through
+/// `compare_registers_eq`, which returns `false` for mismatched types instead
+/// of erroring like the ordering comparisons do. It is kept here as a
+/// regression guard: the three-valued-logic change must not break it.
+#[test]
+fn null_never_equals_a_value() -> Result<(), Box> {
+ assert_query(
+ &[nullable_csv()],
+ "SELECT name FROM nullable WHERE score = 90",
+ "name\n\
+ Alice\nEve",
+ )
+}
+
+/// awk-like coercion: a string that parses as a number compares numerically.
+/// Matches what `arithmetic_op` already does, so `x + '1'` and `x > '1'` agree.
+#[test]
+fn string_number_comparison_coerces() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name FROM employees WHERE salary > '60000'",
+ "name\n\
+ Alice\nCharlie\nDavid\nFrank\nGrace",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Group 6: ORDER BY expressiveness
+// ---------------------------------------------------------------------------
+
+/// `build_sort_spec` maps the ORDER BY key into the *projected* column list
+/// and errors if it is absent, so you cannot sort by a column you don't select.
+#[test]
+fn order_by_column_not_in_select_list() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name FROM employees ORDER BY salary DESC",
+ "name\n\
+ Grace\nDavid\nFrank\nAlice\nCharlie\nEve\nBob\nHenry",
+ )
+}
+
+/// ORDER BY an arbitrary expression.
+#[test]
+fn order_by_expression() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name FROM employees ORDER BY salary * -1 LIMIT 3",
+ "name\n\
+ Grace\nDavid\nFrank",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Group 7: DML and FROM expressiveness
+// ---------------------------------------------------------------------------
+
+/// `compile_expr` had no `Identifier` arm, so a column could not appear on the
+/// right-hand side of a SET. Verified through writeback rather than a trailing
+/// SELECT, because `UPDATE ...; SELECT ...` in one invocation does not emit the
+/// SELECT's rows -- see `update_then_select_emits_the_select`.
+#[test]
+fn update_set_from_column() -> Result<(), Box> {
+ assert_after_write(
+ employees_csv(),
+ "employees",
+ "UPDATE employees SET salary = salary + 1",
+ "id,name,age,salary,department,role\n\
+ 1,Alice,30,70001,Engineering,Developer\n\
+ 2,Bob,25,55001,Marketing,Specialist\n\
+ 3,Charlie,35,65001,Engineering,Manager\n\
+ 4,David,40,80001,Sales,Director\n\
+ 5,Eve,28,60001,Marketing,Analyst\n\
+ 6,Frank,32,75001,Engineering,Developer\n\
+ 7,Grace,45,90001,Sales,Director\n\
+ 8,Henry,22,45001,HR,Intern",
+ )
+}
+
+/// An expression RHS combined with WHERE.
+#[test]
+fn update_set_expression_with_where() -> Result<(), Box> {
+ assert_after_write(
+ employees_csv(),
+ "employees",
+ "UPDATE employees SET salary = salary * 2 WHERE department = 'HR'",
+ "id,name,age,salary,department,role\n\
+ 1,Alice,30,70000,Engineering,Developer\n\
+ 2,Bob,25,55000,Marketing,Specialist\n\
+ 3,Charlie,35,65000,Engineering,Manager\n\
+ 4,David,40,80000,Sales,Director\n\
+ 5,Eve,28,60000,Marketing,Analyst\n\
+ 6,Frank,32,75000,Engineering,Developer\n\
+ 7,Grace,45,90000,Sales,Director\n\
+ 8,Henry,22,90000,HR,Intern",
+ )
+}
+
+/// A DML statement followed by a SELECT must run the SELECT against the
+/// updated table and emit its rows. Statements used to share one bytecode
+/// program and one result buffer, and modifications were applied only after
+/// the whole program finished, so the SELECT saw the pre-update table.
+#[test]
+fn update_then_select_emits_the_select() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "UPDATE employees SET salary = 1 WHERE id = 1; \
+ SELECT name, salary FROM employees WHERE id = 1",
+ "name,salary\nAlice,1",
+ )
+}
+
+/// UPDATE must leave rows where it found them.
+///
+/// It is implemented as delete-plus-insert, and the insert appends, so any row
+/// it touches jumps to the end of the table. This is not merely cosmetic: with
+/// --write the reordering is persisted to the user's file.
+#[test]
+fn update_preserves_row_order() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "UPDATE employees SET salary = 1 WHERE id = 1; SELECT id, salary FROM employees LIMIT 3",
+ "id,salary\n1,1\n2,55000\n3,65000",
+ )
+}
+
+/// Derived tables parse fine in sqlparser 0.36; `TableFactor::Derived` simply
+/// has no handler anywhere in the compiler.
+#[test]
+fn derived_table_in_from() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name FROM (SELECT name FROM employees) t LIMIT 2",
+ "name\n\
+ Alice\nBob",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Group 8: window functions
+// ---------------------------------------------------------------------------
+
+/// Defect 8: an aggregate OVER (PARTITION BY x) with no ORDER BY must return
+/// the *partition total* on every row, not a running total.
+#[test]
+fn partition_aggregate_returns_partition_total() -> Result<(), Box> {
+ assert_query(
+ &[employees_csv()],
+ "SELECT name, SUM(salary) OVER (PARTITION BY department) FROM employees",
+ "name,SUM\n\
+ Alice,210000\n\
+ Charlie,210000\n\
+ Frank,210000\n\
+ Henry,45000\n\
+ Bob,115000\n\
+ Eve,115000\n\
+ David,170000\n\
+ Grace,170000",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Group 9: aggregate + GROUP BY interaction on the departments fixture
+// ---------------------------------------------------------------------------
+
+/// The suite's existing `test_group_by_with_order_by` asserts this ordering and
+/// **passes by coincidence**: alphabetical department order happens to equal
+/// avg_salary-descending order in this fixture. Pinned here as an exact-output
+/// test so the coincidence cannot hide a regression again.
+#[test]
+fn group_by_order_by_avg_is_not_coincidental() -> Result<(), Box> {
+ assert_query(
+ &[departments_csv()],
+ "SELECT department, COUNT(*) FROM departments GROUP BY department \
+ ORDER BY COUNT(*) ASC, department DESC",
+ "department,COUNT\n\
+ Sales,2\n\
+ Marketing,2\n\
+ Engineering,3",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Group 10: multi-statement output
+// ---------------------------------------------------------------------------
+
+/// Each statement in a `stmt; stmt` script must print its own result set with
+/// its own header. Currently the rows of every statement are concatenated and
+/// printed under the LAST statement's schema, so
+/// `SELECT name ...; SELECT COUNT(*) ...` emits the header `COUNT` above
+/// `Alice`, `Charlie`, `3`.
+#[test]
+fn multi_statement_output_has_per_statement_headers() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/people.csv"],
+ "SELECT name FROM people WHERE age > 30; SELECT COUNT(*) FROM people",
+ "name\n\
+ Alice\n\
+ Charlie\n\
+ COUNT\n\
+ 3",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Group 11: aggregates over joins
+// ---------------------------------------------------------------------------
+
+/// An aggregate over a join is rejected: the join projection resolver accepts
+/// only column references, so `COUNT(*)` never reaches the aggregate path.
+#[test]
+fn aggregate_over_join() -> Result<(), Box> {
+ assert_query(
+ &[users_csv(), orders_csv()],
+ "SELECT COUNT(*) FROM users JOIN orders ON users.id = orders.user_id",
+ "COUNT\n5",
+ )
+}
diff --git a/tests/golden/mod.rs b/tests/golden/mod.rs
new file mode 100644
index 0000000..966f757
--- /dev/null
+++ b/tests/golden/mod.rs
@@ -0,0 +1,2291 @@
+//! Characterization ("golden") tests: exact-output pins on behaviour that
+//! sqawk currently gets RIGHT.
+//!
+//! Generated from real output, then reviewed. These exist to be a semantic
+//! oracle for refactors -- most immediately the direct sqlparser 0.36 -> 0.62
+//! jump, which has no bisection handle and ~130 mechanical edit sites.
+//!
+//! Unlike the rest of the suite, these assert the COMPLETE stdout: every row,
+//! in order. If a refactor drops a row, reorders output, renames a header, or
+//! changes a value, exactly one of these fails and names the query.
+//!
+//! Known-wrong behaviour is NOT pinned here -- it lives in `tests/defects/`.
+
+use crate::helpers::assert_query;
+
+#[test]
+fn select_star() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT * FROM employees",
+ concat!(
+ "id,name,age,salary,department,role",
+ "\n1,Alice,30,70000,Engineering,Developer",
+ "\n2,Bob,25,55000,Marketing,Specialist",
+ "\n3,Charlie,35,65000,Engineering,Manager",
+ "\n4,David,40,80000,Sales,Director",
+ "\n5,Eve,28,60000,Marketing,Analyst",
+ "\n6,Frank,32,75000,Engineering,Developer",
+ "\n7,Grace,45,90000,Sales,Director",
+ "\n8,Henry,22,45000,HR,Intern",
+ ),
+ )
+}
+
+#[test]
+fn select_columns() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, salary FROM employees",
+ concat!(
+ "name,salary",
+ "\nAlice,70000",
+ "\nBob,55000",
+ "\nCharlie,65000",
+ "\nDavid,80000",
+ "\nEve,60000",
+ "\nFrank,75000",
+ "\nGrace,90000",
+ "\nHenry,45000",
+ ),
+ )
+}
+
+#[test]
+fn select_literal() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT 1",
+ concat!("col0", "\n1",),
+ )
+}
+
+#[test]
+fn where_eq() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE department = 'Engineering'",
+ concat!("name", "\nAlice", "\nCharlie", "\nFrank",),
+ )
+}
+
+#[test]
+fn where_gt() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE salary > 70000",
+ concat!("name", "\nDavid", "\nFrank", "\nGrace",),
+ )
+}
+
+#[test]
+fn where_and() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE salary > 60000 AND department = 'Engineering'",
+ concat!("name", "\nAlice", "\nCharlie", "\nFrank",),
+ )
+}
+
+#[test]
+fn where_or() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE department = 'HR' OR department = 'Sales'",
+ concat!("name", "\nDavid", "\nGrace", "\nHenry",),
+ )
+}
+
+#[test]
+fn where_ne() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE department <> 'Engineering'",
+ concat!("name", "\nBob", "\nDavid", "\nEve", "\nGrace", "\nHenry",),
+ )
+}
+
+#[test]
+fn where_like() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE name LIKE 'A%'",
+ concat!("name", "\nAlice",),
+ )
+}
+
+#[test]
+fn where_not_like() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE name NOT LIKE 'A%'",
+ concat!(
+ "name",
+ "\nBob",
+ "\nCharlie",
+ "\nDavid",
+ "\nEve",
+ "\nFrank",
+ "\nGrace",
+ "\nHenry",
+ ),
+ )
+}
+
+#[test]
+fn where_ilike() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE name ILIKE 'a%'",
+ concat!("name", "\nAlice",),
+ )
+}
+
+#[test]
+fn where_between() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE age BETWEEN 30 AND 40",
+ concat!("name", "\nAlice", "\nCharlie", "\nDavid", "\nFrank",),
+ )
+}
+
+#[test]
+fn where_not_between() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE age NOT BETWEEN 30 AND 40",
+ concat!("name", "\nBob", "\nEve", "\nGrace", "\nHenry",),
+ )
+}
+
+#[test]
+fn where_in_list() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE department IN ('HR', 'Sales')",
+ concat!("name", "\nDavid", "\nGrace", "\nHenry",),
+ )
+}
+
+#[test]
+fn where_not_in_list() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE department NOT IN ('HR', 'Sales')",
+ concat!("name", "\nAlice", "\nBob", "\nCharlie", "\nEve", "\nFrank",),
+ )
+}
+
+#[test]
+fn where_is_null() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/nullable.csv"],
+ "SELECT id FROM nullable WHERE score IS NULL",
+ concat!("id", "\n2",),
+ )
+}
+
+#[test]
+fn where_is_not_null() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/nullable.csv"],
+ "SELECT id FROM nullable WHERE score IS NOT NULL",
+ concat!("id", "\n1", "\n3", "\n4", "\n5",),
+ )
+}
+
+#[test]
+fn order_by_asc() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, salary FROM employees ORDER BY salary ASC",
+ concat!(
+ "name,salary",
+ "\nHenry,45000",
+ "\nBob,55000",
+ "\nEve,60000",
+ "\nCharlie,65000",
+ "\nAlice,70000",
+ "\nFrank,75000",
+ "\nDavid,80000",
+ "\nGrace,90000",
+ ),
+ )
+}
+
+#[test]
+fn order_by_desc() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, salary FROM employees ORDER BY salary DESC",
+ concat!(
+ "name,salary",
+ "\nGrace,90000",
+ "\nDavid,80000",
+ "\nFrank,75000",
+ "\nAlice,70000",
+ "\nCharlie,65000",
+ "\nEve,60000",
+ "\nBob,55000",
+ "\nHenry,45000",
+ ),
+ )
+}
+
+#[test]
+fn order_by_multi() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT department, name FROM employees ORDER BY department ASC, name DESC",
+ concat!(
+ "department,name",
+ "\nEngineering,Frank",
+ "\nEngineering,Charlie",
+ "\nEngineering,Alice",
+ "\nHR,Henry",
+ "\nMarketing,Eve",
+ "\nMarketing,Bob",
+ "\nSales,Grace",
+ "\nSales,David",
+ ),
+ )
+}
+
+#[test]
+fn limit() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees LIMIT 3",
+ concat!("name", "\nAlice", "\nBob", "\nCharlie",),
+ )
+}
+
+#[test]
+fn limit_offset() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees LIMIT 3 OFFSET 2",
+ concat!("name", "\nCharlie", "\nDavid", "\nEve",),
+ )
+}
+
+#[test]
+fn order_by_limit() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, salary FROM employees ORDER BY salary DESC LIMIT 3",
+ concat!(
+ "name,salary",
+ "\nGrace,90000",
+ "\nDavid,80000",
+ "\nFrank,75000",
+ ),
+ )
+}
+
+#[test]
+fn distinct() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT DISTINCT department FROM employees",
+ concat!(
+ "department",
+ "\nEngineering",
+ "\nMarketing",
+ "\nSales",
+ "\nHR",
+ ),
+ )
+}
+
+#[test]
+fn distinct_star() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/duplicates.csv"],
+ "SELECT DISTINCT name, department FROM duplicates",
+ concat!(
+ "name,department",
+ "\nAlice,Engineering",
+ "\nBob,Marketing",
+ "\nCharlie,Engineering",
+ "\nDave,Finance",
+ "\nEve,HR",
+ "\nFrank,Sales",
+ ),
+ )
+}
+
+#[test]
+fn agg_count() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT COUNT(*) FROM employees",
+ concat!("COUNT", "\n8",),
+ )
+}
+
+#[test]
+fn agg_sum() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT SUM(salary) FROM employees",
+ concat!("SUM", "\n540000",),
+ )
+}
+
+#[test]
+fn agg_avg() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT AVG(salary) FROM employees",
+ concat!("AVG", "\n67500",),
+ )
+}
+
+#[test]
+fn agg_min_max() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT MIN(salary), MAX(salary) FROM employees",
+ concat!("MIN,MAX", "\n45000,90000",),
+ )
+}
+
+#[test]
+fn agg_with_where() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT COUNT(*) FROM employees WHERE department = 'Engineering'",
+ concat!("COUNT", "\n3",),
+ )
+}
+
+#[test]
+fn agg_alias() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT COUNT(*) AS n, SUM(salary) AS total FROM employees",
+ concat!("n,total", "\n8,540000",),
+ )
+}
+
+#[test]
+fn group_by_basic() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT department, COUNT(*) FROM employees GROUP BY department",
+ concat!(
+ "department,COUNT",
+ "\nEngineering,3",
+ "\nHR,1",
+ "\nMarketing,2",
+ "\nSales,2",
+ ),
+ )
+}
+
+#[test]
+fn group_by_sum() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT department, SUM(salary) FROM employees GROUP BY department",
+ concat!(
+ "department,SUM",
+ "\nEngineering,210000",
+ "\nHR,45000",
+ "\nMarketing,115000",
+ "\nSales,170000",
+ ),
+ )
+}
+
+#[test]
+fn group_by_having() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT department, COUNT(*) FROM employees GROUP BY department HAVING COUNT(*) > 1",
+ concat!(
+ "department,COUNT",
+ "\nEngineering,3",
+ "\nMarketing,2",
+ "\nSales,2",
+ ),
+ )
+}
+
+#[test]
+fn join_on() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/users.csv", "tests/data/orders.csv"],
+ "SELECT users.name, orders.id FROM users JOIN orders ON users.id = orders.user_id",
+ concat!(
+ "users.name,orders.id",
+ "\nJohn,101",
+ "\nJohn,103",
+ "\nJohn,105",
+ "\nJane,102",
+ "\nJane,104",
+ ),
+ )
+}
+
+#[test]
+fn join_inner() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/users.csv", "tests/data/orders.csv"],
+ "SELECT users.name, orders.date FROM users INNER JOIN orders ON users.id = orders.user_id",
+ concat!(
+ "users.name,orders.date",
+ "\nJohn,2023-01-15",
+ "\nJohn,2023-02-10",
+ "\nJohn,2023-03-05",
+ "\nJane,2023-01-20",
+ "\nJane,2023-02-25",
+ ),
+ )
+}
+
+#[test]
+fn join_left() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/users.csv", "tests/data/orders.csv"],
+ "SELECT users.name, orders.id FROM users LEFT JOIN orders ON users.id = orders.user_id",
+ concat!(
+ "users.name,orders.id",
+ "\nJohn,101",
+ "\nJohn,103",
+ "\nJohn,105",
+ "\nJane,102",
+ "\nJane,104",
+ ),
+ )
+}
+
+#[test]
+fn join_right() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/users.csv", "tests/data/orders.csv"],
+ "SELECT users.name, orders.id FROM users RIGHT JOIN orders ON users.id = orders.user_id",
+ concat!(
+ "users.name,orders.id",
+ "\nJohn,101",
+ "\nJane,102",
+ "\nJohn,103",
+ "\nJane,104",
+ "\nJohn,105",
+ ),
+ )
+}
+
+#[test]
+fn join_full() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/users.csv", "tests/data/orders.csv"],
+ "SELECT users.name, orders.id FROM users FULL JOIN orders ON users.id = orders.user_id",
+ concat!(
+ "users.name,orders.id",
+ "\nJohn,101",
+ "\nJohn,103",
+ "\nJohn,105",
+ "\nJane,102",
+ "\nJane,104",
+ ),
+ )
+}
+
+#[test]
+fn join_cross() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/users.csv", "tests/data/orders.csv"],
+ "SELECT users.name, orders.id FROM users CROSS JOIN orders",
+ concat!(
+ "users.name,orders.id",
+ "\nJohn,101",
+ "\nJohn,102",
+ "\nJohn,103",
+ "\nJohn,104",
+ "\nJohn,105",
+ "\nJane,101",
+ "\nJane,102",
+ "\nJane,103",
+ "\nJane,104",
+ "\nJane,105",
+ ),
+ )
+}
+
+#[test]
+fn join_implicit() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/users.csv", "tests/data/orders.csv"],
+ "SELECT users.name, orders.id FROM users, orders WHERE users.id = orders.user_id",
+ concat!(
+ "users.name,orders.id",
+ "\nJohn,101",
+ "\nJohn,103",
+ "\nJohn,105",
+ "\nJane,102",
+ "\nJane,104",
+ ),
+ )
+}
+
+#[test]
+fn join_three() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/users.csv", "tests/data/orders.csv", "tests/data/products.csv"],
+ "SELECT users.name, products.name FROM users, orders, products WHERE users.id = orders.user_id AND orders.product_id = products.product_id",
+ concat!(
+ "users.name,products.name",
+ "\nJohn,Laptop",
+ "\nJohn,Headphones",
+ "\nJohn,Keyboard",
+ "\nJane,Phone",
+ "\nJane,Monitor",
+ ),
+ )
+}
+
+#[test]
+fn union() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT department FROM employees WHERE age < 26 UNION SELECT department FROM employees WHERE age > 44",
+ concat!(
+ "department",
+ "\nMarketing",
+ "\nHR",
+ "\nSales",
+ ),
+ )
+}
+
+#[test]
+fn union_all() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT department FROM employees WHERE age < 26 UNION ALL SELECT department FROM employees WHERE age > 44",
+ concat!(
+ "department",
+ "\nMarketing",
+ "\nHR",
+ "\nSales",
+ ),
+ )
+}
+
+#[test]
+fn intersect() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT department FROM employees WHERE age < 31 INTERSECT SELECT department FROM employees WHERE age > 31",
+ concat!(
+ "department",
+ "\nEngineering",
+ ),
+ )
+}
+
+#[test]
+fn except() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT department FROM employees WHERE age < 31 EXCEPT SELECT department FROM employees WHERE age > 31",
+ concat!(
+ "department",
+ "\nMarketing",
+ "\nHR",
+ ),
+ )
+}
+
+#[test]
+fn subquery_scalar() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE salary = (SELECT MAX(salary) FROM employees)",
+ concat!("name", "\nGrace",),
+ )
+}
+
+#[test]
+fn subquery_in() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/users.csv", "tests/data/orders.csv"],
+ "SELECT name FROM users WHERE id IN (SELECT user_id FROM orders)",
+ concat!("name", "\nJohn", "\nJane",),
+ )
+}
+
+#[test]
+fn subquery_exists() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/users.csv", "tests/data/orders.csv"],
+ "SELECT name FROM users WHERE EXISTS (SELECT 1 FROM orders WHERE orders.user_id = users.id)",
+ concat!(
+ "name",
+ "\nJohn",
+ "\nJane",
+ ),
+ )
+}
+
+#[test]
+fn str_upper() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT UPPER(name) FROM employees",
+ concat!(
+ "UPPER(name)",
+ "\nALICE",
+ "\nBOB",
+ "\nCHARLIE",
+ "\nDAVID",
+ "\nEVE",
+ "\nFRANK",
+ "\nGRACE",
+ "\nHENRY",
+ ),
+ )
+}
+
+#[test]
+fn str_lower() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT LOWER(name) FROM employees",
+ concat!(
+ "LOWER(name)",
+ "\nalice",
+ "\nbob",
+ "\ncharlie",
+ "\ndavid",
+ "\neve",
+ "\nfrank",
+ "\ngrace",
+ "\nhenry",
+ ),
+ )
+}
+
+#[test]
+fn str_length() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT LENGTH(name) FROM employees",
+ concat!(
+ "LENGTH(name)",
+ "\n5",
+ "\n3",
+ "\n7",
+ "\n5",
+ "\n3",
+ "\n5",
+ "\n5",
+ "\n5",
+ ),
+ )
+}
+
+#[test]
+fn str_substr() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT SUBSTR(name, 1, 3) FROM employees",
+ concat!(
+ "SUBSTR(name)",
+ "\nAli",
+ "\nBob",
+ "\nCha",
+ "\nDav",
+ "\nEve",
+ "\nFra",
+ "\nGra",
+ "\nHen",
+ ),
+ )
+}
+
+#[test]
+fn str_replace() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT REPLACE(department, 'Engineering', 'Eng') FROM employees",
+ concat!(
+ "REPLACE(department)",
+ "\nEng",
+ "\nMarketing",
+ "\nEng",
+ "\nSales",
+ "\nMarketing",
+ "\nEng",
+ "\nSales",
+ "\nHR",
+ ),
+ )
+}
+
+#[test]
+fn str_trim() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/strings.csv"],
+ "SELECT TRIM(padded_text) FROM strings",
+ concat!(
+ "expr",
+ "\ntrimme",
+ "\nneeds space",
+ "\nwhitespace",
+ "\nextra",
+ "\npadding",
+ ),
+ )
+}
+
+#[test]
+fn str_concat() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT CONCAT(name, department) FROM employees",
+ concat!(
+ "CONCAT(name)",
+ "\nAliceEngineering",
+ "\nBobMarketing",
+ "\nCharlieEngineering",
+ "\nDavidSales",
+ "\nEveMarketing",
+ "\nFrankEngineering",
+ "\nGraceSales",
+ "\nHenryHR",
+ ),
+ )
+}
+
+#[test]
+fn str_left_right() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT LEFT(name, 2), RIGHT(name, 2) FROM employees",
+ concat!(
+ "LEFT(name),RIGHT(name)",
+ "\nAl,ce",
+ "\nBo,ob",
+ "\nCh,ie",
+ "\nDa,id",
+ "\nEv,ve",
+ "\nFr,nk",
+ "\nGr,ce",
+ "\nHe,ry",
+ ),
+ )
+}
+
+#[test]
+fn math_abs() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/boundaries.csv"],
+ // NOTE: ABS(i64::MIN) has no representable result. sqawk saturates to
+ // i64::MAX rather than erroring or wrapping. Pinned so any change to
+ // that choice is deliberate, not so that it is endorsed.
+ "SELECT ABS(value) FROM boundaries",
+ concat!(
+ "ABS(value)",
+ "\n0",
+ "\n1",
+ "\n1",
+ "\n9223372036854775807",
+ "\n9223372036854775807",
+ ),
+ )
+}
+
+#[test]
+fn math_round() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT ROUND(salary) FROM employees",
+ concat!(
+ "ROUND(salary)",
+ "\n70000",
+ "\n55000",
+ "\n65000",
+ "\n80000",
+ "\n60000",
+ "\n75000",
+ "\n90000",
+ "\n45000",
+ ),
+ )
+}
+
+#[test]
+fn arith_add() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, salary + 1000 FROM employees",
+ concat!(
+ "name,expr",
+ "\nAlice,71000",
+ "\nBob,56000",
+ "\nCharlie,66000",
+ "\nDavid,81000",
+ "\nEve,61000",
+ "\nFrank,76000",
+ "\nGrace,91000",
+ "\nHenry,46000",
+ ),
+ )
+}
+
+#[test]
+fn arith_mul() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, salary * 2 FROM employees",
+ concat!(
+ "name,expr",
+ "\nAlice,140000",
+ "\nBob,110000",
+ "\nCharlie,130000",
+ "\nDavid,160000",
+ "\nEve,120000",
+ "\nFrank,150000",
+ "\nGrace,180000",
+ "\nHenry,90000",
+ ),
+ )
+}
+
+#[test]
+fn arith_div() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, salary / 1000 FROM employees",
+ concat!(
+ "name,expr",
+ "\nAlice,70",
+ "\nBob,55",
+ "\nCharlie,65",
+ "\nDavid,80",
+ "\nEve,60",
+ "\nFrank,75",
+ "\nGrace,90",
+ "\nHenry,45",
+ ),
+ )
+}
+
+#[test]
+fn arith_mod() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, age % 7 FROM employees",
+ concat!(
+ "name,expr",
+ "\nAlice,2",
+ "\nBob,4",
+ "\nCharlie,0",
+ "\nDavid,5",
+ "\nEve,0",
+ "\nFrank,4",
+ "\nGrace,3",
+ "\nHenry,1",
+ ),
+ )
+}
+
+#[test]
+fn arith_unary_minus() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, -salary FROM employees",
+ concat!(
+ "name,expr",
+ "\nAlice,-70000",
+ "\nBob,-55000",
+ "\nCharlie,-65000",
+ "\nDavid,-80000",
+ "\nEve,-60000",
+ "\nFrank,-75000",
+ "\nGrace,-90000",
+ "\nHenry,-45000",
+ ),
+ )
+}
+
+#[test]
+fn coalesce() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/nullable.csv"],
+ "SELECT id, COALESCE(score, 0) FROM nullable",
+ concat!(
+ "id,COALESCE(score)",
+ "\n1,90",
+ "\n2,0",
+ "\n3,70",
+ "\n4,50",
+ "\n5,90",
+ ),
+ )
+}
+
+#[test]
+fn nullif() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, NULLIF(department, 'HR') FROM employees",
+ concat!(
+ "name,NULLIF(department)",
+ "\nAlice,Engineering",
+ "\nBob,Marketing",
+ "\nCharlie,Engineering",
+ "\nDavid,Sales",
+ "\nEve,Marketing",
+ "\nFrank,Engineering",
+ "\nGrace,Sales",
+ "\nHenry,NULL",
+ ),
+ )
+}
+
+#[test]
+fn win_row_number() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, ROW_NUMBER() OVER (ORDER BY salary DESC) FROM employees",
+ concat!(
+ "name,ROW_NUMBER",
+ "\nGrace,1",
+ "\nDavid,2",
+ "\nFrank,3",
+ "\nAlice,4",
+ "\nCharlie,5",
+ "\nEve,6",
+ "\nBob,7",
+ "\nHenry,8",
+ ),
+ )
+}
+
+#[test]
+fn win_rank() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, RANK() OVER (ORDER BY department) FROM employees",
+ concat!(
+ "name,RANK",
+ "\nAlice,1",
+ "\nCharlie,1",
+ "\nFrank,1",
+ "\nHenry,4",
+ "\nBob,5",
+ "\nEve,5",
+ "\nDavid,7",
+ "\nGrace,7",
+ ),
+ )
+}
+
+#[test]
+fn win_dense_rank() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, DENSE_RANK() OVER (ORDER BY department) FROM employees",
+ concat!(
+ "name,DENSE_RANK",
+ "\nAlice,1",
+ "\nCharlie,1",
+ "\nFrank,1",
+ "\nHenry,2",
+ "\nBob,3",
+ "\nEve,3",
+ "\nDavid,4",
+ "\nGrace,4",
+ ),
+ )
+}
+
+#[test]
+fn win_lag() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, LAG(salary) OVER (ORDER BY salary) FROM employees",
+ concat!(
+ "name,LAG",
+ "\nHenry,NULL",
+ "\nBob,45000",
+ "\nEve,55000",
+ "\nCharlie,60000",
+ "\nAlice,65000",
+ "\nFrank,70000",
+ "\nDavid,75000",
+ "\nGrace,80000",
+ ),
+ )
+}
+
+#[test]
+fn win_lead() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, LEAD(salary) OVER (ORDER BY salary) FROM employees",
+ concat!(
+ "name,LEAD",
+ "\nHenry,55000",
+ "\nBob,60000",
+ "\nEve,65000",
+ "\nCharlie,70000",
+ "\nAlice,75000",
+ "\nFrank,80000",
+ "\nDavid,90000",
+ "\nGrace,NULL",
+ ),
+ )
+}
+
+#[test]
+fn case_in_where() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE CASE WHEN age > 30 THEN 1 ELSE 0 END = 1",
+ concat!("name", "\nCharlie", "\nDavid", "\nFrank", "\nGrace",),
+ )
+}
+
+#[test]
+fn cast_in_where() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE CAST(salary AS INTEGER) > 70000",
+ concat!("name", "\nDavid", "\nFrank", "\nGrace",),
+ )
+}
+
+#[test]
+fn alias_column() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name AS who, salary AS pay FROM employees",
+ concat!(
+ "who,pay",
+ "\nAlice,70000",
+ "\nBob,55000",
+ "\nCharlie,65000",
+ "\nDavid,80000",
+ "\nEve,60000",
+ "\nFrank,75000",
+ "\nGrace,90000",
+ "\nHenry,45000",
+ ),
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Outer joins WITH unmatched rows on both sides.
+//
+// The users/orders fixtures cannot exercise NULL-fill: every user has an order
+// and every order has a user, so LEFT/RIGHT/FULL all degenerate to INNER there.
+// customers/purchases has an unmatched row on each side (Cara has no purchase;
+// purchase 13 references customer 99), so these pin the actual outer-join
+// semantics.
+// ---------------------------------------------------------------------------
+
+#[test]
+fn outer_join_inner_baseline() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/customers.csv", "tests/data/purchases.csv"],
+ "SELECT customers.name, purchases.item FROM customers INNER JOIN purchases ON customers.id = purchases.customer_id",
+ concat!(
+ "customers.name,purchases.item",
+ "\nAnn,Book",
+ "\nAnn,Pen",
+ "\nBen,Desk",
+ ),
+ )
+}
+
+#[test]
+fn outer_join_left_nullfill() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/customers.csv", "tests/data/purchases.csv"],
+ "SELECT customers.name, purchases.item FROM customers LEFT JOIN purchases ON customers.id = purchases.customer_id",
+ concat!(
+ "customers.name,purchases.item",
+ "\nAnn,Book",
+ "\nAnn,Pen",
+ "\nBen,Desk",
+ "\nCara,NULL",
+ ),
+ )
+}
+
+#[test]
+fn outer_join_right_nullfill() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/customers.csv", "tests/data/purchases.csv"],
+ "SELECT customers.name, purchases.item FROM customers RIGHT JOIN purchases ON customers.id = purchases.customer_id",
+ concat!(
+ "customers.name,purchases.item",
+ "\nAnn,Book",
+ "\nAnn,Pen",
+ "\nBen,Desk",
+ "\nNULL,Orphan",
+ ),
+ )
+}
+
+#[test]
+fn outer_join_full_nullfill() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/customers.csv", "tests/data/purchases.csv"],
+ "SELECT customers.name, purchases.item FROM customers FULL JOIN purchases ON customers.id = purchases.customer_id",
+ concat!(
+ "customers.name,purchases.item",
+ "\nAnn,Book",
+ "\nAnn,Pen",
+ "\nBen,Desk",
+ "\nCara,NULL",
+ "\nNULL,Orphan",
+ ),
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Error paths.
+//
+// Only two tests in the pre-existing suite assert a failure at all, so nothing
+// pinned which inputs must be REJECTED. These do -- and they matter especially
+// after an AST upgrade, where an unhandled node can quietly start being
+// accepted (or a supported one start being refused).
+// ---------------------------------------------------------------------------
+
+use crate::helpers::assert_after_write;
+use crate::helpers::assert_query_fails;
+
+#[test]
+fn err_unknown_column() -> Result<(), Box> {
+ assert_query_fails(
+ &["tests/data/employees.csv"],
+ "SELECT nosuchcolumn FROM employees",
+ "not found",
+ )
+}
+
+#[test]
+fn err_unknown_table() -> Result<(), Box> {
+ assert_query_fails(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM nosuchtable",
+ "not found",
+ )
+}
+
+#[test]
+fn err_syntax() -> Result<(), Box> {
+ assert_query_fails(
+ &["tests/data/employees.csv"],
+ "SELECT FROM WHERE",
+ "Failed to execute SQL",
+ )
+}
+
+#[test]
+fn err_group_by_all_rejected() -> Result<(), Box> {
+ // GROUP BY ALL parses but carries no key list. It must be rejected rather
+ // than silently treated as "no GROUP BY".
+ assert_query_fails(
+ &["tests/data/employees.csv"],
+ "SELECT department, COUNT(*) FROM employees GROUP BY ALL",
+ "GROUP BY ALL is not supported",
+ )
+}
+
+#[test]
+fn err_order_by_all_rejected() -> Result<(), Box> {
+ // Likewise ORDER BY ALL must not be silently ignored. Under the dialect
+ // sqawk parses with, `ALL` here comes through as an ordinary identifier
+ // rather than as OrderByKind::All, so this is rejected as an unknown
+ // column instead of by the explicit guard in compile_query. Either route
+ // is acceptable; silently returning unsorted rows is not, which is what
+ // this test pins.
+ assert_query_fails(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees ORDER BY ALL",
+ "not found",
+ )
+}
+
+#[test]
+fn err_semi_join_rejected() -> Result<(), Box> {
+ // Join kinds sqawk cannot execute must produce a clear error, not a
+ // silently wrong join. This is the arm that a sqlparser variant rename
+ // would otherwise slip past.
+ assert_query_fails(
+ &["tests/data/users.csv", "tests/data/orders.csv"],
+ "SELECT users.name FROM users LEFT SEMI JOIN orders ON users.id = orders.user_id",
+ "not supported",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// WHERE combined with LIMIT/OFFSET.
+//
+// The suite tested WHERE and tested LIMIT, but never the two together -- and
+// the interaction was broken: the WHERE-false jump was computed as a fixed
+// offset from the test, which landed on the LIMIT counter's DecrJumpZero. Every
+// row that FAILED the filter still decremented the limit, so LIMIT counted rows
+// scanned rather than rows returned, silently truncating results.
+// ---------------------------------------------------------------------------
+
+#[test]
+fn where_with_limit_not_reached() -> Result<(), Box> {
+ // Limit is larger than the match count: every match must be returned.
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE department = 'Engineering' LIMIT 5",
+ "name\nAlice\nCharlie\nFrank",
+ )
+}
+
+#[test]
+fn where_with_limit_late_matches() -> Result<(), Box> {
+ // Both matches come after three non-matching rows, so a limit that counts
+ // scanned rows would drop Grace.
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE department = 'Sales' LIMIT 5",
+ "name\nDavid\nGrace",
+ )
+}
+
+#[test]
+fn where_with_limit_reached() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE department = 'Engineering' LIMIT 2",
+ "name\nAlice\nCharlie",
+ )
+}
+
+#[test]
+fn where_with_limit_offset() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE department = 'Engineering' LIMIT 5 OFFSET 1",
+ "name\nCharlie\nFrank",
+ )
+}
+
+#[test]
+fn where_with_offset_only() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE age > 25 OFFSET 2",
+ "name\nDavid\nEve\nFrank\nGrace",
+ )
+}
+
+/// DISTINCT rides on the aggregate function type as a flag bit, so it composes
+/// with every aggregate rather than only COUNT. repeats.csv holds 10,10,20,30,30.
+#[test]
+fn aggregate_distinct_variants() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/repeats.csv"],
+ "SELECT SUM(v), SUM(DISTINCT v), COUNT(v), COUNT(DISTINCT v), \
+ MIN(DISTINCT v), MAX(DISTINCT v), AVG(DISTINCT v) FROM repeats",
+ "SUM,SUM,COUNT,COUNT,MIN,MAX,AVG\n100,60,5,3,10,30,20",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// Expression position symmetry.
+//
+// SQL draws no distinction between an expression used as a filter and the same
+// expression used as a projected value, but sqawk had three parallel
+// expression compilers and each supported a different subset. These pin both
+// positions for the same node kinds so the two cannot drift apart again.
+// ---------------------------------------------------------------------------
+
+#[test]
+fn case_as_projected_value() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, CASE WHEN age > 30 THEN 'old' ELSE 'young' END FROM employees LIMIT 3",
+ "name,expr\nAlice,young\nBob,young\nCharlie,old",
+ )
+}
+
+#[test]
+fn cast_as_projected_value() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT CAST(salary AS TEXT) FROM employees LIMIT 2",
+ "expr\n70000\n55000",
+ )
+}
+
+#[test]
+fn is_null_as_projected_value() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/nullable.csv"],
+ "SELECT id, score IS NULL FROM nullable",
+ "id,expr\n1,0\n2,1\n3,0\n4,0\n5,0",
+ )
+}
+
+#[test]
+fn like_as_projected_value() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, name LIKE 'A%' FROM employees LIMIT 2",
+ "name,expr\nAlice,1\nBob,0",
+ )
+}
+
+#[test]
+fn between_as_projected_value() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name, age BETWEEN 30 AND 40 FROM employees LIMIT 3",
+ "name,expr\nAlice,1\nBob,0\nCharlie,1",
+ )
+}
+
+#[test]
+fn trim_as_filter() -> Result<(), Box> {
+ // TRIM/CEIL/FLOOR were handled only in the projection path, so they worked
+ // in a SELECT list and failed in WHERE -- the mirror image of SUBSTR.
+ assert_query(
+ &["tests/data/strings.csv"],
+ "SELECT id FROM strings WHERE TRIM(padded_text) = 'trimme'",
+ "id\n1",
+ )
+}
+
+#[test]
+fn ceil_as_filter() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE CEIL(salary) > 85000",
+ "name\nGrace",
+ )
+}
+
+#[test]
+fn concat_operator_as_filter() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT name FROM employees WHERE name || '!' = 'Alice!'",
+ "name\nAlice",
+ )
+}
+
+// ---------------------------------------------------------------------------
+// GROUP BY interactions that were previously unrepresentable.
+// ---------------------------------------------------------------------------
+
+#[test]
+fn group_by_with_where_and_having() -> Result<(), Box> {
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT department, COUNT(*) FROM employees WHERE age > 25 GROUP BY department \
+ HAVING COUNT(*) > 1",
+ // age > 25 drops Bob(25) and Henry(22), leaving Marketing with one
+ // row, which HAVING then removes.
+ "department,COUNT\nEngineering,3\nSales,2",
+ )
+}
+
+#[test]
+fn group_by_three_columns() -> Result<(), Box> {
+ // Every key column must participate; a difference in the last one still
+ // starts a new group.
+ assert_query(
+ &["tests/data/employees.csv"],
+ "SELECT department, role, COUNT(*) FROM employees GROUP BY department, role",
+ "department,role,COUNT\n\
+ Engineering,Developer,2\n\
+ Engineering,Manager,1\n\
+ HR,Intern,1\n\
+ Marketing,Analyst,1\n\
+ Marketing,Specialist,1\n\
+ Sales,Director,2",
+ )
+}
+
+#[test]
+fn group_by_with_where_excluding_a_whole_group() -> Result<(), Box