Co-Authored-By: Claude Sonnet 5.5 <[email protected]> Claude-Session: https://claude.ai/code/session_01JXVi2GMQ7bR5EmbUFdDj7N
72 lines
2.3 KiB
Java
72 lines
2.3 KiB
Java
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();
|
|
}
|
|
}
|