use supervisor agent for complex chat
This commit is contained in:
@@ -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))));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user