fix: avoid low-confidence supervisor retry by default

This commit is contained in:
zhuyongxin
2026-07-03 15:32:09 +08:00
parent b0f288ae36
commit 5b827fe90e
2 changed files with 118 additions and 38 deletions
@@ -104,6 +104,9 @@ public class ChatService {
@Value("${verifier.low-confidence-threshold:0.5}") @Value("${verifier.low-confidence-threshold:0.5}")
private double verifierLowConfidenceThreshold; private double verifierLowConfidenceThreshold;
@Value("${chat.complex.retry-on-low-confidence:false}")
private boolean retryOnLowConfidence;
/** 多 Agent Chat 的 Prompt */ /** 多 Agent Chat 的 Prompt */
private String chatPlannerPrompt; private String chatPlannerPrompt;
private String chatExecutorPrompt; private String chatExecutorPrompt;
@@ -211,16 +214,19 @@ public class ChatService {
/** /**
* 动态构建方法工具数组 * 动态构建方法工具数组
* 根据 cls.mock-enabled 决定是否包含 QueryLogsTools * 根据已注入的 Bean 暴露本地工具,避免 mock/真实模式下漏注入。
*/ */
public Object[] buildMethodToolsArray() { public Object[] buildMethodToolsArray() {
List<Object> methodTools = new ArrayList<>();
methodTools.add(dateTimeTools);
methodTools.add(lookupKnowledgeTool);
if (queryLogsTools != null) { if (queryLogsTools != null) {
// Mock 模式:包含 QueryLogsTools methodTools.add(queryLogsTools);
return new Object[]{dateTimeTools, lookupKnowledgeTool};
} else {
// 真实模式:不包含 QueryLogsTools(由 MCP 提供日志查询功能)
return new Object[]{dateTimeTools, lookupKnowledgeTool, queryMetricsTools};
} }
if (queryMetricsTools != null) {
methodTools.add(queryMetricsTools);
}
return methodTools.toArray();
} }
/** /**
@@ -436,7 +442,10 @@ public class ChatService {
break; break;
} }
if (finalDecision.groundednessScore() >= verifierLowConfidenceThreshold || round == 2) { boolean shouldRetry = retryOnLowConfidence
&& finalDecision.groundednessScore() < verifierLowConfidenceThreshold
&& round < 2;
if (!shouldRetry) {
answer = buildLowConfidenceOutput(answer, finalDecision); answer = buildLowConfidenceOutput(answer, finalDecision);
persistVerifierEvaluation(session, finalDecision, round); persistVerifierEvaluation(session, finalDecision, round);
break; break;
@@ -2,6 +2,7 @@ package com.superbiz.agent.service;
import com.superbiz.agent.agent.tool.DateTimeTools; import com.superbiz.agent.agent.tool.DateTimeTools;
import com.superbiz.agent.agent.tool.QueryLogsTools; 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.AgentStep;
import com.superbiz.agent.domain.entity.DiagnosisSession; import com.superbiz.agent.domain.entity.DiagnosisSession;
import com.superbiz.agent.repository.AgentStepRepository; import com.superbiz.agent.repository.AgentStepRepository;
@@ -23,6 +24,7 @@ import java.util.Optional;
import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicInteger;
import static org.junit.jupiter.api.Assertions.assertEquals; 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.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.anyString; import static org.mockito.ArgumentMatchers.anyString;
@@ -33,9 +35,81 @@ class ChatServiceSupervisorAgentTest {
@Test @Test
void executeChatComplexInvokesSupervisorFlow() throws Exception { void executeChatComplexInvokesSupervisorFlow() throws Exception {
ChatService chatService = new ChatService(); ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel(); ScriptedChatModel chatModel = new ScriptedChatModel();
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);
}
@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(),
"supervisor-low-confidence-session"
);
assertTrue(result.answer().contains("EXECUTOR_FINAL_ANSWER"));
assertEquals(List.of("chat_planner", "chat_executor", "chat_verifier", "FINISH"), chatModel.decisions);
}
@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); DiagnosisSessionRepository diagnosisSessionRepository = mock(DiagnosisSessionRepository.class);
when(diagnosisSessionRepository.findBySessionId(anyString())).thenReturn(Optional.empty()); when(diagnosisSessionRepository.findBySessionId(anyString())).thenReturn(Optional.empty());
when(diagnosisSessionRepository.save(any(DiagnosisSession.class))).thenAnswer(invocation -> invocation.getArgument(0)); when(diagnosisSessionRepository.save(any(DiagnosisSession.class))).thenAnswer(invocation -> invocation.getArgument(0));
@@ -75,19 +149,7 @@ class ChatServiceSupervisorAgentTest {
ReflectionTestUtils.setField(chatService, "chatPlannerPrompt", "PLANNER_TEST_PROMPT"); ReflectionTestUtils.setField(chatService, "chatPlannerPrompt", "PLANNER_TEST_PROMPT");
ReflectionTestUtils.setField(chatService, "chatExecutorPrompt", "EXECUTOR_TEST_PROMPT"); ReflectionTestUtils.setField(chatService, "chatExecutorPrompt", "EXECUTOR_TEST_PROMPT");
ReflectionTestUtils.setField(chatService, "chatVerifierPrompt", "VERIFIER_TEST_PROMPT"); ReflectionTestUtils.setField(chatService, "chatVerifierPrompt", "VERIFIER_TEST_PROMPT");
return chatService;
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 static final class ScriptedChatModel implements ChatModel {
@@ -96,21 +158,10 @@ class ChatServiceSupervisorAgentTest {
private int decisionIndex; private int decisionIndex;
private String promptText = ""; private String promptText = "";
private boolean sawVerifierPrompt; private boolean sawVerifierPrompt;
private final String verifierOutput;
@Override private ScriptedChatModel() {
public ChatResponse call(Prompt prompt) { this("""
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", "verdict": "PASS",
"groundedness_score": 1.0, "groundedness_score": 1.0,
@@ -126,7 +177,27 @@ class ChatServiceSupervisorAgentTest {
], ],
"rationale": "scripted pass" "rationale": "scripted pass"
} }
"""; """);
}
private ScriptedChatModel(String verifierOutput) {
this.verifierOutput = verifierOutput;
}
@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 = verifierOutput;
} else { } else {
text = "UNEXPECTED_PROMPT"; text = "UNEXPECTED_PROMPT";
} }