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\"}";
}
}
}