From b0f288ae3628c759a5f0ce4f97f4c422024b038e Mon Sep 17 00:00:00 2001 From: zhuyongxin Date: Fri, 3 Jul 2026 14:10:55 +0800 Subject: [PATCH] use supervisor agent for complex chat --- .../superbiz/agent/service/ChatService.java | 92 ++++-------- .../ChatServiceSupervisorAgentTest.java | 136 ++++++++++++++++++ 2 files changed, 166 insertions(+), 62 deletions(-) create mode 100644 src/test/java/com/superbiz/agent/service/ChatServiceSupervisorAgentTest.java diff --git a/src/main/java/com/superbiz/agent/service/ChatService.java b/src/main/java/com/superbiz/agent/service/ChatService.java index e1731fe..26bf469 100644 --- a/src/main/java/com/superbiz/agent/service/ChatService.java +++ b/src/main/java/com/superbiz/agent/service/ChatService.java @@ -397,11 +397,25 @@ public class ChatService { .subAgents(List.of(planner, executor, verifier)) .build(); - String plannerPlan = callAgent(planner, buildPlannerInput(question, retryContext), config); - answer = callAgent(executor, buildExecutorInput(question, plannerPlan, retryContext), config); + String supervisorInput = buildSupervisorInput(question, retryContext); + Optional stateOptional = supervisor.invoke(supervisorInput, config); + if (stateOptional.isEmpty()) { + finalDecision = buildVerifierFallbackDecision(round, "supervisor 未返回有效状态"); + answer = buildLowConfidenceOutput(answer, finalDecision); + persistVerifierEvaluation(session, finalDecision, round); + break; + } + + String plannerPlan = extractStateText(stateOptional, "planner_plan"); + answer = extractStateText(stateOptional, "executor_feedback"); VerifierContextHolder.setExecutorFinalAnswer(answer); - String verifierOutput = callAgent(verifier, "VERIFY", config); + String verifierOutput = extractStateText(stateOptional, "verifier_output"); finalDecision = parseVerifierDecision(verifierOutput, round); + logger.debug("Supervisor round {} finished: plannerPlanLength={}, answerLength={}, verifierOutputLength={}", + round, + plannerPlan != null ? plannerPlan.length() : 0, + answer != null ? answer.length() : 0, + verifierOutput != null ? verifierOutput.length() : 0); if (finalDecision == null) { finalDecision = buildVerifierFallbackDecision(round, "verifier_output 缺失或无法解析"); @@ -530,10 +544,6 @@ public class ChatService { .build(); } - private String callAgent(ReactAgent agent, String input, RunnableConfig config) throws GraphRunnerException { - return agent.call(input, config).getText(); - } - private String resolveSessionId(String requestedSessionId) { if (requestedSessionId != null && !requestedSessionId.isBlank()) { return requestedSessionId; @@ -558,21 +568,14 @@ public class ChatService { return session; } - private String buildPlannerInput(String question, String retryContext) { - if (retryContext == null || retryContext.isBlank()) { - return question; - } - return question + "\n\n--- 补充约束 ---\n" + retryContext; - } - - private String buildExecutorInput(String question, String plannerPlan, String retryContext) { - StringBuilder input = new StringBuilder(question); - if (plannerPlan != null && !plannerPlan.isBlank()) { - input.append("\n\n--- planner_plan ---\n").append(plannerPlan); - } + private String buildSupervisorInput(String question, String retryContext) { + StringBuilder input = new StringBuilder(); + input.append("请按 supervisor 系统提示完成本轮 Planner -> Executor -> Verifier 编排。\n\n"); + input.append("--- 用户问题 ---\n").append(question); if (retryContext != null && !retryContext.isBlank()) { input.append("\n\n--- retry_context ---\n").append(retryContext); } + input.append("\n\n完成 verifier 后立即 FINISH,不要额外生成最终答案。"); return input.toString(); } @@ -674,55 +677,20 @@ public class ChatService { """.formatted(round); } - private String buildRoundInput(String question, String retryContext) { - if (retryContext == null || retryContext.isBlank()) { - return question; - } - return question + "\n\n--- 补充约束 ---\n" + retryContext; - } - - private String extractExecutorAnswer(Optional stateOptional) { + private String extractStateText(Optional stateOptional, String key) { if (stateOptional.isEmpty()) { return null; } - return stateOptional.get().value("executor_feedback") - .filter(AssistantMessage.class::isInstance) - .map(AssistantMessage.class::cast) - .map(AssistantMessage::getText) + return stateOptional.get().value(key) + .map(value -> { + if (value instanceof AssistantMessage assistantMessage) { + return assistantMessage.getText(); + } + return String.valueOf(value); + }) .orElse(null); } - private VerifierDecision parseVerifierDecision(Optional stateOptional, int round) { - if (stateOptional.isEmpty()) { - return null; - } - - Optional verifierOutput = stateOptional.get().value("verifier_output") - .filter(AssistantMessage.class::isInstance) - .map(AssistantMessage.class::cast); - - if (verifierOutput.isEmpty() || verifierOutput.get().getText() == null || verifierOutput.get().getText().isBlank()) { - return null; - } - - try { - JsonNode root = objectMapper.readTree(verifierOutput.get().getText()); - List> factsChecked = parseFactsChecked(root.path("facts_checked")); - - return new VerifierDecision( - root.path("verdict").asText("LOW_CONFID"), - root.path("groundedness_score").asDouble(0.0), - root.path("critical_fact_count").asInt(0), - factsChecked, - root.path("rationale").asText(""), - round - ); - } catch (Exception e) { - logger.error("解析 verifier_output 失败: {}", verifierOutput.get().getText(), e); - return null; - } - } - private void persistVerifierEvaluation(DiagnosisSession session, VerifierDecision decision, int round) { if (decision == null) { return; diff --git a/src/test/java/com/superbiz/agent/service/ChatServiceSupervisorAgentTest.java b/src/test/java/com/superbiz/agent/service/ChatServiceSupervisorAgentTest.java new file mode 100644 index 0000000..99c6652 --- /dev/null +++ b/src/test/java/com/superbiz/agent/service/ChatServiceSupervisorAgentTest.java @@ -0,0 +1,136 @@ +package com.superbiz.agent.service; + +import com.superbiz.agent.agent.tool.DateTimeTools; +import com.superbiz.agent.agent.tool.QueryLogsTools; +import com.superbiz.agent.domain.entity.AgentStep; +import com.superbiz.agent.domain.entity.DiagnosisSession; +import com.superbiz.agent.repository.AgentStepRepository; +import com.superbiz.agent.repository.DiagnosisSessionRepository; +import com.superbiz.agent.tool.LookupKnowledgeTool; +import com.superbiz.agent.tool.RetrievedDocTracker; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.chat.model.Generation; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.tool.ToolCallback; +import org.springframework.test.util.ReflectionTestUtils; + +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.concurrent.atomic.AtomicInteger; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +class ChatServiceSupervisorAgentTest { + + @Test + void executeChatComplexInvokesSupervisorFlow() throws Exception { + ChatService chatService = new ChatService(); + ScriptedChatModel chatModel = new ScriptedChatModel(); + + DiagnosisSessionRepository diagnosisSessionRepository = mock(DiagnosisSessionRepository.class); + when(diagnosisSessionRepository.findBySessionId(anyString())).thenReturn(Optional.empty()); + when(diagnosisSessionRepository.save(any(DiagnosisSession.class))).thenAnswer(invocation -> invocation.getArgument(0)); + + AtomicInteger stepId = new AtomicInteger(1); + AgentStepRepository agentStepRepository = mock(AgentStepRepository.class); + when(agentStepRepository.save(any(AgentStep.class))).thenAnswer(invocation -> { + AgentStep step = invocation.getArgument(0); + if (step.getId() == null) { + step.setId((long) stepId.getAndIncrement()); + } + return step; + }); + when(agentStepRepository.findById(any())).thenReturn(Optional.of(new AgentStep())); + when(agentStepRepository.findBySessionIdOrderByStepIndex(anyString())).thenReturn(List.of()); + + EvaluationService evaluationService = mock(EvaluationService.class); + RetrievedDocTracker retrievedDocTracker = mock(RetrievedDocTracker.class); + KnowledgeDomainService knowledgeDomainService = mock(KnowledgeDomainService.class); + when(knowledgeDomainService.buildKnowledgeMap()).thenReturn(""); + ToolTraceSummaryService toolTraceSummaryService = mock(ToolTraceSummaryService.class); + when(toolTraceSummaryService.buildVerifierTraceSummary(anyString(), anyString())).thenReturn(List.of()); + SelfEvaluationMergeService selfEvaluationMergeService = mock(SelfEvaluationMergeService.class); + when(selfEvaluationMergeService.mergeVerifierEvaluation(any(), any())).thenReturn("{}"); + + ReflectionTestUtils.setField(chatService, "dateTimeTools", new DateTimeTools()); + ReflectionTestUtils.setField(chatService, "lookupKnowledgeTool", new LookupKnowledgeTool()); + ReflectionTestUtils.setField(chatService, "queryLogsTools", new QueryLogsTools(mock(ToolInvocationRecorder.class))); + ReflectionTestUtils.setField(chatService, "diagnosisSessionRepository", diagnosisSessionRepository); + ReflectionTestUtils.setField(chatService, "agentStepRepository", agentStepRepository); + ReflectionTestUtils.setField(chatService, "evaluationService", evaluationService); + ReflectionTestUtils.setField(chatService, "retrievedDocTracker", retrievedDocTracker); + ReflectionTestUtils.setField(chatService, "knowledgeDomainService", knowledgeDomainService); + ReflectionTestUtils.setField(chatService, "toolTraceSummaryService", toolTraceSummaryService); + ReflectionTestUtils.setField(chatService, "selfEvaluationMergeService", selfEvaluationMergeService); + ReflectionTestUtils.setField(chatService, "verifierLowConfidenceThreshold", 0.5d); + ReflectionTestUtils.setField(chatService, "chatPlannerPrompt", "PLANNER_TEST_PROMPT"); + ReflectionTestUtils.setField(chatService, "chatExecutorPrompt", "EXECUTOR_TEST_PROMPT"); + ReflectionTestUtils.setField(chatService, "chatVerifierPrompt", "VERIFIER_TEST_PROMPT"); + + ChatService.ChatResult result = chatService.executeChatComplex( + chatModel, + new ToolCallback[0], + "请分析订单支付超时的原因,并给出修复建议", + List.of(), + "supervisor-test-session" + ); + + assertEquals("EXECUTOR_FINAL_ANSWER", result.answer()); + assertEquals("supervisor-test-session", result.sessionId()); + assertEquals(List.of("chat_planner", "chat_executor", "chat_verifier", "FINISH"), chatModel.decisions); + assertTrue(chatModel.sawVerifierPrompt); + } + + private static final class ScriptedChatModel implements ChatModel { + private final List decisionScript = List.of("chat_planner", "chat_executor", "chat_verifier", "FINISH"); + private final java.util.ArrayList decisions = new java.util.ArrayList<>(); + private int decisionIndex; + private String promptText = ""; + private boolean sawVerifierPrompt; + + @Override + public ChatResponse call(Prompt prompt) { + promptText = prompt.getContents(); + String text; + if (promptText.contains("Available options:")) { + text = "{\"agent\":\"" + decisionScript.get(decisionIndex++) + "\"}"; + decisions.add(text.substring(10, text.length() - 2)); + } else if (promptText.contains("PLANNER_TEST_PROMPT")) { + text = "PLANNER_PLAN"; + } else if (promptText.contains("EXECUTOR_TEST_PROMPT")) { + text = "EXECUTOR_FINAL_ANSWER"; + } else if (promptText.contains("VERIFIER_TEST_PROMPT")) { + sawVerifierPrompt = true; + text = """ + { + "verdict": "PASS", + "groundedness_score": 1.0, + "critical_fact_count": 1, + "facts_checked": [ + { + "fact": "executor answer generated", + "is_critical": true, + "verification": "direct_evidence", + "detail": "covered by scripted verifier", + "evidence_refs": [] + } + ], + "rationale": "scripted pass" + } + """; + } else { + text = "UNEXPECTED_PROMPT"; + } + return new ChatResponse(List.of(new Generation(new AssistantMessage(text)))); + } + } +}