Files
spring-ai/text-to-sql/src/main/java/com/ankurm/texttosql/TextToSql.java
T

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();
}
}