|
@@ -0,0 +1,293 @@
|
|
|
|
|
+package com.zsjz.ai.module.agent.tools;
|
|
|
|
|
+
|
|
|
|
|
+import com.baomidou.mybatisplus.core.MybatisConfiguration;
|
|
|
|
|
+import com.baomidou.mybatisplus.core.metadata.TableInfoHelper;
|
|
|
|
|
+import com.baomidou.mybatisplus.extension.plugins.pagination.Page;
|
|
|
|
|
+import com.fasterxml.jackson.databind.JsonNode;
|
|
|
|
|
+import com.fasterxml.jackson.databind.ObjectMapper;
|
|
|
|
|
+import com.zsjz.ai.common.base.BasicColumn;
|
|
|
|
|
+import com.zsjz.ai.common.constants.PathConst;
|
|
|
|
|
+import com.zsjz.ai.common.context.CaseContextHolder;
|
|
|
|
|
+import com.zsjz.ai.common.model.govern.query.GovernTreeQuery;
|
|
|
|
|
+import com.zsjz.ai.common.model.trans.entity.TransRecord;
|
|
|
|
|
+import com.zsjz.ai.module.agent.artifact.ArtifactService;
|
|
|
|
|
+import com.zsjz.ai.module.agent.config.ArtifactProperties;
|
|
|
|
|
+import com.zsjz.ai.module.govern.serivce.GovernTreeService;
|
|
|
|
|
+import io.agentscope.core.message.TextBlock;
|
|
|
|
|
+import io.agentscope.core.message.ToolResultBlock;
|
|
|
|
|
+import io.agentscope.core.message.ToolResultState;
|
|
|
|
|
+import org.apache.ibatis.builder.MapperBuilderAssistant;
|
|
|
|
|
+import org.junit.jupiter.api.AfterEach;
|
|
|
|
|
+import org.junit.jupiter.api.BeforeAll;
|
|
|
|
|
+import org.junit.jupiter.api.DisplayName;
|
|
|
|
|
+import org.junit.jupiter.api.Test;
|
|
|
|
|
+
|
|
|
|
|
+import java.nio.file.Files;
|
|
|
|
|
+import java.nio.file.Path;
|
|
|
|
|
+import java.util.ArrayList;
|
|
|
|
|
+import java.util.LinkedHashMap;
|
|
|
|
|
+import java.util.List;
|
|
|
|
|
+import java.util.Map;
|
|
|
|
|
+import java.util.concurrent.atomic.AtomicReference;
|
|
|
|
|
+
|
|
|
|
|
+import static org.junit.jupiter.api.Assertions.assertEquals;
|
|
|
|
|
+import static org.junit.jupiter.api.Assertions.assertFalse;
|
|
|
|
|
+import static org.junit.jupiter.api.Assertions.assertNotNull;
|
|
|
|
|
+import static org.junit.jupiter.api.Assertions.assertNull;
|
|
|
|
|
+import static org.junit.jupiter.api.Assertions.assertTrue;
|
|
|
|
|
+import static org.mockito.ArgumentMatchers.any;
|
|
|
|
|
+import static org.mockito.Mockito.mock;
|
|
|
|
|
+import static org.mockito.Mockito.when;
|
|
|
|
|
+
|
|
|
|
|
+/**
|
|
|
|
|
+ * {@code query_table_data} 的契约:表清单发现 / 条件白名单校验 / dto 组装 / 结果形态。
|
|
|
|
|
+ *
|
|
|
|
|
+ * <p>治理表查询走的是 {@link GovernTreeService#getTreeTablePage}——工具侧职责只有
|
|
|
|
|
+ * 「组 dto + 校验」,查询本身用 mock 打桩;{@code isRegistered} 用真实的 MyBatis-Plus
|
|
|
|
|
+ * 注册表({@link BeforeAll} 里注册 {@link TransRecord}),不走 spy。
|
|
|
|
|
+ */
|
|
|
|
|
+class GovernTableToolTest {
|
|
|
|
|
+
|
|
|
|
|
+ private static final Long CASE_ID = 952709L;
|
|
|
|
|
+
|
|
|
|
|
+ @BeforeAll
|
|
|
|
|
+ static void initTableInfoCache() {
|
|
|
|
|
+ MybatisConfiguration configuration = new MybatisConfiguration();
|
|
|
|
|
+ MapperBuilderAssistant assistant = new MapperBuilderAssistant(configuration, "");
|
|
|
|
|
+ assistant.setCurrentNamespace("com.zsjz.ai.module.agent.tools.GovernTableToolTest");
|
|
|
|
|
+ TableInfoHelper.initTableInfo(assistant, TransRecord.class);
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ private final GovernTreeService governTreeService = mock(GovernTreeService.class);
|
|
|
|
|
+ private final ArtifactProperties artifactProps = new ArtifactProperties();
|
|
|
|
|
+ private final GovernTableTool tool =
|
|
|
|
|
+ new GovernTableTool(governTreeService, new ArtifactService(artifactProps));
|
|
|
|
|
+
|
|
|
|
|
+ /** 本次测试产生的产物文件(用例结束清理,避免污染 workspace) */
|
|
|
|
|
+ private final List<Path> created = new ArrayList<>();
|
|
|
|
|
+
|
|
|
|
|
+ @AfterEach
|
|
|
|
|
+ void cleanup() {
|
|
|
|
|
+ for (Path file : created) {
|
|
|
|
|
+ try {
|
|
|
|
|
+ Files.deleteIfExists(file);
|
|
|
|
|
+ } catch (Exception ignored) {
|
|
|
|
|
+ // 尽力而为
|
|
|
|
|
+ }
|
|
|
|
|
+ }
|
|
|
|
|
+ created.clear();
|
|
|
|
|
+ // 空目录一并收掉(非空说明是别的用例/真实数据,保留)
|
|
|
|
|
+ try {
|
|
|
|
|
+ Files.deleteIfExists(PathConst.WORKSPACE.resolve(String.valueOf(CASE_ID)).resolve("artifacts"));
|
|
|
|
|
+ Files.deleteIfExists(PathConst.WORKSPACE.resolve(String.valueOf(CASE_ID)));
|
|
|
|
|
+ } catch (Exception ignored) {
|
|
|
|
|
+ // 尽力而为
|
|
|
|
|
+ }
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ private static String text(ToolResultBlock block) {
|
|
|
|
|
+ StringBuilder sb = new StringBuilder();
|
|
|
|
|
+ for (var content : block.getOutput()) {
|
|
|
|
|
+ if (content instanceof TextBlock textBlock) {
|
|
|
|
|
+ sb.append(textBlock.getText());
|
|
|
|
|
+ }
|
|
|
|
|
+ }
|
|
|
|
|
+ return sb.toString();
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ private static JsonNode json(ToolResultBlock block) {
|
|
|
|
|
+ try {
|
|
|
|
|
+ return new ObjectMapper().readTree(text(block));
|
|
|
|
|
+ } catch (Exception e) {
|
|
|
|
|
+ throw new IllegalArgumentException("结果不是 JSON: " + text(block), e);
|
|
|
|
|
+ }
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ /** 让服务端返回一份固定分页结果,并捕获工具组装的 dto */
|
|
|
|
|
+ private AtomicReference<GovernTreeQuery> stubService(List<Object> records, List<BasicColumn> head) {
|
|
|
|
|
+ AtomicReference<GovernTreeQuery> captured = new AtomicReference<>();
|
|
|
|
|
+ when(governTreeService.getTreeTablePage(any(GovernTreeQuery.class))).thenAnswer(inv -> {
|
|
|
|
|
+ captured.set(inv.getArgument(0));
|
|
|
|
|
+ Page<Object> page = new Page<>(1, 100, records.size());
|
|
|
|
|
+ page.setRecords(records);
|
|
|
|
|
+ Map<String, Object> out = new LinkedHashMap<>();
|
|
|
|
|
+ out.put("pages", page);
|
|
|
|
|
+ out.put("head", head);
|
|
|
|
|
+ out.put("detail", false);
|
|
|
|
|
+ return out;
|
|
|
|
|
+ });
|
|
|
|
|
+ return captured;
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ // ==================== 校验 ====================
|
|
|
|
|
+
|
|
|
|
|
+ @Test
|
|
|
|
|
+ @DisplayName("缺 tableName 直接报错(没有表清单发现模式)")
|
|
|
|
|
+ void missingTableNameRejected() {
|
|
|
|
|
+ ToolResultBlock nullQuery = tool.queryTableData(null);
|
|
|
|
|
+ assertEquals(ToolResultState.ERROR, nullQuery.getState(), text(nullQuery));
|
|
|
|
|
+ assertTrue(text(nullQuery).contains("缺少 query"), text(nullQuery));
|
|
|
|
|
+
|
|
|
|
|
+ GovernTableTool.QuerySpec spec = new GovernTableTool.QuerySpec();
|
|
|
|
|
+ ToolResultBlock blankName = tool.queryTableData(spec);
|
|
|
|
|
+ assertEquals(ToolResultState.ERROR, blankName.getState(), text(blankName));
|
|
|
|
|
+ assertTrue(text(blankName).contains("缺少 tableName"), text(blankName));
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ @Test
|
|
|
|
|
+ @DisplayName("未注册的表名被拒,并指回 list_tables")
|
|
|
|
|
+ void unknownTableRejected() {
|
|
|
|
|
+ GovernTableTool.QuerySpec spec = new GovernTableTool.QuerySpec();
|
|
|
|
|
+ spec.tableName = "no_such_table";
|
|
|
|
|
+
|
|
|
|
|
+ ToolResultBlock result = tool.queryTableData(spec);
|
|
|
|
|
+ assertEquals(ToolResultState.ERROR, result.getState(), text(result));
|
|
|
|
|
+ assertTrue(text(result).contains("表不存在"), text(result));
|
|
|
|
|
+ assertTrue(text(result).contains("list_tables"), text(result));
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ @Test
|
|
|
|
|
+ @DisplayName("条件白名单:非法列名 / 比较符 / 缺值 / 连接符逐个被拒")
|
|
|
|
|
+ void conditionValidation() {
|
|
|
|
|
+ // 非法列名(注入面)
|
|
|
|
|
+ assertEquals(ToolResultState.ERROR, errorOf(c -> c.column = "a; drop table x").getState());
|
|
|
|
|
+ assertTrue(text(errorOf(c -> c.column = "a; drop table x")).contains("非法列名"));
|
|
|
|
|
+ // 比较符白名单外
|
|
|
|
|
+ assertTrue(text(errorOf(c -> c.operator = "IN")).contains("不支持的比较符"));
|
|
|
|
|
+ // 缺 value
|
|
|
|
|
+ assertTrue(text(errorOf(c -> {
|
|
|
|
|
+ c.operator = ">=";
|
|
|
|
|
+ c.value = null;
|
|
|
|
|
+ })).contains("缺 value"));
|
|
|
|
|
+ // 连接符白名单(无条件时 conditionRun 不参与拼装,得带上一条条件才会校验)
|
|
|
|
|
+ GovernTableTool.QuerySpec spec = baseSpec();
|
|
|
|
|
+ spec.conditionRun = "XOR";
|
|
|
|
|
+ spec.conditions = List.of(cond("personCardNo", "=", "x"));
|
|
|
|
|
+ assertTrue(text(tool.queryTableData(spec)).contains("conditionRun"));
|
|
|
|
|
+ // 排序列注入
|
|
|
|
|
+ spec = baseSpec();
|
|
|
|
|
+ spec.orderKey = "id; delete from trans_record";
|
|
|
|
|
+ assertTrue(text(tool.queryTableData(spec)).contains("非法排序列"));
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ // ==================== 查询 ====================
|
|
|
|
|
+
|
|
|
|
|
+ @Test
|
|
|
|
|
+ @DisplayName("组 dto:tableName/limit/排序/条件转义齐全,结果带 head 中文列名与预览行")
|
|
|
|
|
+ void queryBuildsDtoAndReturnsRows() {
|
|
|
|
|
+ // LinkedHashMap:列序 = 插入序(Map.of 的遍历序不稳定,会让 columns 顺序抖动)
|
|
|
|
|
+ Map<String, Object> row = new LinkedHashMap<>();
|
|
|
|
|
+ row.put("personCardNo", "62220001");
|
|
|
|
|
+ row.put("transAmount", "100");
|
|
|
|
|
+ AtomicReference<GovernTreeQuery> captured = stubService(
|
|
|
|
|
+ List.of(row),
|
|
|
|
|
+ List.of(new BasicColumn("卡号", "personCardNo"), new BasicColumn("金额", "transAmount")));
|
|
|
|
|
+
|
|
|
|
|
+ GovernTableTool.QuerySpec spec = baseSpec();
|
|
|
|
|
+ spec.conditions = List.of(cond("personCardNo", "=", "O'Brien"),
|
|
|
|
|
+ cond("transAmount", ">=", "100"));
|
|
|
|
|
+ spec.conditionRun = "AND";
|
|
|
|
|
+ spec.orderKey = "transDate";
|
|
|
|
|
+ spec.asc = true;
|
|
|
|
|
+
|
|
|
|
|
+ ToolResultBlock result = tool.queryTableData(spec);
|
|
|
|
|
+ assertEquals(ToolResultState.SUCCESS, result.getState(), text(result));
|
|
|
|
|
+
|
|
|
|
|
+ // dto 组装
|
|
|
|
|
+ GovernTreeQuery dto = captured.get();
|
|
|
|
|
+ assertNotNull(dto);
|
|
|
|
|
+ assertEquals("trans_record", dto.getTableName());
|
|
|
|
|
+ assertEquals(1, dto.getPage());
|
|
|
|
|
+ assertEquals(Integer.MAX_VALUE, dto.getLimit(), "默认 max-rows=0:全量导出不限条数");
|
|
|
|
|
+ assertEquals("transDate", dto.getOrderKey());
|
|
|
|
|
+ assertTrue(dto.hasAsc());
|
|
|
|
|
+ assertNotNull(dto.getConditionSql());
|
|
|
|
|
+ assertEquals("AND", dto.getConditionSql().getConditionRun());
|
|
|
|
|
+ assertEquals(2, dto.getConditionSql().getConditions().size());
|
|
|
|
|
+ assertEquals("O''Brien", dto.getConditionSql().getConditions().get(0).value(),
|
|
|
|
|
+ "值里的单引号必须转义——服务端是裸拼接");
|
|
|
|
|
+ assertEquals("=", dto.getConditionSql().getConditions().get(0).condition());
|
|
|
|
|
+
|
|
|
|
|
+ // 结果形态:table/head 在前,artifact(无案件上下文时省略)之后是 columns/rows/totalRows
|
|
|
|
|
+ JsonNode body = json(result);
|
|
|
|
|
+ assertEquals("trans_record", body.get("table").asText());
|
|
|
|
|
+ assertEquals("卡号", body.get("head").get(0).get("label").asText());
|
|
|
|
|
+ assertEquals("personCardNo", body.get("columns").get(0).get("key").asText());
|
|
|
|
|
+ assertEquals("62220001", body.get("rows").get(0).get("personCardNo").asText());
|
|
|
|
|
+ assertEquals(1, body.get("totalRows").asInt());
|
|
|
|
|
+ assertFalse(body.has("artifact"), "无案件上下文不产文件,artifact 键省略");
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ @Test
|
|
|
|
|
+ @DisplayName("IS NULL 不需要 value;不传条件时 conditionSql 为 null(走服务端默认分支)")
|
|
|
|
|
+ void isNullAndNoConditions() {
|
|
|
|
|
+ AtomicReference<GovernTreeQuery> captured = stubService(List.of(), List.of());
|
|
|
|
|
+
|
|
|
|
|
+ GovernTableTool.QuerySpec spec = baseSpec();
|
|
|
|
|
+ spec.conditions = List.of(cond("otherCardNo", "is null", null));
|
|
|
|
|
+ assertEquals(ToolResultState.SUCCESS, tool.queryTableData(spec).getState(), text(tool.queryTableData(spec)));
|
|
|
|
|
+ assertEquals("IS NULL", captured.get().getConditionSql().getConditions().get(0).condition());
|
|
|
|
|
+
|
|
|
|
|
+ spec = baseSpec();
|
|
|
|
|
+ assertEquals(ToolResultState.SUCCESS, tool.queryTableData(spec).getState());
|
|
|
|
|
+ assertNull(captured.get().getConditionSql());
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ @Test
|
|
|
|
|
+ @DisplayName("行数保险丝:max-rows>0 且总数超限时报错让模型收窄,不静默截断")
|
|
|
|
|
+ void oversizedResultRejected() {
|
|
|
|
|
+ artifactProps.setMaxRows(1);
|
|
|
|
|
+ stubService(List.of(Map.of("a", "1"), Map.of("a", "2")), List.of(new BasicColumn("A", "a")));
|
|
|
|
|
+
|
|
|
|
|
+ ToolResultBlock result = tool.queryTableData(baseSpec());
|
|
|
|
|
+ assertEquals(ToolResultState.ERROR, result.getState(), text(result));
|
|
|
|
|
+ assertTrue(text(result).contains("结果集过大"), text(result));
|
|
|
|
|
+ assertTrue(text(result).contains("收窄"), text(result));
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ @Test
|
|
|
|
|
+ @DisplayName("有案件上下文时产出 query_table_data 的 CSV 产物并回传 artifact 元数据")
|
|
|
|
|
+ void artifactFlowsThrough() {
|
|
|
|
|
+ AtomicReference<GovernTreeQuery> captured = stubService(
|
|
|
|
|
+ List.of(Map.of("personCardNo", "62220001")),
|
|
|
|
|
+ List.of(new BasicColumn("卡号", "personCardNo")));
|
|
|
|
|
+
|
|
|
|
|
+ ToolResultBlock result = CaseContextHolder.callWith(CASE_ID, 1L,
|
|
|
|
|
+ () -> tool.queryTableData(baseSpec()));
|
|
|
|
|
+ assertEquals(ToolResultState.SUCCESS, result.getState(), text(result));
|
|
|
|
|
+ assertNotNull(captured.get());
|
|
|
|
|
+
|
|
|
|
|
+ JsonNode artifact = json(result).get("artifact");
|
|
|
|
|
+ assertNotNull(artifact, "有案件上下文必须回传 artifact");
|
|
|
|
|
+ String id = artifact.get("id").asText();
|
|
|
|
|
+ assertTrue(id.startsWith("query_table_data_"), id);
|
|
|
|
|
+ assertEquals(1, artifact.get("totalRows").asInt());
|
|
|
|
|
+ created.add(PathConst.WORKSPACE
|
|
|
|
|
+ .resolve(String.valueOf(CASE_ID)).resolve("artifacts").resolve(id));
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ // ==================== 辅助 ====================
|
|
|
|
|
+
|
|
|
|
|
+ private GovernTableTool.QuerySpec baseSpec() {
|
|
|
|
|
+ GovernTableTool.QuerySpec spec = new GovernTableTool.QuerySpec();
|
|
|
|
|
+ spec.tableName = "trans_record";
|
|
|
|
|
+ return spec;
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ private GovernTableTool.ConditionSpec cond(String column, String operator, String value) {
|
|
|
|
|
+ GovernTableTool.ConditionSpec c = new GovernTableTool.ConditionSpec();
|
|
|
|
|
+ c.column = column;
|
|
|
|
|
+ c.operator = operator;
|
|
|
|
|
+ c.value = value;
|
|
|
|
|
+ return c;
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ /** 只需要失败文本时的快捷入口 */
|
|
|
|
|
+ private ToolResultBlock errorOf(java.util.function.Consumer<GovernTableTool.ConditionSpec> filler) {
|
|
|
|
|
+ GovernTableTool.QuerySpec spec = baseSpec();
|
|
|
|
|
+ GovernTableTool.ConditionSpec c = cond("personCardNo", "=", "x");
|
|
|
|
|
+ filler.accept(c);
|
|
|
|
|
+ spec.conditions = List.of(c);
|
|
|
|
|
+ return tool.queryTableData(spec);
|
|
|
|
|
+ }
|
|
|
|
|
+}
|