本文目录导读:

我来为您提供几个Java中使用AI生成SQL的实用案例,涵盖不同的实现方式。
案例1:基于规则和模板的SQL生成
import java.util.HashMap;
import java.util.Map;
import java.util.regex.Matcher;
import java.util.regex.Pattern;
public class RuleBasedSQLGenerator {
// 保存表结构和关系的元数据
private Map<String, TableSchema> schemaRegistry = new HashMap<>();
// 自然语言到SQL的规则映射
private Map<String, String> intentPatterns = new HashMap<>();
public RuleBasedSQLGenerator() {
initializeSchema();
initializePatterns();
}
private void initializeSchema() {
// 模拟数据库表结构
TableSchema usersTable = new TableSchema("users");
usersTable.addColumn("id", "BIGINT", "主键");
usersTable.addColumn("name", "VARCHAR(100)", "用户名");
usersTable.addColumn("email", "VARCHAR(200)", "邮箱");
usersTable.addColumn("age", "INT", "年龄");
usersTable.addColumn("created_at", "TIMESTAMP", "创建时间");
usersTable.addColumn("status", "VARCHAR(20)", "状态");
schemaRegistry.put("users", usersTable);
TableSchema ordersTable = new TableSchema("orders");
ordersTable.addColumn("id", "BIGINT", "订单ID");
ordersTable.addColumn("user_id", "BIGINT", "用户ID");
ordersTable.addColumn("amount", "DECIMAL(10,2)", "金额");
ordersTable.addColumn("status", "VARCHAR(20)", "订单状态");
ordersTable.addColumn("created_at", "TIMESTAMP", "创建时间");
schemaRegistry.put("orders", ordersTable);
}
private void initializePatterns() {
intentPatterns.put("查询.*用户.*信息", "SELECT * FROM users WHERE ");
intentPatterns.put("统计.*订单.*金额", "SELECT SUM(amount) FROM orders WHERE ");
intentPatterns.put("查询.*活跃.*用户", "SELECT * FROM users WHERE status = 'active'");
intentPatterns.put("*订单.*数量", "SELECT COUNT(*) FROM orders WHERE DATE(created_at) = CURDATE()");
}
public String generateSQL(String naturalLanguage) {
// 1. 识别意图
String matchedPattern = null;
for (String pattern : intentPatterns.keySet()) {
if (naturalLanguage.contains(pattern.replace(".*", ""))) {
matchedPattern = pattern;
break;
}
}
if (matchedPattern == null) {
return "无法识别的查询请求";
}
// 2. 提取参数
String sqlTemplate = intentPatterns.get(matchedPattern);
// 3. 处理特殊条件
if (naturalLanguage.contains("年龄大于")) {
Pattern p = Pattern.compile("年龄大于(\\d+)");
Matcher m = p.matcher(naturalLanguage);
if (m.find()) {
sqlTemplate = "SELECT * FROM users WHERE age > " + m.group(1);
}
} else if (naturalLanguage.contains("年龄小于")) {
Pattern p = Pattern.compile("年龄小于(\\d+)");
Matcher m = p.matcher(naturalLanguage);
if (m.find()) {
sqlTemplate = "SELECT * FROM users WHERE age < " + m.group(1);
}
}
// 4. 处理排序
if (naturalLanguage.contains("排序")) {
if (naturalLanguage.contains("年龄")) {
sqlTemplate += " ORDER BY age";
if (naturalLanguage.contains("降序") || naturalLanguage.contains("大到小")) {
sqlTemplate += " DESC";
} else {
sqlTemplate += " ASC";
}
}
}
// 5. 处理限制
if (naturalLanguage.contains("前") && naturalLanguage.contains("条")) {
Pattern p = Pattern.compile("前(\\d+)条");
Matcher m = p.matcher(naturalLanguage);
if (m.find()) {
sqlTemplate += " LIMIT " + m.group(1);
}
}
return sqlTemplate;
}
// 表结构类
static class TableSchema {
private String tableName;
private Map<String, ColumnInfo> columns = new HashMap<>();
public TableSchema(String tableName) {
this.tableName = tableName;
}
public void addColumn(String name, String type, String comment) {
columns.put(name, new ColumnInfo(name, type, comment));
}
}
static class ColumnInfo {
String name;
String type;
String comment;
public ColumnInfo(String name, String type, String comment) {
this.name = name;
this.type = type;
this.comment = comment;
}
}
public static void main(String[] args) {
RuleBasedSQLGenerator generator = new RuleBasedSQLGenerator();
// 测试用例
String[] testQueries = {
"查询所有用户信息",
"查询年龄大于25的用户",
"查询年龄小于18的用户并按年龄排序",
"查询前10条活跃用户",
"统计今日订单数量"
};
for (String query : testQueries) {
System.out.println("自然语言: " + query);
System.out.println("生成SQL: " + generator.generateSQL(query));
System.out.println("---");
}
}
}
案例2:使用开源NLP库进行SQL生成
import opennlp.tools.stemmer.PorterStemmer;
import opennlp.tools.tokenize.SimpleTokenizer;
import org.deeplearning4j.text.tokenization.tokenizer.Tokenizer;
import javax.json.Json;
import javax.json.JsonObject;
import java.util.*;
public class NLPSQLGenerator {
private static final String[] TABLE_KEYWORDS = {"用户", "订单", "商品", "分类"};
private static final String[] AGGREGATE_KEYWORDS = {"总数", "平均值", "最大值", "最小值", "总和"};
private static final String[] CONDITION_KEYWORDS = {"大于", "小于", "等于", "包含", "在"};
private Map<String, String> tableSynonyms = new HashMap<>();
private Map<String, String> fieldSynonyms = new HashMap<>();
public NLPSQLGenerator() {
// 初始化同义词映射
tableSynonyms.put("用户", "users");
tableSynonyms.put("会员", "users");
tableSynonyms.put("客户", "users");
tableSynonyms.put("订单", "orders");
tableSynonyms.put("商品", "products");
tableSynonyms.put("产品", "products");
fieldSynonyms.put("名字", "name");
fieldSynonyms.put("名称", "name");
fieldSynonyms.put("邮箱", "email");
fieldSynonyms.put("邮件", "email");
fieldSynonyms.put("年龄", "age");
fieldSynonyms.put("金额", "amount");
fieldSynonyms.put("价格", "price");
fieldSynonyms.put("数量", "quantity");
}
public SQLQuery parseNaturalLanguage(String input) {
SQLQuery query = new SQLQuery();
// 1. 分词
String[] tokens = tokenize(input);
// 2. 识别意图
query.setSelectFields(identifySelectFields(tokens));
// 3. 识别表名
String tableName = identifyTable(tokens);
query.setFromTable(tableName);
// 4. 识别聚合函数
query.setAggregateFunction(identifyAggregate(tokens));
// 5. 识别条件
List<Condition> conditions = identifyConditions(tokens, tableName);
query.setConditions(conditions);
// 6. 识别排序
query.setOrderBy(identifyOrderBy(tokens));
// 7. 识别分组
query.setGroupBy(identifyGroupBy(tokens));
// 8. 生成SQL
query.setSql(generateSQLStatement(query));
return query;
}
private String[] tokenize(String input) {
SimpleTokenizer tokenizer = SimpleTokenizer.INSTANCE;
return tokenizer.tokenize(input);
}
private List<String> identifySelectFields(String[] tokens) {
List<String> fields = new ArrayList<>();
for (String token : tokens) {
String mapped = fieldSynonyms.get(token);
if (mapped != null) {
fields.add(mapped);
}
}
return fields.isEmpty() ? Arrays.asList("*") : fields;
}
private String identifyTable(String[] tokens) {
for (String token : tokens) {
String mapped = tableSynonyms.get(token);
if (mapped != null) {
return mapped;
}
}
return "unknown_table";
}
private String identifyAggregate(String[] tokens) {
for (String token : tokens) {
if (token.equals("总数") || token.equals("数量")) {
return "COUNT";
} else if (token.equals("平均值") || token.equals("平均")) {
return "AVG";
} else if (token.equals("最大值") || token.equals("最大")) {
return "MAX";
} else if (token.equals("最小值") || token.equals("最小")) {
return "MIN";
} else if (token.equals("总和")) {
return "SUM";
}
}
return null;
}
private List<Condition> identifyConditions(String[] tokens, String tableName) {
List<Condition> conditions = new ArrayList<>();
for (int i = 0; i < tokens.length - 1; i++) {
String field = fieldSynonyms.get(tokens[i]);
if (field != null) {
String operator = null;
String value = null;
// 检查各种条件
if (tokens[i + 1].equals("大于")) {
operator = ">";
if (i + 2 < tokens.length) value = tokens[i + 2];
} else if (tokens[i + 1].equals("小于")) {
operator = "<";
if (i + 2 < tokens.length) value = tokens[i + 2];
} else if (tokens[i + 1].equals("等于")) {
operator = "=";
if (i + 2 < tokens.length) value = tokens[i + 2];
}
if (operator != null && value != null) {
conditions.add(new Condition(field, operator, value));
}
}
}
return conditions;
}
private String identifyOrderBy(String[] tokens) {
for (int i = 0; i < tokens.length; i++) {
if (tokens[i].equals("排序") || tokens[i].equals("排序")) {
if (i > 0) {
String field = fieldSynonyms.get(tokens[i - 1]);
if (field != null) {
return field;
}
}
}
}
return null;
}
private String identifyGroupBy(String[] tokens) {
for (int i = 0; i < tokens.length; i++) {
if (tokens[i].equals("分组") || tokens[i].equals("按")) {
if (i + 1 < tokens.length) {
String field = fieldSynonyms.get(tokens[i + 1]);
if (field != null) {
return field;
}
}
}
}
return null;
}
private String generateSQLStatement(SQLQuery query) {
StringBuilder sql = new StringBuilder();
// SELECT
sql.append("SELECT ");
if (query.getAggregateFunction() != null && !query.getSelectFields().isEmpty()) {
sql.append(query.getAggregateFunction())
.append("(")
.append(String.join(", ", query.getSelectFields()))
.append(")");
} else {
sql.append(String.join(", ", query.getSelectFields()));
}
// FROM
sql.append(" FROM ").append(query.getFromTable());
// WHERE
if (!query.getConditions().isEmpty()) {
sql.append(" WHERE ");
List<String> conditionStrings = new ArrayList<>();
for (Condition condition : query.getConditions()) {
conditionStrings.add(condition.toString());
}
sql.append(String.join(" AND ", conditionStrings));
}
// GROUP BY
if (query.getGroupBy() != null) {
sql.append(" GROUP BY ").append(query.getGroupBy());
}
// ORDER BY
if (query.getOrderBy() != null) {
sql.append(" ORDER BY ").append(query.getOrderBy());
}
return sql.toString();
}
// 查询模型类
static class SQLQuery {
private List<String> selectFields = new ArrayList<>();
private String fromTable;
private String aggregateFunction;
private List<Condition> conditions = new ArrayList<>();
private String orderBy;
private String groupBy;
private String sql;
// Getters and Setters
// ... (省略getter/setter方法)
public void setSelectFields(List<String> selectFields) {
this.selectFields = selectFields;
}
public List<String> getSelectFields() {
return selectFields;
}
public void setFromTable(String fromTable) {
this.fromTable = fromTable;
}
public String getFromTable() {
return fromTable;
}
public void setAggregateFunction(String aggregateFunction) {
this.aggregateFunction = aggregateFunction;
}
public String getAggregateFunction() {
return aggregateFunction;
}
public void setConditions(List<Condition> conditions) {
this.conditions = conditions;
}
public List<Condition> getConditions() {
return conditions;
}
public void setOrderBy(String orderBy) {
this.orderBy = orderBy;
}
public String getOrderBy() {
return orderBy;
}
public void setGroupBy(String groupBy) {
this.groupBy = groupBy;
}
public String getGroupBy() {
return groupBy;
}
public void setSql(String sql) {
this.sql = sql;
}
public String getSql() {
return sql;
}
}
static class Condition {
private String field;
private String operator;
private String value;
public Condition(String field, String operator, String value) {
this.field = field;
this.operator = operator;
this.value = value;
}
@Override
public String toString() {
return field + " " + operator + " " + value;
}
}
public static void main(String[] args) {
NLPSQLGenerator generator = new NLPSQLGenerator();
String[] testQueries = {
"查询所有用户信息",
"查询年龄大于25的用户",
"查询用户总数",
"按年龄分组统计用户数量",
"查询金额大于100的订单"
};
for (String query : testQueries) {
System.out.println("输入: " + query);
SQLQuery result = generator.parseNaturalLanguage(query);
System.out.println("生成SQL: " + result.getSql());
System.out.println("---");
}
}
}
案例3:集成AI API的SQL生成
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import okhttp3.*;
import java.io.IOException;
import java.util.concurrent.TimeUnit;
public class AISQLGenerator {
private static final String API_ENDPOINT = "https://api.openai.com/v1/completions";
private static final String API_KEY = "your-api-key-here"; // 替换为实际的API密钥
private final OkHttpClient client;
private final ObjectMapper objectMapper;
public AISQLGenerator() {
this.client = new OkHttpClient.Builder()
.connectTimeout(30, TimeUnit.SECONDS)
.readTimeout(30, TimeUnit.SECONDS)
.build();
this.objectMapper = new ObjectMapper();
}
public String generateSQL(String userQuery) throws IOException {
// 构建提示词
String prompt = buildPrompt(userQuery);
// 构建请求体
String jsonBody = buildRequestBody(prompt);
// 发送请求
String response = sendRequest(jsonBody);
// 解析响应并提取SQL
return extractSQLFromResponse(response);
}
private String buildPrompt(String userQuery) {
StringBuilder prompt = new StringBuilder();
prompt.append("你是一个SQL专家,请将以下自然语言查询转换为SQL语句。\n\n");
prompt.append("数据库表结构:\n");
prompt.append("users 表: id (BIGINT), name (VARCHAR), email (VARCHAR), age (INT), created_at (TIMESTAMP), status (VARCHAR)\n");
prompt.append("orders 表: id (BIGINT), user_id (BIGINT), amount (DECIMAL), status (VARCHAR), created_at (TIMESTAMP)\n");
prompt.append("products 表: id (BIGINT), name (VARCHAR), price (DECIMAL), category (VARCHAR)\n\n");
prompt.append("请只返回SQL语句,不要包含其他解释。\n\n");
prompt.append("用户查询:").append(userQuery);
return prompt.toString();
}
private String buildRequestBody(String prompt) {
try {
JsonNode requestBody = objectMapper.createObjectNode()
.put("model", "text-davinci-003")
.put("prompt", prompt)
.put("max_tokens", 150)
.put("temperature", 0.3)
.put("n", 1);
return objectMapper.writeValueAsString(requestBody);
} catch (Exception e) {
throw new RuntimeException("构建请求体失败", e);
}
}
private String sendRequest(String jsonBody) throws IOException {
RequestBody body = RequestBody.create(
MediaType.parse("application/json"),
jsonBody
);
Request request = new Request.Builder()
.url(API_ENDPOINT)
.addHeader("Authorization", "Bearer " + API_KEY)
.addHeader("Content-Type", "application/json")
.post(body)
.build();
try (Response response = client.newCall(request).execute()) {
if (!response.isSuccessful()) {
throw new IOException("API请求失败: " + response.code());
}
return response.body().string();
}
}
private String extractSQLFromResponse(String response) {
try {
JsonNode jsonNode = objectMapper.readTree(response);
String text = jsonNode.get("choices").get(0).get("text").asText().trim();
// 清理输出,确保只返回SQL
if (text.contains("```sql")) {
text = text.substring(text.indexOf("```sql") + 6);
text = text.substring(0, text.indexOf("```"));
} else if (text.contains("```")) {
text = text.substring(text.indexOf("```") + 3);
text = text.substring(0, text.indexOf("```"));
}
return text.trim();
} catch (Exception e) {
throw new RuntimeException("解析响应失败", e);
}
}
public static void main(String[] args) {
AISQLGenerator generator = new AISQLGenerator();
String[] testQueries = {
"查询最近7天注册的活跃用户",
"统计每个类别的商品数量和平均价格",
"查询订单金额大于1000的用户及其订单详情",
"找出购买商品最多的前10个用户"
};
for (String query : testQueries) {
System.out.println("用户查询: " + query);
try {
String sql = generator.generateSQL(query);
System.out.println("生成SQL: " + sql);
} catch (IOException e) {
System.err.println("错误: " + e.getMessage());
}
System.out.println("---");
}
}
}
案例4:基于Spring Boot的SQL生成服务
import org.springframework.boot.SpringApplication;
import org.springframework.boot.autoconfigure.SpringBootApplication;
import org.springframework.web.bind.annotation.*;
import java.util.*;
@SpringBootApplication
@RestController
@RequestMapping("/api/sql-generator")
public class SQLGeneratorService {
private final SQLGenerationEngine engine;
public SQLGeneratorService() {
this.engine = new SQLGenerationEngine();
}
public static void main(String[] args) {
SpringApplication.run(SQLGeneratorService.class, args);
}
@PostMapping("/generate")
public Map<String, Object> generateSQL(@RequestBody GenerationRequest request) {
Map<String, Object> response = new HashMap<>();
try {
String sql = engine.generate(request.getQuery(), request.getContext());
response.put("success", true);
response.put("sql", sql);
response.put("confidence", calculateConfidence(request.getQuery()));
} catch (Exception e) {
response.put("success", false);
response.put("error", e.getMessage());
}
return response;
}
@PostMapping("/batch-generate")
public List<Map<String, Object>> batchGenerate(@RequestBody List<GenerationRequest> requests) {
List<Map<String, Object>> results = new ArrayList<>();
for (GenerationRequest request : requests) {
Map<String, Object> result = new HashMap<>();
result.put("query", request.getQuery());
try {
String sql = engine.generate(request.getQuery(), request.getContext());
result.put("success", true);
result.put("sql", sql);
} catch (Exception e) {
result.put("success", false);
result.put("error", e.getMessage());
}
results.add(result);
}
return results;
}
@GetMapping("/table-info")
public Map<String, Object> getTableInfo() {
return engine.getTableMetadata();
}
private double calculateConfidence(String query) {
// 简单的置信度计算
int complexity = query.length();
int keywords = countKeywords(query);
double baseConfidence = 0.8;
double complexityFactor = Math.max(0, 1 - complexity / 100.0);
double keywordFactor = keywords * 0.1;
return Math.min(1.0, baseConfidence * complexityFactor + keywordFactor);
}
private int countKeywords(String query) {
String[] keywords = {"SELECT", "WHERE", "JOIN", "GROUP", "ORDER", "HAVING"};
int count = 0;
for (String keyword : keywords) {
if (query.toUpperCase().contains(keyword)) {
count++;
}
}
return count;
}
// 请求体类
static class GenerationRequest {
private String query;
private Map<String, Object> context;
public String getQuery() {
return query;
}
public void setQuery(String query) {
this.query = query;
}
public Map<String, Object> getContext() {
return context;
}
public void setContext(Map<String, Object> context) {
this.context = context;
}
}
// SQL生成引擎
static class SQLGenerationEngine {
private final Map<String, TableMetadata> tableMetadata = new HashMap<>();
public SQLGenerationEngine() {
initializeMetadata();
}
private void initializeMetadata() {
// 初始化表元数据
TableMetadata usersTable = new TableMetadata("users");
usersTable.addField("id", "BIGINT", true, true);
usersTable.addField("name", "VARCHAR(100)", false, false);
usersTable.addField("email", "VARCHAR(200)", false, false);
usersTable.addField("age", "INT", false, false);
usersTable.addField("status", "VARCHAR(20)", false, false);
tableMetadata.put("users", usersTable);
TableMetadata ordersTable = new TableMetadata("orders");
ordersTable.addField("id", "BIGINT", true, true);
ordersTable.addField("user_id", "BIGINT", false, false);
ordersTable.addField("amount", "DECIMAL(10,2)", false, false);
ordersTable.addField("status", "VARCHAR(20)", false, false);
ordersTable.addField("created_at", "TIMESTAMP", false, false);
tableMetadata.put("orders", ordersTable);
}
public String generate(String query, Map<String, Object> context) {
// 简化的SQL生成逻辑
StringBuilder sql = new StringBuilder();
// 识别查询类型
if (query.contains("查询") || query.contains("获取") || query.contains("列出")) {
sql.append("SELECT ");
// 识别字段
sql.append("*");
// 识别表
if (query.contains("用户") || query.contains("会员")) {
sql.append(" FROM users");
// 添加条件
if (query.contains("年龄大于")) {
String age = extractNumber(query, "年龄大于");
sql.append(" WHERE age > ").append(age);
} else if (query.contains("活跃")) {
sql.append(" WHERE status = 'active'");
}
// 添加排序
if (query.contains("排序") || query.contains("顺序")) {
sql.append(" ORDER BY created_at DESC");
}
} else if (query.contains("订单")) {
sql.append(" FROM orders");
// 添加时间条件
if (query.contains("quot;) || query.contains("本月")) {
sql.append(" WHERE created_at >= DATE_SUB(NOW(), INTERVAL 1 MONTH)");
}
// 添加金额条件
if (query.contains("金额大于")) {
String amount = extractNumber(query, "金额大于");
sql.append(" AND amount > ").append(amount);
}
}
// 添加限制
if (query.contains("前") && query.contains("条")) {
String limit = extractNumber(query, "前", "条");
sql.append(" LIMIT ").append(limit);
}
} else if (query.contains("统计") || query.contains("计算")) {
sql.append("SELECT ");
if (query.contains("总数")) {
sql.append("COUNT(*)");
} else if (query.contains("平均值") || query.contains("平均")) {
sql.append("AVG(");
if (query.contains("金额")) {
sql.append("amount");
} else if (query.contains("年龄")) {
sql.append("age");
}
sql.append(")");
}
sql.append(" FROM ");
if (query.contains("用户")) {
sql.append("users");
} else if (query.contains("订单")) {
sql.append("orders");
}
}
return sql.toString();
}
public Map<String, Object> getTableMetadata() {
Map<String, Object> result = new HashMap<>();
for (Map.Entry<String, TableMetadata> entry : tableMetadata.entrySet()) {
result.put(entry.getKey(), entry.getValue());
}
return result;
}
private String extractNumber(String text, String prefix) {
int startIndex = text.indexOf(prefix) + prefix.length();
StringBuilder number = new StringBuilder();
for (int i = startIndex; i < text.length(); i++) {
char c = text.charAt(i);
if (Character.isDigit(c)) {
number.append(c);
} else {
break;
}
}
return number.toString();
}
private String extractNumber(String text, String prefix, String suffix) {
String number = extractNumber(text, prefix);
return number;
}
}
static class TableMetadata {
private String tableName;
private List<FieldMetadata> fields = new ArrayList<>();
public TableMetadata(String tableName) {
this.tableName = tableName;
}
public void addField(String name, String type, boolean isPrimary, boolean isAutoIncrement) {
fields.add(new FieldMetadata(name, type, isPrimary, isAutoIncrement));
}
// Getters
public String getTableName() { return tableName; }
public List<FieldMetadata> getFields() { return fields; }
}
static class FieldMetadata {
private String name;
private String type;
private boolean primary;
private boolean autoIncrement;
public FieldMetadata(String name, String type, boolean primary, boolean autoIncrement) {
this.name = name;
this.type = type;
this.primary = primary;
this.autoIncrement = autoIncrement;
}
// Getters
public String getName() { return name; }
public String getType() { return type; }
public boolean isPrimary() { return primary; }
public boolean isAutoIncrement() { return autoIncrement; }
}
}
使用注意事项
- 数据安全: 在生产环境中使用时,注意SQL注入防护
- 性能优化: 对于复杂查询,考虑使用缓存机制
- 准确性验证: AI生成的SQL需要人工审核,特别是涉及重要数据的操作
- 异常处理: 添加完善的错误处理机制
- 日志记录: 记录生成过程以便调试和优化
这些案例展示了从简单规则到复杂AI集成的不同实现方案,您可以根据实际需求选择合适的方案。