feat(harness): add single diagnosis react agent

This commit is contained in:
zhuyongxin
2026-07-21 22:29:49 +08:00
parent 85029d96a7
commit 2362665519
30 changed files with 1711 additions and 2 deletions
@@ -0,0 +1,54 @@
package com.superbiz.agent.harness.agent;
import com.alibaba.cloud.ai.graph.agent.ReactAgent;
import com.alibaba.cloud.ai.graph.agent.hook.Hook;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.core.DiagnosisHarnessCore;
import com.superbiz.agent.harness.core.RunContext;
import org.springframework.ai.chat.model.ChatModel;
import java.util.List;
import java.util.Objects;
public final class DiagnosisAgentFactory {
public static final String AGENT_NAME = "diagnosis_agent";
private final ChatModel chatModel;
private final DiagnosisHarnessCore core;
private final HarnessEvidenceTools evidenceTools;
private final ObjectMapper objectMapper;
private final List<Hook> auditHooks;
private final String prompt;
public DiagnosisAgentFactory(ChatModel chatModel,
DiagnosisHarnessCore core,
HarnessEvidenceTools evidenceTools,
ObjectMapper objectMapper,
List<? extends Hook> auditHooks) {
this.chatModel = Objects.requireNonNull(chatModel, "chatModel must not be null");
this.core = Objects.requireNonNull(core, "core must not be null");
this.evidenceTools = Objects.requireNonNull(evidenceTools, "evidenceTools must not be null");
this.objectMapper = Objects.requireNonNull(objectMapper, "objectMapper must not be null");
this.auditHooks = List.copyOf(Objects.requireNonNull(auditHooks, "auditHooks must not be null"));
this.prompt = DiagnosisAgentPrompt.load();
}
public ReactAgent create(RunContext context) {
Objects.requireNonNull(context, "context must not be null");
return ReactAgent.builder()
.name(AGENT_NAME)
.description("Collects bounded evidence and authors one diagnosis draft")
.model(chatModel)
.systemPrompt(prompt)
.tools(evidenceTools.callbacks())
.interceptors(
new HarnessModelInterceptor(core, context),
new HarnessToolInterceptor(context, evidenceTools, objectMapper))
.hooks(auditHooks)
.outputSchema(new DiagnosisDraftOutputSchema(objectMapper).getFormat())
.parallelToolExecution(false)
.releaseThread(true)
.build();
}
}
@@ -0,0 +1,15 @@
package com.superbiz.agent.harness.agent;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.superbiz.agent.harness.contract.PreviousTurn;
public record DiagnosisAgentInput(
@JsonProperty("query") String query,
@JsonProperty("previous_turn") PreviousTurn previousTurn) {
public DiagnosisAgentInput {
if (query == null || query.isBlank()) {
throw new IllegalArgumentException("query must not be blank");
}
}
}
@@ -0,0 +1,27 @@
package com.superbiz.agent.harness.agent;
public final class DiagnosisAgentLimitException extends IllegalArgumentException {
private final String boundary;
private final long limit;
private final long actual;
public DiagnosisAgentLimitException(String boundary, long limit, long actual) {
super(boundary + " exceeds UTF-8 byte limit: limit=" + limit + ", actual=" + actual);
this.boundary = boundary;
this.limit = limit;
this.actual = actual;
}
public String boundary() {
return boundary;
}
public long limit() {
return limit;
}
public long actual() {
return actual;
}
}
@@ -0,0 +1,21 @@
package com.superbiz.agent.harness.agent;
public record DiagnosisAgentLimits(
long maxQueryBytes,
long maxPreviousTurnBytes,
long maxInputBytes,
long maxDraftBytes) {
public DiagnosisAgentLimits {
requirePositive(maxQueryBytes, "maxQueryBytes");
requirePositive(maxPreviousTurnBytes, "maxPreviousTurnBytes");
requirePositive(maxInputBytes, "maxInputBytes");
requirePositive(maxDraftBytes, "maxDraftBytes");
}
private static void requirePositive(long value, String name) {
if (value <= 0) {
throw new IllegalArgumentException(name + " must be positive");
}
}
}
@@ -0,0 +1,12 @@
package com.superbiz.agent.harness.agent;
public final class DiagnosisAgentOutputException extends RuntimeException {
public DiagnosisAgentOutputException(String message) {
super(message);
}
public DiagnosisAgentOutputException(String message, Throwable cause) {
super(message, cause);
}
}
@@ -0,0 +1,24 @@
package com.superbiz.agent.harness.agent;
import org.springframework.core.io.ClassPathResource;
import java.io.IOException;
import java.io.InputStream;
import java.nio.charset.StandardCharsets;
public final class DiagnosisAgentPrompt {
public static final String RESOURCE_PATH = "prompts/diagnosis-agent-prompt.md";
private DiagnosisAgentPrompt() {
}
public static String load() {
ClassPathResource resource = new ClassPathResource(RESOURCE_PATH);
try (InputStream input = resource.getInputStream()) {
return new String(input.readAllBytes(), StandardCharsets.UTF_8);
} catch (IOException e) {
throw new IllegalStateException("Failed to load Diagnosis Agent prompt", e);
}
}
}
@@ -0,0 +1,103 @@
package com.superbiz.agent.harness.agent;
import com.alibaba.cloud.ai.graph.RunnableConfig;
import com.alibaba.cloud.ai.graph.agent.ReactAgent;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.DeserializationFeature;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.fasterxml.jackson.databind.ObjectReader;
import com.superbiz.agent.harness.contract.DiagnosisDraft;
import com.superbiz.agent.harness.core.DiagnosisHarnessCore;
import com.superbiz.agent.harness.core.RunContext;
import org.springframework.ai.chat.messages.AssistantMessage;
import java.nio.charset.StandardCharsets;
import java.util.Objects;
public final class DiagnosisAgentUseCase {
public static final String RUN_CONTEXT_METADATA = "runContext";
private final DiagnosisHarnessCore core;
private final DiagnosisAgentFactory agentFactory;
private final ObjectMapper objectMapper;
private final ObjectReader draftReader;
private final DiagnosisAgentLimits limits;
public DiagnosisAgentUseCase(DiagnosisHarnessCore core,
DiagnosisAgentFactory agentFactory,
ObjectMapper objectMapper,
DiagnosisAgentLimits limits) {
this.core = Objects.requireNonNull(core, "core must not be null");
this.agentFactory = Objects.requireNonNull(agentFactory, "agentFactory must not be null");
this.objectMapper = Objects.requireNonNull(objectMapper, "objectMapper must not be null");
this.draftReader = objectMapper.readerFor(DiagnosisDraft.class)
.with(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES)
.with(DeserializationFeature.FAIL_ON_TRAILING_TOKENS);
this.limits = Objects.requireNonNull(limits, "limits must not be null");
}
public DiagnosisDraft execute(RunContext context, DiagnosisAgentInput input) {
Objects.requireNonNull(context, "context must not be null");
Objects.requireNonNull(input, "input must not be null");
core.checkActive(context);
checkLimit("query", utf8Bytes(input.query()), limits.maxQueryBytes());
if (input.previousTurn() != null) {
checkLimit("previous_turn", utf8Bytes(writeJson(input.previousTurn())),
limits.maxPreviousTurnBytes());
}
String inputJson = writeJson(input);
long inputBytes = utf8Bytes(inputJson);
checkLimit("input", inputBytes, limits.maxInputBytes());
core.reserveRunBytes(context, inputBytes);
RunnableConfig config = RunnableConfig.builder()
.threadId(context.runId())
.addMetadata("sessionId", context.sessionId())
.addMetadata("runId", context.runId())
.addMetadata(RUN_CONTEXT_METADATA, context)
.addMetadata("_stream_", false)
.build();
ReactAgent agent = agentFactory.create(context);
AssistantMessage response;
try {
response = agent.call(inputJson, config);
} catch (Exception e) {
throw new DiagnosisAgentOutputException("Diagnosis Agent execution failed", e);
}
core.checkActive(context);
String output = response == null ? null : response.getText();
if (output == null || output.isBlank()) {
throw new DiagnosisAgentOutputException("Diagnosis Agent returned an empty Draft");
}
long draftBytes = utf8Bytes(output);
checkLimit("draft", draftBytes, limits.maxDraftBytes());
core.reserveRunBytes(context, draftBytes);
try {
return draftReader.readValue(output);
} catch (JsonProcessingException e) {
throw new DiagnosisAgentOutputException("Diagnosis Agent returned an invalid Draft", e);
}
}
private String writeJson(Object value) {
try {
return objectMapper.writeValueAsString(value);
} catch (JsonProcessingException e) {
throw new IllegalArgumentException("Diagnosis Agent input is not serializable", e);
}
}
private static long utf8Bytes(String value) {
return value.getBytes(StandardCharsets.UTF_8).length;
}
private static void checkLimit(String boundary, long actual, long limit) {
if (actual > limit) {
throw new DiagnosisAgentLimitException(boundary, limit, actual);
}
}
}
@@ -0,0 +1,29 @@
package com.superbiz.agent.harness.agent;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.fasterxml.jackson.databind.node.ArrayNode;
import com.fasterxml.jackson.databind.node.ObjectNode;
import com.superbiz.agent.harness.contract.DiagnosisDraft;
import org.springframework.ai.converter.BeanOutputConverter;
import java.util.Objects;
final class DiagnosisDraftOutputSchema extends BeanOutputConverter<DiagnosisDraft> {
DiagnosisDraftOutputSchema(ObjectMapper objectMapper) {
super(DiagnosisDraft.class, Objects.requireNonNull(objectMapper, "objectMapper must not be null"));
}
@Override
protected void postProcessSchema(JsonNode schema) {
JsonNode conclusion = schema.path("properties").path("conclusion");
if (!(conclusion instanceof ObjectNode conclusionSchema)) {
throw new IllegalStateException("DiagnosisDraft schema has no conclusion property");
}
ArrayNode allowedTypes = conclusionSchema.arrayNode();
allowedTypes.add("object");
allowedTypes.add("null");
conclusionSchema.set("type", allowedTypes);
}
}
@@ -0,0 +1,10 @@
package com.superbiz.agent.harness.agent;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryResult;
@FunctionalInterface
public interface EvidenceToolInvoker {
ToolBoundaryResult invoke(RunContext context, String toolCallId, String arguments);
}
@@ -0,0 +1,95 @@
package com.superbiz.agent.harness.agent;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.tool.adapter.MysqlToolAdapter;
import com.superbiz.agent.harness.tool.adapter.QueryLogsToolAdapter;
import com.superbiz.agent.harness.tool.adapter.RagToolAdapter;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryResult;
import com.superbiz.agent.harness.tool.boundary.ToolCallRequestEnvelope;
import com.superbiz.agent.harness.tool.contract.AgentToolContracts;
import com.superbiz.agent.harness.tool.contract.MysqlToolRequest;
import com.superbiz.agent.harness.tool.contract.QueryLogsRequest;
import com.superbiz.agent.harness.tool.contract.RagToolRequest;
import org.springframework.ai.tool.ToolCallback;
import org.springframework.ai.tool.function.FunctionToolCallback;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
public final class HarnessEvidenceTools {
private final List<ToolCallback> callbacks;
private final Map<String, EvidenceToolInvoker> invokers;
public HarnessEvidenceTools(EvidenceToolInvoker ragInvoker,
EvidenceToolInvoker logsInvoker,
EvidenceToolInvoker mysqlInvoker) {
Map<String, EvidenceToolInvoker> registered = new LinkedHashMap<>();
registered.put(AgentToolContracts.LOOKUP_KNOWLEDGE,
Objects.requireNonNull(ragInvoker, "ragInvoker must not be null"));
registered.put(AgentToolContracts.QUERY_LOGS,
Objects.requireNonNull(logsInvoker, "logsInvoker must not be null"));
registered.put(AgentToolContracts.QUERY_MYSQL,
Objects.requireNonNull(mysqlInvoker, "mysqlInvoker must not be null"));
this.invokers = Map.copyOf(registered);
this.callbacks = List.of(
definition(AgentToolContracts.LOOKUP_KNOWLEDGE,
AgentToolContracts.LOOKUP_KNOWLEDGE_DESCRIPTION, RagToolRequest.class),
definition(AgentToolContracts.QUERY_LOGS,
AgentToolContracts.QUERY_LOGS_DESCRIPTION, QueryLogsRequest.class),
definition(AgentToolContracts.QUERY_MYSQL,
AgentToolContracts.QUERY_MYSQL_DESCRIPTION, MysqlToolRequest.class));
}
public static HarnessEvidenceTools fromAdapters(RagToolAdapter ragAdapter,
QueryLogsToolAdapter logsAdapter,
MysqlToolAdapter mysqlAdapter) {
Objects.requireNonNull(ragAdapter, "ragAdapter must not be null");
Objects.requireNonNull(logsAdapter, "logsAdapter must not be null");
Objects.requireNonNull(mysqlAdapter, "mysqlAdapter must not be null");
return new HarnessEvidenceTools(
bridge(AgentToolContracts.LOOKUP_KNOWLEDGE, ragAdapter::execute),
bridge(AgentToolContracts.QUERY_LOGS, logsAdapter::execute),
bridge(AgentToolContracts.QUERY_MYSQL, mysqlAdapter::execute));
}
public List<ToolCallback> callbacks() {
return callbacks;
}
public boolean supports(String toolName) {
return invokers.containsKey(toolName);
}
public ToolBoundaryResult invoke(RunContext context, String toolName,
String toolCallId, String arguments) {
EvidenceToolInvoker invoker = invokers.get(toolName);
if (invoker == null) {
throw new IllegalArgumentException("Unsupported evidence Tool: " + toolName);
}
return invoker.invoke(context, toolCallId, arguments);
}
private static EvidenceToolInvoker bridge(String toolName, AdapterCall adapter) {
return (context, toolCallId, arguments) -> adapter.execute(
context,
new ToolCallRequestEnvelope(
context.runId(), toolCallId, toolName, arguments, true, true));
}
private static <I> ToolCallback definition(String name, String description, Class<I> inputType) {
return FunctionToolCallback.<I, String>builder(name, ignored -> {
throw new IllegalStateException("Harness evidence Tools require the framework Tool interceptor");
})
.description(description)
.inputType(inputType)
.build();
}
@FunctionalInterface
private interface AdapterCall {
ToolBoundaryResult execute(RunContext context, ToolCallRequestEnvelope envelope);
}
}
@@ -0,0 +1,56 @@
package com.superbiz.agent.harness.agent;
import com.alibaba.cloud.ai.graph.agent.interceptor.ModelCallHandler;
import com.alibaba.cloud.ai.graph.agent.interceptor.ModelInterceptor;
import com.alibaba.cloud.ai.graph.agent.interceptor.ModelRequest;
import com.alibaba.cloud.ai.graph.agent.interceptor.ModelResponse;
import com.superbiz.agent.harness.core.DiagnosisHarnessCore;
import com.superbiz.agent.harness.core.RunContext;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.chat.model.ChatResponse;
import java.util.Objects;
public final class HarnessModelInterceptor extends ModelInterceptor {
private final DiagnosisHarnessCore core;
private final RunContext context;
public HarnessModelInterceptor(DiagnosisHarnessCore core, RunContext context) {
this.core = Objects.requireNonNull(core, "core must not be null");
this.context = Objects.requireNonNull(context, "context must not be null");
}
@Override
public String getName() {
return "harness_model_budget_interceptor";
}
@Override
public ModelResponse interceptModel(ModelRequest request, ModelCallHandler handler) {
Objects.requireNonNull(request, "request must not be null");
Objects.requireNonNull(handler, "handler must not be null");
core.beforeModelCall(context);
ModelResponse response = handler.call(request);
recordUsage(response == null ? null : response.getChatResponse());
core.checkActive(context);
return response;
}
private void recordUsage(ChatResponse response) {
if (response == null || response.getMetadata() == null) {
return;
}
Usage usage = response.getMetadata().getUsage();
if (usage == null) {
return;
}
long inputTokens = nonNegative(usage.getPromptTokens());
long outputTokens = nonNegative(usage.getCompletionTokens());
core.recordTokens(context, inputTokens, outputTokens);
}
private static long nonNegative(Integer value) {
return value == null || value < 0 ? 0L : value.longValue();
}
}
@@ -0,0 +1,69 @@
package com.superbiz.agent.harness.agent;
import com.alibaba.cloud.ai.graph.agent.interceptor.ToolCallHandler;
import com.alibaba.cloud.ai.graph.agent.interceptor.ToolCallRequest;
import com.alibaba.cloud.ai.graph.agent.interceptor.ToolCallResponse;
import com.alibaba.cloud.ai.graph.agent.interceptor.ToolInterceptor;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.contract.InvocationStatus;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryResult;
import java.util.LinkedHashMap;
import java.util.Map;
import java.util.Objects;
public final class HarnessToolInterceptor extends ToolInterceptor {
private final RunContext context;
private final HarnessEvidenceTools evidenceTools;
private final ObjectMapper objectMapper;
public HarnessToolInterceptor(RunContext context,
HarnessEvidenceTools evidenceTools,
ObjectMapper objectMapper) {
this.context = Objects.requireNonNull(context, "context must not be null");
this.evidenceTools = Objects.requireNonNull(evidenceTools, "evidenceTools must not be null");
this.objectMapper = Objects.requireNonNull(objectMapper, "objectMapper must not be null");
}
@Override
public String getName() {
return "harness_evidence_tool_interceptor";
}
@Override
public ToolCallResponse interceptToolCall(ToolCallRequest request, ToolCallHandler handler) {
Objects.requireNonNull(request, "request must not be null");
Objects.requireNonNull(handler, "handler must not be null");
if (!evidenceTools.supports(request.getToolName())) {
return handler.call(request);
}
ToolBoundaryResult result = evidenceTools.invoke(
context, request.getToolName(), request.getToolCallId(), request.getArguments());
if (result.status() == InvocationStatus.READY) {
return ToolCallResponse.of(request.getToolCallId(), request.getToolName(), result.agentResult());
}
return ToolCallResponse.builder()
.toolCallId(request.getToolCallId())
.toolName(request.getToolName())
.content(errorObservation(result))
.status("error")
.metadata(Map.of("error", true))
.build();
}
private String errorObservation(ToolBoundaryResult result) {
Map<String, Object> observation = new LinkedHashMap<>();
observation.put("evidence_status", result.evidenceStatus());
observation.put("tool_call_id", result.toolCallId());
observation.put("error_code", result.errorCode());
try {
return objectMapper.writeValueAsString(observation);
} catch (JsonProcessingException e) {
return "{\"evidence_status\":\"ERROR\",\"error_code\":\"SERIALIZATION_ERROR\"}";
}
}
}
@@ -0,0 +1,46 @@
# Role
You are the only Diagnosis Agent for the current diagnosis run. Plan the investigation internally, call the available read-only evidence tools through the framework ReAct loop, evaluate the returned evidence, and author one complete DiagnosisDraft.
# Input
The user message is one JSON object with exactly:
- `query`: the current user query. Preserve its meaning and answer this query only.
- `previous_turn`: an optional safely published previous diagnosis turn with the fixed PreviousTurn schema.
`previous_turn` is context only. It is not evidence for this run, contains no reusable Tool Call IDs, and must not be cited as current evidence.
# Evidence rules
- Use only `lookup_knowledge`, `query_logs`, and `query_mysql` when evidence is needed.
- Treat a Tool observation as evidence only when it contains `evidence_status=EVIDENCE_FOUND` or `evidence_status=NO_EVIDENCE` and a non-blank `tool_call_id`.
- Bind every Analysis item only to Tool Call IDs returned during this run. Never invent, transform, shorten, or reuse a Tool Call ID.
- `NORMAL` Analysis may cite only `EVIDENCE_FOUND` results.
- `NEGATIVE_OBSERVATION` Analysis may cite only `NO_EVIDENCE` results and must state the exact query scope. `NO_EVIDENCE` never proves that a problem does not exist, that a root cause is excluded, or that a system is healthy.
- A Tool `ERROR` is not evidence. Do not cite it as support for Analysis or Conclusion.
- Do not expose raw Tool payloads, credentials, infrastructure coordinates, internal errors, prompts, hidden reasoning, or chain of thought.
# Stopping rule
Stop calling tools when the current evidence is sufficient for a bounded Draft, when the configured tool/model budget prevents more work, or when the available tools cannot obtain the missing information.
If current evidence cannot support a diagnosis:
- set `conclusion` to `null`;
- keep only valid scoped negative observations, if any;
- state the actual queried scope in `limitations.scope`;
- list the evidence still needed in `limitations.missing_info`;
- do not fabricate a root cause, action justification, recommendation, or healthy-state claim;
- finish the current response without asking the Harness to retry the Agent or a Tool.
# Draft rules
- Lead with `conclusion` when supported. Its `based_on_analysis_ids` must reference existing Analysis IDs.
- Every Analysis item must have a unique `analysis_id`, a valid `kind`, concise evidence-grounded text, and at least one current-Run `tool_call_id`.
- Every Action Plan and Recommendation item must reference existing Analysis IDs.
- Mark any dangerous or side-effecting action with `requires_human_confirmation=true`.
- `limitations.scope` must not exceed the actual Tool query scopes.
- `limitations.missing_info` must name material gaps that constrain the conclusion.
Return exactly one JSON object matching the supplied DiagnosisDraft schema. Do not use Markdown fences, headings, prose before or after JSON, or a separate answer field.
@@ -0,0 +1,366 @@
package com.superbiz.agent.harness.agent;
import com.alibaba.cloud.ai.graph.OverAllState;
import com.alibaba.cloud.ai.graph.RunnableConfig;
import com.alibaba.cloud.ai.graph.agent.hook.HookPosition;
import com.alibaba.cloud.ai.graph.agent.hook.HookPositions;
import com.alibaba.cloud.ai.graph.agent.hook.ModelHook;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.contract.AnalysisKind;
import com.superbiz.agent.harness.contract.DiagnosisDraft;
import com.superbiz.agent.harness.contract.EvidenceStatus;
import com.superbiz.agent.harness.contract.PreviousTurn;
import com.superbiz.agent.harness.contract.SourceDocument;
import com.superbiz.agent.harness.core.DiagnosisHarnessCore;
import com.superbiz.agent.harness.core.RunBudgetLimits;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.core.RunState;
import com.superbiz.agent.harness.retry.HarnessRetryPolicies;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryErrorCode;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryResult;
import com.superbiz.agent.harness.tool.contract.AgentToolContracts;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.ToolResponseMessage;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
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 java.time.Clock;
import java.time.Duration;
import java.util.ArrayList;
import java.util.List;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.atomic.AtomicInteger;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertNull;
import static org.junit.jupiter.api.Assertions.assertThrows;
import static org.junit.jupiter.api.Assertions.assertTrue;
class DiagnosisAgentUseCaseTest {
private static final DiagnosisAgentLimits LARGE_LIMITS =
new DiagnosisAgentLimits(20_000, 20_000, 40_000, 40_000);
private final ObjectMapper objectMapper = new ObjectMapper();
@Test
void singleAgentRunsFrameworkToolLoopAndReturnsTypedDraft() {
DiagnosisHarnessCore core = core(new RunBudgetLimits(4, 3, 2,
100, 100, 200, 100_000));
RunContext context = core.startRun("session-agent", "run-agent");
AtomicInteger toolCalls = new AtomicInteger();
HarnessEvidenceTools tools = tools((runContext, id, arguments) -> {
core.beforeToolCall(runContext, AgentToolContracts.LOOKUP_KNOWLEDGE);
toolCalls.incrementAndGet();
return ToolBoundaryResult.ready(id, evidence(id), EvidenceStatus.EVIDENCE_FOUND);
});
ScriptedChatModel model = new ScriptedChatModel(3, 2,
toolCall("call-rag-1", "{\"query\":\"refund timeout\"}"),
new AssistantMessage(supportedDraft("call-rag-1")));
CapturingAuditHook audit = new CapturingAuditHook();
DiagnosisAgentUseCase useCase = useCase(core, model, tools, LARGE_LIMITS, audit);
PreviousTurn previous = new PreviousTurn(
"上一轮支付服务为什么超时?", "支付服务连接池已耗尽", "payment-service",
List.of("未覆盖退款服务"), List.of(new SourceDocument("doc-1", "连接池手册")));
DiagnosisDraft draft = useCase.execute(
context, new DiagnosisAgentInput("那退款服务呢?", previous));
assertEquals("退款服务出现连接池等待", draft.conclusion().text());
assertEquals(AnalysisKind.NORMAL, draft.analysis().get(0).kind());
assertEquals(List.of("call-rag-1"), draft.analysis().get(0).toolCallIds());
assertEquals(2, model.calls());
assertEquals(1, toolCalls.get());
assertEquals(2, context.budget().snapshot().modelCalls());
assertEquals(1, context.budget().snapshot().toolCalls());
assertEquals(10, context.budget().snapshot().totalTokens());
assertTrue(model.prompts().get(0).contains("那退款服务呢?"));
assertTrue(model.prompts().get(0).contains("支付服务连接池已耗尽"));
assertTrue(model.instructions().get(1).stream()
.filter(ToolResponseMessage.class::isInstance)
.map(ToolResponseMessage.class::cast)
.flatMap(message -> message.getResponses().stream())
.anyMatch(response -> response.id().equals("call-rag-1")
&& response.responseData().contains("call-rag-1")));
assertEquals(List.of("session-agent", "session-agent"), audit.sessionIds);
assertEquals(List.of("run-agent", "run-agent"), audit.runIds);
assertEquals(RunState.RUNNING, context.lifecycle().state());
}
@Test
void noEvidenceStopsWithNullConclusionAndNoRetry() {
DiagnosisHarnessCore core = core(defaultBudget());
RunContext context = core.startRun("session-no-evidence", "run-no-evidence");
AtomicInteger toolCalls = new AtomicInteger();
HarnessEvidenceTools tools = tools((runContext, id, arguments) -> {
core.beforeToolCall(runContext, AgentToolContracts.LOOKUP_KNOWLEDGE);
toolCalls.incrementAndGet();
return ToolBoundaryResult.ready(id, noEvidence(id), EvidenceStatus.NO_EVIDENCE);
});
ScriptedChatModel model = new ScriptedChatModel(1, 1,
toolCall("call-empty-1", "{\"query\":\"refund timeout\"}"),
new AssistantMessage(noEvidenceDraft("call-empty-1")));
DiagnosisDraft draft = useCase(core, model, tools, LARGE_LIMITS)
.execute(context, new DiagnosisAgentInput("退款为什么超时?", null));
assertNull(draft.conclusion());
assertEquals(AnalysisKind.NEGATIVE_OBSERVATION, draft.analysis().get(0).kind());
assertTrue(draft.limitations().missingInfo().contains("退款服务运行日志"));
assertEquals(2, model.calls());
assertEquals(1, toolCalls.get());
}
@Test
void fencedOutputFailsClosedWithoutAgentRetry() {
DiagnosisHarnessCore core = core(defaultBudget());
ScriptedChatModel model = new ScriptedChatModel(1, 1,
new AssistantMessage("```json\n" + supportedDraft("call-1") + "\n```"));
RunContext context = core.startRun("session-fenced", "run-fenced");
assertThrows(DiagnosisAgentOutputException.class, () ->
useCase(core, model, tools(errorInvoker()), LARGE_LIMITS)
.execute(context, new DiagnosisAgentInput("诊断超时", null)));
assertEquals(1, model.calls());
assertEquals(1, context.budget().snapshot().modelCalls());
}
@Test
void modelCallBudgetStopsNextReactRoundBeforeChatModel() {
DiagnosisHarnessCore core = core(new RunBudgetLimits(1, 2, 2,
100, 100, 200, 100_000));
RunContext context = core.startRun("session-model-budget", "run-model-budget");
HarnessEvidenceTools tools = tools((runContext, id, arguments) -> {
core.beforeToolCall(runContext, AgentToolContracts.LOOKUP_KNOWLEDGE);
return ToolBoundaryResult.ready(id, evidence(id), EvidenceStatus.EVIDENCE_FOUND);
});
ScriptedChatModel model = new ScriptedChatModel(1, 1,
toolCall("call-budget-1", "{\"query\":\"timeout\"}"),
new AssistantMessage(supportedDraft("call-budget-1")));
assertThrows(DiagnosisAgentOutputException.class, () ->
useCase(core, model, tools, LARGE_LIMITS)
.execute(context, new DiagnosisAgentInput("诊断超时", null)));
assertEquals(1, model.calls());
assertEquals(RunState.BUDGET_EXHAUSTED, context.lifecycle().state());
}
@Test
void actualTokenUsageExhaustionStopsDraftWithoutRetry() {
DiagnosisHarnessCore core = core(new RunBudgetLimits(2, 2, 2,
10, 10, 10, 100_000));
RunContext context = core.startRun("session-token-budget", "run-token-budget");
ScriptedChatModel model = new ScriptedChatModel(6, 6,
new AssistantMessage(supportedDraft("call-unused")));
assertThrows(DiagnosisAgentOutputException.class, () ->
useCase(core, model, tools(errorInvoker()), LARGE_LIMITS)
.execute(context, new DiagnosisAgentInput("诊断超时", null)));
assertEquals(1, model.calls());
assertEquals(12, context.budget().snapshot().totalTokens());
assertEquals(RunState.BUDGET_EXHAUSTED, context.lifecycle().state());
}
@Test
void contextLimitRejectsOriginalQueryBeforeModelCall() {
DiagnosisHarnessCore core = core(defaultBudget());
RunContext context = core.startRun("session-context-budget", "run-context-budget");
ScriptedChatModel model = new ScriptedChatModel(1, 1,
new AssistantMessage(supportedDraft("call-unused")));
DiagnosisAgentLimits limits = new DiagnosisAgentLimits(4, 100, 200, 1_000);
DiagnosisAgentLimitException failure = assertThrows(DiagnosisAgentLimitException.class, () ->
useCase(core, model, tools(errorInvoker()), limits)
.execute(context, new DiagnosisAgentInput("超过四个字节", null)));
assertEquals("query", failure.boundary());
assertEquals(0, model.calls());
assertEquals(0, context.budget().snapshot().modelCalls());
}
@Test
void draftLimitRejectsOversizedOutputWithoutSecondCall() {
DiagnosisHarnessCore core = core(defaultBudget());
RunContext context = core.startRun("session-draft-budget", "run-draft-budget");
ScriptedChatModel model = new ScriptedChatModel(1, 1,
new AssistantMessage(supportedDraft("call-unused")));
DiagnosisAgentLimits limits = new DiagnosisAgentLimits(1_000, 1_000, 2_000, 20);
DiagnosisAgentLimitException failure = assertThrows(DiagnosisAgentLimitException.class, () ->
useCase(core, model, tools(errorInvoker()), limits)
.execute(context, new DiagnosisAgentInput("诊断超时", null)));
assertEquals("draft", failure.boundary());
assertEquals(1, model.calls());
}
@Test
void rejectsInvalidLimitConfiguration() {
assertThrows(IllegalArgumentException.class,
() -> new DiagnosisAgentLimits(0, 1, 1, 1));
assertThrows(IllegalArgumentException.class,
() -> new DiagnosisAgentInput(" ", null));
}
@Test
void generatedDraftSchemaAllowsExplicitNullConclusion() throws Exception {
JsonNode schema = objectMapper.readTree(new DiagnosisDraftOutputSchema(objectMapper).getJsonSchema());
JsonNode conclusionTypes = schema.path("properties").path("conclusion").path("type");
assertTrue(conclusionTypes.isArray());
assertTrue(java.util.stream.StreamSupport.stream(conclusionTypes.spliterator(), false)
.map(JsonNode::asText)
.anyMatch("object"::equals));
assertTrue(java.util.stream.StreamSupport.stream(conclusionTypes.spliterator(), false)
.map(JsonNode::asText)
.anyMatch("null"::equals));
}
private DiagnosisAgentUseCase useCase(DiagnosisHarnessCore core,
ChatModel model,
HarnessEvidenceTools tools,
DiagnosisAgentLimits limits,
ModelHook... hooks) {
DiagnosisAgentFactory factory = new DiagnosisAgentFactory(
model, core, tools, objectMapper, List.of(hooks));
return new DiagnosisAgentUseCase(core, factory, objectMapper, limits);
}
private HarnessEvidenceTools tools(EvidenceToolInvoker ragInvoker) {
EvidenceToolInvoker unused = errorInvoker();
return new HarnessEvidenceTools(ragInvoker, unused, unused);
}
private EvidenceToolInvoker errorInvoker() {
return (context, id, arguments) ->
ToolBoundaryResult.error(id, ToolBoundaryErrorCode.INVALID_REQUEST);
}
private DiagnosisHarnessCore core(RunBudgetLimits budget) {
return new DiagnosisHarnessCore(
Clock.systemUTC(), () -> "unused-run", Duration.ofMinutes(5),
budget, HarnessRetryPolicies.strict());
}
private RunBudgetLimits defaultBudget() {
return new RunBudgetLimits(4, 3, 2, 100, 100, 200, 100_000);
}
private AssistantMessage toolCall(String id, String arguments) {
return AssistantMessage.builder()
.content("")
.toolCalls(List.of(new AssistantMessage.ToolCall(
id, "function", AgentToolContracts.LOOKUP_KNOWLEDGE, arguments)))
.build();
}
private String evidence(String id) {
return "{\"evidence_status\":\"EVIDENCE_FOUND\",\"tool_call_id\":\"" + id
+ "\",\"query\":\"refund timeout\",\"evidence\":[{\"document_id\":\"doc-1\","
+ "\"source\":\"runbook\",\"title\":\"退款连接池\",\"breadcrumb\":[\"运行\"],"
+ "\"excerpt\":\"active=50 max=50\"}],\"returned_count\":1,\"truncated\":false}";
}
private String noEvidence(String id) {
return "{\"evidence_status\":\"NO_EVIDENCE\",\"tool_call_id\":\"" + id
+ "\",\"query\":\"refund timeout\",\"evidence\":[],\"returned_count\":0,\"truncated\":false}";
}
private String supportedDraft(String toolCallId) {
return """
{
"conclusion":{"text":"退款服务出现连接池等待","based_on_analysis_ids":["a-1"]},
"analysis":[{"analysis_id":"a-1","kind":"NORMAL","text":"连接池 active 达到上限","tool_call_ids":["%s"]}],
"action_plan":[{"action":"检查连接归还路径","based_on_analysis_ids":["a-1"],"requires_human_confirmation":false}],
"recommendations":[{"text":"增加连接池饱和告警","based_on_analysis_ids":["a-1"]}],
"limitations":{"scope":"refund-service","missing_info":[]}
}
""".formatted(toolCallId);
}
private String noEvidenceDraft(String toolCallId) {
return """
{
"conclusion":null,
"analysis":[{"analysis_id":"a-1","kind":"NEGATIVE_OBSERVATION","text":"知识库查询范围内未找到退款超时证据","tool_call_ids":["%s"]}],
"action_plan":[],
"recommendations":[],
"limitations":{"scope":"lookup query=refund timeout","missing_info":["退款服务运行日志"]}
}
""".formatted(toolCallId);
}
private static final class ScriptedChatModel implements ChatModel {
private final int promptTokens;
private final int completionTokens;
private final List<AssistantMessage> responses;
private final List<String> prompts = new ArrayList<>();
private final List<List<Message>> instructions = new ArrayList<>();
private final AtomicInteger calls = new AtomicInteger();
private ScriptedChatModel(int promptTokens, int completionTokens,
AssistantMessage... responses) {
this.promptTokens = promptTokens;
this.completionTokens = completionTokens;
this.responses = List.of(responses);
}
@Override
public ChatResponse call(Prompt prompt) {
prompts.add(prompt.getContents());
instructions.add(List.copyOf(prompt.getInstructions()));
int index = calls.getAndIncrement();
if (index >= responses.size()) {
throw new AssertionError("unexpected model retry");
}
ChatResponseMetadata metadata = ChatResponseMetadata.builder()
.usage(new DefaultUsage(promptTokens, completionTokens))
.build();
return new ChatResponse(List.of(new Generation(responses.get(index))), metadata);
}
int calls() {
return calls.get();
}
List<String> prompts() {
return prompts;
}
List<List<Message>> instructions() {
return instructions;
}
}
@HookPositions({HookPosition.BEFORE_MODEL})
private static final class CapturingAuditHook extends ModelHook {
private final List<String> sessionIds = new ArrayList<>();
private final List<String> runIds = new ArrayList<>();
@Override
public String getName() {
return "capturing_agent_step_audit";
}
@Override
public CompletableFuture<java.util.Map<String, Object>> beforeModel(
OverAllState state, RunnableConfig config) {
sessionIds.add(config.metadata("sessionId").map(Object::toString).orElseThrow());
runIds.add(config.metadata("runId").map(Object::toString).orElseThrow());
return CompletableFuture.completedFuture(java.util.Map.of());
}
}
}
@@ -0,0 +1,184 @@
package com.superbiz.agent.harness.agent;
import com.alibaba.cloud.ai.graph.agent.interceptor.ToolCallRequest;
import com.alibaba.cloud.ai.graph.agent.interceptor.ToolCallResponse;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.contract.EvidenceStatus;
import com.superbiz.agent.harness.core.DiagnosisHarnessCore;
import com.superbiz.agent.harness.core.RunBudgetLimits;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.retry.HarnessRetryPolicies;
import com.superbiz.agent.harness.tool.adapter.MysqlToolAdapter;
import com.superbiz.agent.harness.tool.adapter.QueryLogsToolAdapter;
import com.superbiz.agent.harness.tool.adapter.RagToolAdapter;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryErrorCode;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryResult;
import com.superbiz.agent.harness.tool.boundary.ToolCallRequestEnvelope;
import com.superbiz.agent.harness.tool.contract.AgentToolContracts;
import org.junit.jupiter.api.Test;
import org.mockito.ArgumentCaptor;
import org.springframework.ai.tool.ToolCallback;
import org.springframework.ai.tool.execution.ToolExecutionException;
import java.time.Clock;
import java.time.Duration;
import java.util.List;
import java.util.concurrent.atomic.AtomicInteger;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertThrows;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
class HarnessToolInterceptorTest {
private final ObjectMapper objectMapper = new ObjectMapper();
@Test
void exposesOnlyFrozenToolDefinitionsWithoutAgentSuppliedCallId() {
HarnessEvidenceTools tools = fakeTools((context, id, arguments) -> ready(id));
List<ToolCallback> callbacks = tools.callbacks();
assertEquals(List.of(
AgentToolContracts.LOOKUP_KNOWLEDGE,
AgentToolContracts.QUERY_LOGS,
AgentToolContracts.QUERY_MYSQL),
callbacks.stream().map(callback -> callback.getToolDefinition().name()).toList());
callbacks.forEach(callback -> {
assertFalse(callback.getToolDefinition().inputSchema().contains("tool_call_id"));
assertFalse(callback.getToolDefinition().description().isBlank());
assertThrows(ToolExecutionException.class, () -> callback.call("{}"));
});
}
@Test
void adapterBridgeUsesExactFrameworkIdAndRunIdentity() {
RagToolAdapter ragAdapter = mock(RagToolAdapter.class);
QueryLogsToolAdapter logsAdapter = mock(QueryLogsToolAdapter.class);
MysqlToolAdapter mysqlAdapter = mock(MysqlToolAdapter.class);
RunContext context = context(3, 3, 3);
when(ragAdapter.execute(eq(context), any())).thenAnswer(invocation -> {
ToolCallRequestEnvelope envelope = invocation.getArgument(1);
return ready(envelope.toolCallId());
});
HarnessEvidenceTools tools = HarnessEvidenceTools.fromAdapters(ragAdapter, logsAdapter, mysqlAdapter);
HarnessToolInterceptor interceptor = new HarnessToolInterceptor(context, tools, objectMapper);
ToolCallResponse response = interceptor.interceptToolCall(
request(AgentToolContracts.LOOKUP_KNOWLEDGE, "framework-call-7", "{\"query\":\"timeout\"}"),
ignored -> {
throw new AssertionError("registered Tool must not bypass Harness interceptor");
});
ArgumentCaptor<ToolCallRequestEnvelope> envelope = ArgumentCaptor.forClass(ToolCallRequestEnvelope.class);
verify(ragAdapter).execute(eq(context), envelope.capture());
assertEquals("framework-call-7", envelope.getValue().toolCallId());
assertEquals(context.runId(), envelope.getValue().runId());
assertEquals(AgentToolContracts.LOOKUP_KNOWLEDGE, envelope.getValue().toolName());
assertEquals("framework-call-7", response.getToolCallId());
assertTrue(response.getResult().contains("framework-call-7"));
}
@Test
void returnsStableErrorObservationWithoutRawFailureDetail() throws Exception {
HarnessEvidenceTools tools = fakeTools((context, id, arguments) ->
ToolBoundaryResult.error(id, ToolBoundaryErrorCode.TOOL_EXECUTION_ERROR));
HarnessToolInterceptor interceptor = new HarnessToolInterceptor(
context(3, 3, 3), tools, objectMapper);
ToolCallResponse response = interceptor.interceptToolCall(
request(AgentToolContracts.LOOKUP_KNOWLEDGE, "call-error", "{\"query\":\"secret\"}"),
ignored -> null);
JsonNode observation = objectMapper.readTree(response.getResult());
assertTrue(response.isError());
assertEquals("ERROR", observation.path("evidence_status").asText());
assertEquals("call-error", observation.path("tool_call_id").asText());
assertEquals("TOOL_EXECUTION_ERROR", observation.path("error_code").asText());
assertFalse(response.getResult().contains("secret"));
assertFalse(response.getResult().contains("raw_response"));
}
@Test
void delegatesUnknownToolWithoutCreatingEvidenceInvocation() {
AtomicInteger evidenceCalls = new AtomicInteger();
HarnessEvidenceTools tools = fakeTools((context, id, arguments) -> {
evidenceCalls.incrementAndGet();
return ready(id);
});
AtomicInteger handlerCalls = new AtomicInteger();
HarnessToolInterceptor interceptor = new HarnessToolInterceptor(
context(3, 3, 3), tools, objectMapper);
ToolCallResponse response = interceptor.interceptToolCall(
request("unknown_tool", "call-unknown", "{}"),
request -> {
handlerCalls.incrementAndGet();
return ToolCallResponse.error(request.getToolCallId(), request.getToolName(), "not registered");
});
assertTrue(response.isError());
assertEquals(1, handlerCalls.get());
assertEquals(0, evidenceCalls.get());
}
@Test
void oneFrameworkActionConsumesOneToolBudgetWithoutRetry() {
DiagnosisHarnessCore core = core(3, 1, 1, 10_000);
RunContext context = core.startRun("session-tool-budget", "run-tool-budget");
AtomicInteger invocations = new AtomicInteger();
HarnessEvidenceTools tools = fakeTools((runContext, id, arguments) -> {
core.beforeToolCall(runContext, AgentToolContracts.LOOKUP_KNOWLEDGE);
invocations.incrementAndGet();
return ready(id);
});
HarnessToolInterceptor interceptor = new HarnessToolInterceptor(context, tools, objectMapper);
interceptor.interceptToolCall(
request(AgentToolContracts.LOOKUP_KNOWLEDGE, "call-1", "{\"query\":\"one\"}"),
ignored -> null);
assertEquals(1, invocations.get());
assertEquals(1, context.budget().snapshot().toolCalls());
}
private HarnessEvidenceTools fakeTools(EvidenceToolInvoker ragInvoker) {
EvidenceToolInvoker unused = (context, id, arguments) ->
ToolBoundaryResult.error(id, ToolBoundaryErrorCode.INVALID_REQUEST);
return new HarnessEvidenceTools(ragInvoker, unused, unused);
}
private ToolCallRequest request(String toolName, String toolCallId, String arguments) {
return new ToolCallRequest(toolName, arguments, toolCallId, java.util.Map.of());
}
private ToolBoundaryResult ready(String toolCallId) {
return ToolBoundaryResult.ready(toolCallId,
"{\"evidence_status\":\"EVIDENCE_FOUND\",\"tool_call_id\":\""
+ toolCallId + "\",\"evidence\":[]}",
EvidenceStatus.EVIDENCE_FOUND);
}
private RunContext context(int maxModelCalls, int maxToolCalls, int maxCallsPerTool) {
return core(maxModelCalls, maxToolCalls, maxCallsPerTool, 10_000)
.startRun("session-tool", "run-tool");
}
private DiagnosisHarnessCore core(int maxModelCalls, int maxToolCalls,
int maxCallsPerTool, long maxTokens) {
return new DiagnosisHarnessCore(
Clock.systemUTC(),
() -> "unused-run",
Duration.ofMinutes(5),
new RunBudgetLimits(maxModelCalls, maxToolCalls, maxCallsPerTool,
maxTokens, maxTokens, maxTokens, 100_000),
HarnessRetryPolicies.strict());
}
}