95 lines
3.8 KiB
Java
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);
|
|
}
|
|
}
|
|
}
|