|
|
@@ -17,6 +17,7 @@ package com.qingjian.module.ai.agent.skill.text2sql;
|
|
|
|
|
|
import com.google.common.collect.Lists;
|
|
|
import com.qingjian.module.ai.agent.skill.text2sql.dialect.PostgreDialect;
|
|
|
+import lombok.extern.slf4j.Slf4j;
|
|
|
import org.noear.snack4.ONode;
|
|
|
import org.noear.solon.Solon;
|
|
|
import org.noear.solon.Utils;
|
|
|
@@ -56,6 +57,7 @@ import java.util.stream.Collectors;
|
|
|
* @author noear
|
|
|
* @since 3.9.1
|
|
|
*/
|
|
|
+@Slf4j
|
|
|
@Preview("3.9.1")
|
|
|
public class Text2SqlSkill extends AbsSkill {
|
|
|
protected final static Logger LOG = LoggerFactory.getLogger(Text2SqlSkill.class);
|
|
|
@@ -74,7 +76,6 @@ public class Text2SqlSkill extends AbsSkill {
|
|
|
protected int maxContextLength = 8000;
|
|
|
protected SchemaMode schemaMode = SchemaMode.FULL;
|
|
|
protected boolean readOnly = true;
|
|
|
- private SqlUtils sqlUtils;
|
|
|
|
|
|
public record ColumnInfo(String name, String en, String type) {
|
|
|
}
|
|
|
@@ -82,13 +83,6 @@ public class Text2SqlSkill extends AbsSkill {
|
|
|
public Text2SqlSkill() {
|
|
|
super();
|
|
|
this.dialect = new PostgreDialect();
|
|
|
- try {
|
|
|
- DynamicDataSource dds = Solon.context().getBean("db1");
|
|
|
- DataSource ds = dds.getDefaultTargetDataSource();
|
|
|
- this.sqlUtils = SqlUtils.of(ds);
|
|
|
- } catch (Exception e) {
|
|
|
- LOG.error("Failed to initialize Text2SqlSkill", e);
|
|
|
- }
|
|
|
}
|
|
|
|
|
|
private void init() {
|
|
|
@@ -120,8 +114,8 @@ public class Text2SqlSkill extends AbsSkill {
|
|
|
* 初始化:识别数据库方言并预加载表元数据
|
|
|
*/
|
|
|
private void initDialectAndMetadata() {
|
|
|
- tableColumnsMap.put("call_record", getCallColumns());
|
|
|
- tableColumnsMap.put("trans_record", getColumnInfos());
|
|
|
+ tableColumnsMap.put("call_record_all", getCallColumns());
|
|
|
+ tableColumnsMap.put("trans_record_all", getColumnInfos());
|
|
|
tableColumnsMap.put("express_info", getExpressColumns());
|
|
|
tableColumnsMap.put("together_live_info", getTogetherLiveInfoColumns());
|
|
|
tableColumnsMap.put("together_flight_info", getTogetherFlightInfoColumns());
|
|
|
@@ -130,8 +124,8 @@ public class Text2SqlSkill extends AbsSkill {
|
|
|
tableColumnsMap.put("flight_ticket_info", getFlightTicketInfoColumns());
|
|
|
tableColumnsMap.put("hotel_stay_info", getHotelStayInfoColumns());
|
|
|
|
|
|
- tableRemarksMap.put("call_record", "通信记录表");
|
|
|
- tableRemarksMap.put("trans_record", "交易记录表");
|
|
|
+ tableRemarksMap.put("call_record_all", "通信记录表");
|
|
|
+ tableRemarksMap.put("trans_record_all", "交易记录表");
|
|
|
tableRemarksMap.put("express_info", "快递收发记录表");
|
|
|
tableRemarksMap.put("together_live_info", "同住记录表");
|
|
|
tableRemarksMap.put("together_flight_info", "同航班记录表");
|
|
|
@@ -140,8 +134,8 @@ public class Text2SqlSkill extends AbsSkill {
|
|
|
tableRemarksMap.put("flight_ticket_info", "飞机票购买记录表");
|
|
|
tableRemarksMap.put("hotel_stay_info", "酒店住宿记录表");
|
|
|
|
|
|
- tableNames.add("call_record");
|
|
|
- tableNames.add("trans_record");
|
|
|
+ tableNames.add("call_record_all");
|
|
|
+ tableNames.add("trans_record_all");
|
|
|
tableNames.add("express_info");
|
|
|
tableNames.add("together_live_info");
|
|
|
tableNames.add("together_flight_info");
|
|
|
@@ -185,7 +179,7 @@ public class Text2SqlSkill extends AbsSkill {
|
|
|
|
|
|
@Override
|
|
|
public String description() {
|
|
|
- return "数据库专家:具备深厚的 DuckDB 方言知识,擅长多表分析。";
|
|
|
+ return "数据库专家:具备深厚的 DuckDB 方言知识,擅长多表分析。\n 找不到其他可用工具函数时,才使用该技能。";
|
|
|
}
|
|
|
|
|
|
@Override
|
|
|
@@ -226,7 +220,7 @@ public class Text2SqlSkill extends AbsSkill {
|
|
|
return sb.toString();
|
|
|
}
|
|
|
|
|
|
- @ToolMapping(name = "execute_sql", description = "执行单条 SELECT 查询语句。")
|
|
|
+ @ToolMapping(name = "execute_sql", description = "无其他可用工具函数时,可使用execute_sql执行单条 SELECT 查询语句进行分析。")
|
|
|
public String executeSql(@Param("sql") String sql) {
|
|
|
if (Assert.isBlank(sql)) return "Error: SQL is empty.";
|
|
|
|
|
|
@@ -243,12 +237,15 @@ public class Text2SqlSkill extends AbsSkill {
|
|
|
cleanSql = dialect.applyPagination(cleanSql, maxRows);
|
|
|
}
|
|
|
try {
|
|
|
-
|
|
|
+ DynamicDataSource dds = Solon.context().getBean("db1");
|
|
|
+ DataSource ds = dds.getDefaultTargetDataSource();
|
|
|
+ SqlUtils sqlUtils = SqlUtils.of(ds);
|
|
|
List<Map> rows = sqlUtils.sql(cleanSql).queryRowList(Map.class);
|
|
|
if (rows == null || rows.isEmpty()) return "Query OK. No data found.";
|
|
|
String json = ONode.serialize(rows);
|
|
|
return json.length() > maxContextLength ? json.substring(0, maxContextLength) + "... [Truncated]" : json;
|
|
|
} catch (SQLException e) {
|
|
|
+ log.error("SQL Error: ", e);
|
|
|
return "SQL Error: " + e.getMessage() + "\nHint: " + dialect.getErrorHint(e);
|
|
|
}
|
|
|
}
|