package flow import ( "ai-agent/workflow/consts/flow" flowDao "ai-agent/workflow/dao/flow" sessionDao "ai-agent/workflow/dao/session" flowDto "ai-agent/workflow/model/dto/flow" sessionDto "ai-agent/workflow/model/dto/session" "ai-agent/workflow/model/entity" "bytes" "context" "encoding/json" "errors" "fmt" "sync" "time" "gitea.redpowerfuture.com/red-future/common/utils" wsCommon "gitea.redpowerfuture.com/red-future/common/websocket" "github.com/cloudwego/eino/compose" "github.com/gogf/gf/v2/frame/g" "github.com/gogf/gf/v2/os/glog" "github.com/gogf/gf/v2/util/gconv" "github.com/google/uuid" ) // ====================== WebSocket 服务器 ====================== func init() { // 工作流消息处理器注册在统一的 SessionWsService(见 ws_server.go)上: // 首次连接仅升级,连接后按消息 type 路由,不再建连时区分普通对话/工作流 SessionWsService.OnMessage("workflow", handleExecute) SessionWsService.OnMessage("workflow_cancel", handleCancel) } // defaultSessionName 工作流执行但查不到流程名时,会话的兜底名称 const defaultSessionName = "工作流执行" // errWorkflowTerminated 前端终止工作流执行时的错误标记(写入 exec_workflow.error_message) var errWorkflowTerminated = "用户已终止执行" // ====================== 进度上报 ====================== type wsProgressCtxKey struct{} // ProgressReporter 节点执行进度回调接口 type ProgressReporter interface { ReportStart(nodeId, nodeName string, nodeIndex, nodeCount int) ReportComplete(nodeId, nodeName string, nodeIndex, nodeCount int) } // GetProgressReporter 从context中获取进度上报器 func GetProgressReporter(ctx context.Context) ProgressReporter { if reporter, ok := ctx.Value(wsProgressCtxKey{}).(ProgressReporter); ok { return reporter } return nil } type wsProgressReporter struct { conn *wsCommon.WsConnection mu sync.Mutex } func (r *wsProgressReporter) ReportStart(nodeId, nodeName string, nodeIndex, nodeCount int) { if r.conn.IsClosed() { return } r.mu.Lock() msg := &wsCommon.WsPushMsg{ Type: "node_start", Message: fmt.Sprintf("开始执行(%d/%d): %s ", nodeIndex, nodeCount, nodeName), Data: map[string]interface{}{ "nodeId": nodeId, "nodeName": nodeName, "nodeIndex": nodeIndex, "nodeCount": nodeCount, }, } r.mu.Unlock() _ = writeJSON(r.conn, msg) } func (r *wsProgressReporter) ReportComplete(nodeId, nodeName string, nodeIndex, nodeCount int) { if r.conn.IsClosed() { return } r.mu.Lock() msg := &wsCommon.WsPushMsg{ Type: "node_complete", Message: fmt.Sprintf("执行完成(%d/%d): %s ", nodeIndex, nodeCount, nodeName), Data: map[string]interface{}{ "nodeId": nodeId, "nodeName": nodeName, "nodeIndex": nodeIndex, "nodeCount": nodeCount, }, } r.mu.Unlock() _ = writeJSON(r.conn, msg) } // ====================== 消息处理 ====================== // handleExecute 处理工作流执行(由 workerPool 异步调用,不阻塞读循环) func handleExecute(ctx context.Context, conn *wsCommon.WsConnection, payload interface{}) { execPayload := new(sessionDto.WebSocketExecWorkflowReq) if err := gconv.Struct(payload, execPayload); err != nil { glog.Errorf(ctx, "工作流执行参数解析失败: %v", err) _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "error", Message: "执行参数解析失败", Error: err.Error()}) return } execCtx, execCancel := context.WithCancel(ctx) // 替换旧 cancel,设入新 cancel if oldCancel := getExecCancel(conn); oldCancel != nil { oldCancel() } conn.SetMeta("execCancel", execCancel) //_ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "ack", Message: "开始执行工作流"}) // 异步执行工作流(直接 goroutine,不依赖上游 workerPool 二次排队) go func() { defer func() { if r := recover(); r != nil { glog.Errorf(execCtx, "workflow panic: %v", r) _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "error", Message: "工作流异常", Error: fmt.Sprintf("%v", r)}) } }() defer conn.SetMeta("execCancel", nil) // 落库用不带取消的 ctx(保留 request 值),保证前端终止/断连后记录仍能写入 saveCtx := context.WithoutCancel(ctx) // 会话落库:前端 sessionId 对应会话已存在则复用,否则按流程名新建 flowName := defaultSessionName if flowUser, e := flowDao.FlowUserDao.Get(saveCtx, &flowDto.GetFlowUserReq{Id: execPayload.FlowId}); e == nil && flowUser != nil && flowUser.FlowName != "" { flowName = flowUser.FlowName } if e := ensureSession(saveCtx, conn.SessionId, flowName); e != nil { glog.Errorf(saveCtx, "工作流会话创建失败: %v", e) _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "error", Message: "工作流会话创建失败", Error: fmt.Sprintf("%v", e)}) } start := time.Now() reporter := &wsProgressReporter{conn: conn} progressCtx := context.WithValue(execCtx, wsProgressCtxKey{}, reporter) _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "ack", Message: fmt.Sprintf("开始执行工作流(共 %d 个节点)", len(execPayload.FlowContent.Nodes))}) execId, err := executeOrResume(progressCtx, conn, execPayload) recordWorkflow(saveCtx, execId, time.Since(start), err) if err != nil { _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "error", Message: "工作流执行失败", Error: err.Error()}) return } // 成功:把本次执行保存的结果文件路径(exec_workflow_result)一并推给前端 _ = writeJSON(conn, &wsCommon.WsPushMsg{ Type: "flow_complete", Message: "工作流执行完成", Data: map[string]interface{}{ "resultFileUrls": workflowResultFileUrls(saveCtx, execId), }, }) }() } // recordWorkflow 把一次工作流执行写入 exec_workflow/exec_workflow_result:运行记录 + 输出文件结果 func recordWorkflow(ctx context.Context, id int64, duration time.Duration, runErr error) { // exec_workflow 状态沿用 1-运行中,2-成功,3-失败;前端结果卡片也只识别 1/2/3 // (4 会误显示为"运行中"),故取消同样记为失败,错误信息写"用户已终止执行" // error_message 存友好提示,error 存原始错误明细 status := flow.FlowExecutionStatusSuccess var errorMessage, errorDetail string if runErr != nil { status = flow.FlowExecutionStatusFailed if errors.Is(runErr, context.Canceled) { errorMessage = errWorkflowTerminated } else { errorMessage = "工作流执行失败" errorDetail = runErr.Error() } } _, err := sessionDao.ExecWorkflowDao.Update(ctx, &sessionDto.UpdateWorkflowReq{ Id: id, Status: status.Code(), Duration: int64(duration.Seconds()), ErrorMessage: errorMessage, Error: errorDetail, }) if err != nil { glog.Errorf(ctx, "exec_workflow 落库失败: %v", err) return } // 执行成功:重新执行复用了同一条记录,OmitEmpty 的 Update 会跳过空 error_message/error, // 需显式清空,避免上一次失败的报错残留 if runErr == nil { if _, err := sessionDao.ExecWorkflowDao.ClearError(ctx, id); err != nil { glog.Errorf(ctx, "exec_workflow 报错信息清空失败: %v", err) } } } // workflowResultFileUrls 查询指定工作流执行保存的结果文件路径(带文件前缀,与 session/get 返回一致) func workflowResultFileUrls(ctx context.Context, execId int64) []string { results, err := sessionDao.ExecWorkflowResultDao.ListByExecId(ctx, execId) if err != nil { glog.Errorf(ctx, "查询工作流结果路径失败: %v", err) return nil } prefix, _ := utils.GetFileAddressPrefix(ctx) urls := make([]string, 0, len(results)) for _, r := range results { if r.ResultFileUrl != "" { urls = append(urls, prefix+r.ResultFileUrl) } } return urls } // handleCancel 取消工作流执行 func handleCancel(ctx context.Context, conn *wsCommon.WsConnection, _ interface{}) { if cancel := getExecCancel(conn); cancel != nil { cancel() } _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "ack", Message: "已取消工作流执行"}) } // ====================== 工具函数 ====================== func getExecCancel(conn *wsCommon.WsConnection) context.CancelFunc { cancel, _ := wsCommon.GetMetaT[context.CancelFunc](conn, "execCancel") return cancel } // writeJSON 业务层写入,委托 WsConnection.WriteJSON(共享 writeMu 写锁) func writeJSON(conn *wsCommon.WsConnection, data interface{}) error { return conn.WriteJSON(data) } // executeOrResume 决策工作流执行方式: // - 同会话+同工作流的最近一次执行失败,且本次传递参数与上次一致 → 断点续跑(reExecute,复用原执行记录,从失败断点继续) // - 其余情况(上次成功 / 上次参数与本次不同 / 无历史记录 / 查询出错)→ 全新执行(execute) func executeOrResume(ctx context.Context, conn *wsCommon.WsConnection, req *sessionDto.WebSocketExecWorkflowReq) (id int64, err error) { lastExec, err := sessionDao.ExecWorkflowDao.GetLatestBySessionAndFlow(ctx, conn.SessionId, req.FlowId) if err != nil { glog.Errorf(ctx, "查询最近工作流执行记录失败: %v", err) return execute(ctx, conn, req) } if lastExec != nil && *lastExec.Status == *flow.FlowExecutionStatusFailed.Code() && flowContentEqual(lastExec.RequestParams, req.FlowContent) { _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "round_start", Message: "运行开始", Data: map[string]interface{}{ "id": lastExec.Id, }}) return reExecute(ctx, lastExec.Id) } return execute(ctx, conn, req) } // flowContentEqual 判断两次工作流参数是否一致(JSON 序列化后字节比对。 // Go struct 按字段声明序序列化、map 键自动排序,同一内容结果确定,可用于参数等价判断) func flowContentEqual(a, b *entity.FlowInfo) bool { if a == nil || b == nil { return a == b } ab, err1 := json.Marshal(a) bb, err2 := json.Marshal(b) if err1 != nil || err2 != nil { return false } return bytes.Equal(ab, bb) } // execute 执行工作流(首次执行;同会话+同工作流最近一次执行为失败状态时复用该记录重新执行,不新建数据) func execute(ctx context.Context, conn *wsCommon.WsConnection, req *sessionDto.WebSocketExecWorkflowReq) (id int64, err error) { var nodeGroupId = uuid.NewString() // 复用失败记录:查询会话+工作流最近一次执行,若为失败状态则复用同一条记录(更新状态+节点组+本次请求参数), // 全新执行(forceNewRun,参数取本次请求),避免重跑新建数据;查询出错按无记录处理走新建 lastExec, qErr := sessionDao.ExecWorkflowDao.GetLatestBySessionAndFlow(ctx, conn.SessionId, req.FlowId) if qErr != nil { glog.Errorf(ctx, "查询最近工作流执行记录失败: %v", qErr) lastExec = nil } if lastExec != nil && *lastExec.Status == *flow.FlowExecutionStatusFailed.Code() { id = lastExec.Id _, err = sessionDao.ExecWorkflowDao.Update(ctx, &sessionDto.UpdateWorkflowReq{ Id: id, NodeGroupId: nodeGroupId, Status: flow.FlowExecutionStatusRunning.Code(), RequestParams: req.FlowContent, }) if err != nil { return } } else { id, err = sessionDao.ExecWorkflowDao.Insert(ctx, &sessionDto.CreateWorkflowReq{ SessionId: conn.SessionId, FlowId: req.FlowId, NodeGroupId: nodeGroupId, Status: flow.FlowExecutionStatusRunning.Code(), RequestParams: req.FlowContent, }) if err != nil { return } } _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "round_start", Message: "运行开始", Data: map[string]interface{}{ "id": id, }}) err = BuildExecution(ctx, true, req.FlowId, id, nodeGroupId, conn.SessionId, req.FlowContent) if err != nil { return } return id, nil } // reExecute 重新执行工作流 func reExecute(ctx context.Context, execWorkflowId int64) (id int64, err error) { flowInfo, err := sessionDao.ExecWorkflowDao.GetById(ctx, execWorkflowId) if err != nil { return } var nodeGroupId = uuid.NewString() _, err = sessionDao.ExecWorkflowDao.Update(ctx, &sessionDto.UpdateWorkflowReq{ Id: flowInfo.Id, NodeGroupId: nodeGroupId, Status: flow.FlowExecutionStatusRunning.Code(), }) if err != nil { return } err = BuildExecution(ctx, false, flowInfo.FlowId, flowInfo.Id, nodeGroupId, flowInfo.SessionId, flowInfo.RequestParams) if err != nil { return } return flowInfo.Id, nil } func BuildExecution(ctx context.Context, forceNewRun bool, flowId, executionId int64, nodeGroupId string, sessionId string, flowContent *entity.FlowInfo) (err error) { // ========================================================================= // 构建执行图 // ========================================================================= var nodeList []entity.FlowNode var runGraph compose.Runnable[any, any] nodeList, runGraph, err = BuildGraphFromFlowContent(ctx, flowContent) if err != nil { return fmt.Errorf("执行工作流失败: %v", err) } // ========================================================================= // 构建 ConfigMap // ========================================================================= nodeInputParams := ExtractFlowNodeFrom(flowContent) configMap := make(map[string]*entity.FlowNode) for _, cfg := range nodeInputParams { configMap[cfg.Id] = cfg } for _, i := range nodeList { configMap[i.Id] = &i } // ========================================================================= // 构建全局执行入参 // ========================================================================= execInput := &flowDto.FlowExecutionInput{ NodeGroupId: nodeGroupId, ExecutionId: executionId, FlowId: flowId, ConfigMap: configMap, SessionId: sessionId, } var opts []compose.Option opts = append(opts, compose.WithCheckPointID(gconv.String(executionId))) if forceNewRun { opts = append(opts, compose.WithForceNewRun()) } _, err = runGraph.Invoke(ctx, execInput, opts...) if err != nil { info, infoOk := compose.ExtractInterruptInfo(err) if infoOk { var errMsg string var errNodeCount int for _, item := range info.InterruptContexts { if item.Info == nil { continue } if g.NewVar(item.Info).IsMap() { errNodeCount++ valMap := gconv.Map(item.Info) errMsg = fmt.Sprintf("%v\n%v", errMsg, fmt.Sprintf("节点:%v, 失败原因:%v", valMap["node"], valMap["error"])) } } if !g.IsEmpty(errMsg) { err = fmt.Errorf("%v个节点,%v", errNodeCount, errMsg) } } return fmt.Errorf("执行工作流失败: %v", err) } // 清理断点数据 _ = flowDao.FlowCheckpointDao.Delete(ctx, gconv.String(executionId)) return }