296 lines
11 KiB
Go
296 lines
11 KiB
Go
package flow
|
|
|
|
import (
|
|
"ai-agent/gateway"
|
|
"ai-agent/workflow/consts/node"
|
|
flowDto "ai-agent/workflow/model/dto/flow"
|
|
"ai-agent/workflow/model/entity"
|
|
"context"
|
|
"fmt"
|
|
"time"
|
|
|
|
"github.com/cloudwego/eino/compose"
|
|
"github.com/gogf/gf/v2/frame/g"
|
|
"github.com/gogf/gf/v2/util/gconv"
|
|
)
|
|
|
|
// BuildGraph 根据 FlowInfo 构建完整的 Eino Graph 拓扑
|
|
func BuildGraph(ctx context.Context, flowContent *entity.FlowInfo) ([]entity.FlowNode, *compose.Graph[any, any]) {
|
|
// 注册自定义合并函数:处理 *flowDto.FlowExecutionInput 类型合并
|
|
compose.RegisterValuesMergeFunc(func(values []*flowDto.FlowExecutionInput) (*flowDto.FlowExecutionInput, error) {
|
|
if len(values) == 0 {
|
|
return nil, nil
|
|
}
|
|
// 首次运行所有并行分支共享同一个 ConfigMap 指针,直接返回 values[0] 即可。
|
|
// 但续跑(ReExecute)时各分支从 checkpoint 反序列化出独立的 ConfigMap 副本,
|
|
// 只返回 values[0] 会丢失其他分支写入的 OutputResult(用户实测:node-8 成功的结果
|
|
// 在汇合节点 node-7 变 null)。以第一个为基底,把其余分支中缺失的节点输出合并进来。
|
|
base := values[0]
|
|
for _, v := range values[1:] {
|
|
if v == nil {
|
|
continue
|
|
}
|
|
for nodeId, cfg := range v.ConfigMap {
|
|
if cfg == nil {
|
|
continue
|
|
}
|
|
baseCfg, ok := base.ConfigMap[nodeId]
|
|
if !ok || baseCfg == nil {
|
|
base.ConfigMap[nodeId] = cfg
|
|
continue
|
|
}
|
|
if len(baseCfg.OutputResult) == 0 && len(cfg.OutputResult) > 0 {
|
|
baseCfg.OutputResult = cfg.OutputResult
|
|
}
|
|
}
|
|
// 合并已执行节点列表(按 NodeId 去重),续跑时被恢复分支的进度不丢失
|
|
for _, en := range v.ExecutedNodes {
|
|
dup := false
|
|
for _, b := range base.ExecutedNodes {
|
|
if b.NodeId == en.NodeId {
|
|
dup = true
|
|
break
|
|
}
|
|
}
|
|
if !dup {
|
|
base.ExecutedNodes = append(base.ExecutedNodes, en)
|
|
}
|
|
}
|
|
}
|
|
return base, nil
|
|
})
|
|
|
|
graph := compose.NewGraph[any, any](
|
|
// 本地状态初始化
|
|
compose.WithGenLocalState(func(ctx context.Context) *flowDto.NodeExecutionState {
|
|
return &flowDto.NodeExecutionState{}
|
|
}),
|
|
)
|
|
|
|
// 注册所有节点
|
|
nodeMap := make(map[string]entity.FlowNode)
|
|
for _, item := range flowContent.Nodes {
|
|
nodeMap[item.Id] = item
|
|
//if item.NodeCode != node.NodeTypeJudge {
|
|
registerNodeToGraph(graph, item)
|
|
//}
|
|
}
|
|
|
|
// 注册开始节点
|
|
if flowContent.StartNodeId != "" {
|
|
_ = graph.AddEdge(compose.START, flowContent.StartNodeId)
|
|
}
|
|
|
|
var nodeList []entity.FlowNode
|
|
originalEndNodes := FindEndNodes(flowContent.StartNodeId, flowContent.Edges)
|
|
for _, endID := range originalEndNodes {
|
|
// 保存结果节点 ID 必须稳定:ReExecute 续跑重建图时复用同一 ID,
|
|
// 否则 checkpoint 里 ConfigMap 存的是上次的旧 ID,续跑时新图按新 ID 查不到配置,
|
|
// 报"节点信息为空"。一个末端节点对应一个保存结果节点,用 endID 派生即唯一且稳定。
|
|
summaryNodeId := fmt.Sprintf("%s_%s", node.NodeTypeSystemSum, endID)
|
|
summaryNode := entity.FlowNode{
|
|
Id: summaryNodeId,
|
|
NodeCode: node.NodeTypeSystemSum,
|
|
Name: node.GetNodeTypeName(node.NodeTypeSystemSum),
|
|
}
|
|
nodeList = append(nodeList, summaryNode)
|
|
flowContent.Nodes = append(flowContent.Nodes, summaryNode)
|
|
|
|
registerNodeToGraph(graph, summaryNode)
|
|
_ = graph.AddEdge(endID, summaryNodeId)
|
|
_ = graph.AddEdge(summaryNodeId, compose.END)
|
|
}
|
|
|
|
// 构建边关系
|
|
upstreamMap := make(map[string][]string)
|
|
edgeMap := make(map[string][]entity.FlowEdge)
|
|
for _, edge := range flowContent.Edges {
|
|
edgeMap[edge.From] = append(edgeMap[edge.From], edge)
|
|
upstreamMap[edge.To] = append(upstreamMap[edge.To], edge.From)
|
|
}
|
|
|
|
// 处理连线 & 分支
|
|
for _, edges := range edgeMap {
|
|
//fromNode := nodeMap[fromNodeID]
|
|
|
|
// 判断节点 → 分支处理
|
|
//if fromNode.NodeCode == node.NodeTypeJudge {
|
|
// branchMap := make(map[string]bool)
|
|
// for _, e := range edges {
|
|
// branchMap[e.To] = true
|
|
// }
|
|
//
|
|
// judgeLambda := func(ctx context.Context, input any) (string, error) {
|
|
// execInput, ok := input.(*flowDto.FlowExecutionInput)
|
|
// if !ok {
|
|
// return "", fmt.Errorf("入参类型错误")
|
|
// }
|
|
//
|
|
// currentConfig := execInput.ConfigMap[fromNodeID]
|
|
// if currentConfig == nil {
|
|
// return "", fmt.Errorf("判断节点%s无配置", fromNodeID)
|
|
// }
|
|
//
|
|
// branchIdNameMap := make(map[string]string)
|
|
// var branchIDs []string
|
|
// for nodeID := range branchMap {
|
|
// branchIDs = append(branchIDs, nodeID)
|
|
// // 从configMap获取分支节点的名称
|
|
// if branchNodeCfg, ok := execInput.ConfigMap[nodeID]; ok {
|
|
// branchIdNameMap[nodeID] = branchNodeCfg.Name
|
|
// } else {
|
|
// branchIdNameMap[nodeID] = "未命名节点" // 兜底
|
|
// }
|
|
// }
|
|
//
|
|
// // 把分支ID-名称映射塞进 ModelConfig,带给意图节点
|
|
// m := make(map[string]interface{})
|
|
// m["branch_ids"] = branchIDs
|
|
// m["branch_id_name_map"] = branchIdNameMap
|
|
// currentConfig.Config = m
|
|
//
|
|
// // 构造 NodeExecutionInput 传入 JudgeLambda
|
|
// nodeExecInput := &flowDto.NodeExecutionInput{
|
|
// Config: currentConfig,
|
|
// Global: execInput,
|
|
// }
|
|
// return JudgeLambda(ctx, nodeExecInput)
|
|
// }
|
|
//
|
|
// _ = graph.AddBranch(upstreamMap[fromNodeID][0], compose.NewGraphBranch(judgeLambda, branchMap))
|
|
// continue
|
|
//}
|
|
|
|
// 普通节点连线
|
|
for _, e := range edges {
|
|
//toNode := nodeMap[e.To]
|
|
//if toNode.NodeCode == node.NodeTypeJudge {
|
|
// continue
|
|
//}
|
|
_ = graph.AddEdge(e.From, e.To)
|
|
}
|
|
}
|
|
return nodeList, graph
|
|
}
|
|
|
|
// BuildGraphFromFlowContent 根据前端保存的工作流JSON,自动构建执行图并编译
|
|
func BuildGraphFromFlowContent(ctx context.Context, flowContent *entity.FlowInfo) ([]entity.FlowNode, compose.Runnable[any, any], error) {
|
|
nodeList, graph := BuildGraph(ctx, flowContent)
|
|
compile, err := graph.Compile(ctx, compose.WithGraphName("auto_build_workflow"), compose.WithCheckPointStore(NewDbCheckPointStore()), compose.WithNodeTriggerMode(compose.AllPredecessor))
|
|
return nodeList, compile, err
|
|
}
|
|
|
|
// registerNodeToGraph 将单个节点注册到图中(包含通用包装逻辑)
|
|
func registerNodeToGraph(graph *compose.Graph[any, any], flowNode entity.FlowNode) {
|
|
// 通用包装:全程入参都是 *FlowExecutionInput
|
|
wrapLambda := func(lambda func(ctx context.Context, input any) (any, error)) func(ctx context.Context, input any) (any, error) {
|
|
return func(ctx context.Context, input any) (any, error) {
|
|
startTime := time.Now()
|
|
|
|
// 构建节点执行入参(含中断恢复)
|
|
execInput, realInput, err := BuildNodeExecutionInput(ctx, input, flowNode)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
flowNodeDesc := flowNode.Desc
|
|
if g.IsEmpty(flowNodeDesc) {
|
|
flowNodeDesc = flowNode.Name
|
|
}
|
|
|
|
// 上报节点执行进度(WebSocket场景下推送进度给前端)
|
|
if reporter := GetProgressReporter(ctx); reporter != nil {
|
|
nodeIndex := len(execInput.ExecutedNodes) + 1
|
|
if IndexOf(execInput.ExecutedNodes, flowNode.Id) != -1 {
|
|
nodeIndex = IndexOf(execInput.ExecutedNodes, flowNode.Id)
|
|
}
|
|
reporter.ReportStart(flowNode.Id, flowNodeDesc, nodeIndex, len(execInput.ConfigMap))
|
|
}
|
|
|
|
// 上传入参到OSS
|
|
ossResult, err := gateway.Upload(ctx, fmt.Sprintf("nodeInput:%v.txt", time.Now().UnixMilli()), gconv.Bytes(gconv.String(realInput)))
|
|
if err != nil {
|
|
return nil, HandleFailedNodeExecution(ctx, execInput, 0, flowNode, err, 0)
|
|
}
|
|
|
|
// 创建节点执行记录
|
|
nodeExecutionId, err := CreateNodeExecutionRecord(ctx, execInput, flowNode, ossResult)
|
|
if err != nil {
|
|
return nil, HandleFailedNodeExecution(ctx, execInput, 0, flowNode, err, 0)
|
|
}
|
|
realInput.NodeExecutionId = nodeExecutionId
|
|
|
|
// 执行节点
|
|
_, err = lambda(ctx, realInput)
|
|
durationMs := time.Since(startTime).Milliseconds()
|
|
|
|
if err != nil {
|
|
// 执行失败处理
|
|
return nil, HandleFailedNodeExecution(ctx, execInput, nodeExecutionId, flowNode, err, durationMs)
|
|
}
|
|
|
|
// 执行成功处理
|
|
if err = HandleSuccessfulNodeExecution(ctx, execInput, realInput, nodeExecutionId, flowNode, durationMs); err != nil {
|
|
return nil, HandleFailedNodeExecution(ctx, execInput, nodeExecutionId, flowNode, err, durationMs)
|
|
}
|
|
|
|
// 上报节点执行进度(WebSocket场景下推送进度给前端)
|
|
if reporter := GetProgressReporter(ctx); reporter != nil {
|
|
nodeIndex := len(execInput.ExecutedNodes)
|
|
if IndexOf(execInput.ExecutedNodes, flowNode.Id) != -1 {
|
|
nodeIndex = IndexOf(execInput.ExecutedNodes, flowNode.Id)
|
|
}
|
|
reporter.ReportComplete(flowNode.Id, flowNodeDesc, nodeIndex, len(execInput.ConfigMap))
|
|
}
|
|
|
|
// 返回整个 execInput,让下一个节点继续用
|
|
return execInput, nil
|
|
}
|
|
}
|
|
|
|
switch flowNode.NodeCode {
|
|
case node.NodeTypeStart:
|
|
_ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(StartLambda)))
|
|
case node.NodeTypeSystemSum:
|
|
_ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(SummaryLambda)))
|
|
case node.NodeTypeModel:
|
|
_ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(ModelLambda)))
|
|
case node.NodeTypeForm:
|
|
_ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(FormLambda)))
|
|
case node.NodeTypeDataMerge:
|
|
_ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(DataMergeLambda)))
|
|
case node.NodeTypeSubFlow:
|
|
_ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(SubFlowLambda)))
|
|
case node.NodeTypeHttp:
|
|
_ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(HttpLambda)))
|
|
case node.NodeTypeScriptTranscribe:
|
|
_ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(ScriptTranscribeLambda)))
|
|
//case node.NodeTypeTextModel:
|
|
// _ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(TextModelLambda)))
|
|
//case node.NodeTypeImageModel:
|
|
// _ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(ImageModelLambda)))
|
|
//case node.NodeTypeVideoModel:
|
|
// _ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(VideoModelLambda)))
|
|
//case node.NodeTypeAudioModel:
|
|
// _ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(AudioModelLambda)))
|
|
//case node.NodeTypeBatchModel:
|
|
// _ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(BatchModelLambda)))
|
|
//case node.NodeTypeDataConversionModel:
|
|
// _ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(DataConversionLambda)))
|
|
//case node.NodeTypeCustomNode:
|
|
// _ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(CustomLambda)))
|
|
//case node.NodeTypeMerge:
|
|
// _ = graph.AddLambdaNode(flowNode.Id, compose.InvokableLambda(wrapLambda(MergeLambda)))
|
|
}
|
|
}
|
|
|
|
// IndexOf 返回元素第一次出现的下标,不存在返回 -1
|
|
func IndexOf(slice []flowDto.ExecutedNode, target string) int {
|
|
for i, v := range slice {
|
|
if v.NodeId == target {
|
|
return i + 1
|
|
}
|
|
}
|
|
return -1
|
|
}
|