409 lines
11 KiB
Go
409 lines
11 KiB
Go
package workflow
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"reflect"
|
|
"strconv"
|
|
"strings"
|
|
|
|
"agent-desk/internal/ai"
|
|
"agent-desk/internal/ai/runtime/internal/impl/retrievers"
|
|
"agent-desk/internal/ai/workflow/dsl"
|
|
workflowregistry "agent-desk/internal/ai/workflow/registry"
|
|
"agent-desk/internal/models"
|
|
"agent-desk/internal/pkg/utils"
|
|
)
|
|
|
|
const maxWorkflowSteps = 128
|
|
|
|
type Input struct {
|
|
Definition dsl.Definition
|
|
Conversation models.Conversation
|
|
UserMessage models.Message
|
|
AIAgent models.AIAgent
|
|
AIConfig models.AIConfig
|
|
}
|
|
|
|
type Result struct {
|
|
Status string
|
|
ReplyText string
|
|
NodePath []string
|
|
PromptTokens int
|
|
CompletionTokens int
|
|
RetrieverCount int
|
|
TraceData string
|
|
}
|
|
|
|
type Executor struct{}
|
|
|
|
func NewExecutor() *Executor {
|
|
return &Executor{}
|
|
}
|
|
|
|
type runState struct {
|
|
input Input
|
|
nodesByID map[string]dsl.Node
|
|
outgoing map[string][]dsl.Edge
|
|
vars map[string]map[string]any
|
|
result Result
|
|
}
|
|
|
|
func (e *Executor) Execute(ctx context.Context, input Input) (*Result, error) {
|
|
state := newRunState(input)
|
|
currentID := strings.TrimSpace(input.Definition.EntryNodeID)
|
|
if currentID == "" {
|
|
return nil, fmt.Errorf("workflow entry node is required")
|
|
}
|
|
for step := 0; step < maxWorkflowSteps; step++ {
|
|
node, ok := state.nodesByID[currentID]
|
|
if !ok {
|
|
return nil, fmt.Errorf("workflow node does not exist: %s", currentID)
|
|
}
|
|
state.result.NodePath = append(state.result.NodePath, node.ID)
|
|
if err := e.executeNode(ctx, state, node); err != nil {
|
|
return nil, err
|
|
}
|
|
if node.Type == workflowregistry.NodeTypeEnd {
|
|
state.result.Status = "completed"
|
|
return &state.result, nil
|
|
}
|
|
nextID, ok, err := state.nextNodeID(node.ID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if !ok {
|
|
state.result.Status = "completed"
|
|
return &state.result, nil
|
|
}
|
|
currentID = nextID
|
|
}
|
|
return nil, fmt.Errorf("workflow exceeded max steps")
|
|
}
|
|
|
|
func newRunState(input Input) *runState {
|
|
state := &runState{
|
|
input: input,
|
|
nodesByID: make(map[string]dsl.Node, len(input.Definition.Nodes)),
|
|
outgoing: make(map[string][]dsl.Edge),
|
|
vars: make(map[string]map[string]any),
|
|
result: Result{
|
|
Status: "started",
|
|
NodePath: make([]string, 0),
|
|
},
|
|
}
|
|
for _, node := range input.Definition.Nodes {
|
|
node.ID = strings.TrimSpace(node.ID)
|
|
node.Type = strings.TrimSpace(node.Type)
|
|
if node.ID != "" {
|
|
state.nodesByID[node.ID] = node
|
|
}
|
|
}
|
|
for _, edge := range input.Definition.Edges {
|
|
state.outgoing[edge.Source] = append(state.outgoing[edge.Source], edge)
|
|
}
|
|
return state
|
|
}
|
|
|
|
func (e *Executor) executeNode(ctx context.Context, state *runState, node dsl.Node) error {
|
|
switch node.Type {
|
|
case workflowregistry.NodeTypeStart:
|
|
state.setNodeVars(node.ID, map[string]any{
|
|
"conversationId": state.input.Conversation.ID,
|
|
"messageId": state.input.UserMessage.ID,
|
|
"aiAgentId": state.input.AIAgent.ID,
|
|
"userMessage": strings.TrimSpace(state.input.UserMessage.Content),
|
|
"knowledgeBaseIds": utils.SplitInt64s(state.input.AIAgent.KnowledgeIDs),
|
|
"conversationState": state.input.Conversation.Status,
|
|
})
|
|
case workflowregistry.NodeTypeKnowledgeRetrieve:
|
|
return e.executeKnowledgeRetrieve(ctx, state, node)
|
|
case workflowregistry.NodeTypeAnswerabilityGate:
|
|
return e.executeAnswerabilityGate(state, node)
|
|
case workflowregistry.NodeTypeCondition:
|
|
state.setNodeVars(node.ID, map[string]any{"matched": true})
|
|
case workflowregistry.NodeTypeLLMReply:
|
|
return e.executeLLMReply(ctx, state, node)
|
|
case workflowregistry.NodeTypeSendReply:
|
|
replyText := strings.TrimSpace(toString(state.resolveInput(node, "replyText")))
|
|
state.result.ReplyText = replyText
|
|
state.setNodeVars(node.ID, map[string]any{
|
|
"sent": replyText != "",
|
|
"replyMessageId": int64(0),
|
|
})
|
|
case workflowregistry.NodeTypeHandoffToHuman:
|
|
reason := strings.TrimSpace(toString(state.resolveInput(node, "reason")))
|
|
replyText := strings.TrimSpace(readStringConfig(node.Config, "replyText"))
|
|
if replyText == "" {
|
|
replyText = "已为你转接人工客服,请稍候。"
|
|
}
|
|
state.result.ReplyText = replyText
|
|
state.setNodeVars(node.ID, map[string]any{
|
|
"handoffId": int64(0),
|
|
"reason": reason,
|
|
})
|
|
case workflowregistry.NodeTypeEnd:
|
|
state.setNodeVars(node.ID, map[string]any{"status": "completed"})
|
|
default:
|
|
return fmt.Errorf("unsupported workflow node type: %s", node.Type)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (e *Executor) executeKnowledgeRetrieve(ctx context.Context, state *runState, node dsl.Node) error {
|
|
query := strings.TrimSpace(toString(state.resolveInput(node, "query")))
|
|
retriever := retrievers.NewKnowledgeRetriever(state.input.AIAgent)
|
|
result, err := retriever.RetrieveContext(ctx, query)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
items := make([]map[string]any, 0, len(result.ContextResults))
|
|
for _, item := range result.ContextResults {
|
|
items = append(items, map[string]any{
|
|
"knowledgeBaseId": item.KnowledgeBaseID,
|
|
"documentId": item.DocumentID,
|
|
"chunkId": item.ChunkID,
|
|
"content": item.Content,
|
|
"score": item.Score,
|
|
})
|
|
}
|
|
state.result.RetrieverCount = len(result.Hits)
|
|
state.setNodeVars(node.ID, map[string]any{
|
|
"items": items,
|
|
"summary": result.ContextText,
|
|
})
|
|
return nil
|
|
}
|
|
|
|
func (e *Executor) executeAnswerabilityGate(state *runState, node dsl.Node) error {
|
|
items := state.resolveInput(node, "knowledgeItems")
|
|
answerability := "unanswerable"
|
|
reason := "no retrieved knowledge items"
|
|
if hasItems(items) {
|
|
answerability = "answerable"
|
|
reason = "retrieved knowledge items are available"
|
|
}
|
|
state.setNodeVars(node.ID, map[string]any{
|
|
"answerability": answerability,
|
|
"reason": reason,
|
|
})
|
|
return nil
|
|
}
|
|
|
|
func (e *Executor) executeLLMReply(ctx context.Context, state *runState, node dsl.Node) error {
|
|
if staticReply := strings.TrimSpace(readStringConfig(node.Config, "staticReply")); staticReply != "" {
|
|
state.setNodeVars(node.ID, map[string]any{"replyText": staticReply})
|
|
return nil
|
|
}
|
|
userPrompt := strings.TrimSpace(toString(state.resolveInput(node, "userMessage")))
|
|
if userPrompt == "" {
|
|
userPrompt = strings.TrimSpace(state.input.UserMessage.Content)
|
|
}
|
|
knowledgeItems := toString(state.resolveInput(node, "knowledgeItems"))
|
|
systemPrompt := strings.TrimSpace(state.input.AIAgent.SystemPrompt)
|
|
if prompt := strings.TrimSpace(readStringConfig(node.Config, "prompt")); prompt != "" {
|
|
systemPrompt = strings.TrimSpace(systemPrompt + "\n\n" + prompt)
|
|
}
|
|
if knowledgeItems != "" {
|
|
userPrompt = userPrompt + "\n\nKnowledge context:\n" + knowledgeItems
|
|
}
|
|
result, err := ai.LLM.ChatWithConfig(ctx, state.input.AIConfig, systemPrompt, userPrompt)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
state.result.PromptTokens += result.PromptTokens
|
|
state.result.CompletionTokens += result.CompletionTokens
|
|
state.setNodeVars(node.ID, map[string]any{"replyText": result.Content})
|
|
return nil
|
|
}
|
|
|
|
func (s *runState) nextNodeID(sourceNodeID string) (string, bool, error) {
|
|
edges := s.outgoing[sourceNodeID]
|
|
if len(edges) == 0 {
|
|
return "", false, nil
|
|
}
|
|
for _, edge := range edges {
|
|
if edge.Condition == nil {
|
|
continue
|
|
}
|
|
matched, err := s.evaluateCondition(edge.Condition)
|
|
if err != nil {
|
|
return "", false, err
|
|
}
|
|
if matched {
|
|
return strings.TrimSpace(edge.Target), true, nil
|
|
}
|
|
}
|
|
for _, edge := range edges {
|
|
if edge.Condition == nil {
|
|
return strings.TrimSpace(edge.Target), true, nil
|
|
}
|
|
}
|
|
return "", false, nil
|
|
}
|
|
|
|
func (s *runState) evaluateCondition(condition *dsl.Condition) (bool, error) {
|
|
if condition == nil {
|
|
return true, nil
|
|
}
|
|
left := s.resolveSelector(condition.Left)
|
|
operator := strings.TrimSpace(condition.Operator)
|
|
if operator == "" && strings.TrimSpace(condition.Expression) != "" {
|
|
return false, fmt.Errorf("free-form workflow condition expressions are not supported")
|
|
}
|
|
switch operator {
|
|
case "eq", "equals":
|
|
return compareString(left, condition.Right) == 0, nil
|
|
case "neq", "not_equals":
|
|
return compareString(left, condition.Right) != 0, nil
|
|
case "contains":
|
|
return strings.Contains(toString(left), toString(condition.Right)), nil
|
|
case "exists":
|
|
return exists(left), nil
|
|
case "not_exists":
|
|
return !exists(left), nil
|
|
case "truthy", "is_true":
|
|
return truthy(left), nil
|
|
case "falsy", "is_false":
|
|
return !truthy(left), nil
|
|
case "gt":
|
|
return compareNumber(left, condition.Right) > 0, nil
|
|
case "gte":
|
|
return compareNumber(left, condition.Right) >= 0, nil
|
|
case "lt":
|
|
return compareNumber(left, condition.Right) < 0, nil
|
|
case "lte":
|
|
return compareNumber(left, condition.Right) <= 0, nil
|
|
default:
|
|
return false, fmt.Errorf("unsupported workflow condition operator: %s", operator)
|
|
}
|
|
}
|
|
|
|
func (s *runState) setNodeVars(nodeID string, values map[string]any) {
|
|
s.vars[nodeID] = values
|
|
}
|
|
|
|
func (s *runState) resolveInput(node dsl.Node, inputName string) any {
|
|
selector, ok := node.Inputs[inputName]
|
|
if !ok {
|
|
return nil
|
|
}
|
|
return s.resolveSelector(&selector)
|
|
}
|
|
|
|
func (s *runState) resolveSelector(selector *dsl.VariableSelector) any {
|
|
if selector == nil {
|
|
return nil
|
|
}
|
|
fields := s.vars[strings.TrimSpace(selector.NodeID)]
|
|
if fields == nil {
|
|
return nil
|
|
}
|
|
return fields[strings.TrimSpace(selector.Field)]
|
|
}
|
|
|
|
func readStringConfig(raw json.RawMessage, key string) string {
|
|
if len(raw) == 0 {
|
|
return ""
|
|
}
|
|
var cfg map[string]any
|
|
if err := json.Unmarshal(raw, &cfg); err != nil {
|
|
return ""
|
|
}
|
|
return toString(cfg[key])
|
|
}
|
|
|
|
func compareString(left any, right any) int {
|
|
return strings.Compare(toString(left), toString(right))
|
|
}
|
|
|
|
func compareNumber(left any, right any) int {
|
|
leftNum := toFloat(left)
|
|
rightNum := toFloat(right)
|
|
switch {
|
|
case leftNum > rightNum:
|
|
return 1
|
|
case leftNum < rightNum:
|
|
return -1
|
|
default:
|
|
return 0
|
|
}
|
|
}
|
|
|
|
func toString(value any) string {
|
|
switch v := value.(type) {
|
|
case nil:
|
|
return ""
|
|
case string:
|
|
return v
|
|
case fmt.Stringer:
|
|
return v.String()
|
|
case []map[string]any:
|
|
buf, _ := json.Marshal(v)
|
|
return string(buf)
|
|
default:
|
|
return strings.TrimSpace(fmt.Sprint(v))
|
|
}
|
|
}
|
|
|
|
func toFloat(value any) float64 {
|
|
switch v := value.(type) {
|
|
case int:
|
|
return float64(v)
|
|
case int64:
|
|
return float64(v)
|
|
case float64:
|
|
return v
|
|
case float32:
|
|
return float64(v)
|
|
case json.Number:
|
|
f, _ := v.Float64()
|
|
return f
|
|
case string:
|
|
f, _ := strconv.ParseFloat(strings.TrimSpace(v), 64)
|
|
return f
|
|
default:
|
|
return 0
|
|
}
|
|
}
|
|
|
|
func truthy(value any) bool {
|
|
switch v := value.(type) {
|
|
case nil:
|
|
return false
|
|
case bool:
|
|
return v
|
|
case string:
|
|
normalized := strings.ToLower(strings.TrimSpace(v))
|
|
return normalized != "" && normalized != "false" && normalized != "0"
|
|
default:
|
|
return !reflect.ValueOf(value).IsZero()
|
|
}
|
|
}
|
|
|
|
func exists(value any) bool {
|
|
if value == nil {
|
|
return false
|
|
}
|
|
switch v := value.(type) {
|
|
case string:
|
|
return strings.TrimSpace(v) != ""
|
|
default:
|
|
return true
|
|
}
|
|
}
|
|
|
|
func hasItems(value any) bool {
|
|
if value == nil {
|
|
return false
|
|
}
|
|
rv := reflect.ValueOf(value)
|
|
switch rv.Kind() {
|
|
case reflect.Array, reflect.Slice, reflect.Map:
|
|
return rv.Len() > 0
|
|
default:
|
|
return exists(value)
|
|
}
|
|
}
|