package com.superbiz.agent.graph.diagnosis; import com.alibaba.cloud.ai.graph.OverAllState; import com.alibaba.cloud.ai.graph.RunnableConfig; import com.alibaba.cloud.ai.graph.action.AsyncNodeActionWithConfig; import com.fasterxml.jackson.databind.ObjectMapper; import java.util.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.Objects; import java.util.concurrent.CompletableFuture; abstract class AbstractDiagnosisAgentNodeAdapter implements AsyncNodeActionWithConfig { protected final ObjectMapper objectMapper; private final String nodeName; private final String retryCountKey; private final DiagnosisAgentInvoker invoker; private final DiagnosisNodeFailureClassifier failureClassifier; AbstractDiagnosisAgentNodeAdapter( String nodeName, String retryCountKey, DiagnosisAgentInvoker invoker, ObjectMapper objectMapper, DiagnosisNodeFailureClassifier failureClassifier) { this.nodeName = Objects.requireNonNull(nodeName, "nodeName"); this.retryCountKey = Objects.requireNonNull(retryCountKey, "retryCountKey"); this.invoker = Objects.requireNonNull(invoker, "invoker"); this.objectMapper = Objects.requireNonNull(objectMapper, "objectMapper"); this.failureClassifier = Objects.requireNonNull( failureClassifier, "failureClassifier"); } @Override public CompletableFuture> apply( OverAllState state, RunnableConfig config) { int attempt = DiagnosisGraphState.intValue(state, retryCountKey) + 1; try { String input = objectMapper.writeValueAsString(projectInput(state)); NodeResult result = parseOutput(state, invoker.invoke(input, config)); return CompletableFuture.completedFuture(toUpdate(result, attempt)); } catch (Throwable failure) { boolean retryable = failureClassifier.classify(failure) == DiagnosisNodeFailureClassifier.FailureKind.RETRYABLE; NodeResult result = failureResult(failure, retryable); return CompletableFuture.completedFuture(toUpdate(result, attempt)); } } protected NodeResult failureResult(Throwable failure, boolean retryable) { return new NodeResult( retryable ? retryableFailureStatus() : nonRetryableFailureStatus(), retryable ? retryableFailureReason() : nonRetryableFailureReason(), Map.of(DiagnosisGraphState.FAILURE_REASON, retryable ? retryableFailureReason() : nonRetryableFailureReason())); } protected abstract Map projectInput(OverAllState state); protected abstract NodeResult parseOutput(OverAllState state, String rawOutput); protected abstract String retryableFailureStatus(); protected abstract String nonRetryableFailureStatus(); protected abstract String retryableFailureReason(); protected abstract String nonRetryableFailureReason(); private Map toUpdate(NodeResult result, int attempt) { Map update = new LinkedHashMap<>(result.values()); update.put(DiagnosisGraphState.ORCHESTRATION_EVENTS, List.of(new OrchestrationEvent( nodeName, result.outcome(), result.reasonCode(), attempt))); return update; } protected record NodeResult( String outcome, String reasonCode, Map values) { protected NodeResult { Objects.requireNonNull(outcome, "outcome"); Objects.requireNonNull(reasonCode, "reasonCode"); values = values == null ? Map.of() : Map.copyOf(values); } } }