From 13903baf61f59c562b97b5db13b2886a0d0e32b8 Mon Sep 17 00:00:00 2001 From: mlogclub Date: Tue, 14 Apr 2026 11:15:07 +0800 Subject: [PATCH] feat: add structured types for graph tool outcomes and enhance candidate extraction logic --- .../impl/callbacks/agent_trace_handler.go | 123 ++++++++++-------- .../callbacks/agent_trace_handler_test.go | 38 ++++++ .../ai/runtime/reply_interrupt_helpers.go | 11 +- 3 files changed, 114 insertions(+), 58 deletions(-) create mode 100644 internal/ai/runtime/internal/impl/callbacks/agent_trace_handler_test.go diff --git a/internal/ai/runtime/internal/impl/callbacks/agent_trace_handler.go b/internal/ai/runtime/internal/impl/callbacks/agent_trace_handler.go index 0f3682c..15f6fb0 100644 --- a/internal/ai/runtime/internal/impl/callbacks/agent_trace_handler.go +++ b/internal/ai/runtime/internal/impl/callbacks/agent_trace_handler.go @@ -27,6 +27,43 @@ type RuntimeTraceHandler struct { toolMetadataBy map[string]ToolMetadata } +type graphAnalyzeConversationResult struct { + RecommendedNextAction string `json:"recommendedNextAction"` + RiskLevel string `json:"riskLevel"` +} + +type graphTriageAnalysisResult struct { + RiskLevel string `json:"riskLevel"` +} + +type graphTriageTicketDraftResult struct { + Ready bool `json:"ready"` +} + +type graphTriageServiceRequestResult struct { + RecommendedAction string `json:"recommendedAction"` + Analysis graphTriageAnalysisResult `json:"analysis"` + TicketDraft *graphTriageTicketDraftResult `json:"ticketDraft"` +} + +type toolSearchArguments struct { + Query string `json:"query"` + RegexPattern string `json:"regex_pattern"` + ToolCode string `json:"toolCode"` +} + +type toolSearchCandidateResult struct { + ToolCode string `json:"toolCode"` +} + +type toolSearchInvokeResult struct { + SelectedTools []string `json:"selectedTools"` +} + +type toolSearchSearchResult struct { + Candidates []toolSearchCandidateResult `json:"candidates"` +} + func NewRuntimeTraceHandler(collector *RuntimeTraceCollector, toolMetadataBy map[string]ToolMetadata) *RuntimeTraceHandler { return &RuntimeTraceHandler{ BaseChatModelAgentMiddleware: &adk.BaseChatModelAgentMiddleware{}, @@ -94,21 +131,22 @@ func parseGraphToolOutcome(toolCode string, result string) (recommendedAction, r if toolCode == "" || strings.TrimSpace(result) == "" { return "", "", false } - payload := make(map[string]any) - if err := json.Unmarshal([]byte(strings.TrimSpace(result)), &payload); err != nil { - return "", "", false - } switch toolCode { case toolx.GraphAnalyzeConversation.Code: - return strings.TrimSpace(readToolSearchString(payload, "recommendedNextAction")), strings.TrimSpace(readToolSearchString(payload, "riskLevel")), false - case toolx.GraphTriageServiceRequest.Code: - recommendedAction = strings.TrimSpace(readToolSearchString(payload, "recommendedAction")) - if analysis, ok := payload["analysis"].(map[string]any); ok { - riskLevel = strings.TrimSpace(readToolSearchString(analysis, "riskLevel")) + var payload graphAnalyzeConversationResult + if err := json.Unmarshal([]byte(strings.TrimSpace(result)), &payload); err != nil { + return "", "", false } - if ticketDraft, ok := payload["ticketDraft"].(map[string]any); ok { - ready, _ := ticketDraft["ready"].(bool) - ticketDraftReady = ready + return strings.TrimSpace(payload.RecommendedNextAction), strings.TrimSpace(payload.RiskLevel), false + case toolx.GraphTriageServiceRequest.Code: + var payload graphTriageServiceRequestResult + if err := json.Unmarshal([]byte(strings.TrimSpace(result)), &payload); err != nil { + return "", "", false + } + recommendedAction = strings.TrimSpace(payload.RecommendedAction) + riskLevel = strings.TrimSpace(payload.Analysis.RiskLevel) + if payload.TicketDraft != nil { + ticketDraftReady = payload.TicketDraft.Ready } return recommendedAction, riskLevel, ticketDraftReady default: @@ -163,12 +201,12 @@ func previewToolText(text string, limit int) string { func (h *RuntimeTraceHandler) buildToolSearchTraceItem(argumentsInJSON string, result string, runErr error) ToolSearchTraceItem { item := ToolSearchTraceItem{Status: "ok"} - args := parseToolArguments(argumentsInJSON) - item.Query = strings.TrimSpace(firstNonBlank( - readToolSearchString(args, "query"), - readToolSearchString(args, "regex_pattern"), - )) - item.TargetToolCode = strings.TrimSpace(readToolSearchString(args, "toolCode")) + var args toolSearchArguments + if strings.TrimSpace(argumentsInJSON) != "" { + _ = json.Unmarshal([]byte(argumentsInJSON), &args) + } + item.Query = strings.TrimSpace(firstNonBlank(args.Query, args.RegexPattern)) + item.TargetToolCode = strings.TrimSpace(args.ToolCode) item.TargetServerCode, item.TargetToolName = toolx.SplitMCPToolCode(item.TargetToolCode) if item.TargetToolCode != "" { item.Action = "invoke" @@ -180,26 +218,10 @@ func (h *RuntimeTraceHandler) buildToolSearchTraceItem(argumentsInJSON string, r item.ErrorMessage = runErr.Error() return item } - payload := make(map[string]any) - if err := json.Unmarshal([]byte(strings.TrimSpace(result)), &payload); err != nil { - return item - } - item.CandidateToolCodes = h.extractCandidateToolCodes(payload) + item.CandidateToolCodes = h.extractCandidateToolCodes(result) return item } -func readToolSearchString(data map[string]any, key string) string { - if len(data) == 0 { - return "" - } - value, ok := data[key] - if !ok { - return "" - } - text, _ := value.(string) - return text -} - func firstNonBlank(values ...string) string { for _, value := range values { value = strings.TrimSpace(value) @@ -210,30 +232,29 @@ func firstNonBlank(values ...string) string { return "" } -func (h *RuntimeTraceHandler) extractCandidateToolCodes(payload map[string]any) []string { - if len(payload) == 0 { +func (h *RuntimeTraceHandler) extractCandidateToolCodes(result string) []string { + result = strings.TrimSpace(result) + if result == "" { return nil } - if items, ok := payload["selectedTools"].([]any); ok { - return h.extractSelectedToolCodes(items) + var invokePayload toolSearchInvokeResult + if err := json.Unmarshal([]byte(result), &invokePayload); err == nil && len(invokePayload.SelectedTools) > 0 { + return h.extractSelectedToolCodes(invokePayload.SelectedTools) } - if items, ok := payload["candidates"].([]any); ok { - return h.extractCandidateObjects(items) + var searchPayload toolSearchSearchResult + if err := json.Unmarshal([]byte(result), &searchPayload); err == nil && len(searchPayload.Candidates) > 0 { + return extractCandidateObjectCodes(searchPayload.Candidates) } return nil } -func (h *RuntimeTraceHandler) extractSelectedToolCodes(items []any) []string { +func (h *RuntimeTraceHandler) extractSelectedToolCodes(items []string) []string { if len(items) == 0 { return nil } ret := make([]string, 0, len(items)) for _, item := range items { - toolName, ok := item.(string) - if !ok { - continue - } - toolName = strings.TrimSpace(toolName) + toolName := strings.TrimSpace(item) if toolName == "" { continue } @@ -249,17 +270,13 @@ func (h *RuntimeTraceHandler) extractSelectedToolCodes(items []any) []string { return ret } -func (h *RuntimeTraceHandler) extractCandidateObjects(items []any) []string { +func extractCandidateObjectCodes(items []toolSearchCandidateResult) []string { if len(items) == 0 { return nil } ret := make([]string, 0, len(items)) for _, item := range items { - obj, ok := item.(map[string]any) - if !ok { - continue - } - toolCode := strings.TrimSpace(readToolSearchString(obj, "toolCode")) + toolCode := strings.TrimSpace(item.ToolCode) if toolCode == "" { continue } diff --git a/internal/ai/runtime/internal/impl/callbacks/agent_trace_handler_test.go b/internal/ai/runtime/internal/impl/callbacks/agent_trace_handler_test.go new file mode 100644 index 0000000..591dcfc --- /dev/null +++ b/internal/ai/runtime/internal/impl/callbacks/agent_trace_handler_test.go @@ -0,0 +1,38 @@ +package callbacks + +import ( + "testing" + + "cs-agent/internal/pkg/toolx" +) + +func TestParseGraphToolOutcome(t *testing.T) { + action, risk, ready := parseGraphToolOutcome(toolx.GraphAnalyzeConversation.Code, `{"recommendedNextAction":"handoff_to_human","riskLevel":"high"}`) + if action != "handoff_to_human" || risk != "high" || ready { + t.Fatalf("unexpected analyze graph outcome: %q %q %v", action, risk, ready) + } + + action, risk, ready = parseGraphToolOutcome(toolx.GraphTriageServiceRequest.Code, `{"recommendedAction":"prepare_ticket","analysis":{"riskLevel":"medium"},"ticketDraft":{"ready":true}}`) + if action != "prepare_ticket" || risk != "medium" || !ready { + t.Fatalf("unexpected triage graph outcome: %q %q %v", action, risk, ready) + } +} + +func TestExtractCandidateToolCodes(t *testing.T) { + handler := &RuntimeTraceHandler{ + toolMetadataBy: map[string]ToolMetadata{ + "tool_search": {ToolCode: toolx.BuiltinToolSearch.Code, ToolName: toolx.BuiltinToolSearch.Name}, + "foo_model": {ToolCode: "mcp/server/foo", ToolName: "foo"}, + }, + } + + got := handler.extractCandidateToolCodes(`{"selectedTools":["foo_model"]}`) + if len(got) != 1 || got[0] != "mcp/server/foo" { + t.Fatalf("unexpected selectedTools codes: %#v", got) + } + + got = handler.extractCandidateToolCodes(`{"candidates":[{"toolCode":"mcp/server/bar"}]}`) + if len(got) != 1 || got[0] != "mcp/server/bar" { + t.Fatalf("unexpected candidate codes: %#v", got) + } +} diff --git a/internal/ai/runtime/reply_interrupt_helpers.go b/internal/ai/runtime/reply_interrupt_helpers.go index 3f3c392..15f659b 100644 --- a/internal/ai/runtime/reply_interrupt_helpers.go +++ b/internal/ai/runtime/reply_interrupt_helpers.go @@ -9,6 +9,10 @@ import ( svc "cs-agent/internal/services" ) +type interruptMessagePreview struct { + Message string `json:"message"` +} + func buildConversationInterrupt(conversation models.Conversation, message models.Message, aiAgent models.AIAgent, summary *Summary) *models.ConversationInterrupt { if summary == nil { return nil @@ -50,14 +54,11 @@ func extractInterruptMessage(infoPreview string) string { if infoPreview == "" { return "" } - payload := make(map[string]any) + var payload interruptMessagePreview if err := json.Unmarshal([]byte(infoPreview), &payload); err != nil { return "" } - if message, ok := payload["message"].(string); ok { - return strings.TrimSpace(message) - } - return "" + return strings.TrimSpace(payload.Message) } func firstInterruptID(summary *Summary) string {