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
@@ -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))));
}
}
}