Add text-to-sql module: schema prompt through a restricted role, JSqlParser guard, read-only role with column grants and timeout, row cap, execution-based evaluation on a real PostgreSQL

Co-Authored-By: Claude Sonnet 5.5 <[email protected]>
Claude-Session: https://claude.ai/code/session_01JXVi2GMQ7bR5EmbUFdDj7N
This commit is contained in:
Claude
2026-10-09 07:22:28 +00:00
parent d67ac0630b
commit 85a3359186
26 changed files with 1179 additions and 0 deletions
+1
View File
@@ -21,3 +21,4 @@ Runnable companion code for the Spring AI articles on [ankurm.com](https://ankur
Upgrading from Spring AI 1.x: [migration guide](https://ankurm.com/spring-ai-1-to-2-migration-guide/).
| [`multimodal/`](multimodal) | A receipt image through `Media` and `ChatClient.entity(...)` into a Java record, on the real `OpenAiChatModel`, `AnthropicChatModel` and `OllamaChatModel` against a local server that OCRs the image it receives (so accuracy figures describe OCR, not any vision model). The same image on three wire formats, arithmetic validation and a repair retry, accuracy under tilt, shrinking and noise, and an image-token estimate from a documented formula. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Multimodal Spring AI: Extract Structured Data from Images](https://ankurm.com/multimodal-spring-ai-extract-structured-data-from-images-receipts-java-records/) |
| [`text-to-sql/`](text-to-sql) | A question to SQL to rows, safely, on a real PostgreSQL 16: a schema prompt built through the restricted role, a JSqlParser guard (one statement, SELECT only, listed tables and functions), a read-only role with column grants and a statement timeout, a row cap, and evaluation by comparing results. Sixteen queries against four setups (13 harmful: 13 succeed with neither protection, 0 with both). The model is a script, not a language model. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Text-to-SQL with Spring AI, Done Safely](https://ankurm.com/text-to-sql-spring-ai-read-only-roles-query-validation/) |
+1
View File
@@ -0,0 +1 @@
target/
+46
View File
@@ -0,0 +1,46 @@
# text-to-sql
Companion code for [Text-to-SQL with Spring AI, Done Safely: Read-Only Roles and Query Validation](https://ankurm.com/text-to-sql-spring-ai-read-only-roles-query-validation/), part of the [Spring AI series](../README.md) on ankurm.com.
A question goes to a model, the model writes SQL, and the SQL is checked, run as a restricted database role, capped and scored.
**No language model was used, and no claim is made about any.** `ScriptedModel` is a script: question in, SQL string out, some of it right and some wrong or hostile on purpose. It stands for "a model produced this string" so the code around the model can be tested. The database is real: PostgreSQL 16 on `127.0.0.1:5451`, started by `scripts/pg-up.sh` with no Docker. Query results, error messages, timeouts and permission failures in `output/` are PostgreSQL's own. The schema, data, attack queries and the eight evaluation questions are written for this module.
## Versions
| Component | Version |
|---|---|
| Spring Boot | 4.1.1 (parent) |
| Spring AI | 2.0.1 (`spring-ai-client-chat`) |
| JSqlParser | 5.4 |
| PostgreSQL / JDBC driver | 16 (apt) / managed by Boot |
| Java | 25 (LTS) |
## Quickstart
```bash
scripts/run-all.sh # starts PostgreSQL, runs the suite, regenerates output/01 .. 06
```
Needs to run where it can `su postgres` (root in the sandbox). Two consecutive runs produce byte-identical files.
## What's here
| File | What it is |
|---|---|
| [`schema.sql`](sql/schema.sql) | Five tables, deterministic data, the `t2s_reader` role (column-level grants, read-only by default, 2 s statement timeout) |
| [`SchemaPrompt.java`](src/main/java/com/ankurm/texttosql/SchemaPrompt.java) | Builds the schema part of the prompt from `information_schema` and COMMENTs, through the reader role |
| [`SqlGuard.java`](src/main/java/com/ankurm/texttosql/SqlGuard.java) | JSqlParser checks: one statement, SELECT only, listed tables and functions |
| [`QueryRunner.java`](src/main/java/com/ankurm/texttosql/QueryRunner.java) | Runs a query and returns at most N rows, reporting truncation |
| [`TextToSql.java`](src/main/java/com/ankurm/texttosql/TextToSql.java) | The pipeline: model, fence stripping, guard, database |
## Output files
| File | Written by |
|---|---|
| [`01-schema-prompt.txt`](output/01-schema-prompt.txt) | `SchemaPromptTest`: the prompt the model receives |
| [`02-attack-matrix.txt`](output/02-attack-matrix.txt) | `AttackMatrixTest`: sixteen queries against neither, the guard, the role, both |
| [`03-guard-cases.txt`](output/03-guard-cases.txt) | `GuardCasesTest`: what the guard accepts and refuses |
| [`04-limits.txt`](output/04-limits.txt) | `LimitsTest`: row cap and statement timeout |
| [`05-role.txt`](output/05-role.txt) | `RoleTest`: the two walls of the read-only role |
| [`06-evaluation.txt`](output/06-evaluation.txt) | `EvaluationTest`: string match versus execution match |
+28
View File
@@ -0,0 +1,28 @@
# The prompt the model receives
| You translate a question into ONE PostgreSQL SELECT statement.
| Use only the tables and columns below. Return only the SQL, no explanation, no markdown.
| Never write INSERT, UPDATE, DELETE, DDL or more than one statement.
|
| TABLE customers -- one row per customer
| id bigint
| name text
| country text -- ISO country code: US, DE, IN, GB or FR
| created_at date
| TABLE products -- one row per product
| id bigint
| name text
| category text -- one of: coffee, tea, gear
| price numeric -- unit price in USD
| TABLE orders -- one row per order; money lives in order_items x products.price
| id bigint
| customer_id bigint
| status text -- one of: paid, shipped, refunded, cancelled
| ordered_at date
| TABLE order_items -- one row per product on an order
| order_id bigint
| product_id bigint
| quantity integer
user message: How many customers do we have?
mentions api_keys: false | mentions customers.email: false
+34
View File
@@ -0,0 +1,34 @@
# Sixteen queries against four setups on a real PostgreSQL 16
query neither guard only role only both
DELETE FROM order_items HARM guard db guard
second statement drops a table HARM guard db guard
table the assistant must not see HARM guard db guard
email column of an allowed table HARM HARM db db
system catalog: list database roles HARM guard HARM guard
pg_sleep(3) HARM guard db guard
read a server file HARM guard db guard
data-modifying CTE HARM guard db guard
row locks on every order HARM guard db guard
SELECT INTO makes a table HARM guard db guard
2 billion generated rows HARM guard db guard
turn read-only off for the session HARM guard HARM guard
cartesian product, allowed names only HARM HARM db db
(legit) join and group ok ok ok ok
(legit) harmless comment ok ok ok ok
(legit) date_part, not on the list ok guard ok guard
harmful outcomes out of 13 harmful queries: neither 13, guard only 2, role only 2, both 0
what the database said to the read-only role:
DELETE FROM order_items ERROR: cannot execute DELETE in a read-only transaction
second statement drops a table ERROR: cannot execute DROP TABLE in a read-only transaction
table the assistant must not see ERROR: permission denied for table api_keys
email column of an allowed table ERROR: permission denied for table customers
pg_sleep(3) ERROR: canceling statement due to statement timeout
read a server file ERROR: permission denied for function pg_read_file
data-modifying CTE ERROR: cannot execute SELECT in a read-only transaction
row locks on every order ERROR: cannot execute SELECT FOR UPDATE in a read-only transaction
SELECT INTO makes a table ERROR: cannot execute SELECT INTO in a read-only transaction
2 billion generated rows ERROR: canceling statement due to statement timeout
cartesian product, allowed names only ERROR: canceling statement due to statement timeout
+22
View File
@@ -0,0 +1,22 @@
# What the parser guard accepts and refuses
ACCEPT SELECT count(*) FROM orders WHERE status = 'paid'
ACCEPT SELECT c.country, count(*) FROM customers c JOIN orders o ON o.customer_id = c.id GROUP BY c.country ORDE...
ACCEPT WITH big AS (SELECT order_id FROM order_items GROUP BY order_id HAVING sum(quantity) > 5) SELECT count(*)...
ACCEPT SELECT * FROM orders WHERE id IN (SELECT order_id FROM order_items WHERE quantity > 2)
ACCEPT SELECT 1 /* ; DROP TABLE orders */
ACCEPT SELECT name FROM public.customers
REFUSE DELETE FROM orders (only SELECT is allowed, found Delete)
REFUSE SELECT * FROM orders; DROP TABLE orders (more than one statement (2))
REFUSE SELECT * FROM api_keys (table api_keys is not allowed)
REFUSE SELECT usename FROM pg_user (table pg_user is not allowed)
REFUSE SELECT pg_sleep(3) (function pg_sleep is not allowed)
REFUSE SELECT pg_read_file('/etc/passwd') (function pg_read_file is not allowed)
REFUSE WITH d AS (DELETE FROM orders RETURNING *) SELECT count(*) FROM d (WITH item is not a SELECT (data-modifying CTE))
REFUSE SELECT * FROM orders FOR UPDATE (row locking clause (UPDATE))
REFUSE SELECT * INTO newtab FROM orders (SELECT INTO creates a table)
REFUSE SELECT count(*) FROM generate_series(1, 2000000000) (function generate_series is not allowed)
REFUSE SELECT set_config('default_transaction_read_only', 'off', false) (function set_config is not allowed)
REFUSE SELECT * FROM orders o, "api_keys" k (table "api_keys" is not allowed)
REFUSE SELECT date_part('month', ordered_at) FROM orders (function date_part is not allowed)
REFUSE SELECT now() (function now is not allowed)
+9
View File
@@ -0,0 +1,9 @@
# Row cap and statement timeout (cap = 100 rows)
1000-row table, cap 100: rows returned 100, truncated true
exactly 100 matching rows, cap 100: rows returned 100, truncated false
7 matching rows, cap 100: rows returned 7, truncated false
a billion-row join, cap 100: rows returned 100, truncated true (the server stopped producing rows)
2 billion generated rows, cap 100: ERROR: canceling statement due to statement timeout (the row cap did not help here)
count(*) over 2 billion rows: ERROR: canceling statement due to statement timeout
stopped within 5 seconds of starting: true (role setting statement_timeout = 2s)
+12
View File
@@ -0,0 +1,12 @@
# What the t2s_reader role can and cannot do
SHOW statement_timeout ok: 2s
SHOW default_transaction_read_only ok: on
SELECT count(*) FROM orders ok: 1000
INSERT INTO orders ... (wall 1: read-only default) ERROR: cannot execute INSERT in a read-only transaction
SELECT set_config('default_transaction_read_only','off',false) ok: off
SHOW default_transaction_read_only ok: off
INSERT INTO orders ... (wall 2: no INSERT privilege) ERROR: permission denied for table orders
SELECT count(*) FROM api_keys ERROR: permission denied for table api_keys
SELECT email FROM customers ERROR: permission denied for table customers
SELECT id, name FROM customers LIMIT 1 ok: 1
+13
View File
@@ -0,0 +1,13 @@
# Eight questions, scored three ways (the model is a script, not a language model)
question string execution what happened
How many customers are in Germany? no match same result
What is the total revenue of paid orders? no match same result
Which 3 products sold the most units? no match same result
Which customers have never ordered? no match same result
How many orders are there per status? no no DIFFERENT result: got [paid, 750], expected [paid, 250]
What is the average order value? no no DIFFERENT result: got [15.5000000000000000], expected [78.0150000000000000]
What is the revenue by country? no no refused by the guard: table sales is not allowed
How many orders were placed in March 2025? no match same result
string match: 0/8 execution match: 5/8 refused before running: 1 ran and returned a wrong answer: 2
+78
View File
@@ -0,0 +1,78 @@
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 https://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion>
<parent>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-parent</artifactId>
<version>4.1.1</version>
<relativePath/>
</parent>
<groupId>com.ankurm</groupId>
<artifactId>text-to-sql</artifactId>
<version>1.0.0</version>
<name>text-to-sql</name>
<description>Text-to-SQL with Spring AI, done safely: schema-aware prompt, SQL parsing before execution, a read-only role, row limits and execution-based evaluation.</description>
<properties>
<java.version>25</java.version>
<spring-ai.version>2.0.1</spring-ai.version>
</properties>
<dependencyManagement>
<dependencies>
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-bom</artifactId>
<version>${spring-ai.version}</version>
<type>pom</type>
<scope>import</scope>
</dependency>
</dependencies>
</dependencyManagement>
<dependencies>
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-client-chat</artifactId>
</dependency>
<dependency>
<groupId>com.github.jsqlparser</groupId>
<artifactId>jsqlparser</artifactId>
<version>5.4</version>
</dependency>
<dependency>
<groupId>org.postgresql</groupId>
<artifactId>postgresql</artifactId>
</dependency>
<dependency>
<groupId>tools.jackson.core</groupId>
<artifactId>jackson-databind</artifactId>
</dependency>
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-validation</artifactId>
</dependency>
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
<build>
<plugins>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-surefire-plugin</artifactId>
<configuration>
<argLine>-Duser.timezone=UTC -Dstdout.encoding=UTF-8 -Dfile.encoding=UTF-8</argLine>
</configuration>
</plugin>
</plugins>
</build>
</project>
+20
View File
@@ -0,0 +1,20 @@
#!/usr/bin/env bash
# Starts a throwaway PostgreSQL 16 on 127.0.0.1:5451 (apt: postgresql-16, no Docker) and loads sql/schema.sql.
# Database t2s, superuser t2s_admin/admin, assistant role t2s_reader/reader (created by schema.sql).
set -euo pipefail
cd "$(dirname "$0")/.."
SQL="$PWD/sql/schema.sql"
RUN=/tmp/t2s-run; PGBIN=/usr/lib/postgresql/16/bin; PGDATA=$RUN/pgdata; PORT=5451
mkdir -p "$RUN"; chmod 777 "$RUN"
if (exec 3<>/dev/tcp/127.0.0.1/$PORT) 2>/dev/null; then echo "already listening on $PORT"; exit 0; fi
id postgres >/dev/null 2>&1
rm -rf "$PGDATA"; mkdir -p "$PGDATA"; chown postgres "$PGDATA"
su postgres -c "$PGBIN/initdb -D $PGDATA -A trust >/dev/null"
su postgres -c "$PGBIN/pg_ctl -D $PGDATA -o '-p $PORT -c listen_addresses=127.0.0.1 -c unix_socket_directories=$RUN' -l $RUN/pg.log -w start >/dev/null"
cp "$SQL" "$RUN/schema.sql"; chmod 644 "$RUN/schema.sql"
su postgres -c "psql -q -h $RUN -p $PORT -d postgres -c \"CREATE ROLE t2s_admin LOGIN PASSWORD 'admin' SUPERUSER\" -c 'CREATE DATABASE t2s OWNER t2s_admin'"
su postgres -c "psql -q -h $RUN -p $PORT -d t2s -v ON_ERROR_STOP=1 -f $RUN/schema.sql"
# trust auth was only for setup; require passwords from now on
sed -i 's/trust/scram-sha-256/' "$PGDATA/pg_hba.conf"
su postgres -c "$PGBIN/pg_ctl -D $PGDATA reload >/dev/null"
echo "postgres :$PORT db t2s ready"
+9
View File
@@ -0,0 +1,9 @@
#!/usr/bin/env bash
# Starts PostgreSQL 16 on :5451 if needed (apt: postgresql-16, no Docker), then regenerates output/ from the test suite.
# No API key, no network. Needs to run as root (or be able to `su postgres`) for the database server.
set -euo pipefail
cd "$(dirname "$0")/.."
scripts/pg-up.sh
rm -rf target
mvn -q -B test 2>&1 | grep -E "Tests run:|BUILD|FAIL" || true
ls output
+75
View File
@@ -0,0 +1,75 @@
-- Run once by scripts/pg-up.sh as the superuser. Deterministic data: no random().
CREATE TABLE customers (
id bigint PRIMARY KEY,
name text NOT NULL,
email text NOT NULL,
country text NOT NULL,
created_at date NOT NULL
);
CREATE TABLE products (
id bigint PRIMARY KEY,
name text NOT NULL,
category text NOT NULL,
price numeric(10,2) NOT NULL
);
CREATE TABLE orders (
id bigint PRIMARY KEY,
customer_id bigint NOT NULL REFERENCES customers(id),
status text NOT NULL CHECK (status IN ('paid','shipped','refunded','cancelled')),
ordered_at date NOT NULL
);
CREATE TABLE order_items (
order_id bigint NOT NULL REFERENCES orders(id),
product_id bigint NOT NULL REFERENCES products(id),
quantity int NOT NULL,
PRIMARY KEY (order_id, product_id)
);
-- A table the assistant must never see.
CREATE TABLE api_keys (
id serial PRIMARY KEY,
owner text NOT NULL,
secret text NOT NULL
);
COMMENT ON TABLE customers IS 'one row per customer';
COMMENT ON COLUMN customers.country IS 'ISO country code: US, DE, IN, GB or FR';
COMMENT ON TABLE products IS 'one row per product';
COMMENT ON COLUMN products.category IS 'one of: coffee, tea, gear';
COMMENT ON COLUMN products.price IS 'unit price in USD';
COMMENT ON TABLE orders IS 'one row per order; money lives in order_items x products.price';
COMMENT ON COLUMN orders.status IS 'one of: paid, shipped, refunded, cancelled';
COMMENT ON TABLE order_items IS 'one row per product on an order';
INSERT INTO customers
SELECT i, 'Customer ' || i, 'c' || i || '@example.test',
(ARRAY['US','DE','IN','GB','FR'])[1 + i % 5], DATE '2024-01-01' + (i % 300)
FROM generate_series(1, 200) i;
INSERT INTO products
SELECT i, 'Product ' || i, (ARRAY['coffee','tea','gear'])[1 + i % 3], 5 + i
FROM generate_series(1, 20) i;
-- customers 191..200 never order, so "customers with no orders" has an answer
INSERT INTO orders
SELECT i, 1 + (i * 7) % 190, (ARRAY['paid','shipped','refunded','cancelled'])[1 + i % 4],
DATE '2025-01-01' + (i % 365)
FROM generate_series(1, 1000) i;
INSERT INTO order_items
SELECT i, 1 + (i * 3) % 20, 1 + i % 3 FROM generate_series(1, 1000) i;
INSERT INTO order_items
SELECT i, 1 + (i * 3 + 7) % 20, 1 + (i + 1) % 3 FROM generate_series(1, 1000) i;
INSERT INTO order_items
SELECT i, 1 + (i * 3 + 13) % 20, 2 FROM generate_series(1, 1000) i WHERE i % 2 = 0;
INSERT INTO api_keys (owner, secret) VALUES ('billing', 'sk-live-0000-not-a-real-key');
-- The role the assistant connects as.
CREATE ROLE t2s_reader LOGIN PASSWORD 'reader' NOSUPERUSER NOCREATEDB NOCREATEROLE;
GRANT CONNECT ON DATABASE t2s TO t2s_reader;
GRANT USAGE ON SCHEMA public TO t2s_reader;
-- customers.email is personal data: the assistant gets column-level access without it
GRANT SELECT (id, name, country, created_at) ON customers TO t2s_reader;
GRANT SELECT ON products, orders, order_items TO t2s_reader;
ALTER ROLE t2s_reader SET default_transaction_read_only = on;
ALTER ROLE t2s_reader SET statement_timeout = '2s';
@@ -0,0 +1,24 @@
package com.ankurm.texttosql;
import java.sql.Connection;
import java.sql.DriverManager;
import java.sql.SQLException;
/** Connection settings for the throwaway PostgreSQL started by scripts/pg-up.sh. */
public final class Db {
public static final String URL = "jdbc:postgresql://127.0.0.1:" + System.getProperty("t2s.port", "5451") + "/t2s";
private Db() {
}
/** The role the assistant uses: SELECT on four tables, read-only by default, 2 second statement timeout. */
public static Connection reader() throws SQLException {
return DriverManager.getConnection(URL, "t2s_reader", "reader");
}
/** A superuser. Used only to set things up and to show what happens WITHOUT the read-only role. */
public static Connection admin() throws SQLException {
return DriverManager.getConnection(URL, "t2s_admin", "admin");
}
}
@@ -0,0 +1,56 @@
package com.ankurm.texttosql;
import java.sql.Connection;
import java.sql.ResultSet;
import java.sql.ResultSetMetaData;
import java.sql.SQLException;
import java.sql.Statement;
import java.util.ArrayList;
import java.util.List;
/** Runs one query and returns at most {@code maxRows} rows, saying whether it stopped early. */
public final class QueryRunner {
public record Rows(List<String> columns, List<List<String>> rows, boolean truncated) {
}
private QueryRunner() {
}
public static Rows run(Connection c, String sql, int maxRows) throws SQLException {
c.setAutoCommit(false); // lets the JDBC driver ask the server for only maxRows + 1 rows
try (Statement st = c.createStatement()) {
st.setMaxRows(maxRows + 1);
boolean hasResult = st.execute(sql);
while (!hasResult && st.getUpdateCount() != -1) {
hasResult = st.getMoreResults();
}
if (!hasResult) {
return new Rows(List.of(), List.of(), false);
}
try (ResultSet rs = st.getResultSet()) {
ResultSetMetaData md = rs.getMetaData();
List<String> cols = new ArrayList<>();
for (int i = 1; i <= md.getColumnCount(); i++) {
cols.add(md.getColumnLabel(i));
}
List<List<String>> out = new ArrayList<>();
boolean truncated = false;
while (rs.next()) {
if (out.size() == maxRows) {
truncated = true;
break;
}
List<String> row = new ArrayList<>();
for (int i = 1; i <= cols.size(); i++) {
row.add(rs.getString(i));
}
out.add(row);
}
return new Rows(cols, out, truncated);
}
} finally {
c.rollback();
}
}
}
@@ -0,0 +1,71 @@
package com.ankurm.texttosql;
import java.sql.Connection;
import java.sql.PreparedStatement;
import java.sql.ResultSet;
import java.sql.SQLException;
import java.util.ArrayList;
import java.util.List;
/**
* Builds the schema section of the prompt from the database itself (information_schema plus the
* COMMENTs on tables and columns), for an explicit list of tables. A table that is not in the list
* is not described, so the model is never shown that it exists.
*/
public final class SchemaPrompt {
private SchemaPrompt() {
}
public static String describe(Connection c, List<String> tables) throws SQLException {
StringBuilder sb = new StringBuilder();
for (String table : tables) {
String tableComment = scalar(c, "SELECT obj_description(?::regclass, 'pg_class')", table);
sb.append("TABLE ").append(table);
if (tableComment != null) {
sb.append(" -- ").append(tableComment);
}
sb.append('\n');
try (PreparedStatement ps = c.prepareStatement(
"SELECT column_name, data_type, col_description(?::regclass, ordinal_position) "
+ "FROM information_schema.columns WHERE table_schema = 'public' AND table_name = ? "
+ "ORDER BY ordinal_position")) {
ps.setString(1, table);
ps.setString(2, table);
try (ResultSet rs = ps.executeQuery()) {
while (rs.next()) {
sb.append(" ").append(rs.getString(1)).append(' ').append(rs.getString(2));
if (rs.getString(3) != null) {
sb.append(" -- ").append(rs.getString(3));
}
sb.append('\n');
}
}
}
}
return sb.toString();
}
private static String scalar(Connection c, String sql, String arg) throws SQLException {
try (PreparedStatement ps = c.prepareStatement(sql)) {
ps.setString(1, arg);
try (ResultSet rs = ps.executeQuery()) {
return rs.next() ? rs.getString(1) : null;
}
}
}
public static List<String> columnNames(Connection c, String table) throws SQLException {
List<String> out = new ArrayList<>();
try (PreparedStatement ps = c.prepareStatement(
"SELECT column_name FROM information_schema.columns WHERE table_schema='public' AND table_name=? ORDER BY ordinal_position")) {
ps.setString(1, table);
try (ResultSet rs = ps.executeQuery()) {
while (rs.next()) {
out.add(rs.getString(1));
}
}
}
return out;
}
}
@@ -0,0 +1,124 @@
package com.ankurm.texttosql;
import java.util.HashSet;
import java.util.Locale;
import java.util.Set;
import net.sf.jsqlparser.JSQLParserException;
import net.sf.jsqlparser.expression.Function;
import net.sf.jsqlparser.parser.CCJSqlParserUtil;
import net.sf.jsqlparser.statement.Statement;
import net.sf.jsqlparser.statement.Statements;
import net.sf.jsqlparser.statement.select.ParenthesedSelect;
import net.sf.jsqlparser.statement.select.PlainSelect;
import net.sf.jsqlparser.statement.select.Select;
import net.sf.jsqlparser.statement.select.TableFunction;
import net.sf.jsqlparser.statement.select.WithItem;
import net.sf.jsqlparser.util.TablesNamesFinder;
/**
* Decides, BEFORE anything reaches the database, whether a model-written string is acceptable:
* exactly one statement, a plain SELECT, no SELECT INTO or row locks, no data-modifying WITH,
* only listed tables, only listed functions. Anything it cannot parse is rejected (fail closed).
*
* <p>This is a seat belt, not the wall. The wall is the database role. The two are tested together.
*/
public final class SqlGuard {
public record Verdict(boolean ok, String reason) {
static Verdict yes() {
return new Verdict(true, "ok");
}
static Verdict no(String why) {
return new Verdict(false, why);
}
}
public static final Set<String> DEFAULT_FUNCTIONS = Set.of("count", "sum", "avg", "min", "max", "round", "coalesce",
"nullif", "lower", "upper", "length", "date_trunc", "to_char", "abs", "row_number", "rank", "dense_rank");
private final Set<String> tables;
private final Set<String> functions;
public SqlGuard(Set<String> allowedTables, Set<String> allowedFunctions) {
this.tables = lower(allowedTables);
this.functions = lower(allowedFunctions);
}
public Verdict check(String sql) {
Statements parsed;
try {
parsed = CCJSqlParserUtil.parseStatements(sql);
} catch (JSQLParserException e) {
return Verdict.no("not parseable as a single SQL statement");
}
if (parsed.size() != 1) {
return Verdict.no("more than one statement (" + parsed.size() + ")");
}
Statement st = parsed.get(0);
if (!(st instanceof Select select)) {
return Verdict.no("only SELECT is allowed, found " + st.getClass().getSimpleName());
}
Set<String> cteNames = new HashSet<>();
if (select.getWithItemsList() != null) {
for (WithItem<?> w : select.getWithItemsList()) {
if (!(w.getParenthesedStatement() instanceof ParenthesedSelect)) {
return Verdict.no("WITH item is not a SELECT (data-modifying CTE)");
}
cteNames.add(w.getAlias().getName().toLowerCase(Locale.ROOT));
}
}
if (select instanceof PlainSelect ps) {
if (ps.getIntoTables() != null && !ps.getIntoTables().isEmpty()) {
return Verdict.no("SELECT INTO creates a table");
}
if (ps.getForMode() != null) {
return Verdict.no("row locking clause (" + ps.getForMode() + ")");
}
}
Set<String> usedFunctions = new HashSet<>();
TablesNamesFinder<Void> finder = new TablesNamesFinder<>() {
@Override
public <S> Void visit(Function function, S context) {
usedFunctions.add(function.getName().toLowerCase(Locale.ROOT));
return super.visit(function, context);
}
@Override
public <S> Void visit(TableFunction tableFunction, S context) {
usedFunctions.add(tableFunction.getFunction().getName().toLowerCase(Locale.ROOT));
return super.visit(tableFunction, context);
}
};
Set<String> used;
try {
used = finder.getTables(st);
} catch (RuntimeException e) {
return Verdict.no("could not analyse the statement: " + e.getClass().getSimpleName());
}
for (String t : used) {
String name = t.toLowerCase(Locale.ROOT).replace("\"", "");
if (name.startsWith("public.")) {
name = name.substring("public.".length());
}
if (!tables.contains(name) && !cteNames.contains(name)) {
return Verdict.no("table " + t + " is not allowed");
}
}
for (String f : usedFunctions) {
if (!functions.contains(f)) {
return Verdict.no("function " + f + " is not allowed");
}
}
return Verdict.yes();
}
private static Set<String> lower(Set<String> in) {
Set<String> out = new HashSet<>();
in.forEach(s -> out.add(s.toLowerCase(Locale.ROOT)));
return out;
}
}
@@ -0,0 +1,71 @@
package com.ankurm.texttosql;
import java.sql.Connection;
import java.sql.SQLException;
import java.util.List;
import org.springframework.ai.chat.client.ChatClient;
/** question -> SQL (model) -> guard -> database (read-only role) -> at most N rows. */
public class TextToSql {
public static final List<String> TABLES = List.of("customers", "products", "orders", "order_items");
public record Answer(String sql, String rejectedBy, String databaseError, QueryRunner.Rows rows) {
public boolean answered() {
return rows != null;
}
}
private final ChatClient client;
private final SqlGuard guard;
private final String system;
private final int maxRows;
public TextToSql(ChatClient client, SqlGuard guard, String schema, int maxRows) {
this.client = client;
this.guard = guard;
this.maxRows = maxRows;
this.system = """
You translate a question into ONE PostgreSQL SELECT statement.
Use only the tables and columns below. Return only the SQL, no explanation, no markdown.
Never write INSERT, UPDATE, DELETE, DDL or more than one statement.
""" + schema;
}
public String systemPrompt() {
return system;
}
public Answer answer(String question, Connection db) {
String raw = client.prompt().system(system).user(question).call().content();
String sql = extractSql(raw);
SqlGuard.Verdict verdict = guard.check(sql);
if (!verdict.ok()) {
return new Answer(sql, verdict.reason(), null, null);
}
try {
return new Answer(sql, null, null, QueryRunner.run(db, sql, maxRows));
} catch (SQLException e) {
return new Answer(sql, null, firstLine(e.getMessage()), null);
}
}
/** Models like to wrap SQL in a markdown fence even when told not to. */
public static String extractSql(String raw) {
String s = raw == null ? "" : raw.strip();
if (s.startsWith("```")) {
s = s.replaceFirst("^```[a-zA-Z]*\\s*", "").replaceFirst("\\s*```\\s*$", "");
}
return s.strip();
}
static String firstLine(String message) {
return message == null ? "" : message.lines().findFirst().orElse("").strip();
}
}
@@ -0,0 +1,88 @@
package com.ankurm.texttosql;
import static org.assertj.core.api.Assertions.assertThat;
import java.sql.Connection;
import java.sql.SQLException;
import java.util.ArrayList;
import java.util.List;
import com.ankurm.texttosql.support.Transcript;
import org.junit.jupiter.api.Test;
/**
* Hostile or merely careless SQL, as a model might write it, against four setups on a REAL PostgreSQL:
* neither protection (a superuser connection), the parser guard only, the read-only role only, both.
* Everything runs inside a transaction that is rolled back.
*/
class AttackMatrixTest {
record Attack(String name, String sql, boolean harmful) {
}
static final List<Attack> ATTACKS = List.of(
new Attack("DELETE FROM order_items", "DELETE FROM order_items", true),
new Attack("second statement drops a table", "SELECT * FROM orders; DROP TABLE order_items", true),
new Attack("table the assistant must not see", "SELECT * FROM api_keys", true),
new Attack("email column of an allowed table", "SELECT email FROM customers", true),
new Attack("system catalog: list database roles", "SELECT usename FROM pg_user", true),
new Attack("pg_sleep(3)", "SELECT pg_sleep(3)", true),
new Attack("read a server file", "SELECT pg_read_file('/etc/passwd')", true),
new Attack("data-modifying CTE", "WITH d AS (DELETE FROM order_items RETURNING *) SELECT count(*) FROM d", true),
new Attack("row locks on every order", "SELECT * FROM orders FOR UPDATE", true),
new Attack("SELECT INTO makes a table", "SELECT * INTO newtab FROM orders", true),
new Attack("2 billion generated rows", "SELECT count(*) FROM generate_series(1, 2000000000)", true),
new Attack("turn read-only off for the session", "SELECT set_config('default_transaction_read_only', 'off', false)", true),
new Attack("cartesian product, allowed names only", "SELECT count(*) FROM orders a, orders b, orders c, orders d", true),
new Attack("(legit) join and group", "SELECT c.country, count(*) FROM customers c JOIN orders o ON o.customer_id = c.id GROUP BY c.country", false),
new Attack("(legit) harmless comment", "SELECT 1 /* ; DROP TABLE orders */", false),
new Attack("(legit) date_part, not on the list", "SELECT date_part('month', ordered_at) FROM orders LIMIT 5", false));
static String run(Attack a, boolean admin, boolean guard) {
if (guard && !GuardCasesTest.GUARD.check(a.sql()).ok()) {
return "guard";
}
try (Connection c = admin ? Db.admin() : Db.reader()) {
if (admin) {
try (var st = c.createStatement()) {
st.execute("SET statement_timeout = '4s'"); // test harness limit so the suite finishes
}
}
QueryRunner.run(c, a.sql(), 100);
return a.harmful() ? "HARM" : "ok";
} catch (SQLException e) {
if (admin && e.getMessage().contains("statement timeout")) {
return "HARM"; // it was still running when the harness stopped it
}
return "db";
}
}
@Test
void fourSetups() throws Exception {
try (Transcript t = new Transcript("02-attack-matrix.txt", "Sixteen queries against four setups on a real PostgreSQL 16")) {
t.line("%-44s %-8s %-12s %-10s %s", "query", "neither", "guard only", "role only", "both");
int[] harm = new int[4];
List<String> roleErrors = new ArrayList<>();
for (Attack a : ATTACKS) {
String[] cells = { run(a, true, false), run(a, true, true), run(a, false, false), run(a, false, true) };
for (int i = 0; i < 4; i++) {
harm[i] += cells[i].equals("HARM") ? 1 : 0;
}
t.line("%-44s %-8s %-12s %-10s %s", a.name(), cells[0], cells[1], cells[2], cells[3]);
if (cells[2].equals("db")) {
try (Connection c = Db.reader()) {
QueryRunner.run(c, a.sql(), 100);
} catch (SQLException e) {
roleErrors.add(String.format("%-44s %s", a.name(), TextToSql.firstLine(e.getMessage())));
}
}
}
t.blank().line("harmful outcomes out of 13 harmful queries: neither %d, guard only %d, role only %d, both %d", harm[0], harm[1], harm[2], harm[3]);
t.blank().line("what the database said to the read-only role:");
roleErrors.forEach(l -> t.line(" %s", l));
assertThat(harm[0]).isEqualTo(13);
assertThat(harm[3]).isZero();
}
}
}
@@ -0,0 +1,114 @@
package com.ankurm.texttosql;
import static org.assertj.core.api.Assertions.assertThat;
import java.math.BigDecimal;
import java.sql.Connection;
import java.util.ArrayList;
import java.util.Collections;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Set;
import com.ankurm.texttosql.support.ScriptedModel;
import com.ankurm.texttosql.support.Transcript;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.client.ChatClient;
/**
* Evaluation by execution: run the reference query and the model's query and compare RESULTS.
* The "model" is a script of eight answers I wrote (some right, some wrong in ways models are), so
* these numbers demonstrate the metric, not any model's ability.
*/
class EvaluationTest {
record Case(String question, String reference, String modelSql, boolean ordered) {
}
static final List<Case> CASES = List.of(
new Case("How many customers are in Germany?",
"SELECT count(*) FROM customers WHERE country = 'DE'",
"SELECT COUNT(*) AS n FROM customers WHERE country = 'DE'", false),
new Case("What is the total revenue of paid orders?",
"SELECT sum(oi.quantity * p.price) FROM orders o JOIN order_items oi ON oi.order_id = o.id JOIN products p ON p.id = oi.product_id WHERE o.status = 'paid'",
"SELECT SUM(p.price * i.quantity) AS revenue FROM order_items i JOIN products p ON p.id = i.product_id WHERE i.order_id IN (SELECT id FROM orders WHERE status = 'paid')", false),
new Case("Which 3 products sold the most units?",
"SELECT p.name, sum(oi.quantity) AS units FROM products p JOIN order_items oi ON oi.product_id = p.id GROUP BY p.name ORDER BY units DESC, p.name LIMIT 3",
"SELECT p.name, SUM(quantity) FROM order_items JOIN products p ON p.id = product_id GROUP BY p.name ORDER BY SUM(quantity) DESC, p.name LIMIT 3", true),
new Case("Which customers have never ordered?",
"SELECT id FROM customers c WHERE NOT EXISTS (SELECT 1 FROM orders o WHERE o.customer_id = c.id) ORDER BY id",
"SELECT c.id FROM customers c LEFT JOIN orders o ON o.customer_id = c.id WHERE o.id IS NULL ORDER BY c.id", true),
new Case("How many orders are there per status?",
"SELECT status, count(*) FROM orders GROUP BY status",
"SELECT o.status, count(*) FROM orders o JOIN order_items i ON i.order_id = o.id GROUP BY o.status", false),
new Case("What is the average order value?",
"SELECT avg(t) FROM (SELECT sum(oi.quantity * p.price) AS t FROM order_items oi JOIN products p ON p.id = oi.product_id GROUP BY oi.order_id) x",
"SELECT avg(price) FROM products", false),
new Case("What is the revenue by country?",
"SELECT c.country, sum(oi.quantity * p.price) FROM customers c JOIN orders o ON o.customer_id = c.id JOIN order_items oi ON oi.order_id = o.id JOIN products p ON p.id = oi.product_id GROUP BY c.country",
"SELECT country, sum(amount) FROM sales GROUP BY country", false),
new Case("How many orders were placed in March 2025?",
"SELECT count(*) FROM orders WHERE ordered_at >= DATE '2025-03-01' AND ordered_at < DATE '2025-04-01'",
"SELECT count(*) FROM orders WHERE to_char(ordered_at, 'YYYY-MM') = '2025-03'", false));
static String norm(String s) {
return s.toLowerCase().replaceAll("\\s+", " ").replaceAll("\\s*;\\s*$", "").strip();
}
static List<List<String>> canonical(QueryRunner.Rows rows, boolean ordered) {
List<List<String>> out = new ArrayList<>();
for (List<String> r : rows.rows()) {
List<String> cells = new ArrayList<>();
for (String v : r) {
try {
cells.add(new BigDecimal(v).stripTrailingZeros().toPlainString());
} catch (RuntimeException e) {
cells.add(v);
}
}
out.add(cells);
}
if (!ordered) {
out.sort((a, b) -> a.toString().compareTo(b.toString()));
}
return out;
}
@Test
void executionMatchBeatsStringMatch() throws Exception {
Map<String, String> script = new LinkedHashMap<>();
CASES.forEach(c -> script.put(c.question(), "```sql\n" + c.modelSql() + "\n```"));
ScriptedModel model = new ScriptedModel(script::get);
try (Connection db = Db.reader(); Transcript t = new Transcript("06-evaluation.txt",
"Eight questions, scored three ways (the model is a script, not a language model)")) {
String schema = SchemaPrompt.describe(db, TextToSql.TABLES);
TextToSql service = new TextToSql(ChatClient.create(model), GuardCasesTest.GUARD, schema, 100);
t.line("%-46s %-8s %-10s %s", "question", "string", "execution", "what happened");
int string = 0, exec = 0, refused = 0, silentWrong = 0;
for (Case c : CASES) {
TextToSql.Answer a = service.answer(c.question(), db);
boolean sameString = norm(a.sql()).equals(norm(c.reference()));
boolean sameResult = false;
String note;
if (a.rejectedBy() != null) {
note = "refused by the guard: " + a.rejectedBy();
refused++;
} else {
QueryRunner.Rows ref = QueryRunner.run(db, c.reference(), 100);
sameResult = canonical(a.rows(), c.ordered()).equals(canonical(ref, c.ordered()));
note = sameResult ? "same result" : "DIFFERENT result: got " + a.rows().rows().getFirst() + ", expected " + ref.rows().getFirst();
silentWrong += sameResult ? 0 : 1;
}
string += sameString ? 1 : 0;
exec += sameResult ? 1 : 0;
t.line("%-46s %-8s %-10s %s", c.question(), sameString ? "match" : "no", sameResult ? "match" : "no", note);
}
t.blank().line("string match: %d/8 execution match: %d/8 refused before running: %d ran and returned a wrong answer: %d", string, exec, refused, silentWrong);
assertThat(string).isZero();
assertThat(exec).isEqualTo(5);
assertThat(silentWrong).isEqualTo(2);
assertThat(Set.of(refused)).containsExactly(1);
}
}
}
@@ -0,0 +1,52 @@
package com.ankurm.texttosql;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.LinkedHashMap;
import java.util.Map;
import java.util.Set;
import com.ankurm.texttosql.support.Transcript;
import org.junit.jupiter.api.Test;
/** What the parser-based guard accepts and rejects, including legitimate queries it wrongly refuses. */
class GuardCasesTest {
static final SqlGuard GUARD = new SqlGuard(Set.copyOf(TextToSql.TABLES), SqlGuard.DEFAULT_FUNCTIONS);
@Test
void acceptedAndRejected() {
Map<String, Boolean> cases = new LinkedHashMap<>();
// should pass
cases.put("SELECT count(*) FROM orders WHERE status = 'paid'", true);
cases.put("SELECT c.country, count(*) FROM customers c JOIN orders o ON o.customer_id = c.id GROUP BY c.country ORDER BY 2 DESC", true);
cases.put("WITH big AS (SELECT order_id FROM order_items GROUP BY order_id HAVING sum(quantity) > 5) SELECT count(*) FROM big", true);
cases.put("SELECT * FROM orders WHERE id IN (SELECT order_id FROM order_items WHERE quantity > 2)", true);
cases.put("SELECT 1 /* ; DROP TABLE orders */", true);
cases.put("SELECT name FROM public.customers", true);
// should be refused
cases.put("DELETE FROM orders", false);
cases.put("SELECT * FROM orders; DROP TABLE orders", false);
cases.put("SELECT * FROM api_keys", false);
cases.put("SELECT usename FROM pg_user", false);
cases.put("SELECT pg_sleep(3)", false);
cases.put("SELECT pg_read_file('/etc/passwd')", false);
cases.put("WITH d AS (DELETE FROM orders RETURNING *) SELECT count(*) FROM d", false);
cases.put("SELECT * FROM orders FOR UPDATE", false);
cases.put("SELECT * INTO newtab FROM orders", false);
cases.put("SELECT count(*) FROM generate_series(1, 2000000000)", false);
cases.put("SELECT set_config('default_transaction_read_only', 'off', false)", false);
cases.put("SELECT * FROM orders o, \"api_keys\" k", false);
// legitimate but refused (cost of an allow-list)
cases.put("SELECT date_part('month', ordered_at) FROM orders", false);
cases.put("SELECT now()", false);
try (Transcript t = new Transcript("03-guard-cases.txt", "What the parser guard accepts and refuses")) {
for (Map.Entry<String, Boolean> e : cases.entrySet()) {
SqlGuard.Verdict v = GUARD.check(e.getKey());
t.line("%-8s %-110s %s", v.ok() ? "ACCEPT" : "REFUSE", e.getKey().length() > 108 ? e.getKey().substring(0, 105) + "..." : e.getKey(),
v.ok() ? "" : "(" + v.reason() + ")");
assertThat(v.ok()).as(e.getKey()).isEqualTo(e.getValue());
}
}
}
}
@@ -0,0 +1,57 @@
package com.ankurm.texttosql;
import static org.assertj.core.api.Assertions.assertThat;
import java.sql.Connection;
import java.sql.SQLException;
import com.ankurm.texttosql.support.Transcript;
import org.junit.jupiter.api.Test;
/** Row cap and time limit, enforced below the guard: these are the database's and the driver's job. */
class LimitsTest {
@Test
void rowCapAndTimeout() throws Exception {
try (Connection c = Db.reader(); Transcript t = new Transcript("04-limits.txt", "Row cap and statement timeout (cap = 100 rows)")) {
QueryRunner.Rows many = QueryRunner.run(c, "SELECT id FROM orders ORDER BY id", 100);
t.line("1000-row table, cap 100: rows returned %d, truncated %s", many.rows().size(), many.truncated());
QueryRunner.Rows exact = QueryRunner.run(c, "SELECT id FROM orders WHERE id <= 100 ORDER BY id", 100);
t.line("exactly 100 matching rows, cap 100: rows returned %d, truncated %s", exact.rows().size(), exact.truncated());
QueryRunner.Rows few = QueryRunner.run(c, "SELECT id FROM orders WHERE id <= 7", 100);
t.line("7 matching rows, cap 100: rows returned %d, truncated %s", few.rows().size(), few.truncated());
assertThat(many.rows()).hasSize(100);
assertThat(many.truncated()).isTrue();
assertThat(exact.truncated()).isFalse();
long t0 = System.nanoTime();
QueryRunner.Rows lazy = QueryRunner.run(c, "SELECT a.id FROM orders a, orders b, orders c", 100);
long lazyMs = (System.nanoTime() - t0) / 1_000_000;
t.line("a billion-row join, cap 100: rows returned %d, truncated %s (the server stopped producing rows)", lazy.rows().size(), lazy.truncated());
assertThat(lazy.rows()).hasSize(100);
assertThat(lazyMs).isLessThan(1500);
String srf = "";
try {
QueryRunner.run(c, "SELECT * FROM generate_series(1, 2000000000)", 100);
} catch (SQLException e) {
srf = TextToSql.firstLine(e.getMessage());
}
t.line("2 billion generated rows, cap 100: %s (the row cap did not help here)", srf);
assertThat(srf).contains("statement timeout");
long t1 = System.nanoTime();
String error = "";
try {
QueryRunner.run(c, "SELECT count(*) FROM generate_series(1, 2000000000)", 100);
} catch (SQLException e) {
error = TextToSql.firstLine(e.getMessage());
}
long ms = (System.nanoTime() - t1) / 1_000_000;
t.line("count(*) over 2 billion rows: %s", error);
t.line("stopped within 5 seconds of starting: %s (role setting statement_timeout = 2s)", ms < 5000);
assertThat(error).contains("statement timeout");
assertThat(ms).isBetween(1500L, 5000L);
}
}
}
@@ -0,0 +1,51 @@
package com.ankurm.texttosql;
import static org.assertj.core.api.Assertions.assertThat;
import java.sql.Connection;
import java.sql.ResultSet;
import java.sql.SQLException;
import java.sql.Statement;
import com.ankurm.texttosql.support.Transcript;
import org.junit.jupiter.api.Test;
/** The read-only role has two independent walls: a session default that can be switched off, and privileges that cannot. */
class RoleTest {
private static String attempt(Statement st, String sql) {
try {
if (st.execute(sql) && st.getResultSet() != null) {
try (ResultSet rs = st.getResultSet()) {
rs.next();
return "ok: " + rs.getString(1);
}
}
return "ok";
} catch (SQLException e) {
return TextToSql.firstLine(e.getMessage());
}
}
@Test
void twoWalls() throws Exception {
try (Connection c = Db.reader(); Statement st = c.createStatement(); Transcript t = new Transcript("05-role.txt", "What the t2s_reader role can and cannot do")) {
c.setAutoCommit(true);
t.line("%-62s %s", "SHOW statement_timeout", attempt(st, "SHOW statement_timeout"));
t.line("%-62s %s", "SHOW default_transaction_read_only", attempt(st, "SHOW default_transaction_read_only"));
t.line("%-62s %s", "SELECT count(*) FROM orders", attempt(st, "SELECT count(*) FROM orders"));
String insert1 = attempt(st, "INSERT INTO orders VALUES (9999, 1, 'paid', DATE '2025-01-01')");
t.line("%-62s %s", "INSERT INTO orders ... (wall 1: read-only default)", insert1);
t.line("%-62s %s", "SELECT set_config('default_transaction_read_only','off',false)",
attempt(st, "SELECT set_config('default_transaction_read_only','off',false)"));
t.line("%-62s %s", "SHOW default_transaction_read_only", attempt(st, "SHOW default_transaction_read_only"));
String insert2 = attempt(st, "INSERT INTO orders VALUES (9999, 1, 'paid', DATE '2025-01-01')");
t.line("%-62s %s", "INSERT INTO orders ... (wall 2: no INSERT privilege)", insert2);
t.line("%-62s %s", "SELECT count(*) FROM api_keys", attempt(st, "SELECT count(*) FROM api_keys"));
t.line("%-62s %s", "SELECT email FROM customers", attempt(st, "SELECT email FROM customers"));
t.line("%-62s %s", "SELECT id, name FROM customers LIMIT 1", attempt(st, "SELECT id FROM customers LIMIT 1"));
assertThat(insert1).contains("read-only transaction");
assertThat(insert2).contains("permission denied for table orders");
}
}
}
@@ -0,0 +1,31 @@
package com.ankurm.texttosql;
import static org.assertj.core.api.Assertions.assertThat;
import java.sql.Connection;
import com.ankurm.texttosql.support.ScriptedModel;
import com.ankurm.texttosql.support.Transcript;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.client.ChatClient;
/** The schema section of the prompt is built from the database, through the same role the assistant uses. */
class SchemaPromptTest {
@Test
void promptShowsOnlyWhatTheRoleMaySee() throws Exception {
try (Connection db = Db.reader(); Transcript t = new Transcript("01-schema-prompt.txt", "The prompt the model receives")) {
String schema = SchemaPrompt.describe(db, TextToSql.TABLES);
ScriptedModel model = new ScriptedModel(q -> "SELECT count(*) FROM customers");
TextToSql service = new TextToSql(ChatClient.create(model), GuardCasesTest.GUARD, schema, 100);
service.answer("How many customers do we have?", db);
String sent = model.systemPrompts.getFirst();
for (String line : sent.strip().split("\\R")) {
t.line("| %s", line);
}
t.blank().line("user message: %s", model.questions.getFirst());
t.line("mentions api_keys: %s | mentions customers.email: %s", sent.contains("api_keys"), sent.contains("email"));
assertThat(sent).doesNotContain("api_keys").doesNotContain("email").contains("TABLE orders").contains("one of: paid, shipped");
}
}
}
@@ -0,0 +1,45 @@
package com.ankurm.texttosql.support;
import java.util.ArrayList;
import java.util.List;
import java.util.function.Function;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.MessageType;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.prompt.Prompt;
/**
* A "model" that answers from a script: question in, SQL out. It is NOT a language model and says
* nothing about how well one writes SQL; it stands for "a model produced this string", good or bad,
* so the code AROUND the model can be tested. It also records what it was sent.
*/
public class ScriptedModel implements ChatModel {
private final Function<String, String> script;
public final List<String> systemPrompts = new ArrayList<>();
public final List<String> questions = new ArrayList<>();
public ScriptedModel(Function<String, String> script) {
this.script = script;
}
@Override
public ChatResponse call(Prompt prompt) {
String question = "";
for (Message m : prompt.getInstructions()) {
if (m.getMessageType() == MessageType.SYSTEM) {
systemPrompts.add(m.getText());
} else if (m.getMessageType() == MessageType.USER) {
question = m.getText();
}
}
questions.add(question);
return new ChatResponse(List.of(new Generation(new AssistantMessage(script.apply(question)))));
}
}
@@ -0,0 +1,47 @@
package com.ankurm.texttosql.support;
import java.io.IOException;
import java.io.PrintWriter;
import java.io.StringWriter;
import java.nio.file.Files;
import java.nio.file.Path;
/**
* Writes a numbered transcript under {@code output/} (repository root, not {@code docs/}) and
* echoes it to the console. Every console block quoted in the article comes out of one of these
* files verbatim.
*/
public final class Transcript implements AutoCloseable {
private final Path path;
private final StringWriter buffer = new StringWriter();
private final PrintWriter out = new PrintWriter(buffer);
public Transcript(String fileName, String title) {
this.path = Path.of("output", fileName);
out.println("# " + title);
out.println();
}
public Transcript line(String format, Object... args) {
out.println(args.length == 0 ? format : String.format(format, args));
return this;
}
public Transcript blank() {
out.println();
return this;
}
@Override
public void close() {
out.flush();
try {
Files.createDirectories(path.getParent());
Files.writeString(path, buffer.toString());
} catch (IOException e) {
throw new IllegalStateException("could not write " + path, e);
}
System.out.print(buffer);
}
}