use supervisor agent for complex chat

This commit is contained in:
zhuyongxin
2026-07-03 14:10:55 +08:00
parent 1ff7f09d25
commit b0f288ae36
2 changed files with 166 additions and 62 deletions
@@ -397,11 +397,25 @@ public class ChatService {
.subAgents(List.of(planner, executor, verifier)) .subAgents(List.of(planner, executor, verifier))
.build(); .build();
String plannerPlan = callAgent(planner, buildPlannerInput(question, retryContext), config); String supervisorInput = buildSupervisorInput(question, retryContext);
answer = callAgent(executor, buildExecutorInput(question, plannerPlan, retryContext), config); Optional<OverAllState> 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); VerifierContextHolder.setExecutorFinalAnswer(answer);
String verifierOutput = callAgent(verifier, "VERIFY", config); String verifierOutput = extractStateText(stateOptional, "verifier_output");
finalDecision = parseVerifierDecision(verifierOutput, round); 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) { if (finalDecision == null) {
finalDecision = buildVerifierFallbackDecision(round, "verifier_output 缺失或无法解析"); finalDecision = buildVerifierFallbackDecision(round, "verifier_output 缺失或无法解析");
@@ -530,10 +544,6 @@ public class ChatService {
.build(); .build();
} }
private String callAgent(ReactAgent agent, String input, RunnableConfig config) throws GraphRunnerException {
return agent.call(input, config).getText();
}
private String resolveSessionId(String requestedSessionId) { private String resolveSessionId(String requestedSessionId) {
if (requestedSessionId != null && !requestedSessionId.isBlank()) { if (requestedSessionId != null && !requestedSessionId.isBlank()) {
return requestedSessionId; return requestedSessionId;
@@ -558,21 +568,14 @@ public class ChatService {
return session; return session;
} }
private String buildPlannerInput(String question, String retryContext) { private String buildSupervisorInput(String question, String retryContext) {
if (retryContext == null || retryContext.isBlank()) { StringBuilder input = new StringBuilder();
return question; input.append("请按 supervisor 系统提示完成本轮 Planner -> Executor -> Verifier 编排。\n\n");
} input.append("--- 用户问题 ---\n").append(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);
}
if (retryContext != null && !retryContext.isBlank()) { if (retryContext != null && !retryContext.isBlank()) {
input.append("\n\n--- retry_context ---\n").append(retryContext); input.append("\n\n--- retry_context ---\n").append(retryContext);
} }
input.append("\n\n完成 verifier 后立即 FINISH,不要额外生成最终答案。");
return input.toString(); return input.toString();
} }
@@ -674,55 +677,20 @@ public class ChatService {
""".formatted(round); """.formatted(round);
} }
private String buildRoundInput(String question, String retryContext) { private String extractStateText(Optional<OverAllState> stateOptional, String key) {
if (retryContext == null || retryContext.isBlank()) {
return question;
}
return question + "\n\n--- 补充约束 ---\n" + retryContext;
}
private String extractExecutorAnswer(Optional<OverAllState> stateOptional) {
if (stateOptional.isEmpty()) { if (stateOptional.isEmpty()) {
return null; return null;
} }
return stateOptional.get().value("executor_feedback") return stateOptional.get().value(key)
.filter(AssistantMessage.class::isInstance) .map(value -> {
.map(AssistantMessage.class::cast) if (value instanceof AssistantMessage assistantMessage) {
.map(AssistantMessage::getText) return assistantMessage.getText();
}
return String.valueOf(value);
})
.orElse(null); .orElse(null);
} }
private VerifierDecision parseVerifierDecision(Optional<OverAllState> stateOptional, int round) {
if (stateOptional.isEmpty()) {
return null;
}
Optional<AssistantMessage> 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<Map<String, Object>> 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) { private void persistVerifierEvaluation(DiagnosisSession session, VerifierDecision decision, int round) {
if (decision == null) { if (decision == null) {
return; return;
@@ -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<String> decisionScript = List.of("chat_planner", "chat_executor", "chat_verifier", "FINISH");
private final java.util.ArrayList<String> 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))));
}
}
}