Files
SuperBizAgent-java/src/main/java/com/superbiz/agent/graph/diagnosis/AbstractDiagnosisAgentNodeAdapter.java
T

95 lines
3.8 KiB
Java

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<Map<String, Object>> 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<String, Object> 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<String, Object> toUpdate(NodeResult result, int attempt) {
Map<String, Object> 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<String, Object> values) {
protected NodeResult {
Objects.requireNonNull(outcome, "outcome");
Objects.requireNonNull(reasonCode, "reasonCode");
values = values == null ? Map.of() : Map.copyOf(values);
}
}
}