From 9be449c882dc92973e41636649975ba018e7698c Mon Sep 17 00:00:00 2001 From: mlogclub Date: Mon, 29 Jun 2026 13:58:13 +0800 Subject: [PATCH] feat: add assertions for condition branch and port edge order in workflow tests --- .../ai_agent_workflow_service_test.go | 49 +++++++++++++++++++ internal/services/ai_workflow_service.go | 10 ++-- 2 files changed, 54 insertions(+), 5 deletions(-) diff --git a/internal/services/ai_agent_workflow_service_test.go b/internal/services/ai_agent_workflow_service_test.go index 659086d..94bcdfd 100644 --- a/internal/services/ai_agent_workflow_service_test.go +++ b/internal/services/ai_agent_workflow_service_test.go @@ -92,6 +92,24 @@ func TestAIAgentServiceCreatesDefaultWorkflow(t *testing.T) { assertConditionBranchesHavePortEdges(t, stored, "policy_route_1") assertConditionBranchesHavePortEdges(t, stored, "ticket_confirm_route_1") assertConditionBranchesHavePortEdges(t, stored, "answerability_route_1") + assertConditionBranchOrder(t, stored, "policy_route_1", []string{ + "handoff", + "direct", + "clarify", + "end_conversation", + "ticket", + "knowledge", + "default", + }) + assertConditionPortEdgeOrder(t, stored, "policy_route_1", []string{ + "handoff", + "direct", + "clarify", + "end_conversation", + "ticket", + "knowledge", + "default", + }) } func TestAIWorkflowServiceDefaultAgentWorkflowDefinitionIsValid(t *testing.T) { @@ -317,6 +335,37 @@ func assertConditionBranchesHavePortEdges(t *testing.T, def dsl.Definition, node } } +func assertConditionBranchOrder(t *testing.T, def dsl.Definition, nodeID string, want []string) { + t.Helper() + branches := conditionBranches(t, def, nodeID) + if len(branches) != len(want) { + t.Fatalf("expected %s branch order %v, got %#v", nodeID, want, branches) + } + for index, branch := range branches { + if branch.ID != want[index] { + t.Fatalf("expected %s branch order %v, got branch %d = %s", nodeID, want, index, branch.ID) + } + } +} + +func assertConditionPortEdgeOrder(t *testing.T, def dsl.Definition, nodeID string, want []string) { + t.Helper() + got := make([]string, 0, len(want)) + for _, edge := range def.Edges { + if edge.SourceNodeID == nodeID { + got = append(got, edge.SourcePortID) + } + } + if len(got) != len(want) { + t.Fatalf("expected %s port edge order %v, got %v", nodeID, want, got) + } + for index, sourcePortID := range got { + if sourcePortID != want[index] { + t.Fatalf("expected %s port edge order %v, got edge %d = %s", nodeID, want, index, sourcePortID) + } + } +} + func workflowPortEdgeExists(def dsl.Definition, sourceID string, targetID string, sourcePortID string) bool { for _, edge := range def.Edges { if edge.SourceNodeID == sourceID && edge.TargetNodeID == targetID && edge.SourcePortID == sourcePortID { diff --git a/internal/services/ai_workflow_service.go b/internal/services/ai_workflow_service.go index c8d523b..def3cec 100644 --- a/internal/services/ai_workflow_service.go +++ b/internal/services/ai_workflow_service.go @@ -435,14 +435,16 @@ func defaultAgentWorkflowDefinition() dsl.Definition { }, nil), workflowNode("policy_route_1", workflowregistry.NodeTypeCondition, "策略分流", 1560, 125.5, nil, dsl.ConditionConfig{Branches: []dsl.ConditionBranch{ workflowConditionBranch("handoff", "转人工", "handoff_1", "policy_1", "action", "eq", "handoff_to_human"), - workflowConditionBranch("ticket", "创建工单", "draft_ticket_1", "policy_1", "action", "eq", "prepare_ticket"), - workflowConditionBranch("knowledge", "知识库回复", "retrieve_1", "policy_1", "action", "eq", "retrieve_knowledge"), workflowConditionBranch("direct", "直接回复", "policy_reply_1", "policy_1", "action", "eq", "direct_reply"), workflowConditionBranch("clarify", "追问澄清", "policy_reply_1", "policy_1", "action", "eq", "clarify"), workflowConditionBranch("end_conversation", "结束语", "policy_reply_1", "policy_1", "action", "eq", "end_conversation"), + workflowConditionBranch("ticket", "创建工单", "draft_ticket_1", "policy_1", "action", "eq", "prepare_ticket"), + workflowConditionBranch("knowledge", "知识库回复", "retrieve_1", "policy_1", "action", "eq", "retrieve_knowledge"), {ID: "default", Name: "策略兜底", TargetNodeID: "policy_reply_1", Default: true}, }}), workflowNode("handoff_1", workflowregistry.NodeTypeHandoffToHuman, "转人工", 2020, 0, workflowInputs("reason", "start_1", "userMessage"), nil), + workflowNode("policy_reply_1", workflowregistry.NodeTypeSendReply, "发送策略回复", 4320, 98.5, workflowInputs("replyText", "policy_1", "replyText"), nil), + workflowNode("handoff_end_1", workflowregistry.NodeTypeEnd, "结束", 2480, 0, nil, nil), workflowNode("draft_ticket_1", workflowregistry.NodeTypePrepareTicketDraft, "整理工单草稿", 2020, 379, workflowInputs("issue", "start_1", "userMessage"), nil), workflowNode("ticket_confirm_prompt_1", workflowregistry.NodeTypeLLMReply, "建单确认文案", 2480, 379, workflowInputs("userMessage", "start_1", "userMessage"), map[string]any{"staticReply": "我已整理工单草稿。请回复“确认”创建工单,或回复“取消”放弃。"}), workflowNode("ticket_confirm_1", workflowregistry.NodeTypeHumanConfirm, "确认建单", 2940, 379, workflowInputs("prompt", "ticket_confirm_prompt_1", "replyText"), nil), @@ -457,8 +459,6 @@ func defaultAgentWorkflowDefinition() dsl.Definition { workflowNode("ticket_result_reply_1", workflowregistry.NodeTypeSendReply, "发送建单结果", 4320, 285.5, workflowInputs("replyText", "create_ticket_1", "message"), nil), workflowNode("ticket_cancel_reply_1", workflowregistry.NodeTypeLLMReply, "取消建单提示", 3860, 472.5, workflowInputs("userMessage", "start_1", "userMessage"), map[string]any{"staticReply": "已取消创建工单。你可以继续补充问题,我会继续帮你处理。"}), workflowNode("send_ticket_cancel_1", workflowregistry.NodeTypeSendReply, "发送取消提示", 4320, 472.5, workflowInputs("replyText", "ticket_cancel_reply_1", "replyText"), nil), - workflowNode("policy_reply_1", workflowregistry.NodeTypeSendReply, "发送策略回复", 4320, 98.5, workflowInputs("replyText", "policy_1", "replyText"), nil), - workflowNode("handoff_end_1", workflowregistry.NodeTypeEnd, "结束", 2480, 0, nil, nil), workflowNode("retrieve_1", workflowregistry.NodeTypeKnowledgeRetrieve, "知识检索", 2480, 753, workflowInputs("query", "start_1", "userMessage"), nil), workflowNode("answerability_1", workflowregistry.NodeTypeAnswerabilityGate, "可回答判断", 2940, 753, map[string]dsl.Value{ "userMessage": dsl.RefValue("start_1", "userMessage"), @@ -484,10 +484,10 @@ func defaultAgentWorkflowDefinition() dsl.Definition { workflowEdge("start_1", "understanding_1"), workflowEdge("understanding_1", "policy_1"), workflowEdge("policy_1", "policy_route_1"), + workflowPortEdge("policy_route_1", "handoff_1", "handoff"), workflowPortEdge("policy_route_1", "policy_reply_1", "direct"), workflowPortEdge("policy_route_1", "policy_reply_1", "clarify"), workflowPortEdge("policy_route_1", "policy_reply_1", "end_conversation"), - workflowPortEdge("policy_route_1", "handoff_1", "handoff"), workflowPortEdge("policy_route_1", "draft_ticket_1", "ticket"), workflowPortEdge("policy_route_1", "retrieve_1", "knowledge"), workflowPortEdge("policy_route_1", "policy_reply_1", "default"),