diff --git a/internal/ai/workflow/dsl/types.go b/internal/ai/workflow/dsl/types.go new file mode 100644 index 0000000..1812bed --- /dev/null +++ b/internal/ai/workflow/dsl/types.go @@ -0,0 +1,34 @@ +package dsl + +import "encoding/json" + +type Definition struct { + SchemaVersion int `json:"schemaVersion"` + EntryNodeID string `json:"entryNodeId"` + Nodes []Node `json:"nodes"` + Edges []Edge `json:"edges"` +} + +type Node struct { + ID string `json:"id"` + Type string `json:"type"` + Name string `json:"name"` + Position Position `json:"position"` + Config json.RawMessage `json:"config"` +} + +type Position struct { + X float64 `json:"x"` + Y float64 `json:"y"` +} + +type Edge struct { + ID string `json:"id"` + Source string `json:"source"` + Target string `json:"target"` + Condition *Condition `json:"condition,omitempty"` +} + +type Condition struct { + Expression string `json:"expression"` +} diff --git a/internal/ai/workflow/registry/registry.go b/internal/ai/workflow/registry/registry.go new file mode 100644 index 0000000..13465c7 --- /dev/null +++ b/internal/ai/workflow/registry/registry.go @@ -0,0 +1,33 @@ +package registry + +const ( + NodeTypeStart = "start" + NodeTypeKnowledgeRetrieve = "knowledge_retrieve" + NodeTypeAnswerabilityGate = "answerability_gate" + NodeTypeLLMReply = "llm_reply" + NodeTypeCondition = "condition" + NodeTypeAnalyzeConversation = "analyze_conversation" + NodeTypePrepareTicketDraft = "prepare_ticket_draft" + NodeTypeHumanConfirm = "human_confirm" + NodeTypeCreateTicket = "create_ticket" + NodeTypeHandoffToHuman = "handoff_to_human" + NodeTypeSendReply = "send_reply" + NodeTypeEnd = "end" +) + +func DefaultRegistry() *Registry { + return NewRegistry( + NodeSpec{Type: NodeTypeStart, Title: "Start", Description: "Conversation workflow entry.", RiskLevel: NodeRiskLevelLow}, + NodeSpec{Type: NodeTypeKnowledgeRetrieve, Title: "Knowledge Retrieve", Description: "Retrieve knowledge for the current user message.", RiskLevel: NodeRiskLevelLow}, + NodeSpec{Type: NodeTypeAnswerabilityGate, Title: "Answerability Gate", Description: "Decide whether retrieved knowledge is enough to answer.", RiskLevel: NodeRiskLevelLow}, + NodeSpec{Type: NodeTypeLLMReply, Title: "LLM Reply", Description: "Generate a reply or structured analysis with the configured model.", RiskLevel: NodeRiskLevelMedium}, + NodeSpec{Type: NodeTypeCondition, Title: "Condition", Description: "Route by controlled workflow variables.", RiskLevel: NodeRiskLevelLow}, + NodeSpec{Type: NodeTypeAnalyzeConversation, Title: "Analyze Conversation", Description: "Analyze intent, risk, and recommended next action.", RiskLevel: NodeRiskLevelLow}, + NodeSpec{Type: NodeTypePrepareTicketDraft, Title: "Prepare Ticket Draft", Description: "Build a ticket draft from conversation context.", RiskLevel: NodeRiskLevelMedium}, + NodeSpec{Type: NodeTypeHumanConfirm, Title: "Human Confirm", Description: "Interrupt and wait for explicit user confirmation.", RiskLevel: NodeRiskLevelMedium, Interruptible: true}, + NodeSpec{Type: NodeTypeCreateTicket, Title: "Create Ticket", Description: "Create a ticket from a confirmed draft.", RiskLevel: NodeRiskLevelHigh, RequiresConfirmationPredecessor: true}, + NodeSpec{Type: NodeTypeHandoffToHuman, Title: "Handoff To Human", Description: "Transfer the conversation to human support.", RiskLevel: NodeRiskLevelHigh, RequiresConfirmationPredecessor: true}, + NodeSpec{Type: NodeTypeSendReply, Title: "Send Reply", Description: "Return or commit customer-visible reply text.", RiskLevel: NodeRiskLevelLow}, + NodeSpec{Type: NodeTypeEnd, Title: "End", Description: "End workflow execution.", RiskLevel: NodeRiskLevelLow}, + ) +} diff --git a/internal/ai/workflow/registry/spec.go b/internal/ai/workflow/registry/spec.go new file mode 100644 index 0000000..0f98acb --- /dev/null +++ b/internal/ai/workflow/registry/spec.go @@ -0,0 +1,53 @@ +package registry + +type NodeRiskLevel string + +const ( + NodeRiskLevelLow NodeRiskLevel = "low" + NodeRiskLevelMedium NodeRiskLevel = "medium" + NodeRiskLevelHigh NodeRiskLevel = "high" +) + +type NodeSpec struct { + Type string `json:"type"` + Title string `json:"title"` + Description string `json:"description"` + RiskLevel NodeRiskLevel `json:"riskLevel"` + Interruptible bool `json:"interruptible"` + RequiresConfirmationPredecessor bool `json:"requiresConfirmationPredecessor"` +} + +type Registry struct { + specsByType map[string]NodeSpec + specs []NodeSpec +} + +func NewRegistry(specs ...NodeSpec) *Registry { + ret := &Registry{ + specsByType: make(map[string]NodeSpec, len(specs)), + specs: make([]NodeSpec, 0, len(specs)), + } + for _, spec := range specs { + if spec.Type == "" { + continue + } + ret.specsByType[spec.Type] = spec + ret.specs = append(ret.specs, spec) + } + return ret +} + +func (r *Registry) Get(nodeType string) (NodeSpec, bool) { + if r == nil { + return NodeSpec{}, false + } + spec, ok := r.specsByType[nodeType] + return spec, ok +} + +func (r *Registry) List() []NodeSpec { + if r == nil { + return nil + } + return append([]NodeSpec(nil), r.specs...) +} diff --git a/internal/ai/workflow/validator/validator.go b/internal/ai/workflow/validator/validator.go new file mode 100644 index 0000000..2c9e7e7 --- /dev/null +++ b/internal/ai/workflow/validator/validator.go @@ -0,0 +1,207 @@ +package validator + +import ( + "fmt" + "strings" + + "agent-desk/internal/ai/workflow/dsl" + "agent-desk/internal/ai/workflow/registry" +) + +type Error struct { + Field string `json:"field"` + Message string `json:"message"` +} + +type Result struct { + Valid bool `json:"valid"` + Errors []Error `json:"errors"` +} + +func ValidateDefinition(def dsl.Definition, reg *registry.Registry) Result { + if reg == nil { + reg = registry.DefaultRegistry() + } + v := definitionValidator{ + def: def, + registry: reg, + nodesByID: make(map[string]dsl.Node, len(def.Nodes)), + outgoing: make(map[string][]string), + incoming: make(map[string][]string), + startNodeIDs: make([]string, 0, 1), + endNodeIDs: make([]string, 0, 1), + } + v.validate() + return Result{ + Valid: len(v.errors) == 0, + Errors: v.errors, + } +} + +type definitionValidator struct { + def dsl.Definition + registry *registry.Registry + nodesByID map[string]dsl.Node + outgoing map[string][]string + incoming map[string][]string + startNodeIDs []string + endNodeIDs []string + errors []Error +} + +func (v *definitionValidator) validate() { + v.validateNodes() + v.validateEdges() + v.validateEntry() + v.validateReachability() + v.validateConfirmationGuards() +} + +func (v *definitionValidator) validateNodes() { + for index, node := range v.def.Nodes { + node.ID = strings.TrimSpace(node.ID) + node.Type = strings.TrimSpace(node.Type) + field := fmt.Sprintf("nodes[%d]", index) + if node.ID == "" { + v.addError(field+".id", "node id is required") + continue + } + if _, exists := v.nodesByID[node.ID]; exists { + v.addError(field+".id", "duplicate node id: "+node.ID) + continue + } + v.nodesByID[node.ID] = node + if node.Type == "" { + v.addError(field+".type", "node type is required") + continue + } + if _, ok := v.registry.Get(node.Type); !ok { + v.addError(field+".type", "unknown node type: "+node.Type) + continue + } + switch node.Type { + case registry.NodeTypeStart: + v.startNodeIDs = append(v.startNodeIDs, node.ID) + case registry.NodeTypeEnd: + v.endNodeIDs = append(v.endNodeIDs, node.ID) + } + } + if len(v.startNodeIDs) != 1 { + v.addError("nodes", "workflow must contain exactly one start node") + } + if len(v.endNodeIDs) == 0 { + v.addError("nodes", "workflow must contain at least one end node") + } +} + +func (v *definitionValidator) validateEdges() { + seen := make(map[string]struct{}, len(v.def.Edges)) + for index, edge := range v.def.Edges { + edge.ID = strings.TrimSpace(edge.ID) + edge.Source = strings.TrimSpace(edge.Source) + edge.Target = strings.TrimSpace(edge.Target) + field := fmt.Sprintf("edges[%d]", index) + if edge.ID == "" { + v.addError(field+".id", "edge id is required") + } else if _, exists := seen[edge.ID]; exists { + v.addError(field+".id", "duplicate edge id: "+edge.ID) + } + seen[edge.ID] = struct{}{} + if edge.Source == "" { + v.addError(field+".source", "edge source is required") + } else if _, ok := v.nodesByID[edge.Source]; !ok { + v.addError(field+".source", "edge source node does not exist: "+edge.Source) + } + if edge.Target == "" { + v.addError(field+".target", "edge target is required") + } else if _, ok := v.nodesByID[edge.Target]; !ok { + v.addError(field+".target", "edge target node does not exist: "+edge.Target) + } + if edge.Source != "" && edge.Target != "" { + v.outgoing[edge.Source] = append(v.outgoing[edge.Source], edge.Target) + v.incoming[edge.Target] = append(v.incoming[edge.Target], edge.Source) + } + } +} + +func (v *definitionValidator) validateEntry() { + entryNodeID := strings.TrimSpace(v.def.EntryNodeID) + if entryNodeID == "" { + v.addError("entryNodeId", "entry node id is required") + return + } + entry, ok := v.nodesByID[entryNodeID] + if !ok { + v.addError("entryNodeId", "entry node does not exist: "+entryNodeID) + return + } + if entry.Type != registry.NodeTypeStart { + v.addError("entryNodeId", "entry node must be the start node") + } +} + +func (v *definitionValidator) validateReachability() { + entryNodeID := strings.TrimSpace(v.def.EntryNodeID) + if entryNodeID == "" { + return + } + if _, ok := v.nodesByID[entryNodeID]; !ok { + return + } + reachable := make(map[string]struct{}, len(v.nodesByID)) + queue := []string{entryNodeID} + for len(queue) > 0 { + current := queue[0] + queue = queue[1:] + if _, exists := reachable[current]; exists { + continue + } + reachable[current] = struct{}{} + for _, target := range v.outgoing[current] { + if _, exists := reachable[target]; !exists { + queue = append(queue, target) + } + } + } + for id := range v.nodesByID { + if _, ok := reachable[id]; !ok { + v.addError("nodes", "node is not reachable from entry node: "+id) + } + } +} + +func (v *definitionValidator) validateConfirmationGuards() { + for id, node := range v.nodesByID { + spec, ok := v.registry.Get(node.Type) + if !ok || !spec.RequiresConfirmationPredecessor { + continue + } + if !v.hasConfirmationPredecessor(id, make(map[string]struct{})) { + v.addError("nodes."+id, node.Type+" requires human_confirm before execution") + } + } +} + +func (v *definitionValidator) hasConfirmationPredecessor(nodeID string, visiting map[string]struct{}) bool { + if _, seen := visiting[nodeID]; seen { + return false + } + visiting[nodeID] = struct{}{} + for _, source := range v.incoming[nodeID] { + node, ok := v.nodesByID[source] + if !ok { + continue + } + if node.Type == registry.NodeTypeHumanConfirm { + return true + } + if v.hasConfirmationPredecessor(source, visiting) { + return true + } + } + return false +} + +func (v *definitionValidator) addError(field string, message string) { + v.errors = append(v.errors, Error{Field: field, Message: message}) +} diff --git a/internal/ai/workflow/validator/validator_test.go b/internal/ai/workflow/validator/validator_test.go new file mode 100644 index 0000000..065674f --- /dev/null +++ b/internal/ai/workflow/validator/validator_test.go @@ -0,0 +1,129 @@ +package validator_test + +import ( + "encoding/json" + "strings" + "testing" + + "agent-desk/internal/ai/workflow/dsl" + "agent-desk/internal/ai/workflow/registry" + "agent-desk/internal/ai/workflow/validator" +) + +func TestValidateDefinitionAcceptsMinimalConversationFlow(t *testing.T) { + result := validator.ValidateDefinition(minimalDefinition(), registry.DefaultRegistry()) + + if !result.Valid { + t.Fatalf("expected valid definition, got errors: %#v", result.Errors) + } +} + +func TestValidateDefinitionRejectsMissingStart(t *testing.T) { + def := minimalDefinition() + def.Nodes = []dsl.Node{ + {ID: "reply_1", Type: "send_reply", Config: json.RawMessage(`{"text":"hello"}`)}, + {ID: "end_1", Type: "end"}, + } + + result := validator.ValidateDefinition(def, registry.DefaultRegistry()) + + if result.Valid { + t.Fatalf("expected missing start to be invalid") + } + if !hasValidationMessage(result, "exactly one start node") { + t.Fatalf("expected start error, got %#v", result.Errors) + } +} + +func TestValidateDefinitionRejectsUnknownNodeType(t *testing.T) { + def := minimalDefinition() + def.Nodes = append(def.Nodes, dsl.Node{ID: "unknown_1", Type: "unknown_node"}) + def.Edges = append(def.Edges, dsl.Edge{ID: "e3", Source: "reply_1", Target: "unknown_1"}) + + result := validator.ValidateDefinition(def, registry.DefaultRegistry()) + + if result.Valid { + t.Fatalf("expected unknown node type to be invalid") + } + if !hasValidationMessage(result, "unknown node type") { + t.Fatalf("expected unknown-node error, got %#v", result.Errors) + } +} + +func TestValidateDefinitionRejectsUnguardedCreateTicket(t *testing.T) { + def := dsl.Definition{ + SchemaVersion: 1, + EntryNodeID: "start_1", + Nodes: []dsl.Node{ + {ID: "start_1", Type: "start"}, + {ID: "draft_1", Type: "prepare_ticket_draft"}, + {ID: "create_1", Type: "create_ticket"}, + {ID: "end_1", Type: "end"}, + }, + Edges: []dsl.Edge{ + {ID: "e1", Source: "start_1", Target: "draft_1"}, + {ID: "e2", Source: "draft_1", Target: "create_1"}, + {ID: "e3", Source: "create_1", Target: "end_1"}, + }, + } + + result := validator.ValidateDefinition(def, registry.DefaultRegistry()) + + if result.Valid { + t.Fatalf("expected unguarded create_ticket to be invalid") + } + if !hasValidationMessage(result, "requires human_confirm") { + t.Fatalf("expected confirmation guard error, got %#v", result.Errors) + } +} + +func TestValidateDefinitionAcceptsConfirmedCreateTicket(t *testing.T) { + def := dsl.Definition{ + SchemaVersion: 1, + EntryNodeID: "start_1", + Nodes: []dsl.Node{ + {ID: "start_1", Type: "start"}, + {ID: "draft_1", Type: "prepare_ticket_draft"}, + {ID: "confirm_1", Type: "human_confirm"}, + {ID: "create_1", Type: "create_ticket"}, + {ID: "end_1", Type: "end"}, + }, + Edges: []dsl.Edge{ + {ID: "e1", Source: "start_1", Target: "draft_1"}, + {ID: "e2", Source: "draft_1", Target: "confirm_1"}, + {ID: "e3", Source: "confirm_1", Target: "create_1"}, + {ID: "e4", Source: "create_1", Target: "end_1"}, + }, + } + + result := validator.ValidateDefinition(def, registry.DefaultRegistry()) + + if !result.Valid { + t.Fatalf("expected confirmed create_ticket to be valid, got %#v", result.Errors) + } +} + +func minimalDefinition() dsl.Definition { + return dsl.Definition{ + SchemaVersion: 1, + EntryNodeID: "start_1", + Nodes: []dsl.Node{ + {ID: "start_1", Type: "start"}, + {ID: "reply_1", Type: "send_reply", Config: json.RawMessage(`{"text":"hello"}`)}, + {ID: "end_1", Type: "end"}, + }, + Edges: []dsl.Edge{ + {ID: "e1", Source: "start_1", Target: "reply_1"}, + {ID: "e2", Source: "reply_1", Target: "end_1"}, + }, + } +} + +func hasValidationMessage(result validator.Result, want string) bool { + for _, item := range result.Errors { + if strings.Contains(item.Message, want) { + return true + } + } + return false +}