diff --git a/internal/ai/runtime/graphs/handoff_graph.go b/internal/ai/runtime/graphs/handoff_graph.go index 0a0170b..7fd401a 100644 --- a/internal/ai/runtime/graphs/handoff_graph.go +++ b/internal/ai/runtime/graphs/handoff_graph.go @@ -22,6 +22,10 @@ type HandoffGraphInterruptInfo struct { Message string `json:"message"` } +type handoffGraphArgs struct { + Reason string `json:"reason"` +} + func init() { schema.RegisterName[HandoffGraphState]("cs_agent_handoff_graph_state") schema.RegisterName[HandoffGraphInterruptInfo]("cs_agent_handoff_graph_interrupt_info") @@ -92,13 +96,13 @@ func (g *HandoffGraph) Run(ctx context.Context, argumentsInJSON string) (string, func (g *HandoffGraph) buildReason(argumentsInJSON string) (string, error) { reason := "用户需要转人工支持" - raw := make(map[string]any) + var args handoffGraphArgs if strings.TrimSpace(argumentsInJSON) != "" { - if err := json.Unmarshal([]byte(argumentsInJSON), &raw); err != nil { + if err := json.Unmarshal([]byte(argumentsInJSON), &args); err != nil { return "", fmt.Errorf("invalid handoff arguments: %w", err) } } - if parsed := strings.TrimSpace(graphGetStringValue(raw, "reason")); parsed != "" { + if parsed := strings.TrimSpace(args.Reason); parsed != "" { reason = parsed } return reason, nil @@ -111,15 +115,3 @@ func (g *HandoffGraph) buildConfirmationPrompt(reason string) string { func parseHandoffDecision(value string) ConfirmationDecision { return ParseConfirmationDecision(value) } - -func graphGetStringValue(data map[string]any, key string) string { - if len(data) == 0 { - return "" - } - value, ok := data[key] - if !ok || value == nil { - return "" - } - text, _ := value.(string) - return text -} diff --git a/internal/ai/runtime/graphs/handoff_graph_test.go b/internal/ai/runtime/graphs/handoff_graph_test.go new file mode 100644 index 0000000..b6f469d --- /dev/null +++ b/internal/ai/runtime/graphs/handoff_graph_test.go @@ -0,0 +1,31 @@ +package graphs + +import ( + "testing" + + "cs-agent/internal/models" +) + +func TestHandoffGraphBuildReason(t *testing.T) { + graph := NewHandoffGraph(&models.Conversation{ID: 1}, &models.AIAgent{Name: "AI"}) + + reason, err := graph.buildReason(`{"reason":" 用户需要人工确认 "}`) + if err != nil { + t.Fatalf("buildReason returned error: %v", err) + } + if reason != "用户需要人工确认" { + t.Fatalf("unexpected reason: %q", reason) + } +} + +func TestHandoffGraphBuildReasonFallback(t *testing.T) { + graph := NewHandoffGraph(&models.Conversation{ID: 1}, &models.AIAgent{Name: "AI"}) + + reason, err := graph.buildReason(`{}`) + if err != nil { + t.Fatalf("buildReason returned error: %v", err) + } + if reason != "用户需要转人工支持" { + t.Fatalf("unexpected fallback reason: %q", reason) + } +}