refactor: use sequential agent for chat workflow
This commit is contained in:
@@ -0,0 +1,223 @@
|
||||
package com.superbiz.agent.service;
|
||||
|
||||
import com.superbiz.agent.agent.tool.DateTimeTools;
|
||||
import com.superbiz.agent.agent.tool.QueryLogsTools;
|
||||
import com.superbiz.agent.agent.tool.QueryMetricsTools;
|
||||
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.assertSame;
|
||||
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 ChatServiceSequentialAgentTest {
|
||||
|
||||
@Test
|
||||
void executeChatComplexInvokesSequentialWorkflow() throws Exception {
|
||||
ChatService chatService = createChatService();
|
||||
ScriptedChatModel chatModel = new ScriptedChatModel();
|
||||
|
||||
ChatService.ChatResult result = chatService.executeChatComplex(
|
||||
chatModel,
|
||||
new ToolCallback[0],
|
||||
"请分析订单支付超时的原因,并给出修复建议",
|
||||
List.of(),
|
||||
"sequential-test-session"
|
||||
);
|
||||
|
||||
assertEquals("EXECUTOR_FINAL_ANSWER", result.answer());
|
||||
assertEquals("sequential-test-session", result.sessionId());
|
||||
assertEquals(List.of("chat_planner", "chat_executor", "chat_verifier"), chatModel.agentCalls);
|
||||
assertTrue(chatModel.sawVerifierPrompt);
|
||||
}
|
||||
|
||||
@Test
|
||||
void executeChatComplexDoesNotRetryLowConfidenceByDefault() throws Exception {
|
||||
ChatService chatService = createChatService();
|
||||
ScriptedChatModel chatModel = new ScriptedChatModel("""
|
||||
{
|
||||
"verdict": "LOW_CONFID",
|
||||
"groundedness_score": 0.1,
|
||||
"critical_fact_count": 1,
|
||||
"facts_checked": [
|
||||
{
|
||||
"fact": "missing direct evidence",
|
||||
"is_critical": true,
|
||||
"verification": "no_evidence",
|
||||
"detail": "scripted evidence gap",
|
||||
"evidence_refs": []
|
||||
}
|
||||
],
|
||||
"rationale": "scripted low confidence"
|
||||
}
|
||||
""");
|
||||
|
||||
ChatService.ChatResult result = chatService.executeChatComplex(
|
||||
chatModel,
|
||||
new ToolCallback[0],
|
||||
"请分析订单支付超时的原因,并给出修复建议",
|
||||
List.of(),
|
||||
"sequential-low-confidence-session"
|
||||
);
|
||||
|
||||
assertTrue(result.answer().contains("EXECUTOR_FINAL_ANSWER"));
|
||||
assertEquals(List.of("chat_planner", "chat_executor", "chat_verifier"), chatModel.agentCalls);
|
||||
}
|
||||
|
||||
@Test
|
||||
void executeChatComplexRunsPlannerExecutorVerifierInFixedOrder() throws Exception {
|
||||
ChatService chatService = createChatService();
|
||||
ScriptedChatModel chatModel = new ScriptedChatModel();
|
||||
|
||||
ChatService.ChatResult result = chatService.executeChatComplex(
|
||||
chatModel,
|
||||
new ToolCallback[0],
|
||||
"请分析订单支付超时的原因,并给出修复建议",
|
||||
List.of(),
|
||||
"sequential-workflow-session"
|
||||
);
|
||||
|
||||
assertEquals("EXECUTOR_FINAL_ANSWER", result.answer());
|
||||
assertEquals(List.of("chat_planner", "chat_executor", "chat_verifier"), chatModel.agentCalls);
|
||||
assertTrue(chatModel.sawVerifierPrompt);
|
||||
}
|
||||
|
||||
@Test
|
||||
void buildMethodToolsArrayIncludesLogsAndMetricsWhenAvailable() {
|
||||
ChatService chatService = new ChatService();
|
||||
DateTimeTools dateTimeTools = new DateTimeTools();
|
||||
LookupKnowledgeTool lookupKnowledgeTool = new LookupKnowledgeTool();
|
||||
QueryLogsTools queryLogsTools = new QueryLogsTools(mock(ToolInvocationRecorder.class));
|
||||
QueryMetricsTools queryMetricsTools = new QueryMetricsTools(mock(ToolInvocationRecorder.class));
|
||||
|
||||
ReflectionTestUtils.setField(chatService, "dateTimeTools", dateTimeTools);
|
||||
ReflectionTestUtils.setField(chatService, "lookupKnowledgeTool", lookupKnowledgeTool);
|
||||
ReflectionTestUtils.setField(chatService, "queryLogsTools", queryLogsTools);
|
||||
ReflectionTestUtils.setField(chatService, "queryMetricsTools", queryMetricsTools);
|
||||
|
||||
Object[] methodTools = chatService.buildMethodToolsArray();
|
||||
|
||||
assertEquals(4, methodTools.length);
|
||||
assertSame(dateTimeTools, methodTools[0]);
|
||||
assertSame(lookupKnowledgeTool, methodTools[1]);
|
||||
assertSame(queryLogsTools, methodTools[2]);
|
||||
assertSame(queryMetricsTools, methodTools[3]);
|
||||
}
|
||||
|
||||
private ChatService createChatService() {
|
||||
ChatService chatService = new ChatService();
|
||||
|
||||
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");
|
||||
return chatService;
|
||||
}
|
||||
|
||||
private static final class ScriptedChatModel implements ChatModel {
|
||||
private final java.util.ArrayList<String> agentCalls = new java.util.ArrayList<>();
|
||||
private String promptText = "";
|
||||
private boolean sawVerifierPrompt;
|
||||
private final String verifierOutput;
|
||||
|
||||
private ScriptedChatModel() {
|
||||
this("""
|
||||
{
|
||||
"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"
|
||||
}
|
||||
""");
|
||||
}
|
||||
|
||||
private ScriptedChatModel(String verifierOutput) {
|
||||
this.verifierOutput = verifierOutput;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatResponse call(Prompt prompt) {
|
||||
promptText = prompt.getContents();
|
||||
String text;
|
||||
if (promptText.contains("PLANNER_TEST_PROMPT")) {
|
||||
agentCalls.add("chat_planner");
|
||||
text = "PLANNER_PLAN";
|
||||
} else if (promptText.contains("EXECUTOR_TEST_PROMPT")) {
|
||||
agentCalls.add("chat_executor");
|
||||
text = "EXECUTOR_FINAL_ANSWER";
|
||||
} else if (promptText.contains("VERIFIER_TEST_PROMPT")) {
|
||||
agentCalls.add("chat_verifier");
|
||||
sawVerifierPrompt = true;
|
||||
text = verifierOutput;
|
||||
} else {
|
||||
text = "UNEXPECTED_PROMPT";
|
||||
}
|
||||
return new ChatResponse(List.of(new Generation(new AssistantMessage(text))));
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user