From f0353dc14b55cbbcc531b1dd027ae7e7ac0400ae Mon Sep 17 00:00:00 2001 From: meteor <1839462565@qq.com> Date: Mon, 20 Jul 2026 15:31:36 +0800 Subject: [PATCH 1/6] feat(ai): standardize AI tool results as JSON --- .../domain/api/model/ai/AiToolResult.java | 45 ++++ .../core/impl/ai/AiToolServiceImpl.java | 196 +++++++++++++----- .../core/impl/ai/AiToolServiceImplTest.java | 48 +++++ .../api/adapter/ai/AiChatStreamAdapter.java | 9 + .../web/api/adapter/ai/AiToolAdapter.java | 40 ++-- .../web/api/mcp/adapter/AiToolMcpAdapter.java | 14 +- 6 files changed, 281 insertions(+), 71 deletions(-) create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java new file mode 100644 index 000000000..050a110b0 --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java @@ -0,0 +1,45 @@ +package ai.chat2db.community.domain.api.model.ai; + +import lombok.Data; + +import java.util.ArrayList; +import java.util.List; + +@Data +public class AiToolResult { + + private Boolean success; + + private String tool; + + private String summary; + + private List data = new ArrayList<>(); + + private String errorCode; + + public static AiToolResult success(String tool, String summary, List data) { + AiToolResult result = new AiToolResult(); + result.setSuccess(Boolean.TRUE); + result.setTool(tool); + result.setSummary(summary); + result.setData(copyData(data)); + return result; + } + + public static AiToolResult failure(String tool, String summary, String errorCode) { + AiToolResult result = new AiToolResult(); + result.setSuccess(Boolean.FALSE); + result.setTool(tool); + result.setSummary(summary); + result.setErrorCode(errorCode); + return result; + } + + private static List copyData(List data) { + if (data == null || data.isEmpty()) { + return new ArrayList<>(); + } + return new ArrayList<>(data); + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java index 033a7969d..a2608090a 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java @@ -10,6 +10,7 @@ import ai.chat2db.community.domain.api.model.request.datasource.DbDatabaseQueryAllRequest; import ai.chat2db.community.domain.api.model.request.runtime.DbConnectionContextRequest; import ai.chat2db.community.domain.api.model.PageResponse; +import ai.chat2db.community.domain.api.model.ai.AiToolResult; import ai.chat2db.community.domain.api.model.runtime.ConnectionProfile; import ai.chat2db.community.domain.api.service.db.IDbConnectionContextService; import ai.chat2db.community.domain.api.service.db.IDbDatabaseService; @@ -40,6 +41,8 @@ import ai.chat2db.community.domain.api.model.metadata.TableColumn; import ai.chat2db.community.domain.api.model.metadata.TableIndex; import ai.chat2db.community.domain.api.model.metadata.TableIndexColumn; +import com.alibaba.fastjson2.JSON; +import com.alibaba.fastjson2.JSONWriter; import org.apache.commons.collections4.CollectionUtils; import org.apache.commons.lang3.StringUtils; import org.springframework.beans.factory.annotation.Autowired; @@ -87,27 +90,26 @@ public String listAllDataSources(AiToolContextRequest toolContext) { toolContext, () -> workspaceStorageFacade.listDataSources(queryRequest)); if (Objects.isNull(result)) { - return emitToolResult(toolContext, "list_all_datasources", "Failed to query datasources: unknown error"); + return emitToolFailure(toolContext, "list_all_datasources", + "Failed to query datasources: unknown error", "DATASOURCE_QUERY_FAILED"); } if (CollectionUtils.isEmpty(result.getData())) { - return emitToolResult(toolContext, "list_all_datasources", "No datasources found."); + return emitToolResult(toolContext, "list_all_datasources", "No datasources found.", Collections.emptyList()); } - return emitToolResult(toolContext, "list_all_datasources", result.getData().stream() + List> data = result.getData().stream() .filter(Objects::nonNull) .map(dataSource -> { - List parts = new ArrayList<>(); - parts.add("id=" + dataSource.getId()); - parts.add("name=" + StringUtils.defaultIfBlank(dataSource.getAlias(), "(unnamed)")); - if (StringUtils.isNotBlank(dataSource.getType())) { - parts.add("type=" + dataSource.getType()); - } - if (StringUtils.isNotBlank(dataSource.getEnvType())) { - parts.add("env=" + dataSource.getEnvType()); - } - return String.join("; ", parts); + Map item = new LinkedHashMap<>(); + item.put("id", dataSource.getId()); + item.put("name", StringUtils.defaultIfBlank(dataSource.getAlias(), "(unnamed)")); + item.put("type", dataSource.getType()); + item.put("env", dataSource.getEnvType()); + return item; }) - .collect(Collectors.joining("\n"))); + .collect(Collectors.toList()); + return emitToolResult(toolContext, "list_all_datasources", + "Found " + data.size() + " datasource(s).", data); } public String listAllTables(AiListTablesRequest aiListTablesRequest) { Long dataSourceId = aiListTablesRequest == null ? null : aiListTablesRequest.getDataSourceId(); @@ -127,11 +129,14 @@ public String listAllTables(AiListTablesRequest aiListTablesRequest) { .build(); List result = tableService.queryTables(queryParam); if (CollectionUtils.isEmpty(result)) { - return emitToolResult(toolContext, "list_all_tables", "No tables found."); + return emitToolResult(toolContext, "list_all_tables", "No tables found.", Collections.emptyList()); } - return emitToolResult(toolContext, "list_all_tables", result.stream() - .map(this::formatTableSummary) - .collect(Collectors.joining("\n"))); + List> data = result.stream() + .filter(Objects::nonNull) + .map(this::tableSummaryData) + .collect(Collectors.toList()); + return emitToolResult(toolContext, "list_all_tables", + "Found " + data.size() + " table(s).", data); } finally { connectionContextService.clear(); } @@ -147,15 +152,20 @@ public String listAllDatabases(Long dataSourceId, .build(); List result = databaseService.queryAll(queryParam); if (CollectionUtils.isEmpty(result)) { - return emitToolResult(toolContext, "list_all_databases", "No databases found."); + return emitToolResult(toolContext, "list_all_databases", "No databases found.", Collections.emptyList()); } - return emitToolResult(toolContext, "list_all_databases", result.stream() + List> data = result.stream() + .filter(Objects::nonNull) .map(database -> { - String systemFlag = database.isSystem() ? " [SYSTEM]" : ""; - String comment = StringUtils.isBlank(database.getComment()) ? "" : " - " + database.getComment(); - return StringUtils.defaultString(database.getName(), "(unnamed)") + systemFlag + comment; + Map item = new LinkedHashMap<>(); + item.put("name", StringUtils.defaultString(database.getName(), "(unnamed)")); + item.put("system", database.isSystem()); + item.put("comment", database.getComment()); + return item; }) - .collect(Collectors.joining("\n"))); + .collect(Collectors.toList()); + return emitToolResult(toolContext, "list_all_databases", + "Found " + data.size() + " database(s).", data); } finally { connectionContextService.clear(); } @@ -165,7 +175,8 @@ public String listAllSchemas(String databaseName,Long dataSourceId, ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, databaseName, null); String targetDatabase = StringUtils.defaultIfBlank(databaseName, profile.getDatabaseName()); if (StringUtils.isBlank(targetDatabase)) { - return emitToolResult(toolContext, "list_all_schemas", "databaseName is required for listing schemas."); + return emitToolFailure(toolContext, "list_all_schemas", + "databaseName is required for listing schemas.", "INVALID_ARGUMENT"); } try { connectionContextService.bindProfile(profile); @@ -176,15 +187,20 @@ public String listAllSchemas(String databaseName,Long dataSourceId, .build(); List result = databaseService.querySchema(queryParam); if (CollectionUtils.isEmpty(result)) { - return emitToolResult(toolContext, "list_all_schemas", "No schemas found."); + return emitToolResult(toolContext, "list_all_schemas", "No schemas found.", Collections.emptyList()); } - return emitToolResult(toolContext, "list_all_schemas", result.stream() + List> data = result.stream() + .filter(Objects::nonNull) .map(schema -> { - String systemFlag = schema.isSystem() ? " [SYSTEM]" : ""; - String comment = StringUtils.isBlank(schema.getComment()) ? "" : " - " + schema.getComment(); - return StringUtils.defaultString(schema.getName(), "(unnamed)") + systemFlag + comment; + Map item = new LinkedHashMap<>(); + item.put("name", StringUtils.defaultString(schema.getName(), "(unnamed)")); + item.put("system", schema.isSystem()); + item.put("comment", schema.getComment()); + return item; }) - .collect(Collectors.joining("\n"))); + .collect(Collectors.toList()); + return emitToolResult(toolContext, "list_all_schemas", + "Found " + data.size() + " schema(s).", data); } finally { connectionContextService.clear(); } @@ -198,14 +214,14 @@ public String executeSql(AiExecuteSqlRequest aiExecuteSqlRequest) { AiToolContextRequest toolContext = aiExecuteSqlRequest == null ? null : aiExecuteSqlRequest.getAiToolContextRequest(); if (StringUtils.isBlank(sql)) { - return emitToolResult(toolContext, "execute_sql", "sql is empty."); + return emitToolFailure(toolContext, "execute_sql", "sql is empty.", "INVALID_ARGUMENT"); } ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, databaseName, schemaName); int resolvedPageSize = normalizePageSize(pageSize); String trimmedSql = sql.trim(); String unsafeSqlMessage = buildNonQueryExecutionMessage(trimmedSql, profile); if (StringUtils.isNotBlank(unsafeSqlMessage)) { - return emitToolResult(toolContext, "execute_sql", unsafeSqlMessage); + return emitToolFailure(toolContext, "execute_sql", unsafeSqlMessage, "SQL_REQUIRES_MANUAL_CONFIRMATION"); } boolean operationLogged = false; @@ -229,20 +245,22 @@ public String executeSql(AiExecuteSqlRequest aiExecuteSqlRequest) { sqlOperationLogRecorder.recordListResultAsync(sqlOperationLogListResultRequest); operationLogged = true; if (Objects.isNull(executeResult) || !executeResult.success()) { - return emitToolResult(toolContext, "execute_sql", "SQL execution failed: " - + (Objects.isNull(executeResult) ? "unknown error" : StringUtils.defaultString(executeResult.getErrorMessage()))); + return emitToolFailure(toolContext, "execute_sql", "SQL execution failed: " + + (Objects.isNull(executeResult) ? "unknown error" : StringUtils.defaultString(executeResult.getErrorMessage())), + "SQL_EXECUTION_FAILED"); } if (CollectionUtils.isEmpty(executeResult.getData())) { - return emitToolResult(toolContext, "execute_sql", "SQL executed successfully with no result."); + return emitToolResult(toolContext, "execute_sql", + "SQL executed successfully with no result.", Collections.emptyList()); } - StringBuilder output = new StringBuilder(2048); + List> data = new ArrayList<>(); int index = 1; for (ExecuteResponse item : executeResult.getData()) { - output.append("## Result ").append(index++).append("\n"); - output.append(formatExecuteResponse(item)).append("\n\n"); + data.add(executeResponseData(index++, item)); } - return emitToolResult(toolContext, "execute_sql", output.toString().trim()); + return emitToolResult(toolContext, "execute_sql", + "SQL executed successfully with " + data.size() + " result set(s).", data); } catch (RuntimeException e) { if (!operationLogged) { sqlOperationLogRecorder.recordFailureAsync(trimmedSql, SqlOperationLogSourceEnum.AI_TOOL.name(), e.getMessage()); @@ -260,7 +278,7 @@ public String getTablesSchema(AiGetTablesSchemaRequest aiGetTablesSchemaRequest) AiToolContextRequest toolContext = aiGetTablesSchemaRequest == null ? null : aiGetTablesSchemaRequest.getAiToolContextRequest(); if (CollectionUtils.isEmpty(tableNames)) { - return emitToolResult(toolContext, "get_tables_schema", "tableNames is empty."); + return emitToolFailure(toolContext, "get_tables_schema", "tableNames is empty.", "INVALID_ARGUMENT"); } List normalized = tableNames.stream() @@ -270,27 +288,43 @@ public String getTablesSchema(AiGetTablesSchemaRequest aiGetTablesSchemaRequest) .limit(20) .collect(Collectors.toList()); if (CollectionUtils.isEmpty(normalized)) { - return emitToolResult(toolContext, "get_tables_schema", "tableNames is empty."); + return emitToolFailure(toolContext, "get_tables_schema", "tableNames is empty.", "INVALID_ARGUMENT"); } ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, databaseName, schemaName); try { connectionContextService.bindProfile(profile); - StringBuilder ddlBuilder = new StringBuilder(4096); + List> data = new ArrayList<>(); for (String tableName : normalized) { Table table = fetchDetailedTable(profile, tableName); String ddl = fetchTableDdl(profile, tableName); - ddlBuilder.append(buildRichTableSchema(tableName, ddl, table)) - .append("\n\n"); + String schemaText = buildRichTableSchema(tableName, ddl, table); + Map item = new LinkedHashMap<>(); + item.put("tableName", tableName); + item.put("schema", schemaText); + data.add(item); } - return emitToolResult(toolContext, "get_tables_schema", ddlBuilder.toString()); + return emitToolResult(toolContext, "get_tables_schema", + "Fetched schema for " + data.size() + " table(s).", data); } finally { connectionContextService.clear(); } } - private String emitToolResult(AiToolContextRequest toolContext, String toolName, String content) { - return content; + private String emitToolResult(AiToolContextRequest toolContext, String toolName, String summary, List data) { + return successToolResultJson(toolName, summary, data); + } + + private String emitToolFailure(AiToolContextRequest toolContext, String toolName, String summary, String errorCode) { + return failureToolResultJson(toolName, summary, errorCode); + } + + static String successToolResultJson(String toolName, String summary, List data) { + return JSON.toJSONString(AiToolResult.success(toolName, summary, data), JSONWriter.Feature.WriteNulls); + } + + static String failureToolResultJson(String toolName, String summary, String errorCode) { + return JSON.toJSONString(AiToolResult.failure(toolName, summary, errorCode), JSONWriter.Feature.WriteNulls); } private T invokeWithRequestContext(AiToolContextRequest toolContext, java.util.function.Supplier supplier) { @@ -386,14 +420,12 @@ private Table fetchDetailedTable(ConnectionProfile profile, String tableName) { return table; } - private String formatTableSummary(SimpleTable table) { - StringBuilder builder = new StringBuilder(128); - builder.append(StringUtils.defaultString(table.getName(), "(unnamed)")); - builder.append(" [").append(StringUtils.defaultIfBlank(table.getTableType(), "TABLE")).append("]"); - if (StringUtils.isNotBlank(table.getComment())) { - builder.append(" - ").append(table.getComment()); - } - return builder.toString(); + private Map tableSummaryData(SimpleTable table) { + Map item = new LinkedHashMap<>(); + item.put("name", StringUtils.defaultString(table.getName(), "(unnamed)")); + item.put("type", StringUtils.defaultIfBlank(table.getTableType(), "TABLE")); + item.put("comment", table.getComment()); + return item; } private String buildRichTableSchema(String tableName, String ddl, Table table) { @@ -550,6 +582,58 @@ private String formatExecuteResponse(ExecuteResponse result) { return builder.toString().trim(); } + private Map executeResponseData(int index, ExecuteResponse result) { + Map item = new LinkedHashMap<>(); + item.put("resultIndex", index); + if (Objects.isNull(result)) { + item.put("success", false); + item.put("message", "Empty result."); + item.put("text", "Empty result."); + return item; + } + item.put("success", Boolean.TRUE.equals(result.getSuccess())); + item.put("sqlType", result.getSqlType()); + item.put("durationMs", result.getDuration()); + item.put("updateCount", result.getUpdateCount()); + item.put("message", result.getMessage()); + item.put("description", result.getDescription()); + item.put("hasNextPage", result.getHasNextPage()); + item.put("rowCount", result.getDataList() == null ? 0 : result.getDataList().size()); + item.put("columns", columnNames(result.getHeaderList())); + item.put("rows", rowPreview(result.getHeaderList(), result.getDisplayDataList())); + item.put("text", formatExecuteResponse(result)); + return item; + } + + private List columnNames(List
headers) { + if (CollectionUtils.isEmpty(headers)) { + return Collections.emptyList(); + } + return headers.stream() + .map(header -> StringUtils.defaultIfBlank(header.getName(), header.getColumnName())) + .map(name -> StringUtils.defaultIfBlank(name, "col")) + .collect(Collectors.toList()); + } + + private List> rowPreview(List
headers, List> rows) { + if (CollectionUtils.isEmpty(headers) || CollectionUtils.isEmpty(rows)) { + return Collections.emptyList(); + } + List headerNames = columnNames(headers); + int rowCount = Math.min(rows.size(), MAX_SQL_RESULT_ROWS); + List> result = new ArrayList<>(rowCount); + for (int i = 0; i < rowCount; i++) { + List row = rows.get(i); + Map rowData = new LinkedHashMap<>(); + for (int c = 0; c < headerNames.size(); c++) { + String value = row != null && c < row.size() ? row.get(c) : null; + rowData.put(headerNames.get(c), normalizeCell(value)); + } + result.add(rowData); + } + return result; + } + private ListResult wrapExecuteResults(List results) { ListResult result = ListResult.of(results); if (CollectionUtils.isEmpty(results)) { diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java new file mode 100644 index 000000000..7611e62a3 --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java @@ -0,0 +1,48 @@ +package ai.chat2db.community.domain.core.impl.ai; + +import com.alibaba.fastjson2.JSON; +import com.alibaba.fastjson2.JSONObject; +import org.junit.jupiter.api.Test; + +import java.util.List; +import java.util.Map; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +class AiToolServiceImplTest { + + @Test + void shouldSerializeSuccessfulToolResultAsStandardJson() { + String json = AiToolServiceImpl.successToolResultJson( + "execute_sql", + "SQL executed successfully with 1 result set(s).", + List.of(Map.of("rowCount", 2))); + + JSONObject result = JSON.parseObject(json); + + assertEquals(true, result.getBoolean("success")); + assertEquals("execute_sql", result.getString("tool")); + assertEquals("SQL executed successfully with 1 result set(s).", result.getString("summary")); + assertEquals(1, result.getJSONArray("data").size()); + assertNull(result.get("errorCode")); + assertTrue(json.contains("\"errorCode\":null")); + } + + @Test + void shouldSerializeFailedToolResultAsStandardJson() { + String json = AiToolServiceImpl.failureToolResultJson( + "execute_sql", + "sql is empty.", + "INVALID_ARGUMENT"); + + JSONObject result = JSON.parseObject(json); + + assertEquals(false, result.getBoolean("success")); + assertEquals("execute_sql", result.getString("tool")); + assertEquals("sql is empty.", result.getString("summary")); + assertEquals(0, result.getJSONArray("data").size()); + assertEquals("INVALID_ARGUMENT", result.getString("errorCode")); + } +} diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiChatStreamAdapter.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiChatStreamAdapter.java index 446dfe92d..c7dfa6ce6 100644 --- a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiChatStreamAdapter.java +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiChatStreamAdapter.java @@ -527,6 +527,9 @@ private String buildNl2SqlToolPrompt(boolean hasToolContext, boolean globalDatab - After choosing a datasource, pass dataSourceId and databaseName/schemaName when calling database tools. - Do not use execute_sql to run the final SQL answer. Generate SQL only. - If several datasources or schemas could satisfy the request, make the most reasonable schema-based choice and still return SQL only. + - Tool results are JSON objects with fields: success, tool, summary, data, errorCode. + - If success is false, use summary/errorCode to recover or explain why SQL cannot be generated from live metadata. + - When success is true, consume structured values from data instead of parsing summary text. """; } @@ -545,6 +548,9 @@ private String buildNl2SqlToolPrompt(boolean hasToolContext, boolean globalDatab - Use tools only when they help resolve schema, table, column, or dialect uncertainty. - Do not use execute_sql to run the final SQL answer. Generate SQL only. - The final answer must remain SQL only, even if tools are used. + - Tool results are JSON objects with fields: success, tool, summary, data, errorCode. + - If success is false, use summary/errorCode to recover or explain why SQL cannot be generated from live metadata. + - When success is true, consume structured values from data instead of parsing summary text. """; } @@ -594,6 +600,9 @@ private String buildDatabaseToolPrompt(boolean hasToolContext) { - If a chart, trend, report, aggregation, count, statistics, ranking, or comparison requires live database results, use execute_sql before answering. - Do not answer with only SQL templates, guessed values, or a markdown table when actual query results are required. - Prefer returning CREATE TABLE DDL when available. + - Tool results are JSON objects with fields: success, tool, summary, data, errorCode. + - If success is false, use summary/errorCode to decide whether to retry with narrower arguments or explain the failure. + - When success is true, consume structured values from data instead of parsing summary text. - Never output pseudo tool-call tags or XML-like tool invocation markup in the final answer. """; } diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java index d0aae4124..7db711a28 100644 --- a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java @@ -1,10 +1,13 @@ package ai.chat2db.community.web.api.adapter.ai; +import ai.chat2db.community.domain.api.model.ai.AiToolResult; import ai.chat2db.community.domain.api.model.request.ai.AiExecuteSqlRequest; import ai.chat2db.community.domain.api.model.request.ai.AiGetTablesSchemaRequest; import ai.chat2db.community.domain.api.model.request.ai.AiListTablesRequest; import ai.chat2db.community.domain.api.model.request.ai.AiToolContextRequest; import ai.chat2db.community.web.api.converter.ai.AiToolContextConverter; +import com.alibaba.fastjson2.JSON; +import com.alibaba.fastjson2.JSONWriter; import lombok.extern.slf4j.Slf4j; import org.apache.commons.lang3.StringUtils; import org.springframework.ai.chat.model.ToolContext; @@ -14,6 +17,7 @@ import java.util.List; import java.util.Map; +import java.util.function.Supplier; @Component @Slf4j @@ -30,8 +34,8 @@ public AiToolAdapter(ai.chat2db.community.domain.api.service.ai.IAiToolService a @Tool(name = "list_all_datasources", description = "List available Chat2DB data sources. Use this first when no datasource is selected.") public String listAllDataSources(ToolContext toolContext) { - return emit(toolContext, "list_all_datasources", - aiToolService.listAllDataSources(aiToolContextConverter.toParam(toolContext))); + return invoke(toolContext, "list_all_datasources", + () -> aiToolService.listAllDataSources(aiToolContextConverter.toParam(toolContext))); } @Tool(name = "list_all_tables", description = "List all tables in the connected database with comments and type.") @@ -40,8 +44,8 @@ public String listAllTables( @ToolParam(description = "Optional target database name. If omitted, uses selected database context.", required = false) String databaseName, @ToolParam(description = "Optional target schema name. If omitted, uses selected schema context.", required = false) String schemaName, ToolContext toolContext) { - return emit(toolContext, "list_all_tables", - aiToolService.listAllTables(listTablesRequest(dataSourceId, databaseName, schemaName, + return invoke(toolContext, "list_all_tables", + () -> aiToolService.listAllTables(listTablesRequest(dataSourceId, databaseName, schemaName, aiToolContextConverter.toParam(toolContext)))); } @@ -49,8 +53,8 @@ public String listAllTables( public String listAllDatabases( @ToolParam(description = "Optional datasource id. Required when no datasource is selected in context.", required = false) Long dataSourceId, ToolContext toolContext) { - return emit(toolContext, "list_all_databases", - aiToolService.listAllDatabases(dataSourceId, aiToolContextConverter.toParam(toolContext))); + return invoke(toolContext, "list_all_databases", + () -> aiToolService.listAllDatabases(dataSourceId, aiToolContextConverter.toParam(toolContext))); } @Tool(name = "list_all_schemas", description = "List all schemas in the selected database. If databaseName is empty, uses current database context.") @@ -58,8 +62,8 @@ public String listAllSchemas( @ToolParam(description = "Optional target database name", required = false) String databaseName, @ToolParam(description = "Optional datasource id. Required when no datasource is selected in context.", required = false) Long dataSourceId, ToolContext toolContext) { - return emit(toolContext, "list_all_schemas", - aiToolService.listAllSchemas(databaseName, dataSourceId, aiToolContextConverter.toParam(toolContext))); + return invoke(toolContext, "list_all_schemas", + () -> aiToolService.listAllSchemas(databaseName, dataSourceId, aiToolContextConverter.toParam(toolContext))); } @Tool(name = "execute_sql", description = "Execute SQL in current database context and return concise result (rows for SELECT, update count for DML/DDL).") @@ -70,8 +74,8 @@ public String executeSql( @ToolParam(description = "Optional target database name. If omitted, uses selected database context.", required = false) String databaseName, @ToolParam(description = "Optional target schema name. If omitted, uses selected schema context.", required = false) String schemaName, ToolContext toolContext) { - return emit(toolContext, "execute_sql", - aiToolService.executeSql(executeSqlRequest(sql, pageSize, dataSourceId, databaseName, schemaName, + return invoke(toolContext, "execute_sql", + () -> aiToolService.executeSql(executeSqlRequest(sql, pageSize, dataSourceId, databaseName, schemaName, aiToolContextConverter.toParam(toolContext)))); } @@ -82,8 +86,8 @@ public String getTablesSchema( @ToolParam(description = "Optional target database name. If omitted, uses selected database context.", required = false) String databaseName, @ToolParam(description = "Optional target schema name. If omitted, uses selected schema context.", required = false) String schemaName, ToolContext toolContext) { - return emit(toolContext, "get_tables_schema", - aiToolService.getTablesSchema(tablesSchemaRequest(tableNames, dataSourceId, databaseName, schemaName, + return invoke(toolContext, "get_tables_schema", + () -> aiToolService.getTablesSchema(tablesSchemaRequest(tableNames, dataSourceId, databaseName, schemaName, aiToolContextConverter.toParam(toolContext)))); } @@ -121,6 +125,18 @@ private AiGetTablesSchemaRequest tablesSchemaRequest(List tableNames, Lo return request; } + private String invoke(ToolContext toolContext, String toolName, Supplier action) { + try { + return emit(toolContext, toolName, action.get()); + } catch (Exception e) { + log.error("AI tool call failed, tool={}", toolName, e); + String message = "Tool call failed: " + StringUtils.defaultIfBlank(e.getMessage(), "Unknown error"); + return emit(toolContext, toolName, + JSON.toJSONString(AiToolResult.failure(toolName, message, "TOOL_CALL_FAILED"), + JSONWriter.Feature.WriteNulls)); + } + } + private String emit(ToolContext toolContext, String toolName, String content) { Map payload = AiChatTraceSupport.payload(AiChatTraceSupport.TYPE_TOOL_RESULT); payload.put("name", toolName); diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java index ab97792fc..d51e19f96 100644 --- a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java @@ -2,12 +2,15 @@ import ai.chat2db.community.tools.annotation.NotCliRuntime; +import ai.chat2db.community.domain.api.model.ai.AiToolResult; import ai.chat2db.community.web.api.enums.ai.QuestionTypeEnum; import ai.chat2db.community.web.api.model.request.ai.ChatRequest; import ai.chat2db.community.web.api.adapter.ai.AiChatStreamAdapter; import ai.chat2db.community.web.api.adapter.ai.AiToolAdapter; import ai.chat2db.community.tools.model.Context; import ai.chat2db.community.tools.util.ContextUtils; +import com.alibaba.fastjson2.JSON; +import com.alibaba.fastjson2.JSONWriter; import lombok.extern.slf4j.Slf4j; import org.apache.commons.lang3.StringUtils; import org.springframework.ai.chat.model.ToolContext; @@ -103,7 +106,7 @@ public String text2sql( return aiChatStreamAdapter.chatSync(chatRequest); } catch (Exception e) { log.error("MCP tool call failed, tool=text2sql", e); - return "MCP tool 'text2sql' failed: " + StringUtils.defaultIfBlank(e.getMessage(), "Unknown error"); + return toolFailure("text2sql", e); } } @@ -113,11 +116,16 @@ private String invoke(String toolName, Function action) { return action.apply(toolContext); } catch (Exception e) { log.error("MCP tool call failed, tool={}", toolName, e); - String message = StringUtils.defaultIfBlank(e.getMessage(), "Unknown error"); - return "MCP tool '" + toolName + "' failed: " + message; + return toolFailure(toolName, e); } } + private String toolFailure(String toolName, Exception e) { + String message = "MCP tool call failed: " + StringUtils.defaultIfBlank(e.getMessage(), "Unknown error"); + return JSON.toJSONString(AiToolResult.failure(toolName, message, "MCP_TOOL_CALL_FAILED"), + JSONWriter.Feature.WriteNulls); + } + private ToolContext buildToolContext() { Map context = new LinkedHashMap<>(); Context requestContext = ContextUtils.queryContext(); From 72fc8427552ae0296f8d7da2b587a41235e1edcd Mon Sep 17 00:00:00 2001 From: meteor <1839462565@qq.com> Date: Mon, 20 Jul 2026 20:55:41 +0800 Subject: [PATCH 2/6] refactor(ai): remove tool name from AI tool result payload --- .../community/domain/api/model/ai/AiToolResult.java | 8 ++------ .../domain/core/impl/ai/AiToolServiceImpl.java | 12 ++++++------ .../domain/core/impl/ai/AiToolServiceImplTest.java | 6 ++---- .../web/api/adapter/ai/AiChatStreamAdapter.java | 6 +++--- .../community/web/api/adapter/ai/AiToolAdapter.java | 2 +- .../web/api/mcp/adapter/AiToolMcpAdapter.java | 2 +- 6 files changed, 15 insertions(+), 21 deletions(-) diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java index 050a110b0..e12a943ba 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java @@ -10,27 +10,23 @@ public class AiToolResult { private Boolean success; - private String tool; - private String summary; private List data = new ArrayList<>(); private String errorCode; - public static AiToolResult success(String tool, String summary, List data) { + public static AiToolResult success(String summary, List data) { AiToolResult result = new AiToolResult(); result.setSuccess(Boolean.TRUE); - result.setTool(tool); result.setSummary(summary); result.setData(copyData(data)); return result; } - public static AiToolResult failure(String tool, String summary, String errorCode) { + public static AiToolResult failure(String summary, String errorCode) { AiToolResult result = new AiToolResult(); result.setSuccess(Boolean.FALSE); - result.setTool(tool); result.setSummary(summary); result.setErrorCode(errorCode); return result; diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java index a2608090a..dd9ba4efd 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java @@ -312,19 +312,19 @@ public String getTablesSchema(AiGetTablesSchemaRequest aiGetTablesSchemaRequest) } private String emitToolResult(AiToolContextRequest toolContext, String toolName, String summary, List data) { - return successToolResultJson(toolName, summary, data); + return successToolResultJson(summary, data); } private String emitToolFailure(AiToolContextRequest toolContext, String toolName, String summary, String errorCode) { - return failureToolResultJson(toolName, summary, errorCode); + return failureToolResultJson(summary, errorCode); } - static String successToolResultJson(String toolName, String summary, List data) { - return JSON.toJSONString(AiToolResult.success(toolName, summary, data), JSONWriter.Feature.WriteNulls); + static String successToolResultJson(String summary, List data) { + return JSON.toJSONString(AiToolResult.success(summary, data), JSONWriter.Feature.WriteNulls); } - static String failureToolResultJson(String toolName, String summary, String errorCode) { - return JSON.toJSONString(AiToolResult.failure(toolName, summary, errorCode), JSONWriter.Feature.WriteNulls); + static String failureToolResultJson(String summary, String errorCode) { + return JSON.toJSONString(AiToolResult.failure(summary, errorCode), JSONWriter.Feature.WriteNulls); } private T invokeWithRequestContext(AiToolContextRequest toolContext, java.util.function.Supplier supplier) { diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java index 7611e62a3..dca26143b 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java @@ -16,14 +16,13 @@ class AiToolServiceImplTest { @Test void shouldSerializeSuccessfulToolResultAsStandardJson() { String json = AiToolServiceImpl.successToolResultJson( - "execute_sql", "SQL executed successfully with 1 result set(s).", List.of(Map.of("rowCount", 2))); JSONObject result = JSON.parseObject(json); assertEquals(true, result.getBoolean("success")); - assertEquals("execute_sql", result.getString("tool")); + assertNull(result.get("tool")); assertEquals("SQL executed successfully with 1 result set(s).", result.getString("summary")); assertEquals(1, result.getJSONArray("data").size()); assertNull(result.get("errorCode")); @@ -33,14 +32,13 @@ void shouldSerializeSuccessfulToolResultAsStandardJson() { @Test void shouldSerializeFailedToolResultAsStandardJson() { String json = AiToolServiceImpl.failureToolResultJson( - "execute_sql", "sql is empty.", "INVALID_ARGUMENT"); JSONObject result = JSON.parseObject(json); assertEquals(false, result.getBoolean("success")); - assertEquals("execute_sql", result.getString("tool")); + assertNull(result.get("tool")); assertEquals("sql is empty.", result.getString("summary")); assertEquals(0, result.getJSONArray("data").size()); assertEquals("INVALID_ARGUMENT", result.getString("errorCode")); diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiChatStreamAdapter.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiChatStreamAdapter.java index c7dfa6ce6..e919611b6 100644 --- a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiChatStreamAdapter.java +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiChatStreamAdapter.java @@ -527,7 +527,7 @@ private String buildNl2SqlToolPrompt(boolean hasToolContext, boolean globalDatab - After choosing a datasource, pass dataSourceId and databaseName/schemaName when calling database tools. - Do not use execute_sql to run the final SQL answer. Generate SQL only. - If several datasources or schemas could satisfy the request, make the most reasonable schema-based choice and still return SQL only. - - Tool results are JSON objects with fields: success, tool, summary, data, errorCode. + - Tool results are JSON objects with fields: success, summary, data, errorCode. - If success is false, use summary/errorCode to recover or explain why SQL cannot be generated from live metadata. - When success is true, consume structured values from data instead of parsing summary text. """; @@ -548,7 +548,7 @@ private String buildNl2SqlToolPrompt(boolean hasToolContext, boolean globalDatab - Use tools only when they help resolve schema, table, column, or dialect uncertainty. - Do not use execute_sql to run the final SQL answer. Generate SQL only. - The final answer must remain SQL only, even if tools are used. - - Tool results are JSON objects with fields: success, tool, summary, data, errorCode. + - Tool results are JSON objects with fields: success, summary, data, errorCode. - If success is false, use summary/errorCode to recover or explain why SQL cannot be generated from live metadata. - When success is true, consume structured values from data instead of parsing summary text. """; @@ -600,7 +600,7 @@ private String buildDatabaseToolPrompt(boolean hasToolContext) { - If a chart, trend, report, aggregation, count, statistics, ranking, or comparison requires live database results, use execute_sql before answering. - Do not answer with only SQL templates, guessed values, or a markdown table when actual query results are required. - Prefer returning CREATE TABLE DDL when available. - - Tool results are JSON objects with fields: success, tool, summary, data, errorCode. + - Tool results are JSON objects with fields: success, summary, data, errorCode. - If success is false, use summary/errorCode to decide whether to retry with narrower arguments or explain the failure. - When success is true, consume structured values from data instead of parsing summary text. - Never output pseudo tool-call tags or XML-like tool invocation markup in the final answer. diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java index 7db711a28..af0d70e72 100644 --- a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java @@ -132,7 +132,7 @@ private String invoke(ToolContext toolContext, String toolName, Supplier log.error("AI tool call failed, tool={}", toolName, e); String message = "Tool call failed: " + StringUtils.defaultIfBlank(e.getMessage(), "Unknown error"); return emit(toolContext, toolName, - JSON.toJSONString(AiToolResult.failure(toolName, message, "TOOL_CALL_FAILED"), + JSON.toJSONString(AiToolResult.failure(message, "TOOL_CALL_FAILED"), JSONWriter.Feature.WriteNulls)); } } diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java index d51e19f96..68b217a47 100644 --- a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java @@ -122,7 +122,7 @@ private String invoke(String toolName, Function action) { private String toolFailure(String toolName, Exception e) { String message = "MCP tool call failed: " + StringUtils.defaultIfBlank(e.getMessage(), "Unknown error"); - return JSON.toJSONString(AiToolResult.failure(toolName, message, "MCP_TOOL_CALL_FAILED"), + return JSON.toJSONString(AiToolResult.failure(message, "MCP_TOOL_CALL_FAILED"), JSONWriter.Feature.WriteNulls); } From c954b13ca52bf13e91ffd61327e438ba111f8a57 Mon Sep 17 00:00:00 2001 From: meteor <1839462565@qq.com> Date: Mon, 20 Jul 2026 21:34:08 +0800 Subject: [PATCH 3/6] fix(ai): preserve structured tool result fidelity --- .../core/impl/ai/AiToolServiceImpl.java | 14 ++++----- .../core/impl/ai/AiToolServiceImplTest.java | 21 +++++++++++++ .../web/api/adapter/ai/AiToolAdapter.java | 16 +++++++++- .../web/api/mcp/adapter/AiToolMcpAdapter.java | 10 ++++++- .../web/api/adapter/ai/AiToolAdapterTest.java | 27 +++++++++++++++++ .../api/mcp/adapter/AiToolMcpAdapterTest.java | 30 +++++++++++++++++++ 6 files changed, 109 insertions(+), 9 deletions(-) create mode 100644 chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapterTest.java create mode 100644 chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapterTest.java diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java index dd9ba4efd..d11401186 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java @@ -600,12 +600,12 @@ private Map executeResponseData(int index, ExecuteResponse resul item.put("hasNextPage", result.getHasNextPage()); item.put("rowCount", result.getDataList() == null ? 0 : result.getDataList().size()); item.put("columns", columnNames(result.getHeaderList())); - item.put("rows", rowPreview(result.getHeaderList(), result.getDisplayDataList())); + item.put("rows", rowPreviewRows(result.getHeaderList(), result.getDisplayDataList())); item.put("text", formatExecuteResponse(result)); return item; } - private List columnNames(List
headers) { + static List columnNames(List
headers) { if (CollectionUtils.isEmpty(headers)) { return Collections.emptyList(); } @@ -615,19 +615,19 @@ private List columnNames(List
headers) { .collect(Collectors.toList()); } - private List> rowPreview(List
headers, List> rows) { + static List> rowPreviewRows(List
headers, List> rows) { if (CollectionUtils.isEmpty(headers) || CollectionUtils.isEmpty(rows)) { return Collections.emptyList(); } List headerNames = columnNames(headers); int rowCount = Math.min(rows.size(), MAX_SQL_RESULT_ROWS); - List> result = new ArrayList<>(rowCount); + List> result = new ArrayList<>(rowCount); for (int i = 0; i < rowCount; i++) { List row = rows.get(i); - Map rowData = new LinkedHashMap<>(); + List rowData = new ArrayList<>(headerNames.size()); for (int c = 0; c < headerNames.size(); c++) { String value = row != null && c < row.size() ? row.get(c) : null; - rowData.put(headerNames.get(c), normalizeCell(value)); + rowData.add(value); } result.add(rowData); } @@ -672,7 +672,7 @@ private void appendTabularPreview(StringBuilder builder, List
headers, L } } - private String normalizeCell(String value) { + private static String normalizeCell(String value) { if (value == null) { return "NULL"; } diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java index dca26143b..9a1b28ed8 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java @@ -1,9 +1,11 @@ package ai.chat2db.community.domain.core.impl.ai; +import ai.chat2db.community.domain.api.model.result.Header; import com.alibaba.fastjson2.JSON; import com.alibaba.fastjson2.JSONObject; import org.junit.jupiter.api.Test; +import java.util.Arrays; import java.util.List; import java.util.Map; @@ -43,4 +45,23 @@ void shouldSerializeFailedToolResultAsStandardJson() { assertEquals(0, result.getJSONArray("data").size()); assertEquals("INVALID_ARGUMENT", result.getString("errorCode")); } + + @Test + void shouldKeepRowsPositionBasedWhenColumnNamesAreDuplicated() { + List
headers = List.of( + Header.builder().name("id").build(), + Header.builder().name("id").build(), + Header.builder().name("note").build()); + String longText = "a".repeat(201); + + List> rows = AiToolServiceImpl.rowPreviewRows( + headers, + List.of(Arrays.asList("first-id\nwith\ttab", null, longText))); + + assertEquals(List.of("id", "id", "note"), AiToolServiceImpl.columnNames(headers)); + assertEquals(1, rows.size()); + assertEquals("first-id\nwith\ttab", rows.get(0).get(0)); + assertNull(rows.get(0).get(1)); + assertEquals(longText, rows.get(0).get(2)); + } } diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java index af0d70e72..18ccd7a7d 100644 --- a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java @@ -7,6 +7,7 @@ import ai.chat2db.community.domain.api.model.request.ai.AiToolContextRequest; import ai.chat2db.community.web.api.converter.ai.AiToolContextConverter; import com.alibaba.fastjson2.JSON; +import com.alibaba.fastjson2.JSONObject; import com.alibaba.fastjson2.JSONWriter; import lombok.extern.slf4j.Slf4j; import org.apache.commons.lang3.StringUtils; @@ -140,9 +141,22 @@ private String invoke(ToolContext toolContext, String toolName, Supplier private String emit(ToolContext toolContext, String toolName, String content) { Map payload = AiChatTraceSupport.payload(AiChatTraceSupport.TYPE_TOOL_RESULT); payload.put("name", toolName); - payload.put("content", StringUtils.defaultString(content)); + payload.put("content", traceSummary(content)); AiChatTraceSupport.emit(toolContext, payload); return content; } + static String traceSummary(String content) { + if (StringUtils.isBlank(content)) { + return StringUtils.EMPTY; + } + try { + JSONObject result = JSON.parseObject(content); + String summary = result.getString("summary"); + return StringUtils.defaultIfBlank(summary, content); + } catch (Exception ignored) { + return content; + } + } + } diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java index 68b217a47..67d8f75f5 100644 --- a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java @@ -103,7 +103,10 @@ public String text2sql( chatRequest.setQuestionType(QuestionTypeEnum.NL_2_SQL.getCode()); chatRequest.setEnableTools(Boolean.TRUE); - return aiChatStreamAdapter.chatSync(chatRequest); + String sql = aiChatStreamAdapter.chatSync(chatRequest); + return toolSuccess( + "SQL generated successfully.", + List.of(Map.of("sql", StringUtils.defaultString(sql)))); } catch (Exception e) { log.error("MCP tool call failed, tool=text2sql", e); return toolFailure("text2sql", e); @@ -126,6 +129,11 @@ private String toolFailure(String toolName, Exception e) { JSONWriter.Feature.WriteNulls); } + static String toolSuccess(String summary, List data) { + return JSON.toJSONString(AiToolResult.success(summary, data), + JSONWriter.Feature.WriteNulls); + } + private ToolContext buildToolContext() { Map context = new LinkedHashMap<>(); Context requestContext = ContextUtils.queryContext(); diff --git a/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapterTest.java b/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapterTest.java new file mode 100644 index 000000000..505a83f84 --- /dev/null +++ b/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapterTest.java @@ -0,0 +1,27 @@ +package ai.chat2db.community.web.api.adapter.ai; + +import org.junit.jupiter.api.Test; + +import static org.junit.jupiter.api.Assertions.assertEquals; + +class AiToolAdapterTest { + + @Test + void shouldUseSummaryForStructuredToolTraceContent() { + String content = """ + { + "success": true, + "summary": "SQL executed successfully.", + "data": [{"rows": [["very large raw payload"]]}], + "errorCode": null + } + """; + + assertEquals("SQL executed successfully.", AiToolAdapter.traceSummary(content)); + } + + @Test + void shouldKeepRawTraceContentWhenToolResultIsNotJson() { + assertEquals("plain result", AiToolAdapter.traceSummary("plain result")); + } +} diff --git a/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapterTest.java b/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapterTest.java new file mode 100644 index 000000000..d79afdb7f --- /dev/null +++ b/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapterTest.java @@ -0,0 +1,30 @@ +package ai.chat2db.community.web.api.mcp.adapter; + +import com.alibaba.fastjson2.JSON; +import com.alibaba.fastjson2.JSONObject; +import org.junit.jupiter.api.Test; + +import java.util.List; +import java.util.Map; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +class AiToolMcpAdapterTest { + + @Test + void shouldSerializeText2SqlSuccessAsStandardJson() { + String json = AiToolMcpAdapter.toolSuccess( + "SQL generated successfully.", + List.of(Map.of("sql", "select * from users"))); + + JSONObject result = JSON.parseObject(json); + + assertEquals(true, result.getBoolean("success")); + assertEquals("SQL generated successfully.", result.getString("summary")); + assertEquals("select * from users", result.getJSONArray("data").getJSONObject(0).getString("sql")); + assertNull(result.getString("errorCode")); + assertTrue(json.contains("\"errorCode\":null")); + } +} From 968e8f0625e80bc9b54ed2e8921649c558a82144 Mon Sep 17 00:00:00 2001 From: meteor <1839462565@qq.com> Date: Mon, 20 Jul 2026 21:52:34 +0800 Subject: [PATCH 4/6] fix(ai): refine structured tool result payload --- .../domain/api/service/ai/IAiToolService.java | 24 ++++++++++--------- .../core/impl/ai/AiToolServiceImpl.java | 9 ++++++- .../core/impl/ai/AiToolServiceImplTest.java | 5 ++-- 3 files changed, 24 insertions(+), 14 deletions(-) diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/service/ai/IAiToolService.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/service/ai/IAiToolService.java index 412b5a618..a97d18cc6 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/service/ai/IAiToolService.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/service/ai/IAiToolService.java @@ -7,41 +7,43 @@ /** * Provides AI tool operations over datasource metadata and SQL execution. + * + *

Methods return a serialized AiToolResult JSON string with success, summary, data, and errorCode fields. */ public interface IAiToolService { /** - * Returns datasource context text for AI tools. + * Returns datasource metadata for AI tools. * * @param aiToolContextRequest AI tool context parameters. - * @return datasource context text. + * @return serialized AiToolResult JSON string. */ String listAllDataSources(AiToolContextRequest aiToolContextRequest); /** - * Returns table context text for AI tools. + * Returns table metadata for AI tools. * * @param aiListTablesRequest AI table listing parameters. - * @return table context text. + * @return serialized AiToolResult JSON string. */ String listAllTables(AiListTablesRequest aiListTablesRequest); /** - * Returns database context text for AI tools. + * Returns database metadata for AI tools. * * @param dataSourceId datasource identifier. * @param aiToolContextRequest AI tool context parameters. - * @return database context text. + * @return serialized AiToolResult JSON string. */ String listAllDatabases(Long dataSourceId, AiToolContextRequest aiToolContextRequest); /** - * Returns schema context text for AI tools. + * Returns schema metadata for AI tools. * * @param databaseName database name that scopes the lookup. * @param dataSourceId datasource identifier. * @param aiToolContextRequest AI tool context parameters. - * @return schema context text. + * @return serialized AiToolResult JSON string. */ String listAllSchemas(String databaseName, Long dataSourceId, AiToolContextRequest aiToolContextRequest); @@ -49,15 +51,15 @@ public interface IAiToolService { * Executes SQL for an AI tool request. * * @param aiExecuteSqlRequest AI SQL execution parameters. - * @return SQL execution context text. + * @return serialized AiToolResult JSON string. */ String executeSql(AiExecuteSqlRequest aiExecuteSqlRequest); /** - * Returns table schema context text for AI tools. + * Returns table schema metadata for AI tools. * * @param aiGetTablesSchemaRequest AI table schema lookup parameters. - * @return table schema context text. + * @return serialized AiToolResult JSON string. */ String getTablesSchema(AiGetTablesSchemaRequest aiGetTablesSchemaRequest); } diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java index d11401186..4e5f41f1d 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java @@ -627,7 +627,7 @@ static List> rowPreviewRows(List

headers, List List rowData = new ArrayList<>(headerNames.size()); for (int c = 0; c < headerNames.size(); c++) { String value = row != null && c < row.size() ? row.get(c) : null; - rowData.add(value); + rowData.add(normalizeStructuredCell(value)); } result.add(rowData); } @@ -683,6 +683,13 @@ private static String normalizeCell(String value) { return normalized; } + private static Object normalizeStructuredCell(String value) { + if (value == null) { + return null; + } + return normalizeCell(value); + } + private int normalizePageSize(Integer pageSize) { if (Objects.isNull(pageSize) || pageSize <= 0) { return DEFAULT_SQL_PAGE_SIZE; diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java index 9a1b28ed8..28c592070 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java @@ -60,8 +60,9 @@ void shouldKeepRowsPositionBasedWhenColumnNamesAreDuplicated() { assertEquals(List.of("id", "id", "note"), AiToolServiceImpl.columnNames(headers)); assertEquals(1, rows.size()); - assertEquals("first-id\nwith\ttab", rows.get(0).get(0)); + assertEquals("first-id\\nwith tab", rows.get(0).get(0)); assertNull(rows.get(0).get(1)); - assertEquals(longText, rows.get(0).get(2)); + assertEquals(200, rows.get(0).get(2).toString().length()); + assertTrue(rows.get(0).get(2).toString().endsWith("...")); } } From a1fb0a1acfdc0a6e947bf2844bd715f4b7536dfd Mon Sep 17 00:00:00 2001 From: meteor <1839462565@qq.com> Date: Mon, 20 Jul 2026 22:07:47 +0800 Subject: [PATCH 5/6] fix(ai): preserve raw structured SQL rows --- .../core/impl/ai/AiToolServiceImpl.java | 21 ++++++-------- .../core/impl/ai/AiToolServiceImplTest.java | 29 +++++++++++++++++-- 2 files changed, 35 insertions(+), 15 deletions(-) diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java index 4e5f41f1d..c58179150 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java @@ -550,7 +550,7 @@ private String firstNonBlank(String... values) { return null; } - private String formatExecuteResponse(ExecuteResponse result) { + private static String formatExecuteResponse(ExecuteResponse result) { if (Objects.isNull(result)) { return "Empty result."; } @@ -582,7 +582,7 @@ private String formatExecuteResponse(ExecuteResponse result) { return builder.toString().trim(); } - private Map executeResponseData(int index, ExecuteResponse result) { + static Map executeResponseData(int index, ExecuteResponse result) { Map item = new LinkedHashMap<>(); item.put("resultIndex", index); if (Objects.isNull(result)) { @@ -598,7 +598,11 @@ private Map executeResponseData(int index, ExecuteResponse resul item.put("message", result.getMessage()); item.put("description", result.getDescription()); item.put("hasNextPage", result.getHasNextPage()); - item.put("rowCount", result.getDataList() == null ? 0 : result.getDataList().size()); + int rowCount = result.getDataList() == null ? 0 : result.getDataList().size(); + int previewRowCount = Math.min(rowCount, MAX_SQL_RESULT_ROWS); + item.put("rowCount", rowCount); + item.put("previewRowCount", previewRowCount); + item.put("rowsTruncated", rowCount > previewRowCount); item.put("columns", columnNames(result.getHeaderList())); item.put("rows", rowPreviewRows(result.getHeaderList(), result.getDisplayDataList())); item.put("text", formatExecuteResponse(result)); @@ -627,7 +631,7 @@ static List> rowPreviewRows(List
headers, List List rowData = new ArrayList<>(headerNames.size()); for (int c = 0; c < headerNames.size(); c++) { String value = row != null && c < row.size() ? row.get(c) : null; - rowData.add(normalizeStructuredCell(value)); + rowData.add(value); } result.add(rowData); } @@ -651,7 +655,7 @@ private ListResult wrapExecuteResults(List res return result; } - private void appendTabularPreview(StringBuilder builder, List
headers, List> rows) { + private static void appendTabularPreview(StringBuilder builder, List
headers, List> rows) { List headerNames = headers.stream() .map(header -> StringUtils.defaultIfBlank(header.getName(), header.getColumnName())) .map(name -> StringUtils.defaultIfBlank(name, "col")) @@ -683,13 +687,6 @@ private static String normalizeCell(String value) { return normalized; } - private static Object normalizeStructuredCell(String value) { - if (value == null) { - return null; - } - return normalizeCell(value); - } - private int normalizePageSize(Integer pageSize) { if (Objects.isNull(pageSize) || pageSize <= 0) { return DEFAULT_SQL_PAGE_SIZE; diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java index 28c592070..07197e9a5 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java @@ -1,10 +1,13 @@ package ai.chat2db.community.domain.core.impl.ai; import ai.chat2db.community.domain.api.model.result.Header; +import ai.chat2db.community.domain.api.model.result.ExecuteResponse; +import ai.chat2db.community.domain.api.model.result.ResultCell; import com.alibaba.fastjson2.JSON; import com.alibaba.fastjson2.JSONObject; import org.junit.jupiter.api.Test; +import java.util.ArrayList; import java.util.Arrays; import java.util.List; import java.util.Map; @@ -60,9 +63,29 @@ void shouldKeepRowsPositionBasedWhenColumnNamesAreDuplicated() { assertEquals(List.of("id", "id", "note"), AiToolServiceImpl.columnNames(headers)); assertEquals(1, rows.size()); - assertEquals("first-id\\nwith tab", rows.get(0).get(0)); + assertEquals("first-id\nwith\ttab", rows.get(0).get(0)); assertNull(rows.get(0).get(1)); - assertEquals(200, rows.get(0).get(2).toString().length()); - assertTrue(rows.get(0).get(2).toString().endsWith("...")); + assertEquals(longText, rows.get(0).get(2)); + } + + @Test + void shouldExposeRowPreviewTruncationMetadata() { + Header header = Header.builder().name("id").build(); + List> rows = new ArrayList<>(); + for (int i = 0; i < 51; i++) { + rows.add(List.of(ResultCell.of(String.valueOf(i)))); + } + ExecuteResponse response = ExecuteResponse.builder() + .success(true) + .headerList(List.of(header)) + .dataList(rows) + .build(); + + Map result = AiToolServiceImpl.executeResponseData(1, response); + + assertEquals(51, result.get("rowCount")); + assertEquals(50, result.get("previewRowCount")); + assertEquals(true, result.get("rowsTruncated")); + assertEquals(50, ((List) result.get("rows")).size()); } } From f398d8f91103c623391cc84aea24a16f85bc66b1 Mon Sep 17 00:00:00 2001 From: meteor <1839462565@qq.com> Date: Tue, 21 Jul 2026 21:14:44 +0800 Subject: [PATCH 6/6] refactor(ai-tool): move result envelope to adapter boundary --- .../api/exception/ai/AiToolException.java | 12 + .../ai/AiToolInvalidArgumentException.java | 12 + .../ai/AiToolMetadataQueryException.java | 12 + ...iToolSqlConfirmationRequiredException.java | 12 + .../ai/AiToolSqlExecutionException.java | 12 + .../domain/api/model/ai/AiToolResult.java | 24 +- .../api/model/ai/DataSourceToolData.java | 26 + .../domain/api/model/ai/DatabaseToolData.java | 25 + .../domain/api/model/ai/SchemaToolData.java | 25 + .../domain/api/model/ai/SqlToolData.java | 57 ++ .../api/model/ai/TableSchemaResult.java | 22 + .../api/model/ai/TableSchemaToolData.java | 24 + .../domain/api/model/ai/TableToolData.java | 25 + .../domain/api/model/ai/Text2SqlToolData.java | 13 + .../domain/api/model/result/ResultCell.java | 5 + .../domain/api/service/ai/IAiToolService.java | 34 +- .../core/impl/ai/AiToolServiceImpl.java | 718 ++++++------------ .../core/impl/ai/AiToolServiceImplTest.java | 476 ++++++++++-- .../web/api/adapter/ai/AiToolAdapter.java | 80 +- .../converter/ai/AiToolErrorCodeMapper.java | 28 + .../web/api/converter/ai/AiToolOutput.java | 8 + .../converter/ai/AiToolResultConverter.java | 553 ++++++++++++++ .../converter/ai/AiToolResultSerializer.java | 14 + .../web/api/mcp/adapter/AiToolMcpAdapter.java | 122 ++- .../web/api/adapter/ai/AiToolAdapterTest.java | 172 ++++- .../ai/AiToolResultConverterTest.java | 309 ++++++++ .../api/mcp/adapter/AiToolMcpAdapterTest.java | 155 +++- 27 files changed, 2315 insertions(+), 660 deletions(-) create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolException.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolInvalidArgumentException.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolMetadataQueryException.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolSqlConfirmationRequiredException.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolSqlExecutionException.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/DataSourceToolData.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/DatabaseToolData.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/SchemaToolData.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/SqlToolData.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableSchemaResult.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableSchemaToolData.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableToolData.java create mode 100644 chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/Text2SqlToolData.java create mode 100644 chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolErrorCodeMapper.java create mode 100644 chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolOutput.java create mode 100644 chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolResultConverter.java create mode 100644 chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolResultSerializer.java create mode 100644 chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/converter/ai/AiToolResultConverterTest.java diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolException.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolException.java new file mode 100644 index 000000000..18521cc3a --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolException.java @@ -0,0 +1,12 @@ +package ai.chat2db.community.domain.api.exception.ai; + +public class AiToolException extends RuntimeException { + + protected AiToolException(String message) { + super(message); + } + + protected AiToolException(String message, Throwable cause) { + super(message, cause); + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolInvalidArgumentException.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolInvalidArgumentException.java new file mode 100644 index 000000000..9aac754f1 --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolInvalidArgumentException.java @@ -0,0 +1,12 @@ +package ai.chat2db.community.domain.api.exception.ai; + +public class AiToolInvalidArgumentException extends AiToolException { + + public AiToolInvalidArgumentException(String message) { + super(message); + } + + public AiToolInvalidArgumentException(String message, Throwable cause) { + super(message, cause); + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolMetadataQueryException.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolMetadataQueryException.java new file mode 100644 index 000000000..b2424fbb0 --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolMetadataQueryException.java @@ -0,0 +1,12 @@ +package ai.chat2db.community.domain.api.exception.ai; + +public class AiToolMetadataQueryException extends AiToolException { + + public AiToolMetadataQueryException(String message) { + super(message); + } + + public AiToolMetadataQueryException(String message, Throwable cause) { + super(message, cause); + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolSqlConfirmationRequiredException.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolSqlConfirmationRequiredException.java new file mode 100644 index 000000000..747892d7f --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolSqlConfirmationRequiredException.java @@ -0,0 +1,12 @@ +package ai.chat2db.community.domain.api.exception.ai; + +public class AiToolSqlConfirmationRequiredException extends AiToolException { + + public AiToolSqlConfirmationRequiredException(String message) { + super(message); + } + + public AiToolSqlConfirmationRequiredException(String message, Throwable cause) { + super(message, cause); + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolSqlExecutionException.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolSqlExecutionException.java new file mode 100644 index 000000000..431ccb456 --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/exception/ai/AiToolSqlExecutionException.java @@ -0,0 +1,12 @@ +package ai.chat2db.community.domain.api.exception.ai; + +public class AiToolSqlExecutionException extends AiToolException { + + public AiToolSqlExecutionException(String message) { + super(message); + } + + public AiToolSqlExecutionException(String message, Throwable cause) { + super(message, cause); + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java index e12a943ba..13a8c8056 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/AiToolResult.java @@ -2,40 +2,30 @@ import lombok.Data; -import java.util.ArrayList; -import java.util.List; - @Data -public class AiToolResult { +public class AiToolResult { private Boolean success; private String summary; - private List data = new ArrayList<>(); + private T data; private String errorCode; - public static AiToolResult success(String summary, List data) { - AiToolResult result = new AiToolResult(); + public static AiToolResult success(String summary, T data) { + AiToolResult result = new AiToolResult<>(); result.setSuccess(Boolean.TRUE); result.setSummary(summary); - result.setData(copyData(data)); + result.setData(data); return result; } - public static AiToolResult failure(String summary, String errorCode) { - AiToolResult result = new AiToolResult(); + public static AiToolResult failureWithCode(String errorCode, String summary) { + AiToolResult result = new AiToolResult<>(); result.setSuccess(Boolean.FALSE); result.setSummary(summary); result.setErrorCode(errorCode); return result; } - - private static List copyData(List data) { - if (data == null || data.isEmpty()) { - return new ArrayList<>(); - } - return new ArrayList<>(data); - } } diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/DataSourceToolData.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/DataSourceToolData.java new file mode 100644 index 000000000..722788b9a --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/DataSourceToolData.java @@ -0,0 +1,26 @@ +package ai.chat2db.community.domain.api.model.ai; + +import lombok.AllArgsConstructor; +import lombok.Data; +import lombok.NoArgsConstructor; + +import java.util.ArrayList; +import java.util.List; + +@Data +@NoArgsConstructor +@AllArgsConstructor +public class DataSourceToolData { + + private List datasources = new ArrayList<>(); + + @Data + @NoArgsConstructor + @AllArgsConstructor + public static class Item { + private Long id; + private String name; + private String type; + private String env; + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/DatabaseToolData.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/DatabaseToolData.java new file mode 100644 index 000000000..3b0388fb7 --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/DatabaseToolData.java @@ -0,0 +1,25 @@ +package ai.chat2db.community.domain.api.model.ai; + +import lombok.AllArgsConstructor; +import lombok.Data; +import lombok.NoArgsConstructor; + +import java.util.ArrayList; +import java.util.List; + +@Data +@NoArgsConstructor +@AllArgsConstructor +public class DatabaseToolData { + + private List databases = new ArrayList<>(); + + @Data + @NoArgsConstructor + @AllArgsConstructor + public static class Item { + private String name; + private Boolean system; + private String comment; + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/SchemaToolData.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/SchemaToolData.java new file mode 100644 index 000000000..6d92191be --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/SchemaToolData.java @@ -0,0 +1,25 @@ +package ai.chat2db.community.domain.api.model.ai; + +import lombok.AllArgsConstructor; +import lombok.Data; +import lombok.NoArgsConstructor; + +import java.util.ArrayList; +import java.util.List; + +@Data +@NoArgsConstructor +@AllArgsConstructor +public class SchemaToolData { + + private List schemas = new ArrayList<>(); + + @Data + @NoArgsConstructor + @AllArgsConstructor + public static class Item { + private String name; + private Boolean system; + private String comment; + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/SqlToolData.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/SqlToolData.java new file mode 100644 index 000000000..3fbe887f5 --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/SqlToolData.java @@ -0,0 +1,57 @@ +package ai.chat2db.community.domain.api.model.ai; + +import lombok.AllArgsConstructor; +import lombok.Data; +import lombok.NoArgsConstructor; + +import java.util.ArrayList; +import java.util.List; + +@Data +@NoArgsConstructor +@AllArgsConstructor +public class SqlToolData { + + private List results = new ArrayList<>(); + + @Data + @NoArgsConstructor + @AllArgsConstructor + public static class ResultSet { + private Integer resultIndex; + private Boolean success; + private String sqlType; + private Long durationMs; + private Integer updateCount; + private String message; + private String description; + private Boolean hasNextPage; + private Integer rowCount; + private Integer previewRowCount; + private Boolean rowsTruncated; + private List columns = new ArrayList<>(); + private List> rows = new ArrayList<>(); + private List> rowCellMetadata = new ArrayList<>(); + private String text; + } + + @Data + @NoArgsConstructor + @AllArgsConstructor + public static class CellMetadata { + private Boolean rawValueAvailable; + private Object displayValue; + private Boolean largeValue; + private String largeValueId; + private String valueType; + private Integer sqlType; + private String columnType; + private Long sizeBytes; + private Long sizeChars; + private Long loadedBytes; + private Long loadedChars; + private Boolean truncated; + private String unsupportedReason; + private String rawValueUnavailableReason; + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableSchemaResult.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableSchemaResult.java new file mode 100644 index 000000000..82ea33b82 --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableSchemaResult.java @@ -0,0 +1,22 @@ +package ai.chat2db.community.domain.api.model.ai; + +import ai.chat2db.community.domain.api.model.metadata.Table; +import lombok.AllArgsConstructor; +import lombok.Data; +import lombok.NoArgsConstructor; + +/** + * Typed service use-case result for table schema lookup. + * This is not an AI payload and must not contain JSON, summaries, display text, or transport envelopes. + */ +@Data +@NoArgsConstructor +@AllArgsConstructor +public class TableSchemaResult { + + private String tableName; + + private String ddl; + + private Table table; +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableSchemaToolData.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableSchemaToolData.java new file mode 100644 index 000000000..d151ef8ac --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableSchemaToolData.java @@ -0,0 +1,24 @@ +package ai.chat2db.community.domain.api.model.ai; + +import lombok.AllArgsConstructor; +import lombok.Data; +import lombok.NoArgsConstructor; + +import java.util.ArrayList; +import java.util.List; + +@Data +@NoArgsConstructor +@AllArgsConstructor +public class TableSchemaToolData { + + private List tables = new ArrayList<>(); + + @Data + @NoArgsConstructor + @AllArgsConstructor + public static class Item { + private String tableName; + private String schema; + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableToolData.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableToolData.java new file mode 100644 index 000000000..ec466b859 --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/TableToolData.java @@ -0,0 +1,25 @@ +package ai.chat2db.community.domain.api.model.ai; + +import lombok.AllArgsConstructor; +import lombok.Data; +import lombok.NoArgsConstructor; + +import java.util.ArrayList; +import java.util.List; + +@Data +@NoArgsConstructor +@AllArgsConstructor +public class TableToolData { + + private List tables = new ArrayList<>(); + + @Data + @NoArgsConstructor + @AllArgsConstructor + public static class Item { + private String name; + private String type; + private String comment; + } +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/Text2SqlToolData.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/Text2SqlToolData.java new file mode 100644 index 000000000..b9825e64b --- /dev/null +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/ai/Text2SqlToolData.java @@ -0,0 +1,13 @@ +package ai.chat2db.community.domain.api.model.ai; + +import lombok.AllArgsConstructor; +import lombok.Data; +import lombok.NoArgsConstructor; + +@Data +@NoArgsConstructor +@AllArgsConstructor +public class Text2SqlToolData { + + private String sql; +} diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/result/ResultCell.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/result/ResultCell.java index 660493f22..c4d2b11d6 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/result/ResultCell.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/model/result/ResultCell.java @@ -45,9 +45,14 @@ public Object getRawValue() { return rawValue; } + /** + * Builds a string-only cell where value is the display/serialized value and rawValue is the same stable + * string for structured AI/tool consumers that need lossless positional rows. + */ public static ResultCell of(String value) { return ResultCell.builder() .value(value) + .rawValue(value) .largeValue(false) .truncated(false) .valueType(LargeValueTypeEnum.UNKNOWN.code()) diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/service/ai/IAiToolService.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/service/ai/IAiToolService.java index a97d18cc6..ac93945e8 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/service/ai/IAiToolService.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-api/src/main/java/ai/chat2db/community/domain/api/service/ai/IAiToolService.java @@ -1,14 +1,20 @@ package ai.chat2db.community.domain.api.service.ai; +import ai.chat2db.community.domain.api.model.ai.TableSchemaResult; +import ai.chat2db.community.domain.api.model.metadata.Database; +import ai.chat2db.community.domain.api.model.metadata.Schema; +import ai.chat2db.community.domain.api.model.metadata.SimpleTable; import ai.chat2db.community.domain.api.model.request.ai.AiExecuteSqlRequest; import ai.chat2db.community.domain.api.model.request.ai.AiGetTablesSchemaRequest; import ai.chat2db.community.domain.api.model.request.ai.AiListTablesRequest; import ai.chat2db.community.domain.api.model.request.ai.AiToolContextRequest; +import ai.chat2db.community.domain.api.model.result.ExecuteResponse; +import ai.chat2db.community.domain.api.model.storage.WorkspaceDataSource; + +import java.util.List; /** * Provides AI tool operations over datasource metadata and SQL execution. - * - *

Methods return a serialized AiToolResult JSON string with success, summary, data, and errorCode fields. */ public interface IAiToolService { @@ -16,26 +22,26 @@ public interface IAiToolService { * Returns datasource metadata for AI tools. * * @param aiToolContextRequest AI tool context parameters. - * @return serialized AiToolResult JSON string. + * @return datasource metadata. */ - String listAllDataSources(AiToolContextRequest aiToolContextRequest); + List listAllDataSources(AiToolContextRequest aiToolContextRequest); /** * Returns table metadata for AI tools. * * @param aiListTablesRequest AI table listing parameters. - * @return serialized AiToolResult JSON string. + * @return table metadata. */ - String listAllTables(AiListTablesRequest aiListTablesRequest); + List listAllTables(AiListTablesRequest aiListTablesRequest); /** * Returns database metadata for AI tools. * * @param dataSourceId datasource identifier. * @param aiToolContextRequest AI tool context parameters. - * @return serialized AiToolResult JSON string. + * @return database metadata. */ - String listAllDatabases(Long dataSourceId, AiToolContextRequest aiToolContextRequest); + List listAllDatabases(Long dataSourceId, AiToolContextRequest aiToolContextRequest); /** * Returns schema metadata for AI tools. @@ -43,23 +49,23 @@ public interface IAiToolService { * @param databaseName database name that scopes the lookup. * @param dataSourceId datasource identifier. * @param aiToolContextRequest AI tool context parameters. - * @return serialized AiToolResult JSON string. + * @return schema metadata. */ - String listAllSchemas(String databaseName, Long dataSourceId, AiToolContextRequest aiToolContextRequest); + List listAllSchemas(String databaseName, Long dataSourceId, AiToolContextRequest aiToolContextRequest); /** * Executes SQL for an AI tool request. * * @param aiExecuteSqlRequest AI SQL execution parameters. - * @return serialized AiToolResult JSON string. + * @return SQL execution result sets. */ - String executeSql(AiExecuteSqlRequest aiExecuteSqlRequest); + List executeSql(AiExecuteSqlRequest aiExecuteSqlRequest); /** * Returns table schema metadata for AI tools. * * @param aiGetTablesSchemaRequest AI table schema lookup parameters. - * @return serialized AiToolResult JSON string. + * @return table schema results. */ - String getTablesSchema(AiGetTablesSchemaRequest aiGetTablesSchemaRequest); + List getTablesSchema(AiGetTablesSchemaRequest aiGetTablesSchemaRequest); } diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java index c58179150..fba38dde0 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/main/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImpl.java @@ -1,6 +1,11 @@ package ai.chat2db.community.domain.core.impl.ai; import ai.chat2db.community.domain.api.enums.parser.SqlTypeEnum; +import ai.chat2db.community.domain.api.exception.ai.AiToolException; +import ai.chat2db.community.domain.api.exception.ai.AiToolInvalidArgumentException; +import ai.chat2db.community.domain.api.exception.ai.AiToolMetadataQueryException; +import ai.chat2db.community.domain.api.exception.ai.AiToolSqlConfirmationRequiredException; +import ai.chat2db.community.domain.api.exception.ai.AiToolSqlExecutionException; import ai.chat2db.community.domain.api.model.request.db.DbDlExecuteRequest; import ai.chat2db.community.domain.api.model.request.db.DbSchemaQueryRequest; import ai.chat2db.community.domain.api.model.request.db.DbTablePageQueryRequest; @@ -10,7 +15,7 @@ import ai.chat2db.community.domain.api.model.request.datasource.DbDatabaseQueryAllRequest; import ai.chat2db.community.domain.api.model.request.runtime.DbConnectionContextRequest; import ai.chat2db.community.domain.api.model.PageResponse; -import ai.chat2db.community.domain.api.model.ai.AiToolResult; +import ai.chat2db.community.domain.api.model.ai.TableSchemaResult; import ai.chat2db.community.domain.api.model.runtime.ConnectionProfile; import ai.chat2db.community.domain.api.service.db.IDbConnectionContextService; import ai.chat2db.community.domain.api.service.db.IDbDatabaseService; @@ -33,16 +38,11 @@ import ai.chat2db.community.domain.api.service.ai.IAiToolService; import ai.chat2db.community.domain.api.model.metadata.Database; import ai.chat2db.community.domain.api.model.result.ExecuteResponse; -import ai.chat2db.community.domain.api.model.metadata.ForeignKeyInfo; -import ai.chat2db.community.domain.api.model.result.Header; import ai.chat2db.community.domain.api.model.metadata.Schema; import ai.chat2db.community.domain.api.model.metadata.SimpleTable; import ai.chat2db.community.domain.api.model.metadata.Table; import ai.chat2db.community.domain.api.model.metadata.TableColumn; -import ai.chat2db.community.domain.api.model.metadata.TableIndex; -import ai.chat2db.community.domain.api.model.metadata.TableIndexColumn; -import com.alibaba.fastjson2.JSON; -import com.alibaba.fastjson2.JSONWriter; +import ai.chat2db.community.tools.exception.BusinessException; import org.apache.commons.collections4.CollectionUtils; import org.apache.commons.lang3.StringUtils; import org.springframework.beans.factory.annotation.Autowired; @@ -51,10 +51,7 @@ import java.util.ArrayList; import java.util.Collections; -import java.util.Comparator; -import java.util.LinkedHashMap; import java.util.List; -import java.util.Map; import java.util.Objects; import java.util.Locale; import java.util.stream.Collectors; @@ -77,208 +74,200 @@ public class AiToolServiceImpl implements IAiToolService { private IDbSqlService sqlService; @Autowired private IWorkspaceStorageFacade workspaceStorageFacade; - private static final int DEFAULT_SQL_PAGE_SIZE = 200; - private static final int MAX_SQL_PAGE_SIZE = 500; - private static final int MAX_SQL_RESULT_ROWS = 50; + // Execution page size bounds protect database work; AI preview truncation is handled by the web converter. + private static final int DEFAULT_SQL_EXECUTION_PAGE_SIZE = 200; + private static final int MAX_SQL_EXECUTION_PAGE_SIZE = 500; private static final int MAX_GLOBAL_DATASOURCES = 200; - public String listAllDataSources(AiToolContextRequest toolContext) { + + public List listAllDataSources(AiToolContextRequest toolContext) { DbDataSourcePageQueryRequest queryRequest = new DbDataSourcePageQueryRequest(); queryRequest.setPageNo(1); queryRequest.setPageSize(MAX_GLOBAL_DATASOURCES); - PageResponse result = invokeWithRequestContext( - toolContext, - () -> workspaceStorageFacade.listDataSources(queryRequest)); - if (Objects.isNull(result)) { - return emitToolFailure(toolContext, "list_all_datasources", - "Failed to query datasources: unknown error", "DATASOURCE_QUERY_FAILED"); + PageResponse result; + try { + result = invokeWithRequestContext(toolContext, () -> workspaceStorageFacade.listDataSources(queryRequest)); + } catch (BusinessException e) { + throw new AiToolMetadataQueryException( + "Failed to query datasources: " + StringUtils.defaultString(e.getMessage(), "unknown error"), + e); } - if (CollectionUtils.isEmpty(result.getData())) { - return emitToolResult(toolContext, "list_all_datasources", "No datasources found.", Collections.emptyList()); + if (Objects.isNull(result)) { + throw new AiToolMetadataQueryException("Failed to query datasources: unknown error"); } - - List> data = result.getData().stream() + return result.getData() == null ? Collections.emptyList() : result.getData().stream() .filter(Objects::nonNull) - .map(dataSource -> { - Map item = new LinkedHashMap<>(); - item.put("id", dataSource.getId()); - item.put("name", StringUtils.defaultIfBlank(dataSource.getAlias(), "(unnamed)")); - item.put("type", dataSource.getType()); - item.put("env", dataSource.getEnvType()); - return item; - }) .collect(Collectors.toList()); - return emitToolResult(toolContext, "list_all_datasources", - "Found " + data.size() + " datasource(s).", data); } - public String listAllTables(AiListTablesRequest aiListTablesRequest) { - Long dataSourceId = aiListTablesRequest == null ? null : aiListTablesRequest.getDataSourceId(); - String databaseName = aiListTablesRequest == null ? null : aiListTablesRequest.getDatabaseName(); - String schemaName = aiListTablesRequest == null ? null : aiListTablesRequest.getSchemaName(); - AiToolContextRequest toolContext = aiListTablesRequest == null ? null : aiListTablesRequest.getAiToolContextRequest(); - ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, databaseName, schemaName); - try { - connectionContextService.bindProfile(profile); - DbTablePageQueryRequest queryParam = DbTablePageQueryRequest.builder() - .dataSourceId(profile.getDataSourceId()) - .databaseName(profile.getDatabaseName()) - .schemaName(profile.getSchemaName()) - .pageNo(1) - .pageSize(500) - .refresh(false) - .build(); - List result = tableService.queryTables(queryParam); - if (CollectionUtils.isEmpty(result)) { - return emitToolResult(toolContext, "list_all_tables", "No tables found.", Collections.emptyList()); + + public List listAllTables(AiListTablesRequest aiListTablesRequest) { + AiListTablesRequest request = requireRequest(aiListTablesRequest, "listAllTables request"); + Long dataSourceId = request.getDataSourceId(); + String databaseName = request.getDatabaseName(); + String schemaName = request.getSchemaName(); + AiToolContextRequest toolContext = request.getAiToolContextRequest(); + return invokeWithRequestContext(toolContext, () -> { + ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, databaseName, schemaName); + try { + connectionContextService.bindProfile(profile); + DbTablePageQueryRequest queryParam = DbTablePageQueryRequest.builder() + .dataSourceId(profile.getDataSourceId()) + .databaseName(profile.getDatabaseName()) + .schemaName(profile.getSchemaName()) + .pageNo(1) + .pageSize(500) + .refresh(false) + .build(); + List result = tableService.queryTables(queryParam); + return result == null ? Collections.emptyList() : result.stream() + .filter(Objects::nonNull) + .collect(Collectors.toList()); + } catch (AiToolException e) { + throw e; + } catch (BusinessException e) { + throw new AiToolMetadataQueryException( + "Failed to query tables: " + StringUtils.defaultString(e.getMessage(), "unknown error"), + e); + } finally { + connectionContextService.clear(); } - List> data = result.stream() - .filter(Objects::nonNull) - .map(this::tableSummaryData) - .collect(Collectors.toList()); - return emitToolResult(toolContext, "list_all_tables", - "Found " + data.size() + " table(s).", data); - } finally { - connectionContextService.clear(); - } + }); } - public String listAllDatabases(Long dataSourceId, + + public List listAllDatabases(Long dataSourceId, AiToolContextRequest toolContext) { - ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, null, null); - try { - connectionContextService.bindProfile(profile); - DbDatabaseQueryAllRequest queryParam = DbDatabaseQueryAllRequest.builder() - .dataSourceId(profile.getDataSourceId()) - .refresh(false) - .build(); - List result = databaseService.queryAll(queryParam); - if (CollectionUtils.isEmpty(result)) { - return emitToolResult(toolContext, "list_all_databases", "No databases found.", Collections.emptyList()); + return invokeWithRequestContext(toolContext, () -> { + ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, null, null); + try { + connectionContextService.bindProfile(profile); + DbDatabaseQueryAllRequest queryParam = DbDatabaseQueryAllRequest.builder() + .dataSourceId(profile.getDataSourceId()) + .refresh(false) + .build(); + List result = databaseService.queryAll(queryParam); + return result == null ? Collections.emptyList() : result.stream() + .filter(Objects::nonNull) + .collect(Collectors.toList()); + } catch (AiToolException e) { + throw e; + } catch (BusinessException e) { + throw new AiToolMetadataQueryException( + "Failed to query databases: " + StringUtils.defaultString(e.getMessage(), "unknown error"), + e); + } finally { + connectionContextService.clear(); } - List> data = result.stream() - .filter(Objects::nonNull) - .map(database -> { - Map item = new LinkedHashMap<>(); - item.put("name", StringUtils.defaultString(database.getName(), "(unnamed)")); - item.put("system", database.isSystem()); - item.put("comment", database.getComment()); - return item; - }) - .collect(Collectors.toList()); - return emitToolResult(toolContext, "list_all_databases", - "Found " + data.size() + " database(s).", data); - } finally { - connectionContextService.clear(); - } + }); } - public String listAllSchemas(String databaseName,Long dataSourceId, + + public List listAllSchemas(String databaseName, Long dataSourceId, AiToolContextRequest toolContext) { - ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, databaseName, null); - String targetDatabase = StringUtils.defaultIfBlank(databaseName, profile.getDatabaseName()); - if (StringUtils.isBlank(targetDatabase)) { - return emitToolFailure(toolContext, "list_all_schemas", - "databaseName is required for listing schemas.", "INVALID_ARGUMENT"); - } - try { - connectionContextService.bindProfile(profile); - DbSchemaQueryRequest queryParam = DbSchemaQueryRequest.builder() - .dataSourceId(profile.getDataSourceId()) - .dataBaseName(targetDatabase) - .refresh(false) - .build(); - List result = databaseService.querySchema(queryParam); - if (CollectionUtils.isEmpty(result)) { - return emitToolResult(toolContext, "list_all_schemas", "No schemas found.", Collections.emptyList()); + return invokeWithRequestContext(toolContext, () -> { + ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, databaseName, null); + String targetDatabase = StringUtils.defaultIfBlank(databaseName, profile.getDatabaseName()); + if (StringUtils.isBlank(targetDatabase)) { + throw new AiToolInvalidArgumentException("databaseName is required for listing schemas."); } - List> data = result.stream() - .filter(Objects::nonNull) - .map(schema -> { - Map item = new LinkedHashMap<>(); - item.put("name", StringUtils.defaultString(schema.getName(), "(unnamed)")); - item.put("system", schema.isSystem()); - item.put("comment", schema.getComment()); - return item; - }) - .collect(Collectors.toList()); - return emitToolResult(toolContext, "list_all_schemas", - "Found " + data.size() + " schema(s).", data); - } finally { - connectionContextService.clear(); - } + try { + connectionContextService.bindProfile(profile); + DbSchemaQueryRequest queryParam = DbSchemaQueryRequest.builder() + .dataSourceId(profile.getDataSourceId()) + .dataBaseName(targetDatabase) + .refresh(false) + .build(); + List result = databaseService.querySchema(queryParam); + return result == null ? Collections.emptyList() : result.stream() + .filter(Objects::nonNull) + .collect(Collectors.toList()); + } catch (AiToolException e) { + throw e; + } catch (BusinessException e) { + throw new AiToolMetadataQueryException( + "Failed to query schemas: " + StringUtils.defaultString(e.getMessage(), "unknown error"), + e); + } finally { + connectionContextService.clear(); + } + }); } - public String executeSql(AiExecuteSqlRequest aiExecuteSqlRequest) { - String sql = aiExecuteSqlRequest == null ? null : aiExecuteSqlRequest.getSql(); - Integer pageSize = aiExecuteSqlRequest == null ? null : aiExecuteSqlRequest.getPageSize(); - Long dataSourceId = aiExecuteSqlRequest == null ? null : aiExecuteSqlRequest.getDataSourceId(); - String databaseName = aiExecuteSqlRequest == null ? null : aiExecuteSqlRequest.getDatabaseName(); - String schemaName = aiExecuteSqlRequest == null ? null : aiExecuteSqlRequest.getSchemaName(); - AiToolContextRequest toolContext = aiExecuteSqlRequest == null ? null : aiExecuteSqlRequest.getAiToolContextRequest(); + + public List executeSql(AiExecuteSqlRequest aiExecuteSqlRequest) { + AiExecuteSqlRequest request = requireRequest(aiExecuteSqlRequest, "executeSql request"); + String sql = request.getSql(); + Integer pageSize = request.getPageSize(); + Long dataSourceId = request.getDataSourceId(); + String databaseName = request.getDatabaseName(); + String schemaName = request.getSchemaName(); + AiToolContextRequest toolContext = request.getAiToolContextRequest(); if (StringUtils.isBlank(sql)) { - return emitToolFailure(toolContext, "execute_sql", "sql is empty.", "INVALID_ARGUMENT"); + throw new AiToolInvalidArgumentException("sql is empty."); } - ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, databaseName, schemaName); - int resolvedPageSize = normalizePageSize(pageSize); String trimmedSql = sql.trim(); - String unsafeSqlMessage = buildNonQueryExecutionMessage(trimmedSql, profile); - if (StringUtils.isNotBlank(unsafeSqlMessage)) { - return emitToolFailure(toolContext, "execute_sql", unsafeSqlMessage, "SQL_REQUIRES_MANUAL_CONFIRMATION"); - } - - boolean operationLogged = false; - try { - connectionContextService.bindProfile(profile); - DbDlExecuteRequest executeParam = new DbDlExecuteRequest(); - executeParam.setSql(trimmedSql); - executeParam.setSingle(true); - executeParam.setDataSourceId(profile.getDataSourceId()); - executeParam.setDatabaseName(profile.getDatabaseName()); - executeParam.setSchemaName(profile.getSchemaName()); - executeParam.setPageNo(1); - executeParam.setPageSize(resolvedPageSize); - executeParam.setPageSizeAll(false); - executeParam.setErrorContinue(false); - - ListResult executeResult = wrapExecuteResults(dlTemplateService.execute(executeParam)); - OpsSqlOperationLogListResultRequest sqlOperationLogListResultRequest = OpsSqlOperationLogListResultRequest.of( - trimmedSql, executeResult.getSuccess(), executeResult.getErrorMessage(), executeResult.getData(), - SqlOperationLogSourceEnum.AI_TOOL.name()); - sqlOperationLogRecorder.recordListResultAsync(sqlOperationLogListResultRequest); - operationLogged = true; - if (Objects.isNull(executeResult) || !executeResult.success()) { - return emitToolFailure(toolContext, "execute_sql", "SQL execution failed: " - + (Objects.isNull(executeResult) ? "unknown error" : StringUtils.defaultString(executeResult.getErrorMessage())), - "SQL_EXECUTION_FAILED"); - } - if (CollectionUtils.isEmpty(executeResult.getData())) { - return emitToolResult(toolContext, "execute_sql", - "SQL executed successfully with no result.", Collections.emptyList()); + return invokeWithRequestContext(toolContext, () -> { + ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, databaseName, schemaName); + int resolvedPageSize = normalizePageSize(pageSize); + String unsafeSqlMessage = buildNonQueryExecutionMessage(trimmedSql, profile); + if (StringUtils.isNotBlank(unsafeSqlMessage)) { + throw new AiToolSqlConfirmationRequiredException(unsafeSqlMessage); } - List> data = new ArrayList<>(); - int index = 1; - for (ExecuteResponse item : executeResult.getData()) { - data.add(executeResponseData(index++, item)); - } - return emitToolResult(toolContext, "execute_sql", - "SQL executed successfully with " + data.size() + " result set(s).", data); - } catch (RuntimeException e) { - if (!operationLogged) { - sqlOperationLogRecorder.recordFailureAsync(trimmedSql, SqlOperationLogSourceEnum.AI_TOOL.name(), e.getMessage()); + boolean operationLogged = false; + try { + connectionContextService.bindProfile(profile); + DbDlExecuteRequest executeParam = new DbDlExecuteRequest(); + executeParam.setSql(trimmedSql); + executeParam.setSingle(true); + executeParam.setDataSourceId(profile.getDataSourceId()); + executeParam.setDatabaseName(profile.getDatabaseName()); + executeParam.setSchemaName(profile.getSchemaName()); + executeParam.setPageNo(1); + executeParam.setPageSize(resolvedPageSize); + executeParam.setPageSizeAll(false); + executeParam.setErrorContinue(false); + + ListResult executeResult = wrapExecuteResults(dlTemplateService.execute(executeParam)); + OpsSqlOperationLogListResultRequest sqlOperationLogListResultRequest = OpsSqlOperationLogListResultRequest.of( + trimmedSql, executeResult.getSuccess(), executeResult.getErrorMessage(), executeResult.getData(), + SqlOperationLogSourceEnum.AI_TOOL.name()); + sqlOperationLogRecorder.recordListResultAsync(sqlOperationLogListResultRequest); + operationLogged = true; + if (Objects.isNull(executeResult) || !executeResult.success()) { + throw new AiToolSqlExecutionException("SQL execution failed: " + + (Objects.isNull(executeResult) ? "unknown error" : StringUtils.defaultString(executeResult.getErrorMessage())), + null); + } + return executeResult.getData() == null ? Collections.emptyList() : executeResult.getData(); + } catch (AiToolException e) { + throw e; + } catch (BusinessException e) { + if (!operationLogged) { + sqlOperationLogRecorder.recordFailureAsync(trimmedSql, SqlOperationLogSourceEnum.AI_TOOL.name(), e.getMessage()); + } + throw new AiToolSqlExecutionException( + "SQL execution failed: " + StringUtils.defaultString(e.getMessage(), "unknown error"), + e); + } catch (RuntimeException e) { + if (!operationLogged) { + sqlOperationLogRecorder.recordFailureAsync(trimmedSql, SqlOperationLogSourceEnum.AI_TOOL.name(), e.getMessage()); + } + throw e; + } finally { + connectionContextService.clear(); } - throw e; - } finally { - connectionContextService.clear(); - } + }); } - public String getTablesSchema(AiGetTablesSchemaRequest aiGetTablesSchemaRequest) { - List tableNames = aiGetTablesSchemaRequest == null ? null : aiGetTablesSchemaRequest.getTableNames(); - Long dataSourceId = aiGetTablesSchemaRequest == null ? null : aiGetTablesSchemaRequest.getDataSourceId(); - String databaseName = aiGetTablesSchemaRequest == null ? null : aiGetTablesSchemaRequest.getDatabaseName(); - String schemaName = aiGetTablesSchemaRequest == null ? null : aiGetTablesSchemaRequest.getSchemaName(); - AiToolContextRequest toolContext = aiGetTablesSchemaRequest == null ? null : aiGetTablesSchemaRequest.getAiToolContextRequest(); + + public List getTablesSchema(AiGetTablesSchemaRequest aiGetTablesSchemaRequest) { + AiGetTablesSchemaRequest request = requireRequest(aiGetTablesSchemaRequest, "getTablesSchema request"); + List tableNames = request.getTableNames(); + Long dataSourceId = request.getDataSourceId(); + String databaseName = request.getDatabaseName(); + String schemaName = request.getSchemaName(); + AiToolContextRequest toolContext = request.getAiToolContextRequest(); if (CollectionUtils.isEmpty(tableNames)) { - return emitToolFailure(toolContext, "get_tables_schema", "tableNames is empty.", "INVALID_ARGUMENT"); + throw new AiToolInvalidArgumentException("tableNames is empty."); } List normalized = tableNames.stream() @@ -288,43 +277,30 @@ public String getTablesSchema(AiGetTablesSchemaRequest aiGetTablesSchemaRequest) .limit(20) .collect(Collectors.toList()); if (CollectionUtils.isEmpty(normalized)) { - return emitToolFailure(toolContext, "get_tables_schema", "tableNames is empty.", "INVALID_ARGUMENT"); + throw new AiToolInvalidArgumentException("tableNames is empty."); } - ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, databaseName, schemaName); - try { - connectionContextService.bindProfile(profile); - List> data = new ArrayList<>(); - for (String tableName : normalized) { - Table table = fetchDetailedTable(profile, tableName); - String ddl = fetchTableDdl(profile, tableName); - String schemaText = buildRichTableSchema(tableName, ddl, table); - Map item = new LinkedHashMap<>(); - item.put("tableName", tableName); - item.put("schema", schemaText); - data.add(item); + return invokeWithRequestContext(toolContext, () -> { + ConnectionProfile profile = requireScopedConnectInfo(toolContext, dataSourceId, databaseName, schemaName); + try { + connectionContextService.bindProfile(profile); + List data = new ArrayList<>(); + for (String tableName : normalized) { + Table table = fetchDetailedTable(profile, tableName); + String ddl = fetchTableDdl(profile, tableName); + data.add(new TableSchemaResult(tableName, ddl, table)); + } + return data; + } catch (AiToolException e) { + throw e; + } catch (BusinessException e) { + throw new AiToolMetadataQueryException( + "Failed to query table schema: " + StringUtils.defaultString(e.getMessage(), "unknown error"), + e); + } finally { + connectionContextService.clear(); } - return emitToolResult(toolContext, "get_tables_schema", - "Fetched schema for " + data.size() + " table(s).", data); - } finally { - connectionContextService.clear(); - } - } - - private String emitToolResult(AiToolContextRequest toolContext, String toolName, String summary, List data) { - return successToolResultJson(summary, data); - } - - private String emitToolFailure(AiToolContextRequest toolContext, String toolName, String summary, String errorCode) { - return failureToolResultJson(summary, errorCode); - } - - static String successToolResultJson(String summary, List data) { - return JSON.toJSONString(AiToolResult.success(summary, data), JSONWriter.Feature.WriteNulls); - } - - static String failureToolResultJson(String summary, String errorCode) { - return JSON.toJSONString(AiToolResult.failure(summary, errorCode), JSONWriter.Feature.WriteNulls); + }); } private T invokeWithRequestContext(AiToolContextRequest toolContext, java.util.function.Supplier supplier) { @@ -420,224 +396,6 @@ private Table fetchDetailedTable(ConnectionProfile profile, String tableName) { return table; } - private Map tableSummaryData(SimpleTable table) { - Map item = new LinkedHashMap<>(); - item.put("name", StringUtils.defaultString(table.getName(), "(unnamed)")); - item.put("type", StringUtils.defaultIfBlank(table.getTableType(), "TABLE")); - item.put("comment", table.getComment()); - return item; - } - - private String buildRichTableSchema(String tableName, String ddl, Table table) { - - StringBuilder builder = new StringBuilder(2048); - builder.append("-- TABLE: ").append(tableName).append("\n"); - builder.append("/* physical schema */\n"); - builder.append(StringUtils.defaultIfBlank(ddl, "-- schema unavailable")); - - String primaryKeys = formatPrimaryKeys(table); - if (StringUtils.isNotBlank(primaryKeys)) { - builder.append("\n\n").append(primaryKeys); - } - - String indexes = formatIndexes(table); - if (StringUtils.isNotBlank(indexes)) { - builder.append("\n\n").append(indexes); - } - - String foreignKeys = formatForeignKeys(table); - if (StringUtils.isNotBlank(foreignKeys)) { - builder.append("\n\n").append(foreignKeys); - } - - return builder.toString(); - } - - private String formatPrimaryKeys(Table table) { - if (table == null || CollectionUtils.isEmpty(table.getColumnList())) { - return null; - } - List primaryKeys = table.getColumnList().stream() - .filter(column -> Boolean.TRUE.equals(column.getPrimaryKey())) - .sorted(Comparator.comparingInt(TableColumn::getPrimaryKeyOrder)) - .toList(); - if (CollectionUtils.isEmpty(primaryKeys)) { - return null; - } - List lines = new ArrayList<>(); - lines.add("/* primary keys */"); - lines.add(primaryKeys.stream() - .map(TableColumn::getName) - .collect(Collectors.joining(", "))); - return String.join("\n", lines); - } - - private String formatIndexes(Table table) { - if (table == null || CollectionUtils.isEmpty(table.getIndexList())) { - return null; - } - List lines = new ArrayList<>(); - lines.add("/* indexes */"); - for (TableIndex index : table.getIndexList()) { - List columns = index.getColumnList(); - String columnNames = CollectionUtils.isEmpty(columns) - ? "" - : columns.stream() - .sorted(Comparator.comparing(column -> Objects.requireNonNullElse(column.getOrdinalPosition(), (short) 0))) - .map(TableIndexColumn::getColumnName) - .filter(StringUtils::isNotBlank) - .collect(Collectors.joining(", ")); - List parts = new ArrayList<>(); - parts.add("type=" + StringUtils.defaultIfBlank(index.getType(), "INDEX")); - parts.add("unique=" + Boolean.TRUE.equals(index.getUnique())); - if (StringUtils.isNotBlank(index.getMethod())) { - parts.add("method=" + index.getMethod()); - } - if (StringUtils.isNotBlank(index.getComment())) { - parts.add("comment=" + index.getComment()); - } - lines.add("- " + StringUtils.defaultIfBlank(index.getName(), "(unnamed)") - + (StringUtils.isNotBlank(columnNames) ? " (" + columnNames + ")" : "") - + " | " + String.join("; ", parts)); - } - return lines.size() > 1 ? String.join("\n", lines) : null; - } - - private String formatForeignKeys(Table table) { - if (table == null || CollectionUtils.isEmpty(table.getForeignKeyList())) { - return null; - } - Map> grouped = new LinkedHashMap<>(); - for (ForeignKeyInfo foreignKey : table.getForeignKeyList()) { - String key = firstNonBlank(foreignKey.getFkName(), - foreignKey.getFkTableName() + "->" + foreignKey.getPkTableName()); - grouped.computeIfAbsent(key, ignored -> new ArrayList<>()).add(foreignKey); - } - - List lines = new ArrayList<>(); - lines.add("/* foreign keys */"); - for (Map.Entry> entry : grouped.entrySet()) { - List fkList = entry.getValue().stream() - .sorted(Comparator.comparingInt(item -> item.getKeySeq())) - .toList(); - String fkColumns = fkList.stream() - .map(ForeignKeyInfo::getFkColumnName) - .filter(StringUtils::isNotBlank) - .collect(Collectors.joining(", ")); - String pkTable = fkList.stream() - .map(ForeignKeyInfo::getPkTableName) - .filter(StringUtils::isNotBlank) - .findFirst() - .orElse("(unknown)"); - String pkColumns = fkList.stream() - .map(ForeignKeyInfo::getPkColumnName) - .filter(StringUtils::isNotBlank) - .collect(Collectors.joining(", ")); - lines.add("- " + entry.getKey() + ": (" + fkColumns + ") -> " + pkTable + "(" + pkColumns + ")"); - } - return lines.size() > 1 ? String.join("\n", lines) : null; - } - - private String firstNonBlank(String... values) { - if (values == null) { - return null; - } - for (String value : values) { - if (StringUtils.isNotBlank(value)) { - return value; - } - } - return null; - } - - private static String formatExecuteResponse(ExecuteResponse result) { - if (Objects.isNull(result)) { - return "Empty result."; - } - StringBuilder builder = new StringBuilder(1024); - builder.append("success: ").append(Boolean.TRUE.equals(result.getSuccess())).append("\n"); - if (StringUtils.isNotBlank(result.getSqlType())) { - builder.append("sqlType: ").append(result.getSqlType()).append("\n"); - } - if (Objects.nonNull(result.getDuration())) { - builder.append("durationMs: ").append(result.getDuration()).append("\n"); - } - if (Objects.nonNull(result.getUpdateCount())) { - builder.append("updateCount: ").append(result.getUpdateCount()).append("\n"); - } - if (StringUtils.isNotBlank(result.getMessage())) { - builder.append("message: ").append(result.getMessage()).append("\n"); - } - if (StringUtils.isNotBlank(result.getDescription())) { - builder.append("description: ").append(result.getDescription()).append("\n"); - } - if (CollectionUtils.isNotEmpty(result.getHeaderList()) && CollectionUtils.isNotEmpty(result.getDataList())) { - builder.append("rows: ").append(result.getDataList().size()); - if (Objects.nonNull(result.getHasNextPage())) { - builder.append(", hasNextPage: ").append(result.getHasNextPage()); - } - builder.append("\n"); - appendTabularPreview(builder, result.getHeaderList(), result.getDisplayDataList()); - } - return builder.toString().trim(); - } - - static Map executeResponseData(int index, ExecuteResponse result) { - Map item = new LinkedHashMap<>(); - item.put("resultIndex", index); - if (Objects.isNull(result)) { - item.put("success", false); - item.put("message", "Empty result."); - item.put("text", "Empty result."); - return item; - } - item.put("success", Boolean.TRUE.equals(result.getSuccess())); - item.put("sqlType", result.getSqlType()); - item.put("durationMs", result.getDuration()); - item.put("updateCount", result.getUpdateCount()); - item.put("message", result.getMessage()); - item.put("description", result.getDescription()); - item.put("hasNextPage", result.getHasNextPage()); - int rowCount = result.getDataList() == null ? 0 : result.getDataList().size(); - int previewRowCount = Math.min(rowCount, MAX_SQL_RESULT_ROWS); - item.put("rowCount", rowCount); - item.put("previewRowCount", previewRowCount); - item.put("rowsTruncated", rowCount > previewRowCount); - item.put("columns", columnNames(result.getHeaderList())); - item.put("rows", rowPreviewRows(result.getHeaderList(), result.getDisplayDataList())); - item.put("text", formatExecuteResponse(result)); - return item; - } - - static List columnNames(List

headers) { - if (CollectionUtils.isEmpty(headers)) { - return Collections.emptyList(); - } - return headers.stream() - .map(header -> StringUtils.defaultIfBlank(header.getName(), header.getColumnName())) - .map(name -> StringUtils.defaultIfBlank(name, "col")) - .collect(Collectors.toList()); - } - - static List> rowPreviewRows(List
headers, List> rows) { - if (CollectionUtils.isEmpty(headers) || CollectionUtils.isEmpty(rows)) { - return Collections.emptyList(); - } - List headerNames = columnNames(headers); - int rowCount = Math.min(rows.size(), MAX_SQL_RESULT_ROWS); - List> result = new ArrayList<>(rowCount); - for (int i = 0; i < rowCount; i++) { - List row = rows.get(i); - List rowData = new ArrayList<>(headerNames.size()); - for (int c = 0; c < headerNames.size(); c++) { - String value = row != null && c < row.size() ? row.get(c) : null; - rowData.add(value); - } - result.add(rowData); - } - return result; - } - private ListResult wrapExecuteResults(List results) { ListResult result = ListResult.of(results); if (CollectionUtils.isEmpty(results)) { @@ -655,43 +413,18 @@ private ListResult wrapExecuteResults(List res return result; } - private static void appendTabularPreview(StringBuilder builder, List
headers, List> rows) { - List headerNames = headers.stream() - .map(header -> StringUtils.defaultIfBlank(header.getName(), header.getColumnName())) - .map(name -> StringUtils.defaultIfBlank(name, "col")) - .collect(Collectors.toList()); - builder.append(String.join("\t", headerNames)).append("\n"); - int rowCount = Math.min(rows.size(), MAX_SQL_RESULT_ROWS); - for (int i = 0; i < rowCount; i++) { - List row = rows.get(i); - List normalized = new ArrayList<>(headerNames.size()); - for (int c = 0; c < headerNames.size(); c++) { - String value = c < row.size() ? row.get(c) : ""; - normalized.add(normalizeCell(value)); - } - builder.append(String.join("\t", normalized)).append("\n"); - } - if (rows.size() > rowCount) { - builder.append("... ").append(rows.size() - rowCount).append(" more rows not shown."); - } - } - - private static String normalizeCell(String value) { - if (value == null) { - return "NULL"; - } - String normalized = value.replace("\n", "\\n").replace("\r", "\\r").replace("\t", " "); - if (normalized.length() > 200) { - return normalized.substring(0, 197) + "..."; + private int normalizePageSize(Integer pageSize) { + if (Objects.isNull(pageSize) || pageSize <= 0) { + return DEFAULT_SQL_EXECUTION_PAGE_SIZE; } - return normalized; + return Math.min(pageSize, MAX_SQL_EXECUTION_PAGE_SIZE); } - private int normalizePageSize(Integer pageSize) { - if (Objects.isNull(pageSize) || pageSize <= 0) { - return DEFAULT_SQL_PAGE_SIZE; + private T requireRequest(T request, String requestName) { + if (request == null) { + throw new AiToolInvalidArgumentException(requestName + " is required."); } - return Math.min(pageSize, MAX_SQL_PAGE_SIZE); + return request; } private String buildNonQueryExecutionMessage(String sql, ConnectionProfile profile) { @@ -782,7 +515,7 @@ private ConnectionProfile requireScopedConnectInfo( Long dataSourceId, String databaseName, String schemaName) { - ConnectionProfile contextProfile = resolveConnectionProfile(toolContext); + ConnectionProfile contextProfile = extractConnectionProfile(toolContext); Long resolvedDataSourceId = dataSourceId; String resolvedDatabaseName = databaseName; String resolvedSchemaName = schemaName; @@ -796,39 +529,40 @@ private ConnectionProfile requireScopedConnectInfo( if (StringUtils.isBlank(resolvedSchemaName) && contextProfile != null) { resolvedSchemaName = contextProfile.getSchemaName(); } + if (Objects.isNull(resolvedDataSourceId) && toolContext != null) { + resolvedDataSourceId = toolContext.getDataSourceId(); + } + if (StringUtils.isBlank(resolvedDatabaseName) && toolContext != null) { + resolvedDatabaseName = toolContext.getDatabaseName(); + } + if (StringUtils.isBlank(resolvedSchemaName) && toolContext != null) { + resolvedSchemaName = toolContext.getSchemaName(); + } if (Objects.nonNull(resolvedDataSourceId)) { final Long scopedDataSourceId = resolvedDataSourceId; final String scopedDatabaseName = resolvedDatabaseName; final String scopedSchemaName = resolvedSchemaName; - return invokeWithRequestContext(toolContext, - () -> buildProfile(scopedDataSourceId, scopedDatabaseName, scopedSchemaName)); + try { + return invokeWithRequestContext(toolContext, + () -> buildProfile(scopedDataSourceId, scopedDatabaseName, scopedSchemaName)); + } catch (AiToolException e) { + throw e; + } catch (BusinessException e) { + throw new AiToolMetadataQueryException( + "Failed to resolve database connection context: " + + StringUtils.defaultString(e.getMessage(), "unknown error"), + e); + } } - throw new IllegalArgumentException( + throw new AiToolInvalidArgumentException( "No database connection context found. Call list_all_datasources first, then provide dataSourceId/databaseName."); } - private ConnectionProfile resolveConnectionProfile(AiToolContextRequest toolContext) { + private ConnectionProfile extractConnectionProfile(AiToolContextRequest toolContext) { if (Objects.isNull(toolContext)) { return null; } - if (toolContext.getConnectionProfile() != null) { - return toolContext.getConnectionProfile(); - } - Long dataSourceId = toolContext.getDataSourceId(); - if (Objects.isNull(dataSourceId)) { - return null; - } - String databaseName = toolContext.getDatabaseName(); - String schemaName = toolContext.getSchemaName(); - return buildProfile(dataSourceId, databaseName, schemaName); - } - - private ConnectionProfile requireConnectInfo(AiToolContextRequest toolContext) { - ConnectionProfile profile = resolveConnectionProfile(toolContext); - if (Objects.nonNull(profile)) { - return profile; - } - throw new IllegalArgumentException("No database connection context found. Provide dataSourceId/databaseName."); + return toolContext.getConnectionProfile(); } private ConnectionProfile buildProfile(Long dataSourceId, String databaseName, String schemaName) { diff --git a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java index 07197e9a5..b383567d3 100644 --- a/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java +++ b/chat2db-community-server/chat2db-community-domain/chat2db-community-domain-core/src/test/java/ai/chat2db/community/domain/core/impl/ai/AiToolServiceImplTest.java @@ -1,91 +1,453 @@ package ai.chat2db.community.domain.core.impl.ai; -import ai.chat2db.community.domain.api.model.result.Header; +import ai.chat2db.community.domain.api.exception.ai.AiToolInvalidArgumentException; +import ai.chat2db.community.domain.api.exception.ai.AiToolMetadataQueryException; +import ai.chat2db.community.domain.api.exception.ai.AiToolSqlConfirmationRequiredException; +import ai.chat2db.community.domain.api.exception.ai.AiToolSqlExecutionException; +import ai.chat2db.community.domain.api.model.ai.TableSchemaResult; +import ai.chat2db.community.domain.api.model.metadata.Database; +import ai.chat2db.community.domain.api.model.metadata.SimpleTable; +import ai.chat2db.community.domain.api.model.metadata.Table; +import ai.chat2db.community.domain.api.model.request.ai.AiToolContextRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiExecuteSqlRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiGetTablesSchemaRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiListTablesRequest; +import ai.chat2db.community.domain.api.model.request.db.DbDlExecuteRequest; import ai.chat2db.community.domain.api.model.result.ExecuteResponse; -import ai.chat2db.community.domain.api.model.result.ResultCell; -import com.alibaba.fastjson2.JSON; -import com.alibaba.fastjson2.JSONObject; +import ai.chat2db.community.domain.api.model.runtime.ConnectionProfile; +import ai.chat2db.community.domain.api.model.sql.SimpleSqlStatement; +import ai.chat2db.community.domain.api.service.db.IDbConnectionContextService; +import ai.chat2db.community.domain.api.service.db.IDbDatabaseService; +import ai.chat2db.community.domain.api.service.db.IDbDlTemplateService; +import ai.chat2db.community.domain.api.service.db.IDbSqlService; +import ai.chat2db.community.domain.api.service.db.IDbTableService; +import ai.chat2db.community.domain.api.service.ops.IOpsSqlOperationLogService; +import ai.chat2db.community.tools.exception.BusinessException; +import ai.chat2db.community.tools.model.Context; +import ai.chat2db.community.tools.util.ContextUtils; import org.junit.jupiter.api.Test; -import java.util.ArrayList; -import java.util.Arrays; +import java.lang.reflect.Field; +import java.lang.reflect.InvocationHandler; +import java.lang.reflect.Proxy; +import java.util.Collections; import java.util.List; -import java.util.Map; +import java.util.concurrent.atomic.AtomicInteger; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertNull; -import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertThrows; class AiToolServiceImplTest { @Test - void shouldSerializeSuccessfulToolResultAsStandardJson() { - String json = AiToolServiceImpl.successToolResultJson( - "SQL executed successfully with 1 result set(s).", - List.of(Map.of("rowCount", 2))); + void executeSqlRejectsBlankSqlAsStrongException() { + AiExecuteSqlRequest request = new AiExecuteSqlRequest(); + request.setSql(" "); - JSONObject result = JSON.parseObject(json); + assertThrows(AiToolInvalidArgumentException.class, + () -> new AiToolServiceImpl().executeSql(request)); + } + + @Test + void getTablesSchemaRejectsEmptyTableNamesAsStrongException() { + AiGetTablesSchemaRequest request = new AiGetTablesSchemaRequest(); + request.setTableNames(List.of()); - assertEquals(true, result.getBoolean("success")); - assertNull(result.get("tool")); - assertEquals("SQL executed successfully with 1 result set(s).", result.getString("summary")); - assertEquals(1, result.getJSONArray("data").size()); - assertNull(result.get("errorCode")); - assertTrue(json.contains("\"errorCode\":null")); + assertThrows(AiToolInvalidArgumentException.class, + () -> new AiToolServiceImpl().getTablesSchema(request)); } @Test - void shouldSerializeFailedToolResultAsStandardJson() { - String json = AiToolServiceImpl.failureToolResultJson( - "sql is empty.", - "INVALID_ARGUMENT"); + void dtoRequestMethodsRejectNullRequestAsInvalidArgument() { + AiToolServiceImpl service = new AiToolServiceImpl(); - JSONObject result = JSON.parseObject(json); + AiToolInvalidArgumentException listTablesException = assertThrows( + AiToolInvalidArgumentException.class, + () -> service.listAllTables(null)); + AiToolInvalidArgumentException executeSqlException = assertThrows( + AiToolInvalidArgumentException.class, + () -> service.executeSql(null)); + AiToolInvalidArgumentException tableSchemaException = assertThrows( + AiToolInvalidArgumentException.class, + () -> service.getTablesSchema(null)); - assertEquals(false, result.getBoolean("success")); - assertNull(result.get("tool")); - assertEquals("sql is empty.", result.getString("summary")); - assertEquals(0, result.getJSONArray("data").size()); - assertEquals("INVALID_ARGUMENT", result.getString("errorCode")); + assertEquals("listAllTables request is required.", listTablesException.getMessage()); + assertEquals("executeSql request is required.", executeSqlException.getMessage()); + assertEquals("getTablesSchema request is required.", tableSchemaException.getMessage()); } @Test - void shouldKeepRowsPositionBasedWhenColumnNamesAreDuplicated() { - List
headers = List.of( - Header.builder().name("id").build(), - Header.builder().name("id").build(), - Header.builder().name("note").build()); - String longText = "a".repeat(201); + void listAllTablesReturnsStrongTypedTablesAndClearsConnectionContext() { + Fixture fixture = new Fixture(); + SimpleTable table = SimpleTable.builder() + .name("orders") + .tableType("BASE TABLE") + .comment("business orders") + .build(); + inject(fixture.service, "connectionContextService", fixture.connectionContextService()); + inject(fixture.service, "tableService", proxy(IDbTableService.class, (proxy, method, args) -> { + if ("queryTables".equals(method.getName())) { + return List.of(table); + } + return defaultValue(method.getReturnType()); + })); + + AiListTablesRequest request = new AiListTablesRequest(); + request.setDataSourceId(1L); + request.setDatabaseName("app"); - List> rows = AiToolServiceImpl.rowPreviewRows( - headers, - List.of(Arrays.asList("first-id\nwith\ttab", null, longText))); + List result = fixture.service.listAllTables(request); - assertEquals(List.of("id", "id", "note"), AiToolServiceImpl.columnNames(headers)); - assertEquals(1, rows.size()); - assertEquals("first-id\nwith\ttab", rows.get(0).get(0)); - assertNull(rows.get(0).get(1)); - assertEquals(longText, rows.get(0).get(2)); + assertEquals(1, result.size()); + assertSame(table, result.get(0)); + assertEquals(1, fixture.bindProfileCalls.get()); + assertEquals(1, fixture.clearCalls.get()); } @Test - void shouldExposeRowPreviewTruncationMetadata() { - Header header = Header.builder().name("id").build(); - List> rows = new ArrayList<>(); - for (int i = 0; i < 51; i++) { - rows.add(List.of(ResultCell.of(String.valueOf(i)))); - } + void listAllDatabasesReturnsStrongTypedDatabasesAndClearsConnectionContext() { + Fixture fixture = new Fixture(); + Database database = Database.builder() + .name("app") + .comment("application database") + .build(); + inject(fixture.service, "connectionContextService", fixture.connectionContextService()); + inject(fixture.service, "databaseService", proxy(IDbDatabaseService.class, (proxy, method, args) -> { + if ("queryAll".equals(method.getName())) { + return List.of(database); + } + return defaultValue(method.getReturnType()); + })); + + List result = fixture.service.listAllDatabases(1L, null); + + assertEquals(1, result.size()); + assertSame(database, result.get(0)); + assertEquals(1, fixture.bindProfileCalls.get()); + assertEquals(1, fixture.clearCalls.get()); + } + + @Test + void getTablesSchemaReturnsStrongTypedUseCaseResultAndClearsConnectionContext() { + Fixture fixture = new Fixture(); + Table table = Table.builder() + .name("orders") + .databaseName("app") + .build(); + inject(fixture.service, "connectionContextService", fixture.connectionContextService()); + inject(fixture.service, "tableService", tableService(table, "create table orders(id bigint)", null)); + + List result = fixture.service.getTablesSchema(tablesSchemaRequest("orders")); + + assertEquals(1, result.size()); + assertEquals("orders", result.get(0).getTableName()); + assertEquals("create table orders(id bigint)", result.get(0).getDdl()); + assertSame(table, result.get(0).getTable()); + assertEquals(1, fixture.bindProfileCalls.get()); + assertEquals(1, fixture.clearCalls.get()); + } + + @Test + void getTablesSchemaThrowsMetadataExceptionAndClearsConnectionContextWhenMetadataQueryFails() { + Fixture fixture = new Fixture(); + inject(fixture.service, "connectionContextService", fixture.connectionContextService()); + inject(fixture.service, "tableService", tableService(null, null, new BusinessException("metadata failed"))); + + AiToolMetadataQueryException exception = assertThrows( + AiToolMetadataQueryException.class, + () -> fixture.service.getTablesSchema(tablesSchemaRequest("orders"))); + + assertEquals(true, exception.getMessage().contains("metadata failed")); + assertEquals(1, fixture.bindProfileCalls.get()); + assertEquals(1, fixture.clearCalls.get()); + } + + @Test + void getTablesSchemaPropagatesUnexpectedRuntimeExceptionAndClearsConnectionContext() { + Fixture fixture = new Fixture(); + IllegalStateException failure = new IllegalStateException("bug failed"); + inject(fixture.service, "connectionContextService", fixture.connectionContextService()); + inject(fixture.service, "tableService", tableService(null, null, failure)); + + IllegalStateException exception = assertThrows( + IllegalStateException.class, + () -> fixture.service.getTablesSchema(tablesSchemaRequest("orders"))); + + assertSame(failure, exception); + assertEquals(1, fixture.bindProfileCalls.get()); + assertEquals(1, fixture.clearCalls.get()); + } + + @Test + void executeSqlReturnsStrongTypedResultCleansConnectionContextAndRecordsOperationLog() { + Fixture fixture = new Fixture(); ExecuteResponse response = ExecuteResponse.builder() - .success(true) - .headerList(List.of(header)) - .dataList(rows) + .success(Boolean.TRUE) + .sqlType("SELECT") + .message("ok") .build(); + injectExecuteDependencies(fixture, List.of(response), null); - Map result = AiToolServiceImpl.executeResponseData(1, response); + List result = fixture.service.executeSql(executeRequest("select 1")); + + assertEquals(1, result.size()); + assertSame(response, result.get(0)); + assertEquals(1, fixture.bindProfileCalls.get()); + assertEquals(1, fixture.clearCalls.get()); + assertEquals(1, fixture.recordListResultCalls.get()); + assertEquals(0, fixture.recordFailureCalls.get()); + } + + @Test + void executeSqlRunsProfileBindExecutionAndLogInsideToolRequestContext() { + Fixture fixture = new Fixture(); + Context requestContext = Context.builder() + .organizationId(99L) + .token("request-token") + .build(); + AiToolContextRequest toolContext = new AiToolContextRequest(); + toolContext.setDataSourceId(1L); + toolContext.setDatabaseName("app"); + toolContext.setRequestContext(requestContext); + AiExecuteSqlRequest request = new AiExecuteSqlRequest(); + request.setSql("select 1"); + request.setAiToolContextRequest(toolContext); + ExecuteResponse response = ExecuteResponse.builder() + .success(Boolean.TRUE) + .sqlType("SELECT") + .build(); + injectExecuteDependencies(fixture, List.of(response), null); + + List result = fixture.service.executeSql(request); + + assertEquals(1, result.size()); + assertEquals(1, fixture.buildProfileCalls.get()); + assertSame(requestContext, fixture.buildProfileContext); + assertSame(requestContext, fixture.bindProfileContext); + assertSame(requestContext, fixture.executeContext); + assertSame(requestContext, fixture.recordListResultContext); + assertNull(ContextUtils.queryContext()); + } + + @Test + void executeSqlThrowsConfirmationExceptionBeforeBindingOrLoggingNonQuerySql() { + Fixture fixture = new Fixture(); + inject(fixture.service, "connectionContextService", fixture.connectionContextService()); + inject(fixture.service, "sqlService", sqlService("UPDATE")); + inject(fixture.service, "sqlOperationLogRecorder", fixture.sqlOperationLogService()); + + AiToolSqlConfirmationRequiredException exception = assertThrows( + AiToolSqlConfirmationRequiredException.class, + () -> fixture.service.executeSql(executeRequest("update users set name = 'x'"))); + + assertEquals(true, exception.getMessage().contains("Non-query SQL cannot be auto-executed")); + assertEquals(0, fixture.bindProfileCalls.get()); + assertEquals(0, fixture.clearCalls.get()); + assertEquals(0, fixture.recordListResultCalls.get()); + assertEquals(0, fixture.recordFailureCalls.get()); + } - assertEquals(51, result.get("rowCount")); - assertEquals(50, result.get("previewRowCount")); - assertEquals(true, result.get("rowsTruncated")); - assertEquals(50, ((List) result.get("rows")).size()); + @Test + void executeSqlThrowsExecutionExceptionRecordsFailureAndClearsContextWhenExecutorFails() { + Fixture fixture = new Fixture(); + injectExecuteDependencies(fixture, null, new BusinessException("driver failed")); + + AiToolSqlExecutionException exception = assertThrows( + AiToolSqlExecutionException.class, + () -> fixture.service.executeSql(executeRequest("select 1"))); + + assertEquals(true, exception.getMessage().contains("driver failed")); + assertEquals(1, fixture.bindProfileCalls.get()); + assertEquals(1, fixture.clearCalls.get()); + assertEquals(0, fixture.recordListResultCalls.get()); + assertEquals(1, fixture.recordFailureCalls.get()); + } + + @Test + void executeSqlPropagatesUnexpectedRuntimeExceptionRecordsFailureAndClearsContext() { + Fixture fixture = new Fixture(); + IllegalStateException failure = new IllegalStateException("driver bug"); + injectExecuteDependencies(fixture, null, failure); + + IllegalStateException exception = assertThrows( + IllegalStateException.class, + () -> fixture.service.executeSql(executeRequest("select 1"))); + + assertSame(failure, exception); + assertEquals(1, fixture.bindProfileCalls.get()); + assertEquals(1, fixture.clearCalls.get()); + assertEquals(0, fixture.recordListResultCalls.get()); + assertEquals(1, fixture.recordFailureCalls.get()); + } + + @Test + void executeSqlThrowsExecutionExceptionAndRecordsListResultWhenResultSetReportsFailure() { + Fixture fixture = new Fixture(); + ExecuteResponse failed = ExecuteResponse.builder() + .success(Boolean.FALSE) + .message("syntax error") + .description("bad sql") + .build(); + injectExecuteDependencies(fixture, List.of(failed), null); + + AiToolSqlExecutionException exception = assertThrows( + AiToolSqlExecutionException.class, + () -> fixture.service.executeSql(executeRequest("select broken"))); + + assertEquals(true, exception.getMessage().contains("syntax error")); + assertEquals(1, fixture.clearCalls.get()); + assertEquals(1, fixture.recordListResultCalls.get()); + assertEquals(0, fixture.recordFailureCalls.get()); + } + + private static void injectExecuteDependencies(Fixture fixture, List responses, RuntimeException failure) { + inject(fixture.service, "connectionContextService", fixture.connectionContextService()); + inject(fixture.service, "sqlService", sqlService("SELECT")); + inject(fixture.service, "sqlOperationLogRecorder", fixture.sqlOperationLogService()); + inject(fixture.service, "dlTemplateService", proxy(IDbDlTemplateService.class, (proxy, method, args) -> { + if ("execute".equals(method.getName())) { + fixture.executedSql = ((DbDlExecuteRequest) args[0]).getSql(); + fixture.executeContext = ContextUtils.queryContext(); + if (failure != null) { + throw failure; + } + return responses; + } + return defaultValue(method.getReturnType()); + })); + } + + private static IDbSqlService sqlService(String sqlType) { + return proxy(IDbSqlService.class, (proxy, method, args) -> { + if ("parseStatements".equals(method.getName())) { + SimpleSqlStatement statement = new SimpleSqlStatement(); + statement.setSql((String) args[0]); + statement.setSqlType(sqlType); + return List.of(statement); + } + return defaultValue(method.getReturnType()); + }); + } + + private static AiExecuteSqlRequest executeRequest(String sql) { + AiExecuteSqlRequest request = new AiExecuteSqlRequest(); + request.setSql(sql); + request.setDataSourceId(1L); + request.setDatabaseName("app"); + return request; + } + + private static AiGetTablesSchemaRequest tablesSchemaRequest(String tableName) { + AiGetTablesSchemaRequest request = new AiGetTablesSchemaRequest(); + request.setTableNames(List.of(tableName)); + request.setDataSourceId(1L); + request.setDatabaseName("app"); + return request; + } + + private static IDbTableService tableService(Table table, String ddl, RuntimeException failure) { + return proxy(IDbTableService.class, (proxy, method, args) -> { + if ("query".equals(method.getName())) { + if (failure != null) { + throw failure; + } + return table; + } + if ("showCreateTable".equals(method.getName())) { + return ddl; + } + return defaultValue(method.getReturnType()); + }); + } + + private static void inject(Object target, String fieldName, Object value) { + try { + Field field = target.getClass().getDeclaredField(fieldName); + field.setAccessible(true); + field.set(target, value); + } catch (ReflectiveOperationException e) { + throw new IllegalStateException(e); + } + } + + @SuppressWarnings("unchecked") + private static T proxy(Class type, InvocationHandler handler) { + return (T) Proxy.newProxyInstance(type.getClassLoader(), new Class[]{type}, handler); + } + + private static Object defaultValue(Class returnType) { + if (returnType == Void.TYPE) { + return null; + } + if (returnType == Boolean.TYPE) { + return false; + } + if (returnType == Integer.TYPE) { + return 0; + } + if (returnType == Long.TYPE) { + return 0L; + } + if (List.class.isAssignableFrom(returnType)) { + return Collections.emptyList(); + } + return null; + } + + private static class Fixture { + private final AiToolServiceImpl service = new AiToolServiceImpl(); + private final AtomicInteger buildProfileCalls = new AtomicInteger(); + private final AtomicInteger bindProfileCalls = new AtomicInteger(); + private final AtomicInteger clearCalls = new AtomicInteger(); + private final AtomicInteger recordListResultCalls = new AtomicInteger(); + private final AtomicInteger recordFailureCalls = new AtomicInteger(); + private String executedSql; + private Context buildProfileContext; + private Context bindProfileContext; + private Context executeContext; + private Context recordListResultContext; + private Context recordFailureContext; + + private IDbConnectionContextService connectionContextService() { + return proxy(IDbConnectionContextService.class, (proxy, method, args) -> { + if ("buildProfile".equals(method.getName())) { + buildProfileCalls.incrementAndGet(); + buildProfileContext = ContextUtils.queryContext(); + ConnectionProfile profile = new ConnectionProfile(); + profile.setDataSourceId(1L); + profile.setDatabaseName("app"); + profile.setDbType("MYSQL"); + return profile; + } + if ("bindProfile".equals(method.getName())) { + bindProfileContext = ContextUtils.queryContext(); + bindProfileCalls.incrementAndGet(); + return null; + } + if ("clear".equals(method.getName())) { + clearCalls.incrementAndGet(); + return null; + } + return defaultValue(method.getReturnType()); + }); + } + + private IOpsSqlOperationLogService sqlOperationLogService() { + return proxy(IOpsSqlOperationLogService.class, (proxy, method, args) -> { + if ("recordListResultAsync".equals(method.getName())) { + recordListResultContext = ContextUtils.queryContext(); + recordListResultCalls.incrementAndGet(); + return null; + } + if ("recordFailureAsync".equals(method.getName())) { + recordFailureContext = ContextUtils.queryContext(); + recordFailureCalls.incrementAndGet(); + return null; + } + return defaultValue(method.getReturnType()); + }); + } } } diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java index 18ccd7a7d..9c62fb267 100644 --- a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapter.java @@ -1,16 +1,17 @@ package ai.chat2db.community.web.api.adapter.ai; +import ai.chat2db.community.domain.api.exception.ai.AiToolException; import ai.chat2db.community.domain.api.model.ai.AiToolResult; import ai.chat2db.community.domain.api.model.request.ai.AiExecuteSqlRequest; import ai.chat2db.community.domain.api.model.request.ai.AiGetTablesSchemaRequest; import ai.chat2db.community.domain.api.model.request.ai.AiListTablesRequest; import ai.chat2db.community.domain.api.model.request.ai.AiToolContextRequest; import ai.chat2db.community.web.api.converter.ai.AiToolContextConverter; -import com.alibaba.fastjson2.JSON; -import com.alibaba.fastjson2.JSONObject; -import com.alibaba.fastjson2.JSONWriter; +import ai.chat2db.community.web.api.converter.ai.AiToolErrorCodeMapper; +import ai.chat2db.community.web.api.converter.ai.AiToolOutput; +import ai.chat2db.community.web.api.converter.ai.AiToolResultConverter; +import ai.chat2db.community.web.api.converter.ai.AiToolResultSerializer; import lombok.extern.slf4j.Slf4j; -import org.apache.commons.lang3.StringUtils; import org.springframework.ai.chat.model.ToolContext; import org.springframework.ai.tool.annotation.Tool; import org.springframework.ai.tool.annotation.ToolParam; @@ -26,17 +27,27 @@ public class AiToolAdapter { private final ai.chat2db.community.domain.api.service.ai.IAiToolService aiToolService; private final AiToolContextConverter aiToolContextConverter; + private final AiToolResultConverter aiToolResultConverter; + private final AiToolResultSerializer aiToolResultSerializer; + private final AiToolErrorCodeMapper aiToolErrorCodeMapper; public AiToolAdapter(ai.chat2db.community.domain.api.service.ai.IAiToolService aiToolService, - AiToolContextConverter aiToolContextConverter) { + AiToolContextConverter aiToolContextConverter, + AiToolResultConverter aiToolResultConverter, + AiToolResultSerializer aiToolResultSerializer, + AiToolErrorCodeMapper aiToolErrorCodeMapper) { this.aiToolService = aiToolService; this.aiToolContextConverter = aiToolContextConverter; + this.aiToolResultConverter = aiToolResultConverter; + this.aiToolResultSerializer = aiToolResultSerializer; + this.aiToolErrorCodeMapper = aiToolErrorCodeMapper; } @Tool(name = "list_all_datasources", description = "List available Chat2DB data sources. Use this first when no datasource is selected.") public String listAllDataSources(ToolContext toolContext) { return invoke(toolContext, "list_all_datasources", - () -> aiToolService.listAllDataSources(aiToolContextConverter.toParam(toolContext))); + () -> aiToolResultConverter.fromDataSources( + aiToolService.listAllDataSources(aiToolContextConverter.toParam(toolContext)))); } @Tool(name = "list_all_tables", description = "List all tables in the connected database with comments and type.") @@ -46,8 +57,9 @@ public String listAllTables( @ToolParam(description = "Optional target schema name. If omitted, uses selected schema context.", required = false) String schemaName, ToolContext toolContext) { return invoke(toolContext, "list_all_tables", - () -> aiToolService.listAllTables(listTablesRequest(dataSourceId, databaseName, schemaName, - aiToolContextConverter.toParam(toolContext)))); + () -> aiToolResultConverter.fromTables( + aiToolService.listAllTables(listTablesRequest(dataSourceId, databaseName, schemaName, + aiToolContextConverter.toParam(toolContext))))); } @Tool(name = "list_all_databases", description = "List all databases for the connected data source.") @@ -55,7 +67,8 @@ public String listAllDatabases( @ToolParam(description = "Optional datasource id. Required when no datasource is selected in context.", required = false) Long dataSourceId, ToolContext toolContext) { return invoke(toolContext, "list_all_databases", - () -> aiToolService.listAllDatabases(dataSourceId, aiToolContextConverter.toParam(toolContext))); + () -> aiToolResultConverter.fromDatabases( + aiToolService.listAllDatabases(dataSourceId, aiToolContextConverter.toParam(toolContext)))); } @Tool(name = "list_all_schemas", description = "List all schemas in the selected database. If databaseName is empty, uses current database context.") @@ -64,7 +77,9 @@ public String listAllSchemas( @ToolParam(description = "Optional datasource id. Required when no datasource is selected in context.", required = false) Long dataSourceId, ToolContext toolContext) { return invoke(toolContext, "list_all_schemas", - () -> aiToolService.listAllSchemas(databaseName, dataSourceId, aiToolContextConverter.toParam(toolContext))); + () -> aiToolResultConverter.fromSchemas( + aiToolService.listAllSchemas(databaseName, dataSourceId, + aiToolContextConverter.toParam(toolContext)))); } @Tool(name = "execute_sql", description = "Execute SQL in current database context and return concise result (rows for SELECT, update count for DML/DDL).") @@ -76,8 +91,9 @@ public String executeSql( @ToolParam(description = "Optional target schema name. If omitted, uses selected schema context.", required = false) String schemaName, ToolContext toolContext) { return invoke(toolContext, "execute_sql", - () -> aiToolService.executeSql(executeSqlRequest(sql, pageSize, dataSourceId, databaseName, schemaName, - aiToolContextConverter.toParam(toolContext)))); + () -> aiToolResultConverter.fromExecuteResult( + aiToolService.executeSql(executeSqlRequest(sql, pageSize, dataSourceId, databaseName, schemaName, + aiToolContextConverter.toParam(toolContext))))); } @Tool(name = "get_tables_schema", description = "Get CREATE TABLE DDL for specific tables. Returns DDL first, then falls back to structured columns.") @@ -88,8 +104,9 @@ public String getTablesSchema( @ToolParam(description = "Optional target schema name. If omitted, uses selected schema context.", required = false) String schemaName, ToolContext toolContext) { return invoke(toolContext, "get_tables_schema", - () -> aiToolService.getTablesSchema(tablesSchemaRequest(tableNames, dataSourceId, databaseName, schemaName, - aiToolContextConverter.toParam(toolContext)))); + () -> aiToolResultConverter.fromTableSchemas( + aiToolService.getTablesSchema(tablesSchemaRequest(tableNames, dataSourceId, databaseName, schemaName, + aiToolContextConverter.toParam(toolContext))))); } private AiListTablesRequest listTablesRequest(Long dataSourceId, String databaseName, String schemaName, @@ -126,37 +143,28 @@ private AiGetTablesSchemaRequest tablesSchemaRequest(List tableNames, Lo return request; } - private String invoke(ToolContext toolContext, String toolName, Supplier action) { + private String invoke(ToolContext toolContext, String toolName, Supplier> action) { try { - return emit(toolContext, toolName, action.get()); + AiToolOutput output = action.get(); + return emit(toolContext, toolName, AiToolResult.success(output.summary(), output.data())); + } catch (AiToolException e) { + return emit(toolContext, toolName, AiToolResult.failureWithCode( + aiToolErrorCodeMapper.errorCodeFor(e), + e.getMessage())); } catch (Exception e) { log.error("AI tool call failed, tool={}", toolName, e); - String message = "Tool call failed: " + StringUtils.defaultIfBlank(e.getMessage(), "Unknown error"); - return emit(toolContext, toolName, - JSON.toJSONString(AiToolResult.failure(message, "TOOL_CALL_FAILED"), - JSONWriter.Feature.WriteNulls)); + return emit(toolContext, toolName, AiToolResult.failureWithCode( + "TOOL_CALL_FAILED", + "Tool execution failed.")); } } - private String emit(ToolContext toolContext, String toolName, String content) { + private String emit(ToolContext toolContext, String toolName, AiToolResult result) { Map payload = AiChatTraceSupport.payload(AiChatTraceSupport.TYPE_TOOL_RESULT); payload.put("name", toolName); - payload.put("content", traceSummary(content)); + payload.put("content", result.getSummary()); AiChatTraceSupport.emit(toolContext, payload); - return content; - } - - static String traceSummary(String content) { - if (StringUtils.isBlank(content)) { - return StringUtils.EMPTY; - } - try { - JSONObject result = JSON.parseObject(content); - String summary = result.getString("summary"); - return StringUtils.defaultIfBlank(summary, content); - } catch (Exception ignored) { - return content; - } + return aiToolResultSerializer.toJson(result); } } diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolErrorCodeMapper.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolErrorCodeMapper.java new file mode 100644 index 000000000..381172420 --- /dev/null +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolErrorCodeMapper.java @@ -0,0 +1,28 @@ +package ai.chat2db.community.web.api.converter.ai; + +import ai.chat2db.community.domain.api.exception.ai.AiToolException; +import ai.chat2db.community.domain.api.exception.ai.AiToolInvalidArgumentException; +import ai.chat2db.community.domain.api.exception.ai.AiToolMetadataQueryException; +import ai.chat2db.community.domain.api.exception.ai.AiToolSqlConfirmationRequiredException; +import ai.chat2db.community.domain.api.exception.ai.AiToolSqlExecutionException; +import org.springframework.stereotype.Component; + +@Component +public class AiToolErrorCodeMapper { + + public String errorCodeFor(AiToolException e) { + if (e instanceof AiToolInvalidArgumentException) { + return "INVALID_ARGUMENT"; + } + if (e instanceof AiToolSqlConfirmationRequiredException) { + return "SQL_REQUIRES_MANUAL_CONFIRMATION"; + } + if (e instanceof AiToolSqlExecutionException) { + return "SQL_EXECUTION_FAILED"; + } + if (e instanceof AiToolMetadataQueryException) { + return "METADATA_QUERY_FAILED"; + } + return "TOOL_CALL_FAILED"; + } +} diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolOutput.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolOutput.java new file mode 100644 index 000000000..19cdbbacf --- /dev/null +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolOutput.java @@ -0,0 +1,8 @@ +package ai.chat2db.community.web.api.converter.ai; + +/** + * Internal projection result passed from AI tool converters to transport adapters. + * This is not a public tool protocol envelope. + */ +public record AiToolOutput(String summary, T data) { +} diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolResultConverter.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolResultConverter.java new file mode 100644 index 000000000..80ab6f5cf --- /dev/null +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolResultConverter.java @@ -0,0 +1,553 @@ +package ai.chat2db.community.web.api.converter.ai; + +import ai.chat2db.community.domain.api.model.ai.DataSourceToolData; +import ai.chat2db.community.domain.api.model.ai.DatabaseToolData; +import ai.chat2db.community.domain.api.model.ai.SchemaToolData; +import ai.chat2db.community.domain.api.model.ai.SqlToolData; +import ai.chat2db.community.domain.api.model.ai.TableSchemaResult; +import ai.chat2db.community.domain.api.model.ai.TableSchemaToolData; +import ai.chat2db.community.domain.api.model.ai.TableToolData; +import ai.chat2db.community.domain.api.model.ai.Text2SqlToolData; +import ai.chat2db.community.domain.api.model.metadata.Database; +import ai.chat2db.community.domain.api.model.metadata.ForeignKeyInfo; +import ai.chat2db.community.domain.api.model.metadata.Schema; +import ai.chat2db.community.domain.api.model.metadata.SimpleTable; +import ai.chat2db.community.domain.api.model.metadata.Table; +import ai.chat2db.community.domain.api.model.metadata.TableColumn; +import ai.chat2db.community.domain.api.model.metadata.TableIndex; +import ai.chat2db.community.domain.api.model.metadata.TableIndexColumn; +import ai.chat2db.community.domain.api.model.result.ExecuteResponse; +import ai.chat2db.community.domain.api.model.result.Header; +import ai.chat2db.community.domain.api.model.result.ResultCell; +import ai.chat2db.community.domain.api.model.storage.WorkspaceDataSource; +import org.apache.commons.collections4.CollectionUtils; +import org.apache.commons.lang3.StringUtils; +import org.springframework.stereotype.Component; + +import java.sql.Time; +import java.sql.Timestamp; +import java.time.Instant; +import java.time.LocalDate; +import java.time.LocalDateTime; +import java.time.LocalTime; +import java.time.MonthDay; +import java.time.OffsetDateTime; +import java.time.OffsetTime; +import java.time.Year; +import java.time.YearMonth; +import java.time.ZonedDateTime; +import java.time.temporal.TemporalAccessor; +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.Date; +import java.util.IdentityHashMap; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.stream.Collectors; + +@Component +public class AiToolResultConverter { + + // AI preview limit only affects the serialized tool payload; it is not a database execution limit. + private static final int AI_SQL_PREVIEW_ROW_LIMIT = 50; + + public AiToolOutput fromDataSources(List dataSources) { + List items = emptyIfNull(dataSources).stream() + .filter(Objects::nonNull) + .map(dataSource -> new DataSourceToolData.Item( + dataSource.getId(), + StringUtils.defaultIfBlank(dataSource.getAlias(), "(unnamed)"), + dataSource.getType(), + dataSource.getEnvType())) + .collect(Collectors.toList()); + String summary = items.isEmpty() ? "No datasources found." : "Found " + items.size() + " datasource(s)."; + return new AiToolOutput<>(summary, new DataSourceToolData(items)); + } + + public AiToolOutput fromTables(List tables) { + List items = emptyIfNull(tables).stream() + .filter(Objects::nonNull) + .map(table -> new TableToolData.Item( + StringUtils.defaultString(table.getName(), "(unnamed)"), + StringUtils.defaultIfBlank(table.getTableType(), "TABLE"), + table.getComment())) + .collect(Collectors.toList()); + String summary = items.isEmpty() ? "No tables found." : "Found " + items.size() + " table(s)."; + return new AiToolOutput<>(summary, new TableToolData(items)); + } + + public AiToolOutput fromDatabases(List databases) { + List items = emptyIfNull(databases).stream() + .filter(Objects::nonNull) + .map(database -> new DatabaseToolData.Item( + StringUtils.defaultString(database.getName(), "(unnamed)"), + database.isSystem(), + database.getComment())) + .collect(Collectors.toList()); + String summary = items.isEmpty() ? "No databases found." : "Found " + items.size() + " database(s)."; + return new AiToolOutput<>(summary, new DatabaseToolData(items)); + } + + public AiToolOutput fromSchemas(List schemas) { + List items = emptyIfNull(schemas).stream() + .filter(Objects::nonNull) + .map(schema -> new SchemaToolData.Item( + StringUtils.defaultString(schema.getName(), "(unnamed)"), + schema.isSystem(), + schema.getComment())) + .collect(Collectors.toList()); + String summary = items.isEmpty() ? "No schemas found." : "Found " + items.size() + " schema(s)."; + return new AiToolOutput<>(summary, new SchemaToolData(items)); + } + + public AiToolOutput fromExecuteResult(List executeResponses) { + List items = new ArrayList<>(); + int index = 1; + for (ExecuteResponse response : emptyIfNull(executeResponses)) { + items.add(executeResponseData(index++, response)); + } + String summary = items.isEmpty() + ? "SQL executed successfully with no result." + : "SQL executed successfully with " + items.size() + " result set(s)."; + return new AiToolOutput<>(summary, new SqlToolData(items)); + } + + public AiToolOutput fromTableSchemas(List schemaResults) { + List items = emptyIfNull(schemaResults).stream() + .filter(Objects::nonNull) + .map(result -> new TableSchemaToolData.Item( + result.getTableName(), + buildRichTableSchema(result.getTableName(), result.getDdl(), result.getTable()))) + .collect(Collectors.toList()); + String summary = items.isEmpty() + ? "No table schema found." + : "Fetched schema for " + items.size() + " table(s)."; + return new AiToolOutput<>(summary, new TableSchemaToolData(items)); + } + + public AiToolOutput fromText2Sql(String sql) { + return new AiToolOutput<>( + "SQL generated successfully.", + new Text2SqlToolData(StringUtils.defaultString(sql))); + } + + static SqlToolData.ResultSet executeResponseData(int index, ExecuteResponse result) { + SqlToolData.ResultSet item = new SqlToolData.ResultSet(); + item.setResultIndex(index); + if (Objects.isNull(result)) { + item.setSuccess(Boolean.FALSE); + item.setMessage("Empty result."); + item.setText("Empty result."); + return item; + } + item.setSuccess(Boolean.TRUE.equals(result.getSuccess())); + item.setSqlType(result.getSqlType()); + item.setDurationMs(result.getDuration()); + item.setUpdateCount(result.getUpdateCount()); + item.setMessage(result.getMessage()); + item.setDescription(result.getDescription()); + item.setHasNextPage(result.getHasNextPage()); + int rowCount = result.getDataList() == null ? 0 : result.getDataList().size(); + int previewRowCount = Math.min(rowCount, AI_SQL_PREVIEW_ROW_LIMIT); + item.setRowCount(rowCount); + item.setPreviewRowCount(previewRowCount); + item.setRowsTruncated(rowCount > previewRowCount); + item.setColumns(columnNames(result.getHeaderList())); + item.setRows(rowPreviewRows(result.getHeaderList(), result.getDataList())); + item.setRowCellMetadata(rowPreviewCellMetadata(result.getHeaderList(), result.getDataList())); + item.setText(formatExecuteResponse(result)); + return item; + } + + static List columnNames(List
headers) { + if (CollectionUtils.isEmpty(headers)) { + return Collections.emptyList(); + } + return headers.stream() + .map(header -> StringUtils.defaultIfBlank(header.getName(), header.getColumnName())) + .map(name -> StringUtils.defaultIfBlank(name, "col")) + .collect(Collectors.toList()); + } + + static List> rowPreviewRows(List
headers, List> rows) { + if (CollectionUtils.isEmpty(headers) || CollectionUtils.isEmpty(rows)) { + return Collections.emptyList(); + } + List headerNames = columnNames(headers); + int rowCount = Math.min(rows.size(), AI_SQL_PREVIEW_ROW_LIMIT); + List> result = new ArrayList<>(rowCount); + for (int i = 0; i < rowCount; i++) { + List row = rows.get(i); + List rowData = new ArrayList<>(headerNames.size()); + for (int c = 0; c < headerNames.size(); c++) { + ResultCell cell = row != null && c < row.size() ? row.get(c) : null; + rowData.add(cellValue(cell)); + } + result.add(rowData); + } + return result; + } + + static List> rowPreviewCellMetadata(List
headers, List> rows) { + if (CollectionUtils.isEmpty(headers) || CollectionUtils.isEmpty(rows)) { + return Collections.emptyList(); + } + List headerNames = columnNames(headers); + int rowCount = Math.min(rows.size(), AI_SQL_PREVIEW_ROW_LIMIT); + List> result = new ArrayList<>(rowCount); + for (int i = 0; i < rowCount; i++) { + List row = rows.get(i); + List rowData = new ArrayList<>(headerNames.size()); + for (int c = 0; c < headerNames.size(); c++) { + ResultCell cell = row != null && c < row.size() ? row.get(c) : null; + rowData.add(cellMetadata(cell)); + } + result.add(rowData); + } + return result; + } + + private static Object cellValue(ResultCell cell) { + if (cell == null) { + return null; + } + JsonSafeRawValue rawValue = jsonSafeRawValue(cell); + if (rawValue.safe) { + return rawValue.value; + } + return cell.getValue(); + } + + private static SqlToolData.CellMetadata cellMetadata(ResultCell cell) { + if (cell == null || jsonSafeRawValue(cell).safe || isSqlNull(cell)) { + return null; + } + return new SqlToolData.CellMetadata( + Boolean.FALSE, + cell.getValue(), + cell.isLargeValue(), + cell.getLargeValueId(), + cell.getValueType(), + cell.getSqlType(), + cell.getColumnType(), + cell.getSizeBytes(), + cell.getSizeChars(), + cell.getLoadedBytes(), + cell.getLoadedChars(), + cell.isTruncated(), + cell.getUnsupportedReason(), + rawValueUnavailableReason(cell)); + } + + private static JsonSafeRawValue jsonSafeRawValue(ResultCell cell) { + Object rawValue = cell.getRawValue(); + if (rawValue == null || cell.isLargeValue() || cell.isTruncated()) { + return JsonSafeRawValue.unsafe(); + } + return jsonSafeValue(rawValue, new IdentityHashMap<>()); + } + + private static JsonSafeRawValue jsonSafeValue(Object value, IdentityHashMap visiting) { + if (value == null + || value instanceof String + || value instanceof Number + || value instanceof Boolean) { + return JsonSafeRawValue.safe(value); + } + if (value instanceof java.sql.Date date) { + return JsonSafeRawValue.safe(date.toLocalDate().toString()); + } + if (value instanceof Time time) { + return JsonSafeRawValue.safe(time.toLocalTime().toString()); + } + if (value instanceof Timestamp timestamp) { + return JsonSafeRawValue.safe(timestamp.toInstant().toString()); + } + if (value instanceof Date date) { + return JsonSafeRawValue.safe(date.toInstant().toString()); + } + if (value instanceof Instant + || value instanceof LocalDate + || value instanceof LocalTime + || value instanceof LocalDateTime + || value instanceof OffsetDateTime + || value instanceof OffsetTime + || value instanceof ZonedDateTime + || value instanceof Year + || value instanceof YearMonth + || value instanceof MonthDay) { + return JsonSafeRawValue.safe(value.toString()); + } + if (value instanceof Map map) { + if (visiting.containsKey(value)) { + return JsonSafeRawValue.unsafe(); + } + visiting.put(value, Boolean.TRUE); + Map normalized = new LinkedHashMap<>(map.size()); + for (Map.Entry entry : map.entrySet()) { + if (!(entry.getKey() instanceof String key)) { + visiting.remove(value); + return JsonSafeRawValue.unsafe(); + } + JsonSafeRawValue entryValue = jsonSafeValue(entry.getValue(), visiting); + if (!entryValue.safe) { + visiting.remove(value); + return JsonSafeRawValue.unsafe(); + } + normalized.put(key, entryValue.value); + } + visiting.remove(value); + return JsonSafeRawValue.safe(normalized); + } + if (value instanceof List list) { + if (visiting.containsKey(value)) { + return JsonSafeRawValue.unsafe(); + } + visiting.put(value, Boolean.TRUE); + List normalized = new ArrayList<>(list.size()); + for (Object item : list) { + JsonSafeRawValue itemValue = jsonSafeValue(item, visiting); + if (!itemValue.safe) { + visiting.remove(value); + return JsonSafeRawValue.unsafe(); + } + normalized.add(itemValue.value); + } + visiting.remove(value); + return JsonSafeRawValue.safe(normalized); + } + if (value instanceof TemporalAccessor temporalAccessor) { + return JsonSafeRawValue.safe(temporalAccessor.toString()); + } + return JsonSafeRawValue.unsafe(); + } + + private static String rawValueUnavailableReason(ResultCell cell) { + if (cell.isLargeValue()) { + return "LARGE_VALUE"; + } + if (cell.isTruncated()) { + return "TRUNCATED_VALUE"; + } + Object rawValue = cell.getRawValue(); + if (rawValue != null) { + return "UNSAFE_RAW_VALUE:" + rawValue.getClass().getName(); + } + return "RAW_VALUE_NULL"; + } + + private record JsonSafeRawValue(boolean safe, Object value) { + private static JsonSafeRawValue safe(Object value) { + return new JsonSafeRawValue(true, value); + } + + private static JsonSafeRawValue unsafe() { + return new JsonSafeRawValue(false, null); + } + } + + private static boolean isSqlNull(ResultCell cell) { + return cell.getValue() == null + && !cell.isLargeValue() + && !cell.isTruncated() + && cell.getSizeBytes() == null + && cell.getSizeChars() == null + && cell.getLoadedBytes() == null + && cell.getLoadedChars() == null + && StringUtils.isBlank(cell.getLargeValueId()) + && StringUtils.isBlank(cell.getUnsupportedReason()); + } + + private static String formatExecuteResponse(ExecuteResponse result) { + if (Objects.isNull(result)) { + return "Empty result."; + } + StringBuilder builder = new StringBuilder(1024); + builder.append("success: ").append(Boolean.TRUE.equals(result.getSuccess())).append("\n"); + if (StringUtils.isNotBlank(result.getSqlType())) { + builder.append("sqlType: ").append(result.getSqlType()).append("\n"); + } + if (Objects.nonNull(result.getDuration())) { + builder.append("durationMs: ").append(result.getDuration()).append("\n"); + } + if (Objects.nonNull(result.getUpdateCount())) { + builder.append("updateCount: ").append(result.getUpdateCount()).append("\n"); + } + if (StringUtils.isNotBlank(result.getMessage())) { + builder.append("message: ").append(result.getMessage()).append("\n"); + } + if (StringUtils.isNotBlank(result.getDescription())) { + builder.append("description: ").append(result.getDescription()).append("\n"); + } + if (CollectionUtils.isNotEmpty(result.getHeaderList()) && CollectionUtils.isNotEmpty(result.getDataList())) { + builder.append("rows: ").append(result.getDataList().size()); + if (Objects.nonNull(result.getHasNextPage())) { + builder.append(", hasNextPage: ").append(result.getHasNextPage()); + } + builder.append("\n"); + appendTabularPreview(builder, result.getHeaderList(), result.getDisplayDataList()); + } + return builder.toString().trim(); + } + + private static void appendTabularPreview(StringBuilder builder, List
headers, List> rows) { + if (CollectionUtils.isEmpty(rows)) { + return; + } + List headerNames = columnNames(headers); + builder.append(String.join("\t", headerNames)).append("\n"); + int rowCount = Math.min(rows.size(), AI_SQL_PREVIEW_ROW_LIMIT); + for (int i = 0; i < rowCount; i++) { + List row = rows.get(i); + List normalized = new ArrayList<>(headerNames.size()); + for (int c = 0; c < headerNames.size(); c++) { + String value = row != null && c < row.size() ? row.get(c) : null; + normalized.add(normalizeCell(value)); + } + builder.append(String.join("\t", normalized)).append("\n"); + } + if (rows.size() > rowCount) { + builder.append("... ").append(rows.size() - rowCount).append(" more rows not shown."); + } + } + + private static String normalizeCell(String value) { + if (value == null) { + return "NULL"; + } + String normalized = value.replace("\n", "\\n").replace("\r", "\\r").replace("\t", " "); + if (normalized.length() > 200) { + return normalized.substring(0, 197) + "..."; + } + return normalized; + } + + private String buildRichTableSchema(String tableName, String ddl, Table table) { + StringBuilder builder = new StringBuilder(2048); + builder.append("-- TABLE: ").append(tableName).append("\n"); + builder.append("/* physical schema */\n"); + builder.append(StringUtils.defaultIfBlank(ddl, "-- schema unavailable")); + + String primaryKeys = formatPrimaryKeys(table); + if (StringUtils.isNotBlank(primaryKeys)) { + builder.append("\n\n").append(primaryKeys); + } + + String indexes = formatIndexes(table); + if (StringUtils.isNotBlank(indexes)) { + builder.append("\n\n").append(indexes); + } + + String foreignKeys = formatForeignKeys(table); + if (StringUtils.isNotBlank(foreignKeys)) { + builder.append("\n\n").append(foreignKeys); + } + + return builder.toString(); + } + + private String formatPrimaryKeys(Table table) { + if (table == null || CollectionUtils.isEmpty(table.getColumnList())) { + return null; + } + List primaryKeys = table.getColumnList().stream() + .filter(column -> Boolean.TRUE.equals(column.getPrimaryKey())) + .sorted(Comparator.comparingInt(column -> Objects.requireNonNullElse(column.getPrimaryKeyOrder(), 0))) + .toList(); + if (CollectionUtils.isEmpty(primaryKeys)) { + return null; + } + List lines = new ArrayList<>(); + lines.add("/* primary keys */"); + lines.add(primaryKeys.stream() + .map(TableColumn::getName) + .filter(StringUtils::isNotBlank) + .collect(Collectors.joining(", "))); + return String.join("\n", lines); + } + + private String formatIndexes(Table table) { + if (table == null || CollectionUtils.isEmpty(table.getIndexList())) { + return null; + } + List lines = new ArrayList<>(); + lines.add("/* indexes */"); + for (TableIndex index : table.getIndexList()) { + List columns = index.getColumnList(); + String columnNames = CollectionUtils.isEmpty(columns) + ? "" + : columns.stream() + .sorted(Comparator.comparing(column -> Objects.requireNonNullElse(column.getOrdinalPosition(), (short) 0))) + .map(TableIndexColumn::getColumnName) + .filter(StringUtils::isNotBlank) + .collect(Collectors.joining(", ")); + List parts = new ArrayList<>(); + parts.add("type=" + StringUtils.defaultIfBlank(index.getType(), "INDEX")); + parts.add("unique=" + Boolean.TRUE.equals(index.getUnique())); + if (StringUtils.isNotBlank(index.getMethod())) { + parts.add("method=" + index.getMethod()); + } + if (StringUtils.isNotBlank(index.getComment())) { + parts.add("comment=" + index.getComment()); + } + lines.add("- " + StringUtils.defaultIfBlank(index.getName(), "(unnamed)") + + (StringUtils.isNotBlank(columnNames) ? " (" + columnNames + ")" : "") + + " | " + String.join("; ", parts)); + } + return lines.size() > 1 ? String.join("\n", lines) : null; + } + + private String formatForeignKeys(Table table) { + if (table == null || CollectionUtils.isEmpty(table.getForeignKeyList())) { + return null; + } + Map> grouped = new LinkedHashMap<>(); + for (ForeignKeyInfo foreignKey : table.getForeignKeyList()) { + String key = firstNonBlank(foreignKey.getFkName(), + foreignKey.getFkTableName() + "->" + foreignKey.getPkTableName()); + grouped.computeIfAbsent(key, ignored -> new ArrayList<>()).add(foreignKey); + } + + List lines = new ArrayList<>(); + lines.add("/* foreign keys */"); + for (Map.Entry> entry : grouped.entrySet()) { + List fkList = entry.getValue().stream() + .sorted(Comparator.comparingInt(ForeignKeyInfo::getKeySeq)) + .toList(); + String fkColumns = fkList.stream() + .map(ForeignKeyInfo::getFkColumnName) + .filter(StringUtils::isNotBlank) + .collect(Collectors.joining(", ")); + String pkTable = fkList.stream() + .map(ForeignKeyInfo::getPkTableName) + .filter(StringUtils::isNotBlank) + .findFirst() + .orElse("(unknown)"); + String pkColumns = fkList.stream() + .map(ForeignKeyInfo::getPkColumnName) + .filter(StringUtils::isNotBlank) + .collect(Collectors.joining(", ")); + lines.add("- " + entry.getKey() + ": (" + fkColumns + ") -> " + pkTable + "(" + pkColumns + ")"); + } + return lines.size() > 1 ? String.join("\n", lines) : null; + } + + private String firstNonBlank(String... values) { + if (values == null) { + return null; + } + for (String value : values) { + if (StringUtils.isNotBlank(value)) { + return value; + } + } + return null; + } + + private List emptyIfNull(List items) { + return items == null ? Collections.emptyList() : items; + } +} diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolResultSerializer.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolResultSerializer.java new file mode 100644 index 000000000..972336d4a --- /dev/null +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/converter/ai/AiToolResultSerializer.java @@ -0,0 +1,14 @@ +package ai.chat2db.community.web.api.converter.ai; + +import ai.chat2db.community.domain.api.model.ai.AiToolResult; +import com.alibaba.fastjson2.JSON; +import com.alibaba.fastjson2.JSONWriter; +import org.springframework.stereotype.Component; + +@Component +public class AiToolResultSerializer { + + public String toJson(AiToolResult result) { + return JSON.toJSONString(result, JSONWriter.Feature.WriteNulls); + } +} diff --git a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java index 67d8f75f5..9546bd87e 100644 --- a/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java +++ b/chat2db-community-server/chat2db-community-web/src/main/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapter.java @@ -2,17 +2,24 @@ import ai.chat2db.community.tools.annotation.NotCliRuntime; +import ai.chat2db.community.domain.api.exception.ai.AiToolException; import ai.chat2db.community.domain.api.model.ai.AiToolResult; +import ai.chat2db.community.domain.api.model.request.ai.AiExecuteSqlRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiGetTablesSchemaRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiListTablesRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiToolContextRequest; +import ai.chat2db.community.domain.api.service.ai.IAiToolService; import ai.chat2db.community.web.api.enums.ai.QuestionTypeEnum; import ai.chat2db.community.web.api.model.request.ai.ChatRequest; import ai.chat2db.community.web.api.adapter.ai.AiChatStreamAdapter; -import ai.chat2db.community.web.api.adapter.ai.AiToolAdapter; +import ai.chat2db.community.web.api.converter.ai.AiToolContextConverter; +import ai.chat2db.community.web.api.converter.ai.AiToolErrorCodeMapper; +import ai.chat2db.community.web.api.converter.ai.AiToolOutput; +import ai.chat2db.community.web.api.converter.ai.AiToolResultConverter; +import ai.chat2db.community.web.api.converter.ai.AiToolResultSerializer; import ai.chat2db.community.tools.model.Context; import ai.chat2db.community.tools.util.ContextUtils; -import com.alibaba.fastjson2.JSON; -import com.alibaba.fastjson2.JSONWriter; import lombok.extern.slf4j.Slf4j; -import org.apache.commons.lang3.StringUtils; import org.springframework.ai.chat.model.ToolContext; import org.springframework.ai.tool.annotation.Tool; import org.springframework.ai.tool.annotation.ToolParam; @@ -21,25 +28,43 @@ import java.util.LinkedHashMap; import java.util.List; import java.util.Map; -import java.util.function.Function; +import java.util.function.Supplier; @Component @Slf4j @NotCliRuntime public class AiToolMcpAdapter { - private final AiToolAdapter aiToolAdapter; + private final IAiToolService aiToolService; + + private final AiToolContextConverter aiToolContextConverter; + + private final AiToolResultConverter aiToolResultConverter; private final AiChatStreamAdapter aiChatStreamAdapter; - public AiToolMcpAdapter(AiToolAdapter aiToolAdapter, AiChatStreamAdapter aiChatStreamAdapter) { - this.aiToolAdapter = aiToolAdapter; + private final AiToolResultSerializer aiToolResultSerializer; + + private final AiToolErrorCodeMapper aiToolErrorCodeMapper; + + public AiToolMcpAdapter(IAiToolService aiToolService, + AiToolContextConverter aiToolContextConverter, + AiToolResultConverter aiToolResultConverter, + AiChatStreamAdapter aiChatStreamAdapter, + AiToolResultSerializer aiToolResultSerializer, + AiToolErrorCodeMapper aiToolErrorCodeMapper) { + this.aiToolService = aiToolService; + this.aiToolContextConverter = aiToolContextConverter; + this.aiToolResultConverter = aiToolResultConverter; this.aiChatStreamAdapter = aiChatStreamAdapter; + this.aiToolResultSerializer = aiToolResultSerializer; + this.aiToolErrorCodeMapper = aiToolErrorCodeMapper; } @Tool(name = "list_all_datasources", description = "List available Chat2DB datasources. Use this first, then pass the returned dataSourceId to datasource-scoped tools.") public String listAllDataSources() { - return invoke("list_all_datasources", aiToolAdapter::listAllDataSources); + return invoke("list_all_datasources", + () -> aiToolResultConverter.fromDataSources(aiToolService.listAllDataSources(contextRequest()))); } @Tool(name = "list_all_tables", description = "List all tables in a target database. Call list_all_datasources and list_all_databases first, then pass dataSourceId and databaseName explicitly.") @@ -47,13 +72,16 @@ public String listAllTables( @ToolParam(description = "Datasource id returned by list_all_datasources.", required = true) Long dataSourceId, @ToolParam(description = "Target database name returned by list_all_databases.", required = true) String databaseName, @ToolParam(description = "Optional target schema name returned by list_all_schemas.", required = false) String schemaName) { - return invoke("list_all_tables", toolContext -> aiToolAdapter.listAllTables(dataSourceId, databaseName, schemaName, toolContext)); + return invoke("list_all_tables", + () -> aiToolResultConverter.fromTables( + aiToolService.listAllTables(listTablesRequest(dataSourceId, databaseName, schemaName)))); } @Tool(name = "list_all_databases", description = "List all databases available on a target Chat2DB datasource. Call list_all_datasources first, then pass dataSourceId explicitly.") public String listAllDatabases( @ToolParam(description = "Datasource id returned by list_all_datasources.", required = true) Long dataSourceId) { - return invoke("list_all_databases", toolContext -> aiToolAdapter.listAllDatabases(dataSourceId, toolContext)); + return invoke("list_all_databases", + () -> aiToolResultConverter.fromDatabases(aiToolService.listAllDatabases(dataSourceId, contextRequest()))); } @Tool(name = "list_all_schemas", description = "List all schemas in a target database. Call list_all_datasources and list_all_databases first, then pass targetDatabaseName and dataSourceId explicitly.") @@ -62,7 +90,8 @@ public String listAllSchemas( @ToolParam(description = "Datasource id returned by list_all_datasources.", required = true) Long dataSourceId) { return invoke( "list_all_schemas", - toolContext -> aiToolAdapter.listAllSchemas(targetDatabaseName, dataSourceId, toolContext)); + () -> aiToolResultConverter.fromSchemas( + aiToolService.listAllSchemas(targetDatabaseName, dataSourceId, contextRequest()))); } @Tool(name = "execute_sql", description = "Execute SQL against the target database and return a concise result. Pass dataSourceId and databaseName explicitly.") @@ -74,7 +103,8 @@ public String executeSql( @ToolParam(description = "Optional target schema name returned by list_all_schemas.", required = false) String schemaName) { return invoke( "execute_sql", - toolContext -> aiToolAdapter.executeSql(sql, pageSize, dataSourceId, databaseName, schemaName, toolContext)); + () -> aiToolResultConverter.fromExecuteResult( + aiToolService.executeSql(executeSqlRequest(sql, pageSize, dataSourceId, databaseName, schemaName)))); } @Tool(name = "get_tables_schema", description = "Get CREATE TABLE DDL or structured schema for specific tables. Pass dataSourceId and databaseName explicitly.") @@ -85,7 +115,8 @@ public String getTablesSchema( @ToolParam(description = "Optional target schema name returned by list_all_schemas.", required = false) String schemaName) { return invoke( "get_tables_schema", - toolContext -> aiToolAdapter.getTablesSchema(tableNames, dataSourceId, databaseName, schemaName, toolContext)); + () -> aiToolResultConverter.fromTableSchemas( + aiToolService.getTablesSchema(tablesSchemaRequest(tableNames, dataSourceId, databaseName, schemaName)))); } @Tool(name = "text2sql", description = "Convert a natural language question into SQL using Chat2DB's internal AI. Pass datasource context explicitly when targeting a specific database.") @@ -104,19 +135,18 @@ public String text2sql( chatRequest.setEnableTools(Boolean.TRUE); String sql = aiChatStreamAdapter.chatSync(chatRequest); - return toolSuccess( - "SQL generated successfully.", - List.of(Map.of("sql", StringUtils.defaultString(sql)))); + return toolSuccess(aiToolResultConverter.fromText2Sql(sql)); } catch (Exception e) { log.error("MCP tool call failed, tool=text2sql", e); return toolFailure("text2sql", e); } } - private String invoke(String toolName, Function action) { + private String invoke(String toolName, Supplier> action) { try { - ToolContext toolContext = buildToolContext(); - return action.apply(toolContext); + return toolSuccess(action.get()); + } catch (AiToolException e) { + return toolFailure(toolName, e); } catch (Exception e) { log.error("MCP tool call failed, tool={}", toolName, e); return toolFailure(toolName, e); @@ -124,14 +154,54 @@ private String invoke(String toolName, Function action) { } private String toolFailure(String toolName, Exception e) { - String message = "MCP tool call failed: " + StringUtils.defaultIfBlank(e.getMessage(), "Unknown error"); - return JSON.toJSONString(AiToolResult.failure(message, "MCP_TOOL_CALL_FAILED"), - JSONWriter.Feature.WriteNulls); + if (e instanceof AiToolException toolException) { + return aiToolResultSerializer.toJson(AiToolResult.failureWithCode( + aiToolErrorCodeMapper.errorCodeFor(toolException), + toolException.getMessage())); + } + return aiToolResultSerializer.toJson(AiToolResult.failureWithCode( + "TOOL_CALL_FAILED", + "Tool execution failed.")); + } + + private String toolSuccess(AiToolOutput output) { + return aiToolResultSerializer.toJson(AiToolResult.success(output.summary(), output.data())); + } + + private AiToolContextRequest contextRequest() { + return aiToolContextConverter.toParam(buildToolContext()); + } + + private AiListTablesRequest listTablesRequest(Long dataSourceId, String databaseName, String schemaName) { + AiListTablesRequest request = new AiListTablesRequest(); + request.setDataSourceId(dataSourceId); + request.setDatabaseName(databaseName); + request.setSchemaName(schemaName); + request.setAiToolContextRequest(contextRequest()); + return request; + } + + private AiExecuteSqlRequest executeSqlRequest(String sql, Integer pageSize, Long dataSourceId, String databaseName, + String schemaName) { + AiExecuteSqlRequest request = new AiExecuteSqlRequest(); + request.setSql(sql); + request.setPageSize(pageSize); + request.setDataSourceId(dataSourceId); + request.setDatabaseName(databaseName); + request.setSchemaName(schemaName); + request.setAiToolContextRequest(contextRequest()); + return request; } - static String toolSuccess(String summary, List data) { - return JSON.toJSONString(AiToolResult.success(summary, data), - JSONWriter.Feature.WriteNulls); + private AiGetTablesSchemaRequest tablesSchemaRequest(List tableNames, Long dataSourceId, + String databaseName, String schemaName) { + AiGetTablesSchemaRequest request = new AiGetTablesSchemaRequest(); + request.setTableNames(tableNames); + request.setDataSourceId(dataSourceId); + request.setDatabaseName(databaseName); + request.setSchemaName(schemaName); + request.setAiToolContextRequest(contextRequest()); + return request; } private ToolContext buildToolContext() { diff --git a/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapterTest.java b/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapterTest.java index 505a83f84..820b7397a 100644 --- a/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapterTest.java +++ b/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/adapter/ai/AiToolAdapterTest.java @@ -1,27 +1,173 @@ package ai.chat2db.community.web.api.adapter.ai; +import ai.chat2db.community.domain.api.exception.ai.AiToolException; +import ai.chat2db.community.domain.api.exception.ai.AiToolInvalidArgumentException; +import ai.chat2db.community.domain.api.exception.ai.AiToolMetadataQueryException; +import ai.chat2db.community.domain.api.exception.ai.AiToolSqlConfirmationRequiredException; +import ai.chat2db.community.domain.api.exception.ai.AiToolSqlExecutionException; +import ai.chat2db.community.domain.api.model.ai.TableSchemaResult; +import ai.chat2db.community.domain.api.model.metadata.Database; +import ai.chat2db.community.domain.api.model.metadata.Schema; +import ai.chat2db.community.domain.api.model.metadata.SimpleTable; +import ai.chat2db.community.domain.api.model.request.ai.AiExecuteSqlRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiGetTablesSchemaRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiListTablesRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiToolContextRequest; +import ai.chat2db.community.domain.api.model.result.ExecuteResponse; +import ai.chat2db.community.domain.api.model.storage.WorkspaceDataSource; +import ai.chat2db.community.domain.api.service.ai.IAiToolService; +import ai.chat2db.community.web.api.converter.ai.AiToolContextConverter; +import ai.chat2db.community.web.api.converter.ai.AiToolErrorCodeMapper; +import ai.chat2db.community.web.api.converter.ai.AiToolResultConverter; +import ai.chat2db.community.web.api.converter.ai.AiToolResultSerializer; +import com.alibaba.fastjson2.JSON; +import com.alibaba.fastjson2.JSONObject; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.MethodSource; +import org.springframework.ai.chat.model.ToolContext; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.function.Consumer; +import java.util.stream.Stream; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotEquals; class AiToolAdapterTest { + @ParameterizedTest + @MethodSource("typedToolExceptions") + void shouldMapTypedToolExceptionToStableErrorEnvelope(AiToolException exception, String expectedErrorCode) { + AiToolAdapter adapter = adapterWith(new StubAiToolService() { + @Override + public List executeSql(AiExecuteSqlRequest request) { + throw exception; + } + }); + + String json = adapter.executeSql("select", null, 1L, "db", null, null); + JSONObject result = JSON.parseObject(json); + + assertEquals(false, result.getBoolean("success")); + assertEquals(expectedErrorCode, result.getString("errorCode")); + assertEquals(exception.getMessage(), result.getString("summary")); + } + @Test - void shouldUseSummaryForStructuredToolTraceContent() { - String content = """ - { - "success": true, - "summary": "SQL executed successfully.", - "data": [{"rows": [["very large raw payload"]]}], - "errorCode": null - } - """; - - assertEquals("SQL executed successfully.", AiToolAdapter.traceSummary(content)); + void shouldMapPlainIllegalStateExceptionFromServiceToToolCallFailedEnvelope() { + AiToolAdapter adapter = adapterWith(new StubAiToolService() { + @Override + public List executeSql(AiExecuteSqlRequest request) { + throw new IllegalStateException("driver exploded"); + } + }); + + String json = adapter.executeSql("select", null, 1L, "db", null, null); + JSONObject result = JSON.parseObject(json); + + assertEquals(false, result.getBoolean("success")); + assertEquals("TOOL_CALL_FAILED", result.getString("errorCode")); + assertEquals("Tool execution failed.", result.getString("summary")); + assertFalse(result.getString("summary").contains("driver exploded")); } @Test - void shouldKeepRawTraceContentWhenToolResultIsNotJson() { - assertEquals("plain result", AiToolAdapter.traceSummary("plain result")); + void shouldSerializeConvertedStrongResultWithoutServiceJsonRoundTrip() { + AiToolAdapter adapter = adapterWith(new StubAiToolService() { + @Override + public List executeSql(AiExecuteSqlRequest request) { + return Collections.emptyList(); + } + }); + + String json = adapter.executeSql("select 1", null, 1L, "db", null, null); + JSONObject result = JSON.parseObject(json); + + assertEquals(true, result.getBoolean("success")); + assertEquals("SQL executed successfully with no result.", result.getString("summary")); + assertEquals(0, result.getJSONObject("data").getJSONArray("results").size()); + } + + @Test + void shouldEmitTraceContentAsSummaryNotRawJson() { + AiToolAdapter adapter = adapterWith(new StubAiToolService() { + @Override + public List executeSql(AiExecuteSqlRequest request) { + return Collections.emptyList(); + } + }); + List> tracePayloads = new ArrayList<>(); + Map context = new LinkedHashMap<>(); + context.put(AiChatTraceSupport.TRACE_EMITTER_KEY, (Consumer>) tracePayloads::add); + + String json = adapter.executeSql("select 1", null, 1L, "db", null, new ToolContext(context)); + + assertEquals(1, tracePayloads.size()); + Map payload = tracePayloads.get(0); + assertEquals(AiChatTraceSupport.TYPE_TOOL_RESULT, payload.get("type")); + assertEquals("execute_sql", payload.get("name")); + assertEquals("SQL executed successfully with no result.", payload.get("content")); + assertNotEquals(json, payload.get("content")); + assertFalse(String.valueOf(payload.get("content")).contains("\"success\"")); + assertFalse(String.valueOf(payload.get("content")).contains("\"data\"")); + } + + private AiToolAdapter adapterWith(IAiToolService service) { + return new AiToolAdapter( + service, + new AiToolContextConverter(), + new AiToolResultConverter(), + new AiToolResultSerializer(), + new AiToolErrorCodeMapper()); + } + + private static Stream typedToolExceptions() { + return Stream.of( + new Object[] {new AiToolInvalidArgumentException("bad input"), "INVALID_ARGUMENT"}, + new Object[] {new AiToolSqlConfirmationRequiredException("manual confirmation required"), + "SQL_REQUIRES_MANUAL_CONFIRMATION"}, + new Object[] {new AiToolSqlExecutionException("SQL execution failed: syntax error"), + "SQL_EXECUTION_FAILED"}, + new Object[] {new AiToolMetadataQueryException("metadata query failed"), + "METADATA_QUERY_FAILED"}); + } + + private static class StubAiToolService implements IAiToolService { + + @Override + public List listAllDataSources(AiToolContextRequest request) { + return Collections.emptyList(); + } + + @Override + public List listAllTables(AiListTablesRequest request) { + return Collections.emptyList(); + } + + @Override + public List listAllDatabases(Long dataSourceId, AiToolContextRequest request) { + return Collections.emptyList(); + } + + @Override + public List listAllSchemas(String databaseName, Long dataSourceId, AiToolContextRequest request) { + return Collections.emptyList(); + } + + @Override + public List executeSql(AiExecuteSqlRequest request) { + return Collections.emptyList(); + } + + @Override + public List getTablesSchema(AiGetTablesSchemaRequest request) { + return Collections.emptyList(); + } } } diff --git a/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/converter/ai/AiToolResultConverterTest.java b/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/converter/ai/AiToolResultConverterTest.java new file mode 100644 index 000000000..cc65df505 --- /dev/null +++ b/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/converter/ai/AiToolResultConverterTest.java @@ -0,0 +1,309 @@ +package ai.chat2db.community.web.api.converter.ai; + +import ai.chat2db.community.domain.api.model.ai.AiToolResult; +import ai.chat2db.community.domain.api.model.ai.SqlToolData; +import ai.chat2db.community.domain.api.model.result.ExecuteResponse; +import ai.chat2db.community.domain.api.model.result.Header; +import ai.chat2db.community.domain.api.model.result.ResultCell; +import com.alibaba.fastjson2.JSON; +import com.alibaba.fastjson2.JSONArray; +import com.alibaba.fastjson2.JSONObject; +import org.junit.jupiter.api.Test; + +import javax.sql.rowset.serial.SerialBlob; +import javax.sql.rowset.serial.SerialClob; +import java.io.ByteArrayInputStream; +import java.io.StringReader; +import java.math.BigDecimal; +import java.sql.Blob; +import java.sql.Clob; +import java.sql.Timestamp; +import java.time.Instant; +import java.time.LocalDate; +import java.util.ArrayList; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +class AiToolResultConverterTest { + + @Test + void shouldProjectSqlResultToSummaryAndTypedPayloadWithoutEnvelope() { + ExecuteResponse response = ExecuteResponse.builder() + .success(true) + .build(); + + AiToolOutput output = new AiToolResultConverter().fromExecuteResult(List.of(response)); + String json = new AiToolResultSerializer().toJson(AiToolResult.success(output.summary(), output.data())); + JSONObject payload = JSON.parseObject(json); + + assertEquals(true, payload.getBoolean("success")); + assertNull(payload.get("tool")); + assertEquals("SQL executed successfully with 1 result set(s).", payload.getString("summary")); + assertEquals(1, payload.getJSONObject("data").getJSONArray("results").size()); + assertNull(payload.get("errorCode")); + assertTrue(json.contains("\"errorCode\":null")); + } + + @Test + void shouldProjectText2SqlResultThroughConverter() { + AiToolOutput output = new AiToolResultConverter().fromText2Sql("select * from users"); + String json = new AiToolResultSerializer().toJson(AiToolResult.success(output.summary(), output.data())); + JSONObject payload = JSON.parseObject(json); + + assertEquals(true, payload.getBoolean("success")); + assertEquals("SQL generated successfully.", payload.getString("summary")); + assertEquals("select * from users", payload.getJSONObject("data").getString("sql")); + assertNull(payload.getString("errorCode")); + } + + @Test + void shouldSerializeFailedToolResultAsStandardJson() { + String json = new AiToolResultSerializer().toJson(AiToolResult.failureWithCode( + "INVALID_ARGUMENT", + "sql is empty.")); + + JSONObject result = JSON.parseObject(json); + + assertEquals(false, result.getBoolean("success")); + assertNull(result.get("tool")); + assertEquals("sql is empty.", result.getString("summary")); + assertNull(result.get("data")); + assertEquals("INVALID_ARGUMENT", result.getString("errorCode")); + } + + @Test + void shouldKeepRowsPositionBasedAndPreferRawValuesWhenColumnNamesAreDuplicated() { + List
headers = List.of( + Header.builder().name("id").build(), + Header.builder().name("id").build(), + Header.builder().name("note").build()); + String longText = "a".repeat(201); + BigDecimal rawId = new BigDecimal("123.45"); + ResultCell displayFallback = ResultCell.builder() + .value(longText) + .rawValue(null) + .sizeChars(201L) + .loadedChars(201L) + .truncated(true) + .build(); + + List> rows = AiToolResultConverter.rowPreviewRows( + headers, + List.of(List.of( + ResultCell.builder().value("123.45").rawValue(rawId).build(), + ResultCell.of(null), + displayFallback))); + List> metadata = AiToolResultConverter.rowPreviewCellMetadata( + headers, + List.of(List.of( + ResultCell.builder().value("123.45").rawValue(rawId).build(), + ResultCell.of(null), + displayFallback))); + + assertEquals(List.of("id", "id", "note"), AiToolResultConverter.columnNames(headers)); + assertEquals(1, rows.size()); + assertEquals(rawId, rows.get(0).get(0)); + assertNull(rows.get(0).get(1)); + assertEquals(longText, rows.get(0).get(2)); + assertNull(metadata.get(0).get(0)); + assertNull(metadata.get(0).get(1)); + assertEquals(false, metadata.get(0).get(2).getRawValueAvailable()); + assertEquals(true, metadata.get(0).get(2).getTruncated()); + assertEquals(201L, metadata.get(0).get(2).getSizeChars()); + } + + @Test + void shouldTreatStringOnlyResultCellAsRawPreserving() { + ExecuteResponse response = ExecuteResponse.builder() + .success(true) + .headerList(List.of(Header.builder().name("key").build())) + .dataList(List.of(List.of(ResultCell.of("redis-value")))) + .build(); + + AiToolOutput output = new AiToolResultConverter().fromExecuteResult(List.of(response)); + SqlToolData.ResultSet resultSet = output.data().getResults().get(0); + + assertEquals("redis-value", resultSet.getRows().get(0).get(0)); + assertNull(resultSet.getRowCellMetadata().get(0).get(0)); + } + + @Test + void shouldFallbackToDisplayPreviewForTransportUnsafeRawValues() throws Exception { + List
headers = List.of( + Header.builder().name("clob_col").build(), + Header.builder().name("blob_col").build(), + Header.builder().name("bytes_col").build(), + Header.builder().name("reader_col").build(), + Header.builder().name("stream_col").build(), + Header.builder().name("large_col").build(), + Header.builder().name("truncated_col").build(), + Header.builder().name("driver_col").build()); + Clob clob = new SerialClob("large text".toCharArray()); + Blob blob = new SerialBlob(new byte[] {1, 2, 3}); + DriverSpecificValue driverValue = new DriverSpecificValue("POINT(1 2)"); + List row = List.of( + ResultCell.builder().value("clob preview").rawValue(clob).valueType("TEXT").sizeChars(10L).build(), + ResultCell.builder().value("blob preview").rawValue(blob).valueType("BINARY").sizeBytes(3L).build(), + ResultCell.builder().value("bytes preview").rawValue(new byte[] {1, 2, 3}).valueType("BINARY").sizeBytes(3L).build(), + ResultCell.builder().value("reader preview").rawValue(new StringReader("abc")).valueType("TEXT").build(), + ResultCell.builder().value("stream preview").rawValue(new ByteArrayInputStream(new byte[] {1})).valueType("BINARY").build(), + ResultCell.builder().value("large preview").rawValue("raw but large").largeValue(true).valueType("TEXT").sizeChars(20L).build(), + ResultCell.builder().value("truncated preview").rawValue("raw but truncated").truncated(true).valueType("TEXT").loadedChars(9L).sizeChars(20L).build(), + ResultCell.builder().value("driver preview").rawValue(driverValue).valueType("GEOMETRY").build()); + + List> rows = AiToolResultConverter.rowPreviewRows(headers, List.of(row)); + List> metadata = AiToolResultConverter.rowPreviewCellMetadata(headers, List.of(row)); + + assertEquals(List.of( + "clob preview", + "blob preview", + "bytes preview", + "reader preview", + "stream preview", + "large preview", + "truncated preview", + "driver preview"), rows.get(0)); + assertEquals(false, metadata.get(0).get(0).getRawValueAvailable()); + assertEquals("UNSAFE_RAW_VALUE:javax.sql.rowset.serial.SerialClob", + metadata.get(0).get(0).getRawValueUnavailableReason()); + assertEquals("UNSAFE_RAW_VALUE:javax.sql.rowset.serial.SerialBlob", + metadata.get(0).get(1).getRawValueUnavailableReason()); + assertEquals("UNSAFE_RAW_VALUE:[B", metadata.get(0).get(2).getRawValueUnavailableReason()); + assertEquals("LARGE_VALUE", metadata.get(0).get(5).getRawValueUnavailableReason()); + assertEquals("TRUNCATED_VALUE", metadata.get(0).get(6).getRawValueUnavailableReason()); + assertEquals(true, metadata.get(0).get(6).getTruncated()); + assertEquals(20L, metadata.get(0).get(6).getSizeChars()); + assertEquals(9L, metadata.get(0).get(6).getLoadedChars()); + assertEquals(false, metadata.get(0).get(7).getRawValueAvailable()); + assertEquals("UNSAFE_RAW_VALUE:" + DriverSpecificValue.class.getName(), + metadata.get(0).get(7).getRawValueUnavailableReason()); + } + + @Test + void shouldSerializeLargeJdbcRawValuesAsPreviewAndMetadata() throws Exception { + String clobPreview = "[CLOB] 20.00 MB preview"; + String blobPreview = "[BLOB] 5.00 MB preview"; + Clob clob = new SerialClob("large text".toCharArray()); + Blob blob = new SerialBlob(new byte[] {1, 2, 3}); + ExecuteResponse response = ExecuteResponse.builder() + .success(true) + .headerList(List.of( + Header.builder().name("content").build(), + Header.builder().name("payload").build())) + .dataList(List.of(List.of( + ResultCell.builder() + .value(clobPreview) + .rawValue(clob) + .largeValue(true) + .truncated(true) + .valueType("TEXT") + .sizeChars(20L * 1024L * 1024L) + .loadedChars(2048L) + .build(), + ResultCell.builder() + .value(blobPreview) + .rawValue(blob) + .largeValue(true) + .truncated(true) + .valueType("BINARY") + .sizeBytes(5L * 1024L * 1024L) + .loadedBytes(1024L) + .build()))) + .build(); + + AiToolOutput output = new AiToolResultConverter().fromExecuteResult(List.of(response)); + String json = new AiToolResultSerializer().toJson(AiToolResult.success(output.summary(), output.data())); + JSONObject payload = JSON.parseObject(json); + JSONObject resultSet = payload.getJSONObject("data").getJSONArray("results").getJSONObject(0); + JSONArray rows = resultSet.getJSONArray("rows"); + JSONArray metadata = resultSet.getJSONArray("rowCellMetadata"); + + assertFalse(json.contains("\"rawValue\"")); + assertEquals(clobPreview, rows.getJSONArray(0).getString(0)); + assertEquals(blobPreview, rows.getJSONArray(0).getString(1)); + assertEquals(false, metadata.getJSONArray(0).getJSONObject(0).getBoolean("rawValueAvailable")); + assertEquals(true, metadata.getJSONArray(0).getJSONObject(0).getBoolean("largeValue")); + assertEquals(true, metadata.getJSONArray(0).getJSONObject(0).getBoolean("truncated")); + assertEquals("TEXT", metadata.getJSONArray(0).getJSONObject(0).getString("valueType")); + assertEquals(20L * 1024L * 1024L, metadata.getJSONArray(0).getJSONObject(0).getLong("sizeChars")); + assertEquals(2048L, metadata.getJSONArray(0).getJSONObject(0).getLong("loadedChars")); + assertEquals("LARGE_VALUE", metadata.getJSONArray(0).getJSONObject(0).getString("rawValueUnavailableReason")); + assertEquals(false, metadata.getJSONArray(0).getJSONObject(1).getBoolean("rawValueAvailable")); + assertEquals("BINARY", metadata.getJSONArray(0).getJSONObject(1).getString("valueType")); + assertEquals(5L * 1024L * 1024L, metadata.getJSONArray(0).getJSONObject(1).getLong("sizeBytes")); + assertEquals(1024L, metadata.getJSONArray(0).getJSONObject(1).getLong("loadedBytes")); + assertEquals("LARGE_VALUE", metadata.getJSONArray(0).getJSONObject(1).getString("rawValueUnavailableReason")); + } + + @Test + void shouldOnlyExposeWhitelistedJsonSafeRawValues() { + Instant instant = Instant.parse("2026-07-21T10:15:30Z"); + Map jsonObject = new LinkedHashMap<>(); + jsonObject.put("date", LocalDate.of(2026, 7, 21)); + jsonObject.put("values", List.of("ok", 12, true)); + DriverSpecificValue driverValue = new DriverSpecificValue("POINT(1 2)"); + ExecuteResponse response = ExecuteResponse.builder() + .success(true) + .headerList(List.of( + Header.builder().name("amount").build(), + Header.builder().name("created_at").build(), + Header.builder().name("json_doc").build(), + Header.builder().name("geometry").build())) + .dataList(List.of(List.of( + ResultCell.builder().value("123.45").rawValue(new BigDecimal("123.45")).build(), + ResultCell.builder().value("2026-07-21 10:15:30").rawValue(Timestamp.from(instant)).build(), + ResultCell.builder().value("{...}").rawValue(jsonObject).build(), + ResultCell.builder().value("POINT(1 2)").rawValue(driverValue).build()))) + .build(); + + AiToolOutput output = new AiToolResultConverter().fromExecuteResult(List.of(response)); + String json = new AiToolResultSerializer().toJson(AiToolResult.success(output.summary(), output.data())); + JSONObject resultSet = JSON.parseObject(json).getJSONObject("data").getJSONArray("results").getJSONObject(0); + JSONArray row = resultSet.getJSONArray("rows").getJSONArray(0); + JSONArray metadata = resultSet.getJSONArray("rowCellMetadata").getJSONArray(0); + + assertEquals(new BigDecimal("123.45"), row.getBigDecimal(0)); + assertEquals("2026-07-21T10:15:30Z", row.getString(1)); + assertEquals("2026-07-21", row.getJSONObject(2).getString("date")); + assertEquals("ok", row.getJSONObject(2).getJSONArray("values").getString(0)); + assertEquals(12, row.getJSONObject(2).getJSONArray("values").getInteger(1)); + assertEquals(true, row.getJSONObject(2).getJSONArray("values").getBoolean(2)); + assertEquals("POINT(1 2)", row.getString(3)); + assertNull(metadata.get(0)); + assertNull(metadata.get(1)); + assertNull(metadata.get(2)); + assertEquals(false, metadata.getJSONObject(3).getBoolean("rawValueAvailable")); + assertEquals("UNSAFE_RAW_VALUE:" + DriverSpecificValue.class.getName(), + metadata.getJSONObject(3).getString("rawValueUnavailableReason")); + } + + @Test + void shouldExposeRowPreviewTruncationMetadata() { + Header header = Header.builder().name("id").build(); + List> rows = new ArrayList<>(); + for (int i = 0; i < 51; i++) { + rows.add(List.of(ResultCell.of(String.valueOf(i)))); + } + ExecuteResponse response = ExecuteResponse.builder() + .success(true) + .headerList(List.of(header)) + .dataList(rows) + .build(); + + SqlToolData.ResultSet result = AiToolResultConverter.executeResponseData(1, response); + + assertEquals(51, result.getRowCount()); + assertEquals(50, result.getPreviewRowCount()); + assertEquals(true, result.getRowsTruncated()); + assertEquals(50, result.getRows().size()); + } + + private record DriverSpecificValue(String value) { + } +} diff --git a/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapterTest.java b/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapterTest.java index d79afdb7f..12d16ab6a 100644 --- a/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapterTest.java +++ b/chat2db-community-server/chat2db-community-web/src/test/java/ai/chat2db/community/web/api/mcp/adapter/AiToolMcpAdapterTest.java @@ -1,13 +1,34 @@ package ai.chat2db.community.web.api.mcp.adapter; +import ai.chat2db.community.domain.api.exception.ai.AiToolInvalidArgumentException; +import ai.chat2db.community.domain.api.exception.ai.AiToolSqlExecutionException; +import ai.chat2db.community.domain.api.model.ai.TableSchemaResult; +import ai.chat2db.community.domain.api.model.metadata.Database; +import ai.chat2db.community.domain.api.model.metadata.Schema; +import ai.chat2db.community.domain.api.model.metadata.SimpleTable; +import ai.chat2db.community.domain.api.model.request.ai.AiExecuteSqlRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiGetTablesSchemaRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiListTablesRequest; +import ai.chat2db.community.domain.api.model.request.ai.AiToolContextRequest; +import ai.chat2db.community.domain.api.model.result.ExecuteResponse; +import ai.chat2db.community.domain.api.model.storage.WorkspaceDataSource; +import ai.chat2db.community.domain.api.service.ai.IAiToolService; +import ai.chat2db.community.web.api.adapter.ai.AiChatStreamAdapter; +import ai.chat2db.community.web.api.adapter.ai.AiToolAdapter; +import ai.chat2db.community.web.api.converter.ai.AiToolContextConverter; +import ai.chat2db.community.web.api.converter.ai.AiToolErrorCodeMapper; +import ai.chat2db.community.web.api.converter.ai.AiToolResultConverter; +import ai.chat2db.community.web.api.converter.ai.AiToolResultSerializer; +import ai.chat2db.community.web.api.model.request.ai.ChatRequest; import com.alibaba.fastjson2.JSON; import com.alibaba.fastjson2.JSONObject; import org.junit.jupiter.api.Test; +import java.util.Collections; import java.util.List; -import java.util.Map; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertTrue; @@ -15,16 +36,140 @@ class AiToolMcpAdapterTest { @Test void shouldSerializeText2SqlSuccessAsStandardJson() { - String json = AiToolMcpAdapter.toolSuccess( - "SQL generated successfully.", - List.of(Map.of("sql", "select * from users"))); + String json = mcpAdapter(new StubAiToolService(), new StubAiChatStreamAdapter("select * from users")) + .text2sql("list users", 1L, "app", null); JSONObject result = JSON.parseObject(json); assertEquals(true, result.getBoolean("success")); assertEquals("SQL generated successfully.", result.getString("summary")); - assertEquals("select * from users", result.getJSONArray("data").getJSONObject(0).getString("sql")); + assertEquals("select * from users", result.getJSONObject("data").getString("sql")); assertNull(result.getString("errorCode")); assertTrue(json.contains("\"errorCode\":null")); } + + @Test + void shouldKeepDatabaseToolProtocolSameAsRegularAdapter() { + IAiToolService service = new StubAiToolService() { + @Override + public List executeSql(AiExecuteSqlRequest request) { + throw new AiToolSqlExecutionException("SQL execution failed: syntax error"); + } + }; + AiToolAdapter regularAdapter = adapter(service); + AiToolMcpAdapter mcpAdapter = mcpAdapter(service); + + String regular = regularAdapter.executeSql("select broken", null, 1L, "app", null, null); + String mcp = mcpAdapter.executeSql("select broken", null, 1L, "app", null); + + assertEquals(regular, mcp); + } + + @Test + void shouldMapTypedServiceExceptionToStableErrorEnvelope() { + IAiToolService service = new StubAiToolService() { + @Override + public List executeSql(AiExecuteSqlRequest request) { + throw new AiToolInvalidArgumentException("bad mcp input"); + } + }; + AiToolMcpAdapter mcpAdapter = mcpAdapter(service); + + String json = mcpAdapter.executeSql("select 1", null, 1L, "app", null); + JSONObject result = JSON.parseObject(json); + + assertEquals(false, result.getBoolean("success")); + assertEquals("bad mcp input", result.getString("summary")); + assertEquals("INVALID_ARGUMENT", result.getString("errorCode")); + } + + @Test + void shouldMapPlainIllegalStateExceptionFromServiceToToolCallFailedEnvelope() { + IAiToolService service = new StubAiToolService() { + @Override + public List executeSql(AiExecuteSqlRequest request) { + throw new IllegalStateException("driver exploded"); + } + }; + AiToolMcpAdapter mcpAdapter = mcpAdapter(service); + + String json = mcpAdapter.executeSql("select 1", null, 1L, "app", null); + JSONObject result = JSON.parseObject(json); + + assertEquals(false, result.getBoolean("success")); + assertEquals("Tool execution failed.", result.getString("summary")); + assertEquals("TOOL_CALL_FAILED", result.getString("errorCode")); + assertFalse(result.getString("summary").contains("driver exploded")); + } + + private static AiToolAdapter adapter(IAiToolService service) { + return new AiToolAdapter( + service, + new AiToolContextConverter(), + new AiToolResultConverter(), + new AiToolResultSerializer(), + new AiToolErrorCodeMapper()); + } + + private static AiToolMcpAdapter mcpAdapter(IAiToolService service) { + return mcpAdapter(service, null); + } + + private static AiToolMcpAdapter mcpAdapter(IAiToolService service, AiChatStreamAdapter aiChatStreamAdapter) { + return new AiToolMcpAdapter( + service, + new AiToolContextConverter(), + new AiToolResultConverter(), + aiChatStreamAdapter, + new AiToolResultSerializer(), + new AiToolErrorCodeMapper()); + } + + private static class StubAiChatStreamAdapter extends AiChatStreamAdapter { + + private final String sql; + + private StubAiChatStreamAdapter(String sql) { + super(null, null, adapter(new StubAiToolService()), null, null, null, null, null, null); + this.sql = sql; + } + + @Override + public String chatSync(ChatRequest request) { + return sql; + } + } + + private static class StubAiToolService implements IAiToolService { + + @Override + public List listAllDataSources(AiToolContextRequest request) { + return Collections.emptyList(); + } + + @Override + public List listAllTables(AiListTablesRequest request) { + return Collections.emptyList(); + } + + @Override + public List listAllDatabases(Long dataSourceId, AiToolContextRequest request) { + return Collections.emptyList(); + } + + @Override + public List listAllSchemas(String databaseName, Long dataSourceId, AiToolContextRequest request) { + return Collections.emptyList(); + } + + @Override + public List executeSql(AiExecuteSqlRequest request) { + return Collections.emptyList(); + } + + @Override + public List getTablesSchema(AiGetTablesSchemaRequest request) { + return Collections.emptyList(); + } + } }