From ca625396a0ca9d71630b6ca91b98a4dba32c4262 Mon Sep 17 00:00:00 2001 From: mlogclub Date: Mon, 29 Jun 2026 12:00:33 +0800 Subject: [PATCH] feat: enhance condition branch validation by adding port edge checks and updating related functions --- docs | 2 +- internal/ai/workflow/validator/validator.go | 10 ++++--- .../ai/workflow/validator/validator_test.go | 26 +++++++++++++++++-- .../ai_agent_workflow_service_test.go | 21 +++++++++++++++ internal/services/ai_workflow_service.go | 15 ++++++----- 5 files changed, 61 insertions(+), 13 deletions(-) diff --git a/docs b/docs index 5b69ec7..7e5f637 160000 --- a/docs +++ b/docs @@ -1 +1 @@ -Subproject commit 5b69ec75398e40a5b186f6b66efb652229b7053f +Subproject commit 7e5f63704304315e1dd743e02b23f9915a1e8839 diff --git a/internal/ai/workflow/validator/validator.go b/internal/ai/workflow/validator/validator.go index 6035476..037782a 100644 --- a/internal/ai/workflow/validator/validator.go +++ b/internal/ai/workflow/validator/validator.go @@ -285,7 +285,7 @@ func (v *definitionValidator) validateConditions() { } else if _, ok := v.nodesByID[targetNodeID]; !ok { v.addError(branchField+".targetNodeId", "condition branch target node does not exist: "+targetNodeID) } - if !v.hasEdgeTo(strings.TrimSpace(node.ID), targetNodeID) { + if !v.hasConditionBranchEdge(strings.TrimSpace(node.ID), targetNodeID, branchID) { v.addError(branchField+".targetNodeId", "condition branch target must have an outgoing edge: "+targetNodeID) } if branch.Default { @@ -441,12 +441,14 @@ func (v *definitionValidator) hasPath(sourceID string, targetID string, visiting return false } -func (v *definitionValidator) hasEdgeTo(sourceID string, targetID string) bool { - if sourceID == "" || targetID == "" { +func (v *definitionValidator) hasConditionBranchEdge(sourceID string, targetID string, sourcePortID string) bool { + if sourceID == "" || targetID == "" || sourcePortID == "" { return true } for _, edge := range v.def.Edges { - if strings.TrimSpace(edge.SourceNodeID) == sourceID && strings.TrimSpace(edge.TargetNodeID) == targetID { + if strings.TrimSpace(edge.SourceNodeID) == sourceID && + strings.TrimSpace(edge.TargetNodeID) == targetID && + strings.TrimSpace(edge.SourcePortID) == sourcePortID { return true } } diff --git a/internal/ai/workflow/validator/validator_test.go b/internal/ai/workflow/validator/validator_test.go index cbd0586..1dfdabb 100644 --- a/internal/ai/workflow/validator/validator_test.go +++ b/internal/ai/workflow/validator/validator_test.go @@ -176,6 +176,24 @@ func TestValidateDefinitionRejectsConditionBranchTargetWithoutEdge(t *testing.T) } } +func TestValidateDefinitionRejectsConditionBranchTargetWithoutPortEdge(t *testing.T) { + def := conditionDefinition() + def.Edges = []dsl.Edge{ + edge("start_1", "condition_1"), + edge("condition_1", "end_1"), + portEdge("condition_1", "end_1", "default"), + } + + result := validator.ValidateDefinition(def, registry.DefaultRegistry()) + + if result.Valid { + t.Fatalf("expected condition branch target without matching port edge to be invalid") + } + if !hasValidationMessage(result, "condition branch target must have an outgoing edge") { + t.Fatalf("expected branch port edge error, got %#v", result.Errors) + } +} + func TestValidateDefinitionRejectsUnknownConditionVariable(t *testing.T) { def := conditionDefinition() var config dsl.ConditionConfig @@ -240,8 +258,8 @@ func conditionDefinition() dsl.Definition { }, Edges: []dsl.Edge{ edge("start_1", "condition_1"), - edge("condition_1", "end_1"), - edge("condition_1", "end_1"), + portEdge("condition_1", "end_1", "hello"), + portEdge("condition_1", "end_1", "default"), }, } } @@ -263,6 +281,10 @@ func edge(source string, target string) dsl.Edge { return dsl.Edge{SourceNodeID: source, TargetNodeID: target} } +func portEdge(source string, target string, sourcePortID string) dsl.Edge { + return dsl.Edge{SourceNodeID: source, TargetNodeID: target, SourcePortID: sourcePortID} +} + func inputs(name string, value dsl.Value) map[string]dsl.Value { return map[string]dsl.Value{name: value} } diff --git a/internal/services/ai_agent_workflow_service_test.go b/internal/services/ai_agent_workflow_service_test.go index c5d11f8..659086d 100644 --- a/internal/services/ai_agent_workflow_service_test.go +++ b/internal/services/ai_agent_workflow_service_test.go @@ -89,6 +89,9 @@ func TestAIAgentServiceCreatesDefaultWorkflow(t *testing.T) { if !workflowEdgeExists(stored, "create_ticket_1", "ticket_result_reply_1") { t.Fatalf("expected create_ticket to flow into a customer-visible result reply") } + assertConditionBranchesHavePortEdges(t, stored, "policy_route_1") + assertConditionBranchesHavePortEdges(t, stored, "ticket_confirm_route_1") + assertConditionBranchesHavePortEdges(t, stored, "answerability_route_1") } func TestAIWorkflowServiceDefaultAgentWorkflowDefinitionIsValid(t *testing.T) { @@ -305,6 +308,24 @@ func conditionBranches(t *testing.T, def dsl.Definition, nodeID string) []dsl.Co return nil } +func assertConditionBranchesHavePortEdges(t *testing.T, def dsl.Definition, nodeID string) { + t.Helper() + for _, branch := range conditionBranches(t, def, nodeID) { + if !workflowPortEdgeExists(def, nodeID, branch.TargetNodeID, branch.ID) { + t.Fatalf("expected condition branch %s.%s to have port edge to %s", nodeID, branch.ID, branch.TargetNodeID) + } + } +} + +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 { + return true + } + } + return false +} + func workflowEdgeExists(def dsl.Definition, sourceID string, targetID string) bool { for _, edge := range def.Edges { if edge.SourceNodeID == sourceID && edge.TargetNodeID == targetID { diff --git a/internal/services/ai_workflow_service.go b/internal/services/ai_workflow_service.go index c027cec..c8d523b 100644 --- a/internal/services/ai_workflow_service.go +++ b/internal/services/ai_workflow_service.go @@ -434,17 +434,15 @@ func defaultAgentWorkflowDefinition() dsl.Definition { "riskSignals": dsl.RefValue("understanding_1", "riskSignals"), }, nil), workflowNode("policy_route_1", workflowregistry.NodeTypeCondition, "策略分流", 1560, 125.5, nil, dsl.ConditionConfig{Branches: []dsl.ConditionBranch{ - 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("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"), - {ID: "default", Name: "默认澄清", TargetNodeID: "policy_reply_1", Default: true}, + 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"), + {ID: "default", Name: "策略兜底", TargetNodeID: "policy_reply_1", Default: true}, }}), - workflowNode("policy_reply_1", workflowregistry.NodeTypeSendReply, "发送策略回复", 4320, 98.5, workflowInputs("replyText", "policy_1", "replyText"), nil), workflowNode("handoff_1", workflowregistry.NodeTypeHandoffToHuman, "转人工", 2020, 0, workflowInputs("reason", "start_1", "userMessage"), 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), @@ -459,6 +457,8 @@ 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"), @@ -485,9 +485,12 @@ func defaultAgentWorkflowDefinition() dsl.Definition { workflowEdge("understanding_1", "policy_1"), workflowEdge("policy_1", "policy_route_1"), 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"), workflowEdge("policy_reply_1", "end_1"), workflowEdge("handoff_1", "handoff_end_1"), workflowEdge("draft_ticket_1", "ticket_confirm_prompt_1"),