feat(graph): add diagnosis real nodes
This commit is contained in:
+94
@@ -0,0 +1,94 @@
|
||||
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);
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user