feat(harness): add chat application use case

This commit is contained in:
zhuyongxin
2026-07-22 00:57:39 +08:00
parent ee0949d464
commit f8809cb7dd
56 changed files with 2815 additions and 6 deletions
@@ -0,0 +1,7 @@
package com.superbiz.agent.harness.application;
public sealed interface ChatApplicationContent permits
SystemChatContent, KnowledgeContent, DiagnosisContent, FallbackContent {
ChatContentType contentType();
}
@@ -0,0 +1,22 @@
package com.superbiz.agent.harness.application;
import java.util.Objects;
public final class ChatApplicationException extends RuntimeException {
private final ChatFailureCode code;
public ChatApplicationException(ChatFailureCode code, String message) {
super(message);
this.code = Objects.requireNonNull(code, "code must not be null");
}
public ChatApplicationException(ChatFailureCode code, String message, Throwable cause) {
super(message, cause);
this.code = Objects.requireNonNull(code, "code must not be null");
}
public ChatFailureCode code() {
return code;
}
}
@@ -0,0 +1,20 @@
package com.superbiz.agent.harness.application;
public interface ChatApplicationObserver {
void onStarted(ChatRunControl runControl);
void onStatus(ChatApplicationStatus status);
static ChatApplicationObserver noop() {
return new ChatApplicationObserver() {
@Override
public void onStarted(ChatRunControl runControl) {
}
@Override
public void onStatus(ChatApplicationStatus status) {
}
};
}
}
@@ -0,0 +1,17 @@
package com.superbiz.agent.harness.application;
import com.fasterxml.jackson.annotation.JsonProperty;
public record ChatApplicationRequest(
@JsonProperty("query") String query,
@JsonProperty("session_id") String sessionId) {
public ChatApplicationRequest {
if (query == null || query.isBlank()) {
throw new IllegalArgumentException("query must not be blank");
}
if (sessionId != null && sessionId.isBlank()) {
sessionId = null;
}
}
}
@@ -0,0 +1,39 @@
package com.superbiz.agent.harness.application;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.superbiz.agent.harness.contract.IntentType;
import com.superbiz.agent.harness.contract.ReleaseOutcome;
import java.util.Objects;
public record ChatApplicationResult(
@JsonProperty("session_id") String sessionId,
@JsonProperty("run_id") String runId,
@JsonProperty("intent") IntentType intent,
@JsonProperty("outcome") ReleaseOutcome outcome,
@JsonProperty("content_type") ChatContentType contentType,
@JsonProperty("content") ChatApplicationContent content) {
public ChatApplicationResult {
requireText(sessionId, "sessionId");
requireText(runId, "runId");
Objects.requireNonNull(intent, "intent must not be null");
Objects.requireNonNull(outcome, "outcome must not be null");
Objects.requireNonNull(content, "content must not be null");
if (outcome != ReleaseOutcome.SUCCESS && outcome != ReleaseOutcome.FALLBACK) {
throw new IllegalArgumentException("application result must be SUCCESS or FALLBACK");
}
if (contentType != content.contentType()) {
throw new IllegalArgumentException("contentType does not match content");
}
if (outcome == ReleaseOutcome.FALLBACK && contentType != ChatContentType.SAFE_FALLBACK) {
throw new IllegalArgumentException("fallback outcome requires safe fallback content");
}
}
private static void requireText(String value, String name) {
if (value == null || value.isBlank()) {
throw new IllegalArgumentException(name + " must not be blank");
}
}
}
@@ -0,0 +1,20 @@
package com.superbiz.agent.harness.application;
public enum ChatApplicationStatus {
ROUTING("正在识别请求类型"),
SYSTEM_RESPONDING("正在生成回答"),
KNOWLEDGE_SEARCHING("正在查询知识库"),
KNOWLEDGE_ANSWERING("正在整理知识答案"),
DIAGNOSIS_RUNNING("正在收集诊断证据"),
SAFETY_VALIDATING("正在进行安全校验");
private final String message;
ChatApplicationStatus(String message) {
this.message = message;
}
public String message() {
return message;
}
}
@@ -0,0 +1,259 @@
package com.superbiz.agent.harness.application;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.application.persistence.ChatRunStore;
import com.superbiz.agent.harness.application.persistence.RoutingHistory;
import com.superbiz.agent.harness.application.routing.IntentRouterInput;
import com.superbiz.agent.harness.application.routing.IntentRoutingException;
import com.superbiz.agent.harness.contract.IntentType;
import com.superbiz.agent.harness.contract.PreviousTurn;
import com.superbiz.agent.harness.contract.PublishedResult;
import com.superbiz.agent.harness.contract.ReleaseOutcome;
import com.superbiz.agent.harness.core.DiagnosisHarnessCore;
import com.superbiz.agent.harness.core.RunCancellationReason;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.core.RunState;
import com.superbiz.agent.harness.retry.RetryExecutionException;
import java.util.Objects;
import java.util.Optional;
import java.util.function.Supplier;
import java.util.regex.Pattern;
public final class ChatApplicationUseCase {
private static final Pattern SAFE_ID = Pattern.compile("[A-Za-z0-9][A-Za-z0-9._-]{0,63}");
private final DiagnosisHarnessCore core;
private final Supplier<String> sessionIdSupplier;
private final ChatRunStore runStore;
private final IntentRouting router;
private final SystemChatOperation systemChat;
private final KnowledgeQueryOperation knowledgeQuery;
private final DiagnosisOperation diagnosis;
private final ObjectMapper objectMapper;
public ChatApplicationUseCase(DiagnosisHarnessCore core,
Supplier<String> sessionIdSupplier,
ChatRunStore runStore,
IntentRouting router,
SystemChatOperation systemChat,
KnowledgeQueryOperation knowledgeQuery,
DiagnosisOperation diagnosis,
ObjectMapper objectMapper) {
this.core = Objects.requireNonNull(core, "core must not be null");
this.sessionIdSupplier = Objects.requireNonNull(
sessionIdSupplier, "sessionIdSupplier must not be null");
this.runStore = Objects.requireNonNull(runStore, "runStore must not be null");
this.router = Objects.requireNonNull(router, "router must not be null");
this.systemChat = Objects.requireNonNull(systemChat, "systemChat must not be null");
this.knowledgeQuery = Objects.requireNonNull(knowledgeQuery, "knowledgeQuery must not be null");
this.diagnosis = Objects.requireNonNull(diagnosis, "diagnosis must not be null");
this.objectMapper = Objects.requireNonNull(objectMapper, "objectMapper must not be null");
}
public ChatApplicationResult execute(ChatApplicationRequest request) {
return execute(request, ChatApplicationObserver.noop());
}
public ChatApplicationResult execute(ChatApplicationRequest request,
ChatApplicationObserver observer) {
Objects.requireNonNull(request, "request must not be null");
Objects.requireNonNull(observer, "observer must not be null");
String sessionId = resolveSessionId(request.sessionId());
Optional<RoutingHistory> history;
Optional<PreviousTurn> previousTurn;
try {
history = runStore.findLatestRoutingHistory(sessionId);
previousTurn = runStore.findPreviousTurn(sessionId);
} catch (RuntimeException exception) {
throw new ChatApplicationException(
ChatFailureCode.RUN_PERSISTENCE_FAILED,
"无法读取会话上下文,请稍后重试", exception);
}
RunContext context = core.startRun(sessionId);
long startedNanos = System.nanoTime();
IntentType intent = null;
try {
persistStart(context, request.query());
observer.onStarted(new CoreRunControl(core, context));
observer.onStatus(ChatApplicationStatus.ROUTING);
intent = router.route(context, new IntentRouterInput(
request.query(),
history.map(RoutingHistory::intent).orElse(null),
history.map(RoutingHistory::userQuery).orElse(null)));
persistIntent(context.runId(), intent);
PathResult path = executePath(
intent, context, request.query(), previousTurn.orElse(null), observer);
core.checkActive(context);
core.completeSuccess(context);
String safeJson = write(path.content());
persistFinish(context, intent, path.outcome(), safeJson,
path.publishedResult(), durationMillis(startedNanos));
return new ChatApplicationResult(
context.sessionId(), context.runId(), intent, path.outcome(),
path.content().contentType(), path.content());
} catch (RuntimeException exception) {
ReleaseOutcome terminal = terminalOutcome(context);
try {
runStore.finish(context, intent, terminal, null, null,
durationMillis(startedNanos));
} catch (RuntimeException persistenceFailure) {
exception.addSuppressed(persistenceFailure);
}
throw safeFailure(exception, terminal);
}
}
private PathResult executePath(IntentType intent,
RunContext context,
String query,
PreviousTurn previousTurn,
ChatApplicationObserver observer) {
return switch (intent) {
case SYSTEM_CHAT -> {
observer.onStatus(ChatApplicationStatus.SYSTEM_RESPONDING);
yield new PathResult(
ReleaseOutcome.SUCCESS, systemChat.execute(context, query), null);
}
case KNOWLEDGE_QUERY -> {
observer.onStatus(ChatApplicationStatus.KNOWLEDGE_SEARCHING);
observer.onStatus(ChatApplicationStatus.KNOWLEDGE_ANSWERING);
yield new PathResult(
ReleaseOutcome.SUCCESS, knowledgeQuery.execute(context, query), null);
}
case DIAGNOSIS -> {
DiagnosisExecutionResult result = diagnosis.execute(
context, query, previousTurn, observer::onStatus);
yield new PathResult(result.outcome(), result.content(), result.publishedResult());
}
};
}
private void persistStart(RunContext context, String query) {
try {
runStore.start(context, query);
} catch (RuntimeException exception) {
throw new ChatApplicationException(
ChatFailureCode.RUN_PERSISTENCE_FAILED,
"无法创建诊断运行记录,请稍后重试", exception);
}
}
private void persistIntent(String runId, IntentType intent) {
try {
runStore.markIntent(runId, intent);
} catch (RuntimeException exception) {
throw new ChatApplicationException(
ChatFailureCode.RUN_PERSISTENCE_FAILED,
"无法记录请求类型,请稍后重试", exception);
}
}
private void persistFinish(RunContext context,
IntentType intent,
ReleaseOutcome outcome,
String safeJson,
PublishedResult publishedResult,
int durationMs) {
try {
runStore.finish(context, intent, outcome, safeJson, publishedResult, durationMs);
} catch (RuntimeException exception) {
throw new ChatApplicationException(
ChatFailureCode.RUN_PERSISTENCE_FAILED,
"无法记录运行终态,请稍后重试", exception);
}
}
private ReleaseOutcome terminalOutcome(RunContext context) {
RunState state = context.lifecycle().state();
if (state == RunState.CANCELLED) {
return ReleaseOutcome.CANCELLED;
}
if (!state.isTerminal()) {
core.completeFailure(context, "APPLICATION_EXECUTION_FAILED");
}
return context.lifecycle().state() == RunState.CANCELLED
? ReleaseOutcome.CANCELLED : ReleaseOutcome.FAILED;
}
private ChatApplicationException safeFailure(RuntimeException exception,
ReleaseOutcome terminal) {
if (terminal == ReleaseOutcome.CANCELLED) {
return new ChatApplicationException(
ChatFailureCode.RUN_CANCELLED, "Chat Run was cancelled", exception);
}
if (exception instanceof IntentRoutingException) {
return new ChatApplicationException(
ChatFailureCode.ROUTING_UNAVAILABLE,
"当前暂时无法识别请求类型,请稍后重试", exception);
}
if (exception instanceof ChatApplicationException applicationFailure) {
return applicationFailure;
}
if (exception instanceof RetryExecutionException retry
&& retry.failure() == com.superbiz.agent.harness.retry.RetryFailure.CANCELLED) {
return new ChatApplicationException(
ChatFailureCode.RUN_CANCELLED, "Chat Run was cancelled", exception);
}
return new ChatApplicationException(
ChatFailureCode.INTERNAL_FAILURE, "当前暂时无法处理该请求,请稍后重试", exception);
}
private String resolveSessionId(String requested) {
String value = requested == null ? sessionIdSupplier.get() : requested;
if (value == null || !SAFE_ID.matcher(value).matches()) {
throw new IllegalArgumentException("sessionId is invalid");
}
return value;
}
private String write(ChatApplicationContent content) {
try {
return objectMapper.writeValueAsString(content);
} catch (JsonProcessingException exception) {
throw new ChatApplicationException(
ChatFailureCode.INTERNAL_FAILURE, "Public content is not serializable", exception);
}
}
private static int durationMillis(long startedNanos) {
long millis = Math.max(0L, (System.nanoTime() - startedNanos) / 1_000_000L);
return millis >= Integer.MAX_VALUE ? Integer.MAX_VALUE : (int) millis;
}
private record PathResult(
ReleaseOutcome outcome,
ChatApplicationContent content,
PublishedResult publishedResult) {
}
private static final class CoreRunControl implements ChatRunControl {
private final DiagnosisHarnessCore core;
private final RunContext context;
private CoreRunControl(DiagnosisHarnessCore core, RunContext context) {
this.core = core;
this.context = context;
}
@Override
public String sessionId() {
return context.sessionId();
}
@Override
public String runId() {
return context.runId();
}
@Override
public boolean cancelClientDisconnect() {
return core.cancel(context, RunCancellationReason.CLIENT_DISCONNECTED);
}
}
}
@@ -0,0 +1,8 @@
package com.superbiz.agent.harness.application;
public enum ChatContentType {
SYSTEM_CHAT,
KNOWLEDGE_ANSWER,
DIAGNOSIS_REPORT,
SAFE_FALLBACK
}
@@ -0,0 +1,11 @@
package com.superbiz.agent.harness.application;
public enum ChatFailureCode {
ROUTING_UNAVAILABLE,
SYSTEM_CHAT_UNAVAILABLE,
KNOWLEDGE_UNAVAILABLE,
DIAGNOSIS_UNAVAILABLE,
RUN_PERSISTENCE_FAILED,
RUN_CANCELLED,
INTERNAL_FAILURE
}
@@ -0,0 +1,10 @@
package com.superbiz.agent.harness.application;
public interface ChatRunControl {
String sessionId();
String runId();
boolean cancelClientDisconnect();
}
@@ -0,0 +1,24 @@
package com.superbiz.agent.harness.application;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.superbiz.agent.harness.contract.SafeFallback;
import com.superbiz.agent.harness.guard.semantic.SemanticDraftView;
import java.util.List;
import java.util.Objects;
public record DiagnosisContent(
@JsonProperty("report") SemanticDraftView report,
@JsonProperty("references") List<SafeFallback.VerifiedSource> references)
implements ChatApplicationContent {
public DiagnosisContent {
Objects.requireNonNull(report, "report must not be null");
references = references == null ? List.of() : List.copyOf(references);
}
@Override
public ChatContentType contentType() {
return ChatContentType.DIAGNOSIS_REPORT;
}
}
@@ -0,0 +1,29 @@
package com.superbiz.agent.harness.application;
import com.superbiz.agent.harness.contract.PublishedResult;
import com.superbiz.agent.harness.contract.ReleaseOutcome;
import java.util.Objects;
public record DiagnosisExecutionResult(
ReleaseOutcome outcome,
ChatApplicationContent content,
PublishedResult publishedResult) {
public DiagnosisExecutionResult {
Objects.requireNonNull(outcome, "outcome must not be null");
Objects.requireNonNull(content, "content must not be null");
if (outcome == ReleaseOutcome.SUCCESS && content.contentType() != ChatContentType.DIAGNOSIS_REPORT) {
throw new IllegalArgumentException("success requires diagnosis report");
}
if (outcome == ReleaseOutcome.FALLBACK && content.contentType() != ChatContentType.SAFE_FALLBACK) {
throw new IllegalArgumentException("fallback requires safe fallback");
}
if (outcome != ReleaseOutcome.SUCCESS && outcome != ReleaseOutcome.FALLBACK) {
throw new IllegalArgumentException("diagnosis execution must be SUCCESS or FALLBACK");
}
if (outcome != ReleaseOutcome.SUCCESS && publishedResult != null) {
throw new IllegalArgumentException("only success can contain PublishedResult");
}
}
}
@@ -0,0 +1,15 @@
package com.superbiz.agent.harness.application;
import com.superbiz.agent.harness.contract.PreviousTurn;
import com.superbiz.agent.harness.core.RunContext;
import java.util.function.Consumer;
@FunctionalInterface
public interface DiagnosisOperation {
DiagnosisExecutionResult execute(RunContext context,
String query,
PreviousTurn previousTurn,
Consumer<ChatApplicationStatus> statusSink);
}
@@ -0,0 +1,19 @@
package com.superbiz.agent.harness.application;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.superbiz.agent.harness.contract.SafeFallback;
import java.util.Objects;
public record FallbackContent(@JsonProperty("fallback") SafeFallback fallback)
implements ChatApplicationContent {
public FallbackContent {
Objects.requireNonNull(fallback, "fallback must not be null");
}
@Override
public ChatContentType contentType() {
return ChatContentType.SAFE_FALLBACK;
}
}
@@ -0,0 +1,11 @@
package com.superbiz.agent.harness.application;
import com.superbiz.agent.harness.application.routing.IntentRouterInput;
import com.superbiz.agent.harness.contract.IntentType;
import com.superbiz.agent.harness.core.RunContext;
@FunctionalInterface
public interface IntentRouting {
IntentType route(RunContext context, IntentRouterInput input);
}
@@ -0,0 +1,26 @@
package com.superbiz.agent.harness.application;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.superbiz.agent.harness.contract.SourceDocument;
import java.util.List;
public record KnowledgeContent(
@JsonProperty("answer") String answer,
@JsonProperty("references") List<SourceDocument> references,
@JsonProperty("limitations") List<String> limitations)
implements ChatApplicationContent {
public KnowledgeContent {
if (answer == null || answer.isBlank()) {
throw new IllegalArgumentException("answer must not be blank");
}
references = references == null ? List.of() : List.copyOf(references);
limitations = limitations == null ? List.of() : List.copyOf(limitations);
}
@Override
public ChatContentType contentType() {
return ChatContentType.KNOWLEDGE_ANSWER;
}
}
@@ -0,0 +1,9 @@
package com.superbiz.agent.harness.application;
import com.superbiz.agent.harness.core.RunContext;
@FunctionalInterface
public interface KnowledgeQueryOperation {
KnowledgeContent execute(RunContext context, String query);
}
@@ -0,0 +1,22 @@
package com.superbiz.agent.harness.application;
import com.fasterxml.jackson.annotation.JsonProperty;
public record SystemChatContent(@JsonProperty("answer") String answer)
implements ChatApplicationContent {
public SystemChatContent {
requireText(answer, "answer");
}
@Override
public ChatContentType contentType() {
return ChatContentType.SYSTEM_CHAT;
}
private static void requireText(String value, String name) {
if (value == null || value.isBlank()) {
throw new IllegalArgumentException(name + " must not be blank");
}
}
}
@@ -0,0 +1,9 @@
package com.superbiz.agent.harness.application;
import com.superbiz.agent.harness.core.RunContext;
@FunctionalInterface
public interface SystemChatOperation {
SystemChatContent execute(RunContext context, String query);
}
@@ -0,0 +1,63 @@
package com.superbiz.agent.harness.application.executor;
import com.superbiz.agent.harness.agent.DiagnosisAgentInput;
import com.superbiz.agent.harness.agent.DiagnosisAgentUseCase;
import com.superbiz.agent.harness.application.ChatApplicationStatus;
import com.superbiz.agent.harness.application.DiagnosisContent;
import com.superbiz.agent.harness.application.DiagnosisExecutionResult;
import com.superbiz.agent.harness.application.DiagnosisOperation;
import com.superbiz.agent.harness.application.FallbackContent;
import com.superbiz.agent.harness.application.persistence.PublishedResultPolicy;
import com.superbiz.agent.harness.contract.DiagnosisDraft;
import com.superbiz.agent.harness.contract.PreviousTurn;
import com.superbiz.agent.harness.contract.PublishedResult;
import com.superbiz.agent.harness.contract.ReleaseOutcome;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.guard.semantic.SemanticDraftView;
import com.superbiz.agent.harness.release.DiagnosisReleaseResult;
import com.superbiz.agent.harness.release.DiagnosisReleaseUseCase;
import java.util.Objects;
import java.util.function.Consumer;
public final class DiagnosisChatExecutor implements DiagnosisOperation {
private final DiagnosisAgentUseCase diagnosisAgent;
private final DiagnosisReleaseUseCase releaseUseCase;
private final PublishedResultPolicy publishedPolicy;
public DiagnosisChatExecutor(DiagnosisAgentUseCase diagnosisAgent,
DiagnosisReleaseUseCase releaseUseCase,
PublishedResultPolicy publishedPolicy) {
this.diagnosisAgent = Objects.requireNonNull(diagnosisAgent, "diagnosisAgent must not be null");
this.releaseUseCase = Objects.requireNonNull(releaseUseCase, "releaseUseCase must not be null");
this.publishedPolicy = Objects.requireNonNull(publishedPolicy, "publishedPolicy must not be null");
}
@Override
public DiagnosisExecutionResult execute(RunContext context,
String query,
PreviousTurn previousTurn,
Consumer<ChatApplicationStatus> statusSink) {
Objects.requireNonNull(statusSink, "statusSink must not be null");
statusSink.accept(ChatApplicationStatus.DIAGNOSIS_RUNNING);
DiagnosisDraft draft = diagnosisAgent.execute(
context, new DiagnosisAgentInput(query, previousTurn));
statusSink.accept(ChatApplicationStatus.SAFETY_VALIDATING);
DiagnosisReleaseResult released = releaseUseCase.execute(context, query, draft);
if (released.outcome() == ReleaseOutcome.FALLBACK) {
return new DiagnosisExecutionResult(
ReleaseOutcome.FALLBACK,
new FallbackContent(released.fallback()),
null);
}
PublishedResult published = publishedPolicy.create(
query, released.draft(), released.verifiedEvidence()).orElse(null);
return new DiagnosisExecutionResult(
ReleaseOutcome.SUCCESS,
new DiagnosisContent(
SemanticDraftView.from(released.draft()),
released.verifiedEvidence().verifiedSources()),
published);
}
}
@@ -0,0 +1,211 @@
package com.superbiz.agent.harness.application.executor;
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.agent.HarnessEvidenceTools;
import com.superbiz.agent.harness.application.ChatApplicationException;
import com.superbiz.agent.harness.application.ChatFailureCode;
import com.superbiz.agent.harness.application.KnowledgeContent;
import com.superbiz.agent.harness.application.KnowledgeQueryOperation;
import com.superbiz.agent.harness.contract.EvidenceStatus;
import com.superbiz.agent.harness.contract.InvocationStatus;
import com.superbiz.agent.harness.contract.KnowledgeAnswerDraft;
import com.superbiz.agent.harness.contract.SourceDocument;
import com.superbiz.agent.harness.core.DiagnosisHarnessCore;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.guard.semantic.GuardModelCall;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryResult;
import com.superbiz.agent.harness.tool.contract.AgentToolContracts;
import com.superbiz.agent.harness.tool.contract.RagEvidence;
import com.superbiz.agent.harness.tool.contract.RagToolRequest;
import com.superbiz.agent.harness.tool.contract.RagToolResult;
import org.springframework.ai.chat.messages.SystemMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import java.nio.charset.StandardCharsets;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.LinkedHashSet;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.Set;
import java.util.function.Supplier;
public final class KnowledgeQueryExecutor implements KnowledgeQueryOperation {
private static final String SYSTEM_PROMPT = """
你只根据提供的有界知识库证据回答原始问题。每个 answer item 必须引用 exact tool_call_id 和实际 document_ids。
不要使用外部知识,不要改写问题,不要输出 markdown fence。返回严格 KnowledgeAnswerDraft JSON。
""";
private final DiagnosisHarnessCore core;
private final HarnessEvidenceTools tools;
private final GuardModelCall modelCall;
private final ObjectMapper objectMapper;
private final ObjectReader ragReader;
private final ObjectReader answerReader;
private final Supplier<String> callIdSupplier;
private final KnowledgeQueryLimits limits;
public KnowledgeQueryExecutor(DiagnosisHarnessCore core,
HarnessEvidenceTools tools,
GuardModelCall modelCall,
ObjectMapper objectMapper,
Supplier<String> callIdSupplier,
KnowledgeQueryLimits limits) {
this.core = Objects.requireNonNull(core, "core must not be null");
this.tools = Objects.requireNonNull(tools, "tools must not be null");
this.modelCall = Objects.requireNonNull(modelCall, "modelCall must not be null");
this.objectMapper = Objects.requireNonNull(objectMapper, "objectMapper must not be null");
this.ragReader = objectMapper.readerFor(RagToolResult.class)
.with(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES)
.with(DeserializationFeature.FAIL_ON_TRAILING_TOKENS);
this.answerReader = objectMapper.readerFor(KnowledgeAnswerDraft.class)
.with(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES)
.with(DeserializationFeature.FAIL_ON_TRAILING_TOKENS);
this.callIdSupplier = Objects.requireNonNull(callIdSupplier, "callIdSupplier must not be null");
this.limits = Objects.requireNonNull(limits, "limits must not be null");
}
@Override
public KnowledgeContent execute(RunContext context, String query) {
if (query == null || query.isBlank()) {
throw new IllegalArgumentException("query must not be blank");
}
String callId = requireId(callIdSupplier.get());
ToolBoundaryResult boundary = tools.invoke(
context,
AgentToolContracts.LOOKUP_KNOWLEDGE,
callId,
write(new RagToolRequest(query)));
if (boundary.status() != InvocationStatus.READY || boundary.agentResult() == null) {
throw unavailable("Knowledge lookup failed", null);
}
RagToolResult rag = readRag(boundary.agentResult());
validateRag(callId, boundary, rag);
if (rag.evidenceStatus() == EvidenceStatus.NO_EVIDENCE) {
return new KnowledgeContent(
"当前知识库中没有找到可用于回答该问题的资料。",
List.of(),
List.of("知识库查询范围内无匹配证据"));
}
String modelInput = write(new KnowledgeModelInput(query, rag));
long bytes = modelInput.getBytes(StandardCharsets.UTF_8).length;
if (bytes > limits.maxModelInputBytes()) {
throw unavailable("Knowledge model input exceeds limit", null);
}
core.reserveRunBytes(context, bytes);
String output = modelCall.call(
context,
new Prompt(List.of(new SystemMessage(SYSTEM_PROMPT), new UserMessage(modelInput))),
limits.modelTimeout(),
limits.maxModelOutputBytes());
KnowledgeAnswerDraft draft = readAnswer(output);
return publish(callId, rag, draft);
}
private KnowledgeContent publish(String callId, RagToolResult rag, KnowledgeAnswerDraft draft) {
if (draft.answerItems().isEmpty()) {
throw unavailable("Knowledge answer has no items", null);
}
Map<String, RagEvidence> evidenceById = new LinkedHashMap<>();
for (RagEvidence evidence : rag.evidence()) {
if (evidence == null || evidence.documentId() == null || evidence.documentId().isBlank()) {
throw unavailable("Knowledge projection is invalid", null);
}
evidenceById.put(evidence.documentId(), evidence);
}
List<String> answers = new ArrayList<>();
Set<String> referencedIds = new LinkedHashSet<>();
for (KnowledgeAnswerDraft.AnswerItem item : draft.answerItems()) {
if (item == null || item.text() == null || item.text().isBlank()
|| !Objects.equals(callId, item.toolCallId())
|| item.documentIds().isEmpty()) {
throw unavailable("Knowledge answer reference is invalid", null);
}
for (String documentId : item.documentIds()) {
if (!evidenceById.containsKey(documentId)) {
throw unavailable("Knowledge answer references unknown document", null);
}
referencedIds.add(documentId);
}
answers.add(item.text());
}
List<SourceDocument> references = referencedIds.stream()
.map(evidenceById::get)
.map(evidence -> new SourceDocument(
evidence.documentId(), firstText(evidence.title(), evidence.source(), evidence.documentId())))
.toList();
return new KnowledgeContent(String.join("\n\n", answers), references, draft.limitations());
}
private void validateRag(String callId, ToolBoundaryResult boundary, RagToolResult rag) {
if (!Objects.equals(callId, boundary.toolCallId())
|| !Objects.equals(callId, rag.toolCallId())
|| rag.evidenceStatus() != boundary.evidenceStatus()
|| rag.returnedCount() != rag.evidence().size()) {
throw unavailable("Knowledge projection identity is invalid", null);
}
if (rag.evidenceStatus() == EvidenceStatus.EVIDENCE_FOUND && rag.evidence().isEmpty()) {
throw unavailable("Knowledge projection has no evidence", null);
}
if (rag.evidenceStatus() == EvidenceStatus.NO_EVIDENCE && !rag.evidence().isEmpty()) {
throw unavailable("No-evidence projection contains evidence", null);
}
}
private RagToolResult readRag(String value) {
try {
return ragReader.readValue(value);
} catch (JsonProcessingException exception) {
throw unavailable("Knowledge projection is invalid", exception);
}
}
private KnowledgeAnswerDraft readAnswer(String value) {
try {
return answerReader.readValue(value);
} catch (JsonProcessingException exception) {
throw unavailable("Knowledge answer is invalid", exception);
}
}
private String write(Object value) {
try {
return objectMapper.writeValueAsString(value);
} catch (JsonProcessingException exception) {
throw unavailable("Knowledge input is not serializable", exception);
}
}
private static String requireId(String value) {
if (value == null || value.isBlank()) {
throw new IllegalArgumentException("call ID must not be blank");
}
return value;
}
private static String firstText(String... values) {
for (String value : values) {
if (value != null && !value.isBlank()) {
return value;
}
}
return "unknown";
}
private static ChatApplicationException unavailable(String message, Throwable cause) {
return new ChatApplicationException(
ChatFailureCode.KNOWLEDGE_UNAVAILABLE, message, cause);
}
private record KnowledgeModelInput(
String query,
RagToolResult evidence) {
}
}
@@ -0,0 +1,20 @@
package com.superbiz.agent.harness.application.executor;
import java.time.Duration;
import java.util.Objects;
public record KnowledgeQueryLimits(
long maxModelInputBytes,
long maxModelOutputBytes,
Duration modelTimeout) {
public KnowledgeQueryLimits {
if (maxModelInputBytes <= 0 || maxModelOutputBytes <= 0) {
throw new IllegalArgumentException("byte limits must be positive");
}
Objects.requireNonNull(modelTimeout, "modelTimeout must not be null");
if (modelTimeout.isZero() || modelTimeout.isNegative()) {
throw new IllegalArgumentException("modelTimeout must be positive");
}
}
}
@@ -0,0 +1,20 @@
package com.superbiz.agent.harness.application.executor;
import java.time.Duration;
import java.util.Objects;
public record SingleTurnExecutorLimits(
long maxInputBytes,
long maxOutputBytes,
Duration timeout) {
public SingleTurnExecutorLimits {
if (maxInputBytes <= 0 || maxOutputBytes <= 0) {
throw new IllegalArgumentException("byte limits must be positive");
}
Objects.requireNonNull(timeout, "timeout must not be null");
if (timeout.isZero() || timeout.isNegative()) {
throw new IllegalArgumentException("timeout must be positive");
}
}
}
@@ -0,0 +1,56 @@
package com.superbiz.agent.harness.application.executor;
import com.superbiz.agent.harness.application.ChatApplicationException;
import com.superbiz.agent.harness.application.ChatFailureCode;
import com.superbiz.agent.harness.application.SystemChatContent;
import com.superbiz.agent.harness.application.SystemChatOperation;
import com.superbiz.agent.harness.core.DiagnosisHarnessCore;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.guard.semantic.GuardModelCall;
import org.springframework.ai.chat.messages.SystemMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Objects;
public final class SystemChatExecutor implements SystemChatOperation {
private static final String SYSTEM_PROMPT = """
你是 SuperBizAgent 的系统助手。只回答产品定位、能力范围、使用方式、简单问候和非业务闲聊。
系统可以查询内部知识文档,也可以对日志和授权只读数据库进行诊断;不会执行写操作。
不要声称已经查询任何实时数据,不要编造系统能力。直接给出简洁回答。
""";
private final DiagnosisHarnessCore core;
private final GuardModelCall modelCall;
private final SingleTurnExecutorLimits limits;
public SystemChatExecutor(DiagnosisHarnessCore core,
GuardModelCall modelCall,
SingleTurnExecutorLimits limits) {
this.core = Objects.requireNonNull(core, "core must not be null");
this.modelCall = Objects.requireNonNull(modelCall, "modelCall must not be null");
this.limits = Objects.requireNonNull(limits, "limits must not be null");
}
@Override
public SystemChatContent execute(RunContext context, String query) {
if (query == null || query.isBlank()) {
throw new IllegalArgumentException("query must not be blank");
}
long bytes = query.getBytes(StandardCharsets.UTF_8).length;
if (bytes > limits.maxInputBytes()) {
throw new ChatApplicationException(
ChatFailureCode.SYSTEM_CHAT_UNAVAILABLE, "System Chat input exceeds limit");
}
core.reserveRunBytes(context, bytes);
String answer = modelCall.call(
context,
new Prompt(List.of(new SystemMessage(SYSTEM_PROMPT), new UserMessage(query))),
limits.timeout(),
limits.maxOutputBytes());
return new SystemChatContent(answer);
}
}
@@ -0,0 +1,27 @@
package com.superbiz.agent.harness.application.persistence;
import com.superbiz.agent.harness.contract.IntentType;
import com.superbiz.agent.harness.contract.PreviousTurn;
import com.superbiz.agent.harness.contract.PublishedResult;
import com.superbiz.agent.harness.contract.ReleaseOutcome;
import com.superbiz.agent.harness.core.RunContext;
import java.util.Optional;
public interface ChatRunStore {
Optional<RoutingHistory> findLatestRoutingHistory(String sessionId);
Optional<PreviousTurn> findPreviousTurn(String sessionId);
void start(RunContext context, String query);
void markIntent(String runId, IntentType intent);
void finish(RunContext context,
IntentType intent,
ReleaseOutcome outcome,
String safeContentJson,
PublishedResult publishedResult,
int durationMs);
}
@@ -0,0 +1,171 @@
package com.superbiz.agent.harness.application.persistence;
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.domain.entity.ChatSession;
import com.superbiz.agent.domain.entity.DiagnosisRun;
import com.superbiz.agent.harness.contract.IntentType;
import com.superbiz.agent.harness.contract.PreviousTurn;
import com.superbiz.agent.harness.contract.PublishedResult;
import com.superbiz.agent.harness.contract.ReleaseOutcome;
import com.superbiz.agent.harness.core.RunBudgetUsage;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.repository.ChatSessionRepository;
import com.superbiz.agent.repository.DiagnosisRunRepository;
import org.springframework.stereotype.Component;
import org.springframework.transaction.annotation.Transactional;
import org.springframework.beans.factory.annotation.Autowired;
import java.time.LocalDateTime;
import java.util.List;
import java.util.Objects;
import java.util.Optional;
@Component
public class JpaChatRunStore implements ChatRunStore {
private final ChatSessionRepository chatSessions;
private final DiagnosisRunRepository runs;
private final ObjectMapper objectMapper;
private final ObjectReader publishedReader;
private final PublishedResultPolicy publishedPolicy;
@Autowired
public JpaChatRunStore(ChatSessionRepository chatSessions,
DiagnosisRunRepository runs,
ObjectMapper objectMapper) {
this(chatSessions, runs, objectMapper,
new PublishedResultPolicy(PreviousTurnLimits.defaults()));
}
public JpaChatRunStore(ChatSessionRepository chatSessions,
DiagnosisRunRepository runs,
ObjectMapper objectMapper,
PublishedResultPolicy publishedPolicy) {
this.chatSessions = Objects.requireNonNull(chatSessions, "chatSessions must not be null");
this.runs = Objects.requireNonNull(runs, "runs must not be null");
this.objectMapper = Objects.requireNonNull(objectMapper, "objectMapper must not be null");
this.publishedReader = objectMapper.readerFor(PublishedResult.class)
.with(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES)
.with(DeserializationFeature.FAIL_ON_TRAILING_TOKENS);
this.publishedPolicy = Objects.requireNonNull(publishedPolicy, "publishedPolicy must not be null");
}
@Override
@Transactional(readOnly = true)
public Optional<RoutingHistory> findLatestRoutingHistory(String sessionId) {
return runs.findFirstBySessionIdAndReleaseOutcomeInOrderByCreatedAtDescIdDesc(
sessionId, List.of(ReleaseOutcome.SUCCESS, ReleaseOutcome.FALLBACK))
.filter(run -> run.getIntent() != null && hasText(run.getQuery()))
.map(run -> new RoutingHistory(run.getIntent(), run.getQuery()));
}
@Override
@Transactional(readOnly = true)
public Optional<PreviousTurn> findPreviousTurn(String sessionId) {
return runs.findFirstBySessionIdAndIntentAndReleaseOutcomeAndPublishedResultIsNotNullOrderByCreatedAtDescIdDesc(
sessionId, IntentType.DIAGNOSIS, ReleaseOutcome.SUCCESS)
.flatMap(this::readPreviousTurn);
}
@Override
@Transactional
public void start(RunContext context, String query) {
Objects.requireNonNull(context, "context must not be null");
if (!hasText(query)) {
throw new IllegalArgumentException("query must not be blank");
}
LocalDateTime now = LocalDateTime.now();
ChatSession session = chatSessions.findBySessionId(context.sessionId())
.orElseGet(() -> ChatSession.builder()
.sessionId(context.sessionId())
.status("ACTIVE")
.build());
session.setStatus("ACTIVE");
session.setLastActiveAt(now);
chatSessions.save(session);
runs.save(DiagnosisRun.builder()
.runId(context.runId())
.sessionId(context.sessionId())
.query(query)
.status("RUNNING")
.agentFlow("CHAT")
.build());
}
@Override
@Transactional
public void markIntent(String runId, IntentType intent) {
DiagnosisRun run = requiredRun(runId);
run.setIntent(Objects.requireNonNull(intent, "intent must not be null"));
runs.save(run);
}
@Override
@Transactional
public void finish(RunContext context, IntentType intent, ReleaseOutcome outcome,
String safeContentJson, PublishedResult publishedResult, int durationMs) {
Objects.requireNonNull(context, "context must not be null");
Objects.requireNonNull(outcome, "outcome must not be null");
DiagnosisRun run = requiredRun(context.runId());
run.setIntent(intent);
run.setReleaseOutcome(outcome);
run.setStatus(status(outcome));
run.setAnswer(safeContentJson);
run.setPublishedResult(outcome == ReleaseOutcome.SUCCESS && intent == IntentType.DIAGNOSIS
&& publishedResult != null ? write(publishedResult) : null);
run.setTotalDurationMs(Math.max(0, durationMs));
RunBudgetUsage usage = context.budget().snapshot();
run.setTotalTokenCount(saturatingInt(usage.totalTokens()));
run.setToolCallCount(saturatingInt(usage.toolCalls()));
runs.save(run);
if (outcome == ReleaseOutcome.SUCCESS || outcome == ReleaseOutcome.FALLBACK) {
chatSessions.findBySessionId(context.sessionId()).ifPresent(session -> {
session.setLastActiveAt(LocalDateTime.now());
session.setMessagePairCount(
session.getMessagePairCount() == null ? 1 : session.getMessagePairCount() + 1);
chatSessions.save(session);
});
}
}
private Optional<PreviousTurn> readPreviousTurn(DiagnosisRun run) {
try {
PublishedResult result = publishedReader.readValue(run.getPublishedResult());
return publishedPolicy.previousTurn(result);
} catch (Exception exception) {
return Optional.empty();
}
}
private DiagnosisRun requiredRun(String runId) {
return runs.findByRunId(runId)
.orElseThrow(() -> new IllegalStateException("Diagnosis Run is missing"));
}
private String write(PublishedResult result) {
try {
return objectMapper.writeValueAsString(result);
} catch (JsonProcessingException exception) {
throw new IllegalStateException("PublishedResult is not serializable", exception);
}
}
private static String status(ReleaseOutcome outcome) {
return switch (outcome) {
case SUCCESS, FALLBACK -> "SUCCESS";
case FAILED -> "FAILED";
case CANCELLED -> "CANCELLED";
};
}
private static int saturatingInt(long value) {
return value >= Integer.MAX_VALUE ? Integer.MAX_VALUE : (int) Math.max(0, value);
}
private static boolean hasText(String value) {
return value != null && !value.isBlank();
}
}
@@ -0,0 +1,24 @@
package com.superbiz.agent.harness.application.persistence;
public record PreviousTurnLimits(
int maxUserQueryChars,
int maxConclusionChars,
int maxScopeChars,
int maxLimitations,
int maxLimitationChars,
int maxSourceDocuments,
int maxDocumentIdChars,
int maxTitleChars) {
public PreviousTurnLimits {
if (maxUserQueryChars <= 0 || maxConclusionChars <= 0 || maxScopeChars <= 0
|| maxLimitations < 0 || maxLimitationChars <= 0
|| maxSourceDocuments < 0 || maxDocumentIdChars <= 0 || maxTitleChars <= 0) {
throw new IllegalArgumentException("PreviousTurn limits are invalid");
}
}
public static PreviousTurnLimits defaults() {
return new PreviousTurnLimits(2_000, 2_000, 1_000, 10, 500, 10, 256, 500);
}
}
@@ -0,0 +1,112 @@
package com.superbiz.agent.harness.application.persistence;
import com.superbiz.agent.harness.contract.DiagnosisDraft;
import com.superbiz.agent.harness.contract.PreviousTurn;
import com.superbiz.agent.harness.contract.PublishedResult;
import com.superbiz.agent.harness.contract.SourceDocument;
import com.superbiz.agent.harness.guard.evidence.VerifiedAnalysisEvidence;
import com.superbiz.agent.harness.guard.evidence.VerifiedEvidence;
import com.superbiz.agent.harness.guard.evidence.VerifiedEvidenceSnapshot;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.Optional;
public final class PublishedResultPolicy {
private final PreviousTurnLimits limits;
public PublishedResultPolicy(PreviousTurnLimits limits) {
this.limits = Objects.requireNonNull(limits, "limits must not be null");
}
public Optional<PublishedResult> create(String query,
DiagnosisDraft draft,
VerifiedEvidenceSnapshot snapshot) {
if (draft == null || draft.conclusion() == null
|| isBlank(query) || isBlank(draft.conclusion().text())
|| draft.limitations() == null || isBlank(draft.limitations().scope())) {
return Optional.empty();
}
return sanitize(new PublishedResult(
query,
draft.conclusion().text(),
draft.limitations().scope(),
draft.limitations().missingInfo(),
sourceDocuments(snapshot)));
}
public Optional<PublishedResult> sanitize(PublishedResult value) {
if (value == null || isBlank(value.userQuery())
|| isBlank(value.publishedConclusion()) || isBlank(value.scope())) {
return Optional.empty();
}
List<String> limitations = value.limitations().stream()
.filter(item -> !isBlank(item))
.limit(limits.maxLimitations())
.map(item -> bound(item, limits.maxLimitationChars()))
.toList();
List<SourceDocument> documents = value.sourceDocuments().stream()
.filter(Objects::nonNull)
.filter(document -> !isBlank(document.documentId()) && !isBlank(document.title()))
.limit(limits.maxSourceDocuments())
.map(document -> new SourceDocument(
bound(document.documentId(), limits.maxDocumentIdChars()),
bound(document.title(), limits.maxTitleChars())))
.toList();
return Optional.of(new PublishedResult(
bound(value.userQuery(), limits.maxUserQueryChars()),
bound(value.publishedConclusion(), limits.maxConclusionChars()),
bound(value.scope(), limits.maxScopeChars()),
limitations,
documents));
}
public Optional<PreviousTurn> previousTurn(PublishedResult value) {
return sanitize(value).map(PreviousTurn::from);
}
private List<SourceDocument> sourceDocuments(VerifiedEvidenceSnapshot snapshot) {
if (snapshot == null) {
return List.of();
}
Map<String, SourceDocument> documents = new LinkedHashMap<>();
for (VerifiedAnalysisEvidence analysis : snapshot.analyses()) {
for (VerifiedEvidence evidence : analysis.evidence()) {
if (!"RAG".equals(evidence.sourceType())) {
continue;
}
String id = stringValue(evidence.values().get("document_id"));
String title = stringValue(evidence.values().get("title"));
if (!isBlank(id)) {
documents.putIfAbsent(id, new SourceDocument(id, firstText(title, evidence.source(), id)));
}
}
}
return new ArrayList<>(documents.values());
}
private static String stringValue(Object value) {
return value == null ? null : String.valueOf(value);
}
private static String firstText(String... values) {
for (String value : values) {
if (!isBlank(value)) {
return value;
}
}
return "unknown";
}
private static String bound(String value, int max) {
return value.length() <= max ? value : value.substring(0, max);
}
private static boolean isBlank(String value) {
return value == null || value.isBlank();
}
}
@@ -0,0 +1,15 @@
package com.superbiz.agent.harness.application.persistence;
import com.superbiz.agent.harness.contract.IntentType;
import java.util.Objects;
public record RoutingHistory(IntentType intent, String userQuery) {
public RoutingHistory {
Objects.requireNonNull(intent, "intent must not be null");
if (userQuery == null || userQuery.isBlank()) {
throw new IllegalArgumentException("userQuery must not be blank");
}
}
}
@@ -0,0 +1,147 @@
package com.superbiz.agent.harness.application.routing;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.contract.IntentType;
import com.superbiz.agent.harness.application.IntentRouting;
import com.superbiz.agent.harness.core.DiagnosisHarnessCore;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.guard.semantic.GuardModelCall;
import com.superbiz.agent.harness.guard.semantic.GuardModelCallException;
import com.superbiz.agent.harness.retry.HarnessRetryExecutor;
import com.superbiz.agent.harness.retry.RetryAttempt;
import com.superbiz.agent.harness.retry.RetryExecutionException;
import com.superbiz.agent.harness.retry.RetryFailure;
import org.springframework.ai.chat.messages.SystemMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import java.nio.charset.StandardCharsets;
import java.time.Duration;
import java.util.HashSet;
import java.util.Iterator;
import java.util.List;
import java.util.Objects;
import java.util.Set;
import java.util.function.Consumer;
public final class IntentRouter implements IntentRouting {
private static final Set<String> OUTPUT_FIELDS = Set.of("intent");
private final DiagnosisHarnessCore core;
private final HarnessRetryExecutor retryExecutor;
private final GuardModelCall modelCall;
private final ObjectMapper objectMapper;
private final IntentRouterLimits limits;
private final Consumer<RetryAttempt> attemptRecorder;
private final String systemPrompt;
public IntentRouter(DiagnosisHarnessCore core,
HarnessRetryExecutor retryExecutor,
GuardModelCall modelCall,
ObjectMapper objectMapper,
IntentRouterLimits limits,
Consumer<RetryAttempt> attemptRecorder) {
this.core = Objects.requireNonNull(core, "core must not be null");
this.retryExecutor = Objects.requireNonNull(retryExecutor, "retryExecutor must not be null");
this.modelCall = Objects.requireNonNull(modelCall, "modelCall must not be null");
this.objectMapper = Objects.requireNonNull(objectMapper, "objectMapper must not be null");
this.limits = Objects.requireNonNull(limits, "limits must not be null");
this.attemptRecorder = Objects.requireNonNull(attemptRecorder, "attemptRecorder must not be null");
this.systemPrompt = IntentRouterPrompt.load();
}
@Override
public IntentType route(RunContext context, IntentRouterInput input) {
Objects.requireNonNull(context, "context must not be null");
Objects.requireNonNull(input, "input must not be null");
String json = serialize(input);
long bytes = utf8Bytes(json);
if (bytes > limits.maxInputBytes()) {
throw new IntentRoutingException(new IllegalArgumentException("Router input exceeded limit"));
}
core.reserveRunBytes(context, bytes);
Prompt prompt = new Prompt(List.of(
new SystemMessage(systemPrompt), new UserMessage(json)));
long started = System.nanoTime();
try {
return retryExecutor.execute(
context,
context.retryPolicies().intentRouter(),
() -> parse(modelCall.call(
context, prompt, remaining(started), limits.maxOutputBytes())),
this::classify,
attemptRecorder);
} catch (RetryExecutionException exception) {
if (exception.failure() == RetryFailure.CANCELLED
|| exception.failure() == RetryFailure.BUDGET_EXHAUSTED) {
throw exception;
}
throw new IntentRoutingException(exception);
}
}
private IntentType parse(String output) {
JsonNode root;
try {
root = objectMapper.readTree(output);
} catch (JsonProcessingException exception) {
throw invalidOutput(exception);
}
if (root == null || !root.isObject() || !fieldNames(root).equals(OUTPUT_FIELDS)
|| !root.path("intent").isTextual()) {
throw invalidOutput(null);
}
try {
return IntentType.valueOf(root.path("intent").asText());
} catch (IllegalArgumentException exception) {
throw invalidOutput(exception);
}
}
private GuardModelCallException invalidOutput(Throwable cause) {
return new GuardModelCallException(
RetryFailure.INVALID_OUTPUT, "Intent Router output is invalid", cause);
}
private RetryFailure classify(Exception exception) {
if (exception instanceof GuardModelCallException failure) {
return switch (failure.failure()) {
case TIMEOUT, TRANSPORT -> failure.failure();
default -> RetryFailure.INVALID_OUTPUT;
};
}
return RetryFailure.UNKNOWN;
}
private Duration remaining(long started) {
long remaining = limits.totalTimeout().toNanos()
- Math.max(0L, System.nanoTime() - started);
if (remaining <= 0) {
throw new GuardModelCallException(
RetryFailure.TIMEOUT, "Intent Router total timeout exhausted");
}
return Duration.ofNanos(Math.min(remaining, limits.perAttemptTimeout().toNanos()));
}
private String serialize(IntentRouterInput input) {
try {
return objectMapper.writeValueAsString(input);
} catch (JsonProcessingException exception) {
throw new IntentRoutingException(exception);
}
}
private static Set<String> fieldNames(JsonNode node) {
Set<String> names = new HashSet<>();
Iterator<String> fields = node.fieldNames();
fields.forEachRemaining(names::add);
return names;
}
private static long utf8Bytes(String value) {
return value.getBytes(StandardCharsets.UTF_8).length;
}
}
@@ -0,0 +1,19 @@
package com.superbiz.agent.harness.application.routing;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.superbiz.agent.harness.contract.IntentType;
public record IntentRouterInput(
@JsonProperty("query") String query,
@JsonProperty("last_intent") IntentType lastIntent,
@JsonProperty("last_user_query") String lastUserQuery) {
public IntentRouterInput {
if (query == null || query.isBlank()) {
throw new IllegalArgumentException("query must not be blank");
}
if (lastUserQuery != null && lastUserQuery.isBlank()) {
lastUserQuery = null;
}
}
}
@@ -0,0 +1,30 @@
package com.superbiz.agent.harness.application.routing;
import java.time.Duration;
import java.util.Objects;
public record IntentRouterLimits(
long maxInputBytes,
long maxOutputBytes,
Duration perAttemptTimeout,
Duration totalTimeout) {
public IntentRouterLimits {
if (maxInputBytes <= 0 || maxOutputBytes <= 0) {
throw new IllegalArgumentException("byte limits must be positive");
}
requirePositive(perAttemptTimeout, "perAttemptTimeout");
requirePositive(totalTimeout, "totalTimeout");
if (totalTimeout.compareTo(perAttemptTimeout) < 0) {
throw new IllegalArgumentException("totalTimeout must not be shorter than perAttemptTimeout");
}
}
private static void requirePositive(Duration value, String name) {
Objects.requireNonNull(value, name + " must not be null");
if (value.isZero() || value.isNegative()) {
throw new IllegalArgumentException(name + " must be positive");
}
value.toNanos();
}
}
@@ -0,0 +1,23 @@
package com.superbiz.agent.harness.application.routing;
import org.springframework.core.io.ClassPathResource;
import java.io.IOException;
import java.io.InputStream;
import java.nio.charset.StandardCharsets;
final class IntentRouterPrompt {
private static final String PATH = "prompts/intent-router-prompt.md";
private IntentRouterPrompt() {
}
static String load() {
try (InputStream input = new ClassPathResource(PATH).getInputStream()) {
return new String(input.readAllBytes(), StandardCharsets.UTF_8);
} catch (IOException exception) {
throw new IllegalStateException("Failed to load Intent Router prompt", exception);
}
}
}
@@ -0,0 +1,8 @@
package com.superbiz.agent.harness.application.routing;
public final class IntentRoutingException extends RuntimeException {
public IntentRoutingException(Throwable cause) {
super("Intent routing is unavailable", cause);
}
}