From d699f7ce14534e4171382e3718938fea9f2f7182 Mon Sep 17 00:00:00 2001 From: qhd <1766646056@qq.com> Date: Thu, 3 Sep 2026 13:22:22 +0800 Subject: [PATCH] =?UTF-8?q?feat(workflow):=20=E5=A2=9E=E5=8A=A0=E5=B7=A5?= =?UTF-8?q?=E4=BD=9C=E6=B5=81=E8=AE=A1=E8=B4=B9=E4=B8=8E=E6=89=A7=E8=A1=8C?= =?UTF-8?q?=E7=94=9F=E5=91=BD=E5=91=A8=E6=9C=9F=E7=AE=A1=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增计费模块:执行开始建单、终态结算/取消/失败处理,支持按条/按秒/按token计费 - 新增执行生命周期跟踪:优雅关停时取消运行中执行并等待落库 - 新增异步任务等待/通知机制(Wait/Notify) - 重构执行记录落库与进度上报,统一失败分类与重试语义 - 重命名文件:async_task.go→async.go、flow_checkpoint_store.go→exec_checkpoint.go、flow_graph_util.go→exec_record.go - 更新 .gitignore 与数据库密码配置 --- .gitignore | 1 + config.yml | 2 +- gateway/model.go | 36 +- gateway/model_stream.go | 3 +- go.mod | 9 +- go.sum | 2 + update.sql | 12 +- workflow/dao/flow/flow_async_task_dao.go | 5 +- workflow/dao/node/node_execution_dao.go | 5 + workflow/model/dto/node/node_execution_dto.go | 2 + .../model/dto/session/exec_workflow_dto.go | 1 + workflow/model/dto/session/session_dto.go | 3 +- workflow/model/entity/exec_workflow.go | 11 +- workflow/model/entity/flow_user.go | 1 + .../service/flow/{async_task.go => async.go} | 46 +- workflow/service/flow/billing.go | 391 +++++++++ ...checkpoint_store.go => exec_checkpoint.go} | 0 workflow/service/flow/exec_hub.go | 7 - workflow/service/flow/exec_lifecycle.go | 159 ++++ workflow/service/flow/exec_progress.go | 46 ++ .../{flow_graph_util.go => exec_record.go} | 136 +++- .../{recover_execution.go => exec_recover.go} | 109 +-- workflow/service/flow/exec_shutdown.go | 70 -- .../flow/{flow_ws_exec.go => exec_ws.go} | 443 +++++------ workflow/service/flow/flow_helper.go | 132 +--- .../{flow_graph_builder.go => graph_build.go} | 138 ++-- workflow/service/flow/lambda_core.go | 252 ++++++ workflow/service/flow/lambda_http.go | 21 + workflow/service/flow/lambda_node.go | 742 ------------------ workflow/service/flow/lambda_savefile.go | 113 +++ .../service/flow/lambda_script_transcribe.go | 3 +- .../service/flow/lambda_segment_resume.go | 3 - workflow/service/flow/lambda_subflow.go | 219 ++++++ workflow/service/flow/lambda_summary.go | 122 +++ workflow/service/flow/lambda_tool.go | 90 +++ workflow/service/flow/lambda_value_source.go | 703 ----------------- .../{lambda_node_util.go => model_call.go} | 235 +----- .../flow/processor/builtin/media/media.go | 24 +- workflow/service/flow/react_ws_exec.go | 2 +- workflow/service/flow/subtitle.go | 166 ++++ workflow/service/flow/values/value_request.go | 140 ++++ workflow/service/flow/values/value_resolve.go | 254 ++++++ workflow/service/flow/values/value_schema.go | 227 ++++++ workflow/service/flow/ws_server.go | 41 - workflow/service/session/session_service.go | 1 + workflow/service/util_service.go | 12 +- 46 files changed, 2690 insertions(+), 2450 deletions(-) rename workflow/service/flow/{async_task.go => async.go} (79%) create mode 100644 workflow/service/flow/billing.go rename workflow/service/flow/{flow_checkpoint_store.go => exec_checkpoint.go} (100%) create mode 100644 workflow/service/flow/exec_lifecycle.go create mode 100644 workflow/service/flow/exec_progress.go rename workflow/service/flow/{flow_graph_util.go => exec_record.go} (53%) rename workflow/service/flow/{recover_execution.go => exec_recover.go} (71%) delete mode 100644 workflow/service/flow/exec_shutdown.go rename workflow/service/flow/{flow_ws_exec.go => exec_ws.go} (58%) rename workflow/service/flow/{flow_graph_builder.go => graph_build.go} (65%) create mode 100644 workflow/service/flow/lambda_core.go create mode 100644 workflow/service/flow/lambda_http.go delete mode 100644 workflow/service/flow/lambda_node.go create mode 100644 workflow/service/flow/lambda_savefile.go create mode 100644 workflow/service/flow/lambda_subflow.go create mode 100644 workflow/service/flow/lambda_summary.go create mode 100644 workflow/service/flow/lambda_tool.go delete mode 100644 workflow/service/flow/lambda_value_source.go rename workflow/service/flow/{lambda_node_util.go => model_call.go} (55%) create mode 100644 workflow/service/flow/subtitle.go create mode 100644 workflow/service/flow/values/value_request.go create mode 100644 workflow/service/flow/values/value_resolve.go create mode 100644 workflow/service/flow/values/value_schema.go delete mode 100644 workflow/service/flow/ws_server.go diff --git a/.gitignore b/.gitignore index 9f7cc75..b1c7a03 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1,3 @@ /.idea/* /.superpowers/ +/docs/superpowers/ diff --git a/config.yml b/config.yml index cf1339e..89cc5e7 100644 --- a/config.yml +++ b/config.yml @@ -29,7 +29,7 @@ database: host: "192.168.0.83" port: "15432" user: "postgres" - pass: "Bjang09@686^*^" + pass: "Q!P@z#M$1@686^*^.." name: "black-deacon" prefix: "black_deacon_" # (可选)表名前缀 role: "master" diff --git a/gateway/model.go b/gateway/model.go index 0a90925..c8359f6 100644 --- a/gateway/model.go +++ b/gateway/model.go @@ -13,6 +13,7 @@ import ( "gitea.redpowerfuture.com/red-future/common/beans" commonHttp "gitea.redpowerfuture.com/red-future/common/http" + "gitea.redpowerfuture.com/red-future/common/utils" gmq "github.com/bjang03/gmq/core/gmq" "github.com/bjang03/gmq/mq" "github.com/bjang03/gmq/types" @@ -80,12 +81,15 @@ type ModelCallReq struct { type ModelCallRes struct { TaskId int64 `json:"id" dc:"任务ID"` State int8 `json:"state" dc:"状态"` + ModelId int64 `json:"modelId" dc:"生效模型ID(引用行=解析后的系统模型ID,计价按此)"` + MediaType string `json:"mediaType" dc:"输入媒体类型(shop词汇: text/audio/video)"` TotalTokens int64 `json:"totalTokens" dc:"总token"` PromptTokens int64 `json:"promptTokens" dc:"输入token"` CompletionTokens int64 `json:"completionTokens" dc:"输出token"` Tools []ModelTool `json:"tools" dc:"工具"` Content map[string]any `json:"content" dc:"内容"` Cost float64 `json:"cost" dc:"费用(元)"` + Duration int64 `json:"duration" dc:"时长(秒)"` ErrorMsg string `json:"errorMsg" dc:"错误消息"` } @@ -98,43 +102,17 @@ type ModelTool struct { } `json:"function"` } -// requestHeaders 透传当前 HTTP 请求头(鉴权/链路信息)。 -// 浏览器 WebSocket 握手无法携带 Authorization 头,前端把 token 放在握手 URL query(?token=)里; -// 若请求头没有 Authorization,则从 query 补回,保证下游(model-gateway → admin-go)能拿到用户 token。 -// 后台恢复续跑无 HTTP 请求时(ctx 携带合成 user),补充 X-User-Info 供下游 GetUserInfo 识别租户。 -func requestHeaders(ctx context.Context) map[string]string { - headers := make(map[string]string) - if r := g.RequestFromCtx(ctx); r != nil { - for k, v := range r.Request.Header { - if len(v) > 0 { - headers[k] = v[0] - } - } - if headers["Authorization"] == "" { - if t := r.Request.URL.Query().Get("token"); t != "" { - headers["Authorization"] = "Bearer " + t - } - } - } - if headers["X-User-Info"] == "" { - if u := ctx.Value("user"); u != nil { - headers["X-User-Info"] = gconv.String(u) - } - } - return headers -} - // ListModelManage 配置列表 func ListModelManage(ctx context.Context, req *ListModelManageReq) (res *ListModelManageRes, err error) { res = new(ListModelManageRes) - err = commonHttp.Get(ctx, "model-gateway/model/manage/listModelManage", requestHeaders(ctx), res, req) + err = commonHttp.Get(ctx, "model-gateway/model/manage/listModelManage", utils.HeadersFromCtx(ctx, utils.HeadersOptions{TokenFromQuery: true}), res, req) return } // GetModelInfoById 查询模型配置 func GetModelInfoById(ctx context.Context, req *GetModelInfoByIdReq) (res *GetModelInfoByIdRes, err error) { res = new(GetModelInfoByIdRes) - err = commonHttp.Get(ctx, "model-gateway/model/manage/getModelManage", requestHeaders(ctx), res, req) + err = commonHttp.Get(ctx, "model-gateway/model/manage/getModelManage", utils.HeadersFromCtx(ctx, utils.HeadersOptions{TokenFromQuery: true}), res, req) return } @@ -163,7 +141,7 @@ func SubmitModelCall(ctx context.Context, modelId int64, responseType model.Resp } res = new(ModelCallRes) - err = commonHttp.Post(ctx, "model-gateway/model/call/modelCall", requestHeaders(ctx), res, &req) + err = commonHttp.Post(ctx, "model-gateway/model/call/modelCall", utils.HeadersFromCtx(ctx, utils.HeadersOptions{TokenFromQuery: true}), res, &req) if err != nil { return nil, "", err } diff --git a/gateway/model_stream.go b/gateway/model_stream.go index b138caf..bb4ccb8 100644 --- a/gateway/model_stream.go +++ b/gateway/model_stream.go @@ -8,6 +8,7 @@ import ( "strings" commonHttp "gitea.redpowerfuture.com/red-future/common/http" + "gitea.redpowerfuture.com/red-future/common/utils" "github.com/gogf/gf/v2/frame/g" "github.com/gogf/gf/v2/util/gconv" @@ -32,7 +33,7 @@ func ModelCallStream(ctx context.Context, modelId int64, sessionId string, reque RequestParams: requestParams, BusinessParams: businessParams, } - body, err := commonHttp.PostStream(ctx, "model-gateway/model/call/modelCallStream", requestHeaders(ctx), &req) + body, err := commonHttp.PostStream(ctx, "model-gateway/model/call/modelCallStream", utils.HeadersFromCtx(ctx, utils.HeadersOptions{TokenFromQuery: true}), &req) if err != nil { return nil, err } diff --git a/go.mod b/go.mod index cb49d4b..b7c3718 100644 --- a/go.mod +++ b/go.mod @@ -3,7 +3,7 @@ module ai-agent go 1.26.0 require ( - gitea.redpowerfuture.com/red-future/common v0.0.32 + gitea.redpowerfuture.com/red-future/common v0.0.33 github.com/bjang03/gmq v0.0.3 github.com/cloudwego/eino v0.9.5 github.com/cloudwego/eino-examples v0.0.0-20260611092511-bd64846fbc1d @@ -12,7 +12,9 @@ require ( github.com/gogf/gf/contrib/nosql/redis/v2 v2.10.2 github.com/gogf/gf/v2 v2.10.2 github.com/google/uuid v1.6.0 + github.com/gorilla/websocket v1.5.4-0.20250319132907-e064f32e3674 github.com/tidwall/gjson v1.18.0 + github.com/tiger1103/gfast-token v1.0.10 ) require ( @@ -56,7 +58,6 @@ require ( github.com/golang/snappy v1.0.0 // indirect github.com/google/flatbuffers v1.12.1 // indirect github.com/goph/emperror v0.17.2 // indirect - github.com/gorilla/websocket v1.5.4-0.20250319132907-e064f32e3674 // indirect github.com/grokify/html-strip-tags-go v0.1.0 // indirect github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.2 // indirect github.com/hashicorp/consul/api v1.26.1 // indirect @@ -103,7 +104,6 @@ require ( github.com/tidwall/match v1.1.1 // indirect github.com/tidwall/pretty v1.2.1 // indirect github.com/tidwall/sjson v1.2.5 // indirect - github.com/tiger1103/gfast-token v1.0.10 // indirect github.com/twitchyliquid64/golang-asm v0.15.1 // indirect github.com/vcaesar/cedar v0.30.0 // indirect github.com/vmihailenco/msgpack v4.0.4+incompatible // indirect @@ -133,3 +133,6 @@ require ( google.golang.org/protobuf v1.36.8 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect ) + +// 本地消费未发布的新 common(redislock/rediscount 下沉 utils 后的新增方法);发 tag 后移除 +replace gitea.redpowerfuture.com/red-future/common => ../common diff --git a/go.sum b/go.sum index bf6f001..650fd1c 100644 --- a/go.sum +++ b/go.sum @@ -3,6 +3,7 @@ cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMT cloud.google.com/go/compute/metadata v0.7.0/go.mod h1:j5MvL9PprKL39t166CoB1uVHfQMs4tFQZZcKwksXUjo= gitea.redpowerfuture.com/red-future/common v0.0.31 h1:9H8nL5Drazcv7Hs9d4j+cXhaB+7uOllIUqEOyZy1Eao= gitea.redpowerfuture.com/red-future/common v0.0.31/go.mod h1:xPU7aaMxn8rtNnWc2LDUXZL+IkaUkpQeLgflqw9FvdU= +gitea.redpowerfuture.com/red-future/common v0.0.32 h1:O3iZrbPddHD8MQ4WacJ/eKOpSThgwkB34jlb0R+k/D8= gitea.redpowerfuture.com/red-future/common v0.0.32/go.mod h1:xPU7aaMxn8rtNnWc2LDUXZL+IkaUkpQeLgflqw9FvdU= github.com/Azure/azure-sdk-for-go/sdk/azcore v1.17.0/go.mod h1:XCW7KnZet0Opnr7HccfUw1PLc4CjHqpcaxW8DHklNkQ= github.com/Azure/azure-sdk-for-go/sdk/internal v1.10.0/go.mod h1:iZDifYGJTIgIIkYRNWPENUnqx6bJ2xnSDFI2tjwZNuY= @@ -43,6 +44,7 @@ github.com/bgentry/speakeasy v0.1.0/go.mod h1:+zsyZBPWlz7T6j88CTgSN5bM796AkVf0kB github.com/bitly/go-simplejson v0.5.0/go.mod h1:cXHtHw4XUPsvGaxgjIAn8PhEWG9NfngEKAMDJEczWVA= github.com/bjang03/gmq v0.0.2 h1:3CcVorDXYoRIN65bbzwRuUxzkBCkEpHWmKHOkfXzUo0= github.com/bjang03/gmq v0.0.2/go.mod h1:Y7TwWGuV4Cw97WUDaM7x+NC4kyFx1z44WAvNwJV3HV8= +github.com/bjang03/gmq v0.0.3 h1:Yn9GZP1okOc8uh0f/1FFTooV5/mbO4pKrkcK9mTMjok= github.com/bjang03/gmq v0.0.3/go.mod h1:Y7TwWGuV4Cw97WUDaM7x+NC4kyFx1z44WAvNwJV3HV8= github.com/bluele/gcache v0.0.2/go.mod h1:m15KV+ECjptwSPxKhOhQoAFQVtUFjTVkc3H8o0t/fp0= github.com/bmatcuk/doublestar/v4 v4.10.0/go.mod h1:xBQ8jztBU6kakFMg+8WGxn0c6z1fTSPVIjEY1Wr7jzc= diff --git a/update.sql b/update.sql index d0d64d4..57527f0 100644 --- a/update.sql +++ b/update.sql @@ -889,12 +889,14 @@ COMMENT ON COLUMN black_deacon_flow_segment_result.video_url IS '已生成成功 --------------------pgsql创建black_deacon_flow_segment_result表语句--------------------------- --------------------工作流重试兜底:exec_workflow 新增重试/心跳列--------------------- -ALTER TABLE black_deacon_exec_workflow ADD COLUMN IF NOT EXISTS retryable SMALLINT NOT NULL DEFAULT 0; -- 1=程序报错可重试,0=用户取消 +ALTER TABLE black_deacon_exec_workflow ADD COLUMN IF NOT EXISTS retryable SMALLINT NOT NULL DEFAULT 0; -- 0=终局不重试(用户取消/计费门禁拦截),1=可重试(程序报错/关停中断/超时等) ALTER TABLE black_deacon_exec_workflow ADD COLUMN IF NOT EXISTS retry_count INTEGER NOT NULL DEFAULT 0; -- 已重试次数 ALTER TABLE black_deacon_exec_workflow ADD COLUMN IF NOT EXISTS last_heartbeat BIGINT NOT NULL DEFAULT 0; -- 最后心跳(毫秒时间戳) COMMENT ON COLUMN black_deacon_exec_workflow.retryable IS '是否可重试:0-用户取消,1-程序报错'; COMMENT ON COLUMN black_deacon_exec_workflow.retry_count IS '已重试次数'; COMMENT ON COLUMN black_deacon_exec_workflow.last_heartbeat IS '最后心跳时间(毫秒时间戳)'; +ALTER TABLE black_deacon_exec_workflow ADD COLUMN IF NOT EXISTS charge_order_id BIGINT NOT NULL DEFAULT 0; -- 关联计费单ID(shop-user-trade pricing,0=未建单) +COMMENT ON COLUMN black_deacon_exec_workflow.charge_order_id IS '关联计费单ID(shop-user-trade pricing,0=未建单)'; --------------------工作流重试兜底:flow_async_task 统一异步任务表(Task 2 使用)--------------------- CREATE TABLE IF NOT EXISTS black_deacon_flow_async_task ( @@ -919,3 +921,11 @@ CREATE TABLE IF NOT EXISTS black_deacon_flow_async_task ( CREATE UNIQUE INDEX IF NOT EXISTS uk_async_task_exec_node_seg ON black_deacon_flow_async_task(execution_id, node_id, segment_index); CREATE INDEX IF NOT EXISTS idx_async_task_tenant_id ON black_deacon_flow_async_task(tenant_id); CREATE INDEX IF NOT EXISTS idx_async_task_deleted_at ON black_deacon_flow_async_task(deleted_at); + +--------------------崩溃恢复补全用户:exec_workflow.user_id(恢复续跑补 X-User-Info 过 model-gateway 最低余额门禁)--------------------- +ALTER TABLE black_deacon_exec_workflow ADD COLUMN IF NOT EXISTS user_id BIGINT NOT NULL DEFAULT 0; -- 执行用户ID(数字;creator 仅 userName,恢复续跑须按此补全用户 Id) +COMMENT ON COLUMN black_deacon_exec_workflow.user_id IS '执行用户ID(数字,恢复续跑补 X-User-Info 用)'; + +--------------------终局回填业务实扣:exec_workflow.actual_amount(用户钱包实际扣除金额,元;区别于 total_fee=模型按次费用合计)--------------------- +ALTER TABLE black_deacon_exec_workflow ADD COLUMN IF NOT EXISTS actual_amount NUMERIC(15,2) NOT NULL DEFAULT 0; -- 业务扣费(settle/cancel 结算实收;失败/未结算=0) +COMMENT ON COLUMN black_deacon_exec_workflow.actual_amount IS '业务扣费(用户钱包实际扣除金额,元;结算/取消回填实收,失败/未结算=0,区别于 total_fee=模型按次费用合计)'; diff --git a/workflow/dao/flow/flow_async_task_dao.go b/workflow/dao/flow/flow_async_task_dao.go index f72cb4f..106963e 100644 --- a/workflow/dao/flow/flow_async_task_dao.go +++ b/workflow/dao/flow/flow_async_task_dao.go @@ -91,7 +91,10 @@ func (d *flowAsyncTaskDao) DeleteByKey(ctx context.Context, execId int64, nodeId return err } -// DeleteByExecution 清理指定执行的异步任务缓存(exec 成功后调用,与段清理同处) +// DeleteByExecution 清理指定执行的异步任务缓存。统一清理策略(Task 11 方案A,无周期兜底): +// 仅在两处 exec 级删除点调用——① 工作流执行成功后(BuildExecution 尾部,与 checkpoint/段清理同处); +// ② 同一条 exec 以"全新跑"重开(forceNewRun 起跑前,参数已变,旧异步结果必须作废防误复用)。 +// 失败/取消/重试耗尽一律保留产物,供同参数手动续跑(reExecute)复用;不作节点级或周期清扫。 func (d *flowAsyncTaskDao) DeleteByExecution(ctx context.Context, execId int64) error { const physicalTable = "black_deacon_flow_async_task" _, err := gfdb.DB(ctx, public.DbNameBlackDeacon). diff --git a/workflow/dao/node/node_execution_dao.go b/workflow/dao/node/node_execution_dao.go index c054c18..7150f6d 100644 --- a/workflow/dao/node/node_execution_dao.go +++ b/workflow/dao/node/node_execution_dao.go @@ -87,6 +87,11 @@ func (d *nodeExecutionDao) ListByFlowExecutionId(ctx context.Context, req *nodeD model := gfdb.DB(ctx, public.DbNameBlackDeacon).Model(ctx, public.TableNameNodeExecution).NoTenantId(ctx).Fields(fields).OmitEmpty() model.Where(entity.NodeExecutionCol.FlowExecutionId, req.FlowExecutionId) model.Where(entity.NodeExecutionCol.NodeGroupId, req.NodeGroupId) + if req.CreatedAtFrom != nil { + // 结算按订单收敛:只聚合订单创建后产生的节点记录(重跑开新单,created_at 各自独立, + // 避免把已终局运行(已结算扣费)的用量计入本次订单) + model.WhereGTE(entity.NodeExecutionCol.CreatedAt, *req.CreatedAtFrom) + } model.Where(entity.NodeExecutionCol.Status, req.Status) model.Where(entity.NodeExecutionCol.NodeId, req.NodeId) model.OrderAsc(entity.NodeExecutionCol.CreatedAt) diff --git a/workflow/model/dto/node/node_execution_dto.go b/workflow/model/dto/node/node_execution_dto.go index 4480966..387c6c9 100644 --- a/workflow/model/dto/node/node_execution_dto.go +++ b/workflow/model/dto/node/node_execution_dto.go @@ -7,6 +7,7 @@ import ( "gitea.redpowerfuture.com/red-future/common/beans" "github.com/gogf/gf/v2/frame/g" + "github.com/gogf/gf/v2/os/gtime" ) // CreateNodeExecutionReq 创建节点执行记录请求 @@ -64,6 +65,7 @@ type ListNodeExecutionByFlowReq struct { NodeGroupId string `json:"nodeGroupId"` NodeId string `json:"nodeId"` Status node.NodeExecutionStatus `json:"status"` + CreatedAtFrom *gtime.Time `json:"createdAtFrom" dc:"结算按订单收敛:只返回创建时间>=该值的记录(重跑开新单按各自 created_at 隔离)"` } // NodeExecutionResp 节点执行记录响应 diff --git a/workflow/model/dto/session/exec_workflow_dto.go b/workflow/model/dto/session/exec_workflow_dto.go index 932cdc5..c62ee35 100644 --- a/workflow/model/dto/session/exec_workflow_dto.go +++ b/workflow/model/dto/session/exec_workflow_dto.go @@ -6,6 +6,7 @@ import ( ) type CreateWorkflowReq struct { + UserId int64 `json:"userId" description:"执行用户ID(数字,恢复续跑补 X-User-Info 用)"` SessionId string `json:"sessionId" description:"所属会话ID"` FlowId int64 `json:"flowId" description:"工作流ID"` NodeGroupId string `json:"nodeGroupId" description:"节点组ID"` diff --git a/workflow/model/dto/session/session_dto.go b/workflow/model/dto/session/session_dto.go index 787dd68..b5ab121 100644 --- a/workflow/model/dto/session/session_dto.go +++ b/workflow/model/dto/session/session_dto.go @@ -83,7 +83,8 @@ type VOSessionInfoResult struct { ResultFileUrl string `json:"resultFileUrl" description:"结果文件路径(供预览/下载)"` ResultContent string `json:"resultContent" description:"结果文件内容(服务端已读取,前端直接展示)"` TotalTokens int `json:"totalTokens" dc:"总token消耗"` - TotalFee float64 `json:"totalFee" dc:"总费用"` + TotalFee float64 `json:"totalFee" dc:"模型扣费合计(各模型按次费用)"` + ActualAmount float64 `json:"actualAmount" dc:"业务扣费(用户钱包实际扣除,元)"` ErrorMsg string `json:"errorMsg" dc:"错误信息(友好提示)"` Error string `json:"error" dc:"错误明细(原始错误)"` CreatedAt *gtime.Time `json:"createdAt" dc:"创建时间"` diff --git a/workflow/model/entity/exec_workflow.go b/workflow/model/entity/exec_workflow.go index 18622ac..fee3d3f 100644 --- a/workflow/model/entity/exec_workflow.go +++ b/workflow/model/entity/exec_workflow.go @@ -9,6 +9,7 @@ import ( // ExecWorkflow 执行工作流 type ExecWorkflow struct { beans.SQLBaseDO `orm:",inherit"` + UserId int64 `orm:"user_id" json:"userId" description:"执行用户ID(数字,恢复续跑补 X-User-Info 用;creator 是 userName)"` SessionId string `orm:"session_id" json:"sessionId" description:"会话ID"` FlowId int64 `orm:"flow_id" json:"flowId" description:"工作流ID"` NodeGroupId string `orm:"node_group_id" json:"nodeGroupId" description:"节点组ID"` @@ -17,15 +18,18 @@ type ExecWorkflow struct { Status flow.FlowExecutionStatus `orm:"status" json:"status" description:"状态:1-运行中,2-成功,3-失败"` TotalTokens int `orm:"total_tokens" json:"totalTokens" description:"总token消耗"` TotalFee float64 `orm:"total_fee" json:"totalFee" description:"总费用"` + ActualAmount float64 `orm:"actual_amount" json:"actualAmount" description:"业务扣费(用户钱包实际扣除金额,元;结算/取消回填实收,失败/未结算=0,区别于 total_fee=模型按次费用合计)"` ErrorMessage string `orm:"error_message" json:"errorMessage" description:"错误信息(友好提示)"` Error string `orm:"error" json:"error" description:"错误明细(原始错误)"` - Retryable int `orm:"retryable" json:"retryable" description:"是否可重试:0-用户取消,1-程序报错"` + Retryable int `orm:"retryable" json:"retryable" description:"是否可重试:0-终局不重试(用户取消/计费门禁拦截),1-可重试(程序报错/关停中断/超时等)"` RetryCount int `orm:"retry_count" json:"retryCount" description:"已重试次数"` + ChargeOrderId int64 `orm:"charge_order_id" json:"chargeOrderId" description:"关联计费单ID(shop-user-trade pricing,0=未建单)"` LastHeartbeat int64 `orm:"last_heartbeat" json:"lastHeartbeat" description:"最后心跳时间(毫秒时间戳)"` } type execWorkflowCol struct { beans.SQLBaseCol + UserId string SessionId string FlowId string NodeGroupId string @@ -34,15 +38,18 @@ type execWorkflowCol struct { Status string TotalTokens string TotalFee string + ActualAmount string ErrorMessage string Error string Retryable string RetryCount string + ChargeOrderId string LastHeartbeat string } var ExecWorkflowCol = execWorkflowCol{ SQLBaseCol: beans.DefSQLBaseCol, + UserId: "user_id", SessionId: "session_id", FlowId: "flow_id", NodeGroupId: "node_group_id", @@ -51,9 +58,11 @@ var ExecWorkflowCol = execWorkflowCol{ Status: "status", TotalTokens: "total_tokens", TotalFee: "total_fee", + ActualAmount: "actual_amount", ErrorMessage: "error_message", Error: "error", Retryable: "retryable", RetryCount: "retry_count", + ChargeOrderId: "charge_order_id", LastHeartbeat: "last_heartbeat", } diff --git a/workflow/model/entity/flow_user.go b/workflow/model/entity/flow_user.go index 86cb830..198a4fe 100644 --- a/workflow/model/entity/flow_user.go +++ b/workflow/model/entity/flow_user.go @@ -12,6 +12,7 @@ type FlowInfo struct { StartNodeId string `json:"startNodeId"` Nodes []FlowNode `json:"nodes"` Edges []FlowEdge `json:"edges"` + ChargeMode string `json:"chargeMode"` // 计费方式:per_item/per_second/per_token(缺省 per_item,对齐 shop-user-trade consts/pricing) } type FlowNode struct { diff --git a/workflow/service/flow/async_task.go b/workflow/service/flow/async.go similarity index 79% rename from workflow/service/flow/async_task.go rename to workflow/service/flow/async.go index 62d9855..bc99e8c 100644 --- a/workflow/service/flow/async_task.go +++ b/workflow/service/flow/async.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "errors" + "sync" "time" "github.com/gogf/gf/v2/frame/g" @@ -14,6 +15,47 @@ import ( "ai-agent/workflow/model/entity" ) +// 全局等待任务回调的工具 +var ( + asyncMu sync.Mutex + asyncTasks = make(map[string]chan any) +) + +// Wait 阻塞等待回调结果 +// 调用后会一直卡住,直到 Notify 唤醒 或 超时/取消 +func Wait(ctx context.Context, taskId string) (any, error) { + asyncMu.Lock() + ch := make(chan any, 1) + asyncTasks[taskId] = ch + asyncMu.Unlock() + + defer close(ch) + for { + select { + case result := <-ch: + return result, nil + case <-ctx.Done(): + asyncMu.Lock() + delete(asyncTasks, taskId) + asyncMu.Unlock() + return nil, ctx.Err() + } + } +} + +// Notify 回调时调用,唤醒等待的任务 +func Notify(taskId string, result any) { + asyncMu.Lock() + defer asyncMu.Unlock() + + ch, exist := asyncTasks[taskId] + if !exist { + return + } + ch <- result + delete(asyncTasks, taskId) +} + // asyncRecoverWaitTimeout in-flight 行重订阅等结果的超时上限。 // 任务可能仍执行中(消息未发布)或提交即失败(不会发布消息);超时视为结果未知,清记录重提。 const asyncRecoverWaitTimeout = 10 * time.Minute @@ -44,9 +86,11 @@ func asyncCallAction(rec *entity.FlowAsyncTask) asyncAction { } } -// AsyncModelCallWithRecovery 统一异步模型调用入口(spec §5.4): +// AsyncModelCallWithRecovery 统一异步模型调用入口: // 提交时把 task_id/msg_topic 落库 flow_async_task,崩溃后重订阅 msg_topic 拿回已完成结果复用,不重复调用。 // 同步模型直接走 gateway.ModelCallResult,不落库(无恢复语义)。 +// 注意:本函数是节点内阻塞调用(WaitModelCallResult 等回调),不是独立并发触发方, +// 不参与 exec 并发仲裁(谁抢到执行权谁跑)——仲裁语义见《工作流执行并发仲裁设计.md》。 func AsyncModelCallWithRecovery(ctx context.Context, execId int64, nodeId string, segIdx int, modelId int64, responseType model.ResponseType, sessionId string, requestParams map[string]any, businessParams map[string]any) (*gateway.ModelCallRes, error) { if responseType == nil || *responseType != *model.ResponseTypeAsync.Code() { return gateway.ModelCallResult(ctx, modelId, responseType, sessionId, requestParams, businessParams) diff --git a/workflow/service/flow/billing.go b/workflow/service/flow/billing.go new file mode 100644 index 0000000..6dedd00 --- /dev/null +++ b/workflow/service/flow/billing.go @@ -0,0 +1,391 @@ +package flow + +import ( + "context" + "errors" + "fmt" + + commonHttp "gitea.redpowerfuture.com/red-future/common/http" + "gitea.redpowerfuture.com/red-future/common/utils" + "github.com/gogf/gf/v2/os/glog" + "github.com/gogf/gf/v2/os/gtime" + "github.com/gogf/gf/v2/util/gconv" + "github.com/google/uuid" + + nodeDao "ai-agent/workflow/dao/node" + sessionDao "ai-agent/workflow/dao/session" + nodeDto "ai-agent/workflow/model/dto/node" + "ai-agent/workflow/model/entity" +) + +// ====================== 计费本地 DTO(shop-user-trade pricing,独立 module 不可 import,JSON 对齐) ====================== + +type pricingOpenOrderReq struct { + UserId int64 `json:"userId"` + SubjectType string `json:"subjectType"` + SubjectID string `json:"subjectId"` + ChargeMode string `json:"chargeMode"` + BizOrderNo string `json:"bizOrderNo"` +} + +type pricingGetConfigReq struct { + SubjectType string `json:"subjectType"` + SubjectID string `json:"subjectId"` +} + +type pricingGetConfigRes struct { + Enabled int `json:"enabled"` +} + +type pricingChargeOrderInfo struct { + ID int64 `json:"id"` + Status int `json:"status"` // 1已建单 2已结算 3已失败 + ActualAmount float64 `json:"actualAmount"` // 实收/实扣(元):settle/cancel 响应回填,回写 exec_workflow.actual_amount + CreatedAt string `json:"createdAt"` + ChargeMode string `json:"chargeMode"` // per_item / per_token +} + +// pricingSettleReq Settle 与 Cancel 同形状(OrderId + Usage) +type pricingSettleReq struct { + OrderId int64 `json:"orderId"` + Usage map[string]any `json:"usage"` +} + +type pricingFailReq struct { + OrderId int64 `json:"orderId"` + Reason string `json:"reason"` +} + +// 计价对象/模式常量(对齐 shop-user-trade consts/pricing) +const ( + pricingSubjectWorkflow = "workflow" + pricingChargeModePerItem = "per_item" + pricingChargeModePerSecond = "per_second" + pricingChargeModePerToken = "per_token" + pricingOrderStatusCreated = 1 +) + +// errBillingGateBlocked 计费门禁拦截(余额不足/钱包不可用/费率非法/未配置计价/用户缺失/per_second 无视频模型):执行终局失败, +// 不进进程内重试、不落 recoverable(shouldRetry 与 handleExecute 分类均排除)。 +var errBillingGateBlocked = errors.New("计费门禁拦截") + +// pricingURL 组装 shop-user-trade 计价接口地址。 +// 跨服务路由前缀 /pricing/controller/ 由 common http.RouteRegister 按 controller struct 名推导 +// (pricingController → pricing/controller),GoFrame doSetHandler 恒 prefix+uri 拼接。 +func pricingURL(sub string) string { + return "shop-user-trade/pricing/controller/" + sub +} + +// ====================== 建单(执行开始,不动钱) ====================== + +// openBillingOrder 工作流执行开始建计费单(不动钱)。 +// 幂等键 bizOrderNo=wf:{execId};复用终态 execId 重跑(原单已结算/失败)→ 开新单 wf:{execId}:{uuid8}, +// 新单 ID 落 exec_workflow.charge_order_id,结算据此定位。 +// 门禁(余额>=min_balance、钱包须存在)失败 → errBillingGateBlocked,执行终局失败不重试。 +// 计费是工作流执行的前置条件:用户缺失/计价未配置/预检失败/per_second 无视频模型 → 终局失败,不免费跑。 +func openBillingOrder(ctx context.Context, execId int64, flowContent *entity.FlowInfo) error { + user, err := utils.GetUserInfo(ctx) + if err != nil || user == nil || user.Id == 0 { + return fmt.Errorf("%w: 取不到用户 %v", errBillingGateBlocked, err) + } + // config/get 预检 = 服务存活探针 + 启用开关: + // 预检失败(服务宕机/路由不通)→ 终局失败; + // 预检通过(服务确认在、计价已开)后再调 open_order,其失败即为业务错误(余额/钱包/费率)→ 阻塞, + // 确保 shop-user-trade 宕机时工作流执行也被阻断而非免费跑。 + var cfg pricingGetConfigRes + if err = commonHttp.Get(ctx, pricingURL("config/get"), utils.HeadersFromCtx(ctx, utils.HeadersOptions{TokenFromQuery: true}), &cfg, + "subjectType", pricingSubjectWorkflow, "subjectId", pricingSubjectWorkflow); err != nil { + return fmt.Errorf("%w: 计价配置查询失败 %v", errBillingGateBlocked, err) + } + if cfg.Enabled != 1 { + return fmt.Errorf("%w: 计价未启用", errBillingGateBlocked) + } + chargeMode := "" + if flowContent != nil && flowContent.ChargeMode != "" { + chargeMode = flowContent.ChargeMode + } else { + return fmt.Errorf("%w: 工作流执行未选择计费模式", errBillingGateBlocked) + } + // per_second 按秒计费以生成视频总时长为基础,工作流须含视频模型节点,否则配置非法 → 终局失败 + if chargeMode == pricingChargeModePerSecond && !flowHasVideoModel(ctx, flowContent) { + return fmt.Errorf("%w: per_second 计费须工作流包含视频模型", errBillingGateBlocked) + } + info, err := openPricingOrder(ctx, int64(user.Id), chargeMode, fmt.Sprintf("wf:%d", execId)) + if err != nil { + return err + } + orderId := info.ID + if info.Status != pricingOrderStatusCreated { + // 复用终态 execId:原单已结算/失败,开新单让本次运行独立计费 + info, err = openPricingOrder(ctx, int64(user.Id), chargeMode, + fmt.Sprintf("wf:%d:%s", execId, uuid.NewString()[:8])) + if err != nil { + return err + } + orderId = info.ID + } + if err = sessionDao.ExecWorkflowDao.UpdateMap(ctx, execId, map[string]any{ + entity.ExecWorkflowCol.ChargeOrderId: orderId, + }); err != nil { + // 结算时按 bizOrderNo=wf:{execId} 兜底定位,不阻塞 + glog.Errorf(ctx, "工作流计费:记录 charge_order_id 失败 execId=%d: %v", execId, err) + } + return nil +} + +// flowHasVideoModel 工作流是否包含视频模型节点(per_second 计费前置条件)。 +// 遍历节点,按 modelId 去重后经 isVideoModel 查模型类型,任一为视频模型即 true; +// flowContent 缺失/查模型失败按无视频模型处理(per_second 被拦截,fail-closed)。 +func flowHasVideoModel(ctx context.Context, flowContent *entity.FlowInfo) bool { + if flowContent == nil { + return false + } + seen := make(map[int64]struct{}) + for i := range flowContent.Nodes { + modelId := flowContent.Nodes[i].ModelConfig.ModelId + if modelId <= 0 { + continue + } + if _, ok := seen[modelId]; ok { + continue + } + seen[modelId] = struct{}{} + if isVideoModel(ctx, modelId) { + return true + } + } + return false +} + +// openPricingOrder 调 shop-user-trade 建单(幂等:同 subject+bizOrderNo 返回既有单) +func openPricingOrder(ctx context.Context, userId int64, chargeMode, bizOrderNo string) (*pricingChargeOrderInfo, error) { + info := new(pricingChargeOrderInfo) + err := commonHttp.Post(ctx, pricingURL("open_order"), utils.HeadersFromCtx(ctx, utils.HeadersOptions{TokenFromQuery: true}), info, &pricingOpenOrderReq{ + UserId: userId, + SubjectType: pricingSubjectWorkflow, + SubjectID: pricingSubjectWorkflow, + ChargeMode: chargeMode, + BizOrderNo: bizOrderNo, + }) + if err != nil { + return nil, fmt.Errorf("%w: %v", errBillingGateBlocked, err) + } + if info.ID == 0 { + return nil, fmt.Errorf("%w: 未返回计费单ID", errBillingGateBlocked) + } + return info, nil +} + +// ====================== 结算(终态) ====================== + +// settleBilling 工作流终态计费(recordWorkflow 汇聚全部执行路径后调用): +// 成功→Settle 实收;用户取消→Cancel 按已消耗实收;永久失败(retryable=0/重试耗尽)→Fail 不扣费; +// 可恢复失败→跳过(订单留 CREATED,恢复续跑后结算/失败)。 +// 全部计费调用错误仅记日志,不拖垮工作流终态落库。 +func settleBilling(ctx context.Context, execId int64, runErr error, retryable, retryCount *int) { + exec, err := sessionDao.ExecWorkflowDao.GetById(ctx, execId) + if err != nil || exec == nil { + glog.Errorf(ctx, "工作流计费:查询执行失败,跳过结算 execId=%d: %v", execId, err) + return + } + orderId := exec.ChargeOrderId + if orderId == 0 { + // 兜底:charge_order_id 未落库(UpdateMap 失败)时按 bizOrderNo=wf:{execId} 定位;查不到→跳过 + info, e := getPricingOrder(ctx, "", fmt.Sprintf("wf:%d", execId)) + if e != nil || info == nil { + return + } + orderId = info.ID + } + usage, full, err := workflowChargeUsage(ctx, exec, orderId, errors.Is(runErr, context.Canceled)) + if err != nil { + glog.Errorf(ctx, "工作流计费:计算用量失败,跳过结算 execId=%d: %v", execId, err) + return + } + // 终局实扣金额:settle/cancel 由 shop 结算响应回填(元);fail / 可恢复=0(未扣费) + var actual float64 + switch { + case runErr == nil: + actual, _ = callSettlePricing(ctx, orderId, usage, pricingURL("settle")) + case errors.Is(runErr, context.Canceled): + actual, _ = callSettlePricing(ctx, orderId, usage, pricingURL("cancel")) // Cancel 按已消耗实收 + case retryable != nil && *retryable == 0: + callFailPricing(ctx, orderId, runErr.Error()) + case retryCount != nil && *retryCount >= execMaxRetryCount: + callFailPricing(ctx, orderId, runErr.Error()) // 重试耗尽 → 永久失败 + default: + // 可恢复失败(retryable=1 且未耗尽/关停中断):订单留 CREATED,恢复续跑后结算 + } + // 终局回填 exec_workflow(成功/取消/失败/可恢复统一落库,前端与对账均看此行): + // 模型消耗(total_tokens/total_fee,本次运行节点 token_info 按订单窗口聚合)+ 业务实扣 + // (actual_amount=settle/cancel 实收金额,失败/可恢复未结算=0)。失败/取消路径 SummaryLambda 不跑, + // total 列此前恒为空,此处按同一窗口补齐(与结算口径一致,不混入上一运行残留); + // 可恢复失败也先落已消耗,续跑成功后同一订单收敛重算覆盖。错误仅记日志不拖垮终态落库。 + if full != nil { + if err := sessionDao.ExecWorkflowDao.UpdateMap(ctx, execId, map[string]any{ + entity.ExecWorkflowCol.TotalTokens: full.TotalTokens, + entity.ExecWorkflowCol.TotalFee: full.TotalFee, + entity.ExecWorkflowCol.ActualAmount: actual, + }); err != nil { + glog.Errorf(ctx, "exec_workflow 回填消耗/实扣失败 execId=%d: %v", execId, err) + } + } +} + +// workflowChargeUsage 计算工作流结算用量并返回全量聚合(full,含 TotalTokens/TotalFee 供终局回填 exec 行): +// per_token → feeByModel(各模型按次已消耗费用,结算侧按此合计实收——每次模型调用在发生时已由 +// shop /calc 计价,含「不足1分按1分」的按次兜底与调用时媒体/费率快照,不再按聚合 token 重算, +// 避免两笔 0.01 合并重算成 0.01); +// per_item / per_second(其余模式)→ durationSec = 本次生成视频总时长 +// (各视频模型节点记录里模型返回时长之和)。仅生成视频的工作流按时间计费,其余模型调用按条/token 计费。 +// 用户取消(forCancel=true)时 per_item/per_second 在时长之外补收已消耗 token:把非视频节点的按次费用 +// 一并上报(视频节点消耗已由时长计价覆盖,排除防双计)——中途取消通常无视频产出(durationSec=0), +// 但文本/分析节点可能已完成并消耗了 token,须按已消耗补收而非按 0 计。 +func workflowChargeUsage(ctx context.Context, exec *entity.ExecWorkflow, orderId int64, forCancel bool) (usage map[string]any, full *nodeUsageAgg, err error) { + info, err := getPricingOrder(ctx, gconv.String(orderId), "") + if err != nil { + return nil, nil, err + } + // 按订单收敛:只聚合订单创建后产生的节点记录。每次重跑(上单已终态)开新单,created_at 各自独立, + // 隔离「上次已结算运行」与「本次运行」——否则重跑后取消会把上一次已扣费的 token 一起再扣。 + // shop 返回 gtime.String() 无时区本地墙钟,与 node_execution.created_at(timestamp without tz)同格式可比。 + var createdAtFrom *gtime.Time + if info.CreatedAt != "" { + createdAtFrom = gtime.NewFromStr(info.CreatedAt) + if createdAtFrom == nil || createdAtFrom.IsZero() { + return nil, nil, fmt.Errorf("计费单创建时间解析失败: %s", info.CreatedAt) + } + } + full, nonVideo, err := workflowNodeUsage(ctx, exec, createdAtFrom) + if err != nil { + return nil, nil, err + } + switch { + case info.ChargeMode == pricingChargeModePerToken: + // per_token:按次已记录费用结算(含视频模型——per_token 无时长计价,视频模型按自身按次费用计收) + return tokenUsageMap(full), full, nil + case forCancel: + // per_item/per_second 取消补收:时长(full,通常 0)+ 非视频节点按次已消耗费用 + // (nonVideo 排除视频节点:其消耗已由时长计价覆盖,避免双计) + usage := tokenUsageMap(nonVideo) + usage["durationSec"] = full.DurationSec + return usage, full, nil + default: + // 正常结算走时长(视频产出按 per_item/per_second 档位/秒价),不附 token/费用拆分 + return map[string]any{"durationSec": full.DurationSec}, full, nil + } +} + +// nodeUsageAgg 本次执行聚合出的结算用量:FeeByModel(按生效(系统)模型 id 的按次已记录费用合计, +// shop 实收依据)+ DurationSec(时长,供 per_item/per_second 用)。逐模型 token/媒体明细不上报—— +// 每调用一条留在 node_execution.token_info(model_id/total_tokens/prompt_tokens/completion_tokens/ +// media_type/total_fee),订单层按需可从明细再聚合,不再冗余携带。 +type nodeUsageAgg struct { + // FeeByModel 各模型窗口内已消耗费用合计 = 节点 token_info.total_fee 求和。total_fee 本身是 + // 该节点内各次模型调用(每次经 shop /calc 计价,calculator.charge 内部已 ceilFen,「不足1分按1分」 + // 按调用次生效)费用之和 → 此处合计即「按次已记录费用」。结算侧按此实收,不再按聚合 token 重算。 + FeeByModel map[string]float64 + // DurationSec 各视频节点生成视频总时长(per_item/per_second 计价依据) + DurationSec float64 + // TotalTokens / TotalFee 全量节点消耗合计(模型消耗 token / 模型按次费用), + // 终局回填 exec_workflow.total_tokens/total_fee(区别于钱包实扣 ActualAmount)。 + TotalTokens int64 + TotalFee float64 +} + +// workflowNodeUsage 聚合本次执行(FlowExecutionId + 订单创建时间下界)下各节点执行记录写入的用量。 +// 返回两组聚合: +// - full:全部节点记录——per_token 结算(视频模型也按自身按次费用计收,无时长计价)与时长累计用; +// - nonVideo:排除 total_duration>0 的节点(视频节点生成时长,per_item/per_second 取消补收时其消耗 +// 已由时长计价覆盖,不再按次补收,避免双计)。 +// +// 聚合内容:feeByModel = 各模型 total_fee 求和(total_fee = 该节点内各次模型调用经 shop /calc 计价 +// (已 ceilFen)的费用之和 → 按次已记录费用,shop 实收依据);durationSec = 各视频节点生成视频总时长 +// (模型返回,total_duration)累加。逐模型 token/媒体明细留在 node_execution.token_info,订单层不上报。 +// +// 收敛到当前运行:重跑复用同一 exec 记录与节点组(检查点恢复的 SavedFlowInput 携带旧 node_group_id, +// exec_workflow 表也无该列持久化),node_group_id 无法区分运行;改按 created_at >= 订单创建时间过滤—— +// 每次重跑开新单,各自 created_at 隔离本次运行消耗,避免把已结算的上一次运行用量一起计入。 +func workflowNodeUsage(ctx context.Context, exec *entity.ExecWorkflow, createdAtFrom *gtime.Time) (full, nonVideo *nodeUsageAgg, err error) { + records, _, err := nodeDao.NodeExecutionDao.ListByFlowExecutionId(ctx, &nodeDto.ListNodeExecutionByFlowReq{ + FlowExecutionId: exec.Id, + CreatedAtFrom: createdAtFrom, + }) + if err != nil { + return nil, nil, err + } + full = &nodeUsageAgg{FeeByModel: make(map[string]float64)} + nonVideo = &nodeUsageAgg{FeeByModel: make(map[string]float64)} + for _, rec := range records { + for _, ti := range rec.TokenInfo { + addNodeUsageEntry(full, ti) + // total_duration>0 = 视频节点产出了生成时长(lambda 只有视频模型节点累加并落库该字段), + // 该节点消耗由时长计价覆盖,排除出 nonVideo(per_item/per_second 取消补收按此计,防双计) + if gconv.Float64(ti["total_duration"]) <= 0 { + addNodeUsageEntry(nonVideo, ti) + } + } + } + return full, nonVideo, nil +} + +// addNodeUsageEntry 把单条 token_info(节点一次模型调用或调用汇总)累加进聚合。 +// gconv.String 兼容字符串与 JSONB 数字(float64)两种 model_id 写入,避免 (string) 断言丢弃条目。 +func addNodeUsageEntry(agg *nodeUsageAgg, ti map[string]any) { + modelID := gconv.String(ti["model_id"]) + agg.DurationSec += gconv.Float64(ti["total_duration"]) + agg.TotalTokens += gconv.Int64(ti["total_tokens"]) + agg.TotalFee += gconv.Float64(ti["total_fee"]) + if modelID == "" { + return // 无模型 id(异常条目)不计费用(feeByModel 按模型键) + } + agg.FeeByModel[modelID] += gconv.Float64(ti["total_fee"]) +} + +// tokenUsageMap 按次费用口径组装结算用量(token 部分):仅 feeByModel(shop 实收依据, +// per_token 无建单快照)。逐模型 token/媒体明细留在 node_execution.token_info,订单层不上报。 +func tokenUsageMap(agg *nodeUsageAgg) map[string]any { + return map[string]any{ + "feeByModel": agg.FeeByModel, + } +} + +// getPricingOrder 查计费单:id>0 按ID,否则按 subjectType+subjectId+bizOrderNo +func getPricingOrder(ctx context.Context, id, bizOrderNo string) (*pricingChargeOrderInfo, error) { + info := new(pricingChargeOrderInfo) + var data []any + if id != "" { + data = []any{"id", id} + } else { + data = []any{"subjectType", pricingSubjectWorkflow, "subjectId", pricingSubjectWorkflow, "bizOrderNo", bizOrderNo} + } + err := commonHttp.Get(ctx, pricingURL("order"), utils.HeadersFromCtx(ctx, utils.HeadersOptions{TokenFromQuery: true}), info, data...) + if err != nil { + return nil, err + } + if info.ID == 0 { + return nil, errors.New("计费单不存在") + } + return info, nil +} + +// callSettlePricing Settle(成功)或 Cancel(用户取消按已消耗实收),返回实收金额(元)。 +// shop settle/cancel 响应体是结算后的 ChargeOrderInfo(actualAmount=本次实扣),幂等重复结算返回既有金额; +// 调用失败返回 0 并记日志(此时无法确知钱包是否已扣,actual_amount 置 0 保守,不臆造金额)。 +func callSettlePricing(ctx context.Context, orderId int64, usage map[string]any, url string) (actual float64, err error) { + info := new(pricingChargeOrderInfo) + if err := commonHttp.Post(ctx, url, utils.HeadersFromCtx(ctx, utils.HeadersOptions{TokenFromQuery: true}), info, + &pricingSettleReq{OrderId: orderId, Usage: usage}); err != nil { + glog.Errorf(ctx, "工作流计费:结算失败 orderId=%d: %v", orderId, err) + return 0, err + } + return info.ActualAmount, nil +} + +// callFailPricing Fail(不扣费) +func callFailPricing(ctx context.Context, orderId int64, reason string) { + if err := commonHttp.Post(ctx, pricingURL("fail"), utils.HeadersFromCtx(ctx, utils.HeadersOptions{TokenFromQuery: true}), &pricingChargeOrderInfo{}, + &pricingFailReq{OrderId: orderId, Reason: reason}); err != nil { + glog.Errorf(ctx, "工作流计费:失败处理失败 orderId=%d: %v", orderId, err) + } +} diff --git a/workflow/service/flow/flow_checkpoint_store.go b/workflow/service/flow/exec_checkpoint.go similarity index 100% rename from workflow/service/flow/flow_checkpoint_store.go rename to workflow/service/flow/exec_checkpoint.go diff --git a/workflow/service/flow/exec_hub.go b/workflow/service/flow/exec_hub.go index a9f8d79..a8894f5 100644 --- a/workflow/service/flow/exec_hub.go +++ b/workflow/service/flow/exec_hub.go @@ -71,13 +71,6 @@ func registerHubIfAbsent(sessionId string, flowId int64, hub *execHub) (existing return hub, true } -// getExecHub 返回已注册的 hub(无则 nil) -func getExecHub(sessionId string, flowId int64) *execHub { - hubRegMu.Lock() - defer hubRegMu.Unlock() - return hubReg[hubKey(sessionId, flowId)] -} - // ====================== 订阅 / 发布 ====================== func (h *execHub) Subscribe(conn *wsCommon.WsConnection) { diff --git a/workflow/service/flow/exec_lifecycle.go b/workflow/service/flow/exec_lifecycle.go new file mode 100644 index 0000000..3a14edd --- /dev/null +++ b/workflow/service/flow/exec_lifecycle.go @@ -0,0 +1,159 @@ +package flow + +import ( + "context" + "errors" + "sync" + "sync/atomic" + "time" + + "ai-agent/workflow/consts/flow" + sessionDao "ai-agent/workflow/dao/session" + "ai-agent/workflow/model/entity" + + "github.com/gogf/gf/v2/frame/g" + "github.com/google/uuid" +) + +// 运行中执行跟踪:优雅关停时取消全部执行(含脱离连接的恢复执行)并等待落完终态再退出, +// 避免"程序停止 → exec 仍卡在 status=1"。 +// +// 背景:WS 执行(handleExecute)的 execCtx 派生自连接 closeCtx,Close() 取消连接即随之中止落终态; +// 但恢复执行(recoverExecution)的 ctx 经 context.WithoutCancel 脱离连接,程序关停时若不显式取消, +// 恢复中的 exec 不会落终态(仍 status=1),重启要等心跳陈旧 60s 才能再捞。这里登记所有运行中执行, +// SetShuttingDown 统一取消、main 等待全部落库后再退出。 +var ( + execRunMu sync.Mutex + execRuns = make(map[string]context.CancelFunc) +) + +// trackExecRun 登记一次运行中执行,返回 finish 在落完终态(recordWorkflow 后)调用。 +// 每次运行唯一 key(重试/断点续跑复用同一 execId 也不冲突)。 +func trackExecRun(cancel context.CancelFunc) (finish func()) { + key := uuid.NewString() + execRunMu.Lock() + execRuns[key] = cancel + execRunMu.Unlock() + return func() { + execRunMu.Lock() + delete(execRuns, key) + execRunMu.Unlock() + } +} + +// cancelAllExecRuns 优雅关停:取消所有运行中执行。恢复执行 ctx 经 WithoutCancel 脱离连接, +// 必须显式取消才能随 WS 执行一起落终态(status=3/retryable=1)。 +func cancelAllExecRuns() { + execRunMu.Lock() + cancels := make([]context.CancelFunc, 0, len(execRuns)) + for _, c := range execRuns { + cancels = append(cancels, c) + } + execRunMu.Unlock() + for _, c := range cancels { + c() + } +} + +// WaitExecRunsDrain 等待所有运行中执行落完终态(限时),供 main 优雅关停收尾后退出进程。 +// 关停后不再启动新执行(execute/reExecute/recoverExecution 顶部有 IsShuttingDown 守卫), +// 因此运行中集合只减不增,轮询安全。 +func WaitExecRunsDrain(timeout time.Duration) { + deadline := time.Now().Add(timeout) + for { + execRunMu.Lock() + n := len(execRuns) + execRunMu.Unlock() + if n == 0 { + return + } + if time.Now().After(deadline) { + g.Log().Warningf(context.Background(), "优雅关停等待执行落库超时,剩余 %d 个执行", n) + return + } + time.Sleep(100 * time.Millisecond) + } +} + +// errExecAlreadyRunning 用户触发时该执行已在运行(本节点/其它节点后台恢复),不新建执行 +var errExecAlreadyRunning = errors.New("工作流正在执行中,不重复执行") + +// errInterruptedByShutdown 程序优雅关停导致连接 ctx 取消时的错误标记(区别于用户主动取消)。 +// 落库为 status=3 + retryable=1,下次启动恢复扫描捞起续跑。 +var errInterruptedByShutdown = errors.New("程序关停中断") + +// shuttingDown 优雅关停标记:程序收到退出信号后置位。 +// 用于区分"程序关停导致的 WS 连接取消"与"用户主动取消",避免前者被误分类为不可重试。 +var shuttingDown atomic.Bool + +// SetShuttingDown 置位优雅关停标记并取消所有运行中执行(main 信号处理在 Close() 前调用)。 +// 恢复执行 ctx 经 WithoutCancel 脱离连接,必须在此显式取消,否则程序关停时恢复中的 exec +// 不会落终态(仍 status=1);取消后 BuildExecution 随之中止、走恢复错误分支写 status=3/retryable=1。 +func SetShuttingDown() { + shuttingDown.Store(true) + cancelAllExecRuns() +} + +// IsShuttingDown 是否处于优雅关停 +func IsShuttingDown() bool { + return shuttingDown.Load() +} + +// shouldRetry 错误分类(retryable 终局语义与 retry_count 预算见《工作流执行并发仲裁设计.md》§5): +// 用户取消不重试;计费门禁拦截不重试;其余程序报错重试 +func shouldRetry(err error) bool { + return err != nil && !errors.Is(err, context.Canceled) && !errors.Is(err, errExecAlreadyRunning) && + !errors.Is(err, errBillingGateBlocked) +} + +// isRecoverable 判定可恢复(恢复侧谓词,见《工作流执行并发仲裁设计.md》§2/§5):僵尸运行中(status=1 且心跳陈旧)或可重试失败(status=3 且 retryable=1 且未耗尽)。 +// 与 ListRecoverable SQL 判定一致;last_heartbeat=0(老数据/默认)视为陈旧。 +func isRecoverable(exec *entity.ExecWorkflow, nowMs int64) bool { + if exec == nil || exec.Status == nil { + return false + } + status := *exec.Status + switch status { + case *flow.FlowExecutionStatusRunning.Code(): + return exec.LastHeartbeat < nowMs-int64(heartbeatStaleAfter/time.Millisecond) + case *flow.FlowExecutionStatusFailed.Code(): + return exec.Retryable == 1 && exec.RetryCount < execMaxRetryCount + default: + return false + } +} + +// startHeartbeat 后台心跳 goroutine:每 30s touch last_heartbeat,返回 stop 函数。 +// 覆盖正常执行与恢复执行,崩溃前最后一次心跳即崩溃近似时间戳(心跳=在跑活体标记/租约语义见《工作流执行并发仲裁设计.md》§3)。 +// onLeaseLost:心跳连续失败达到陈旧阈值(租约丢失)时回调——调用方应取消执行 ctx, +// 使心跳与执行同生共死,避免"心跳已过期但执行还活着"的窗口被其它节点恢复导致双跑。 +func startHeartbeat(ctx context.Context, execId int64, onLeaseLost func()) func() { + stopCh := make(chan struct{}) + done := make(chan struct{}) + go func() { + defer close(done) + ticker := time.NewTicker(heartbeatInterval) + defer ticker.Stop() + var consecutiveFail int + for { + select { + case <-ticker.C: + if err := sessionDao.ExecWorkflowDao.TouchHeartbeat(ctx, execId); err != nil { + g.Log().Warningf(ctx, "心跳落库失败 execId=%d: %v", execId, err) + consecutiveFail++ + if onLeaseLost != nil && consecutiveFail >= heartbeatMaxFail { + g.Log().Errorf(ctx, "心跳连续失败 %d 次,租约丢失,中止执行 execId=%d", consecutiveFail, execId) + onLeaseLost() + } + } else { + consecutiveFail = 0 + } + case <-stopCh: + return + case <-ctx.Done(): + return + } + } + }() + return func() { close(stopCh); <-done } +} diff --git a/workflow/service/flow/exec_progress.go b/workflow/service/flow/exec_progress.go new file mode 100644 index 0000000..b1475ff --- /dev/null +++ b/workflow/service/flow/exec_progress.go @@ -0,0 +1,46 @@ +package flow + +import ( + "context" + + wsCommon "gitea.redpowerfuture.com/red-future/common/websocket" +) + +// ====================== 进度上报 ====================== +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 +} + +// 进度上报由 exec_hub.go 的 execHub 实现(单执行事件中枢,可多连接订阅);wsProgressCtxKey/ProgressReporter/GetProgressReporter 保留。 + +// 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) +} diff --git a/workflow/service/flow/flow_graph_util.go b/workflow/service/flow/exec_record.go similarity index 53% rename from workflow/service/flow/flow_graph_util.go rename to workflow/service/flow/exec_record.go index d752443..38e2671 100644 --- a/workflow/service/flow/flow_graph_util.go +++ b/workflow/service/flow/exec_record.go @@ -1,21 +1,133 @@ package flow import ( - "ai-agent/gateway" - "ai-agent/workflow/consts/node" - nodeDao "ai-agent/workflow/dao/node" - flowDto "ai-agent/workflow/model/dto/flow" - nodeDto "ai-agent/workflow/model/dto/node" - "ai-agent/workflow/model/entity" "context" + "errors" "fmt" "time" + "gitea.redpowerfuture.com/red-future/common/oss" "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" + + "ai-agent/gateway" + "ai-agent/workflow/consts/flow" + "ai-agent/workflow/consts/node" + nodeDao "ai-agent/workflow/dao/node" + sessionDao "ai-agent/workflow/dao/session" + flowDto "ai-agent/workflow/model/dto/flow" + nodeDto "ai-agent/workflow/model/dto/node" + "ai-agent/workflow/model/entity" ) +// ====================== 执行记录落库 ====================== + +// recordExecutionFailure 记录一次失败状态。有 execId 直接更新该记录;拿不到 execId +// (executeOrResume 在创建记录后、返回前 panic,或查询/创建执行记录失败)时, +// 兜底按会话+工作流查最近一条仍处于"运行中"的记录标记为失败,避免前端已报错但记录卡在 Running。 +// 最近记录已是成功/失败状态则不处理(可能是上一次执行的结果,不应误改)。 +func recordExecutionFailure(ctx context.Context, sessionId string, flowId int64, execId int64, runErr error) { + // 兜底失败路径同样携带 retryable 分类(与 handleExecute 一致): + // 用户取消=0;其余(程序关停中断/程序报错)可重试=1,保证任意 status=3 写库都带分类 + var retryable, retryCnt *int + if runErr != nil { + if errors.Is(runErr, context.Canceled) && !IsShuttingDown() { + retryable, retryCnt = intptr(0), intptr(0) + } else { + retryable, retryCnt = intptr(1), intptr(0) + } + } + if !g.IsEmpty(execId) { + recordWorkflow(ctx, execId, 0, runErr, retryable, retryCnt) + return + } + lastExec, err := sessionDao.ExecWorkflowDao.GetLatestBySessionAndFlow(ctx, sessionId, flowId) + if err != nil || lastExec == nil { + glog.Errorf(ctx, "兜底标记失败状态失败: sessionId=%s flowId=%d err=%v", sessionId, flowId, err) + return + } + if lastExec.Status == nil || *lastExec.Status != *flow.FlowExecutionStatusRunning.Code() { + return + } + recordWorkflow(ctx, lastExec.Id, 0, runErr, retryable, retryCnt) +} + +// intptr 取 int 值指针(recordWorkflow 的 retryable/retryCount 参数:nil=不改动) +func intptr(v int) *int { return &v } + +// recordWorkflow 把一次工作流执行写入 exec_workflow/exec_workflow_result:运行记录 + 输出文件结果。 +// retryable/retryCount 非 nil 时随终态同语句原子落库,避免"先 UpdateRetry 再 Update"两步写部分生效 +// 导致 retryable 与 status/error_message 不一致(关停中断场景曾出现 retryable=0 但 message=程序关停中断) +func recordWorkflow(ctx context.Context, id int64, duration time.Duration, runErr error, retryable, retryCount *int) { + // 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 + switch { + case errors.Is(runErr, context.Canceled): + errorMessage = errWorkflowTerminated + case errors.Is(runErr, errInterruptedByShutdown): + errorMessage = "程序关停中断" + default: + errorMessage = "工作流执行失败" + errorDetail = runErr.Error() + } + } + data := map[string]any{ + entity.ExecWorkflowCol.Status: *status.Code(), + } + if d := int64(duration.Seconds()); d != 0 { + data[entity.ExecWorkflowCol.Duration] = d + } + if errorMessage != "" { + data[entity.ExecWorkflowCol.ErrorMessage] = errorMessage + } + if errorDetail != "" { + data[entity.ExecWorkflowCol.Error] = errorDetail + } + if retryable != nil { + data[entity.ExecWorkflowCol.Retryable] = *retryable + data[entity.ExecWorkflowCol.RetryCount] = *retryCount + } + if err := sessionDao.ExecWorkflowDao.UpdateMap(ctx, id, data); err != nil { + glog.Errorf(ctx, "exec_workflow 终态落库失败 execId=%d: %v", id, err) + return + } + // 执行成功:重新执行复用了同一条记录,需显式清空,避免上一次失败的报错残留 + if runErr == nil { + if _, err := sessionDao.ExecWorkflowDao.ClearError(ctx, id); err != nil { + glog.Errorf(ctx, "exec_workflow 报错信息清空失败: %v", err) + } + } + // 工作流计费:终态结算(成功→Settle/用户取消→Cancel/永久失败→Fail/可恢复→跳过)。 + // recordWorkflow 汇聚全部路径(WS/恢复/panic),此处一处接线全覆盖;计费错误仅记日志不拖垮落库 + settleBilling(ctx, id, runErr, retryable, retryCount) +} + +// 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, _ := oss.GetFileAddressPrefix(ctx) + urls := make([]string, 0, len(results)) + for _, r := range results { + if r.ResultFileUrl != "" { + urls = append(urls, prefix+r.ResultFileUrl) + } + } + return urls +} + +// ====================== 节点执行记录 ====================== + // BuildNodeExecutionInput 构建节点执行入参,包含中断恢复逻辑 func BuildNodeExecutionInput(ctx context.Context, input any, flowNode entity.FlowNode) (*flowDto.FlowExecutionInput, *flowDto.NodeExecutionInput, error) { execInput := new(flowDto.FlowExecutionInput) @@ -59,15 +171,6 @@ func BuildNodeExecutionInput(ctx context.Context, input any, flowNode entity.Flo return nil, nil, fmt.Errorf("节点:%v 节点信息为空", flowNode.Name) } - // 聚合输入来源 - //if len(flowNode.InputSource) > 0 { - // for _, inputSource := range currentConfig.InputSource { - // if sourceConfig, ok := configMap[inputSource.NodeId]; ok { - // currentConfig.OutputResult = append(currentConfig.OutputResult, sourceConfig.OutputResult...) - // } - // } - //} - // 构建节点执行入参 realInput := &flowDto.NodeExecutionInput{ Config: currentConfig, @@ -114,9 +217,6 @@ func HandleFailedNodeExecution(ctx context.Context, execInput *flowDto.FlowExecu } } - // 记录失败到已执行列表 - //RecordExecutionResult(execInput, flowNode.Id, node.NodeExecutionStatusFailed.Code()) - // 触发中断 return compose.Interrupt(ctx, map[string]string{ "node": flowNode.Name, diff --git a/workflow/service/flow/recover_execution.go b/workflow/service/flow/exec_recover.go similarity index 71% rename from workflow/service/flow/recover_execution.go rename to workflow/service/flow/exec_recover.go index a5a2f97..d950f53 100644 --- a/workflow/service/flow/recover_execution.go +++ b/workflow/service/flow/exec_recover.go @@ -4,13 +4,9 @@ import ( "context" "errors" "fmt" - "sync/atomic" "time" - "ai-agent/workflow/consts/flow" - flowDao "ai-agent/workflow/dao/flow" sessionDao "ai-agent/workflow/dao/session" - "ai-agent/workflow/model/entity" "gitea.redpowerfuture.com/red-future/common/beans" "gitea.redpowerfuture.com/red-future/common/utils" @@ -32,37 +28,8 @@ const ( execMaxRetryCount = 2 // 整次执行最多自动重试 2 次(共 3 次尝试) ) -// errExecAlreadyRunning 用户触发时该执行已在运行(本节点/其它节点后台恢复),不新建执行 -var errExecAlreadyRunning = errors.New("工作流正在执行中,不重复执行") - -// errInterruptedByShutdown 程序优雅关停导致连接 ctx 取消时的错误标记(区别于用户主动取消)。 -// 落库为 status=3 + retryable=1,下次启动恢复扫描捞起续跑。 -var errInterruptedByShutdown = errors.New("程序关停中断") - -// shuttingDown 优雅关停标记:程序收到退出信号后置位。 -// 用于区分"程序关停导致的 WS 连接取消"与"用户主动取消",避免前者被误分类为不可重试。 -var shuttingDown atomic.Bool - -// SetShuttingDown 置位优雅关停标记并取消所有运行中执行(main 信号处理在 Close() 前调用)。 -// 恢复执行 ctx 经 WithoutCancel 脱离连接,必须在此显式取消,否则程序关停时恢复中的 exec -// 不会落终态(仍 status=1);取消后 BuildExecution 随之中止、走恢复错误分支写 status=3/retryable=1。 -func SetShuttingDown() { - shuttingDown.Store(true) - cancelAllExecRuns() -} - -// IsShuttingDown 是否处于优雅关停 -func IsShuttingDown() bool { - return shuttingDown.Load() -} - -// shouldRetry 错误分类(spec §3):用户取消不重试;其余(模型/DB/网络/超时/panic)程序报错重试 -func shouldRetry(err error) bool { - return err != nil && !errors.Is(err, context.Canceled) && !errors.Is(err, errExecAlreadyRunning) -} - // StartRecoveryLoop 启动恢复扫描:先立即扫一次,再周期扫描。 -// 多节点各自扫描,靠 Redis 锁对同一 exec 抢占去重(spec §8.2/8.3) +// 多节点各自扫描,靠 Redis 锁对同一 exec 抢占去重(Redis 锁只管瞬时互斥、运行期靠 DB 心跳,见《工作流执行并发仲裁设计.md》§1/§2) func StartRecoveryLoop(ctx context.Context) { go func() { scanAndRecover(ctx) @@ -101,7 +68,8 @@ func scanAndRecover(ctx context.Context) { } } -// recoverExecution 统一恢复例程(spec §7):抢锁 → 判定 → 置运行中续跑。 +// recoverExecution 统一恢复例程:抢锁 → 判定(isRecoverable)→ 条件重置抢权 → 置运行中续跑 +// (触发源/抢权闸/输家纪律见《工作流执行并发仲裁设计.md》§1/§2/§4)。 // 两个触发源共用:启动/周期扫描(attach=nil)、executeOrResume 对僵尸行的用户触发(attach 携带连接)。 func recoverExecution(parentCtx context.Context, execId int64, attach *execAttach) { // 优雅关停期间不再启动新的恢复执行(周期扫描已有守卫;用户 executeOrResume 触发路径这里兜底) @@ -165,8 +133,10 @@ func recoverExecution(parentCtx context.Context, execId int64, attach *execAttac // 心跳在条件重置成功后才启动:避免对未抢到重置权的 exec 空跑心跳 saveCtx := context.WithoutCancel(parentCtx) // 恢复无 HTTP 用户,但图节点 lambda 的 INSERT(node_execution/flow_async_task/segment_result)走 - // insertHook 硬性要求 user;用 exec 所属租户合成系统用户 ctx(保留 span),供后续落库使用 - userCtx := context.WithValue(saveCtx, "user", &beans.User{UserName: exec.Creator, TenantId: exec.TenantId}) + // insertHook 硬性要求 user,且节点模型调用外发 model-gateway 带 X-User-Info 要过单次调用最低 + // 余额门禁(须 user.Id>0)——用 exec 所属租户合成系统用户 ctx(Id 取创建时落库的 user_id, + // creator 仅 userName 推不回数字 id;旧记录 user_id=0 时其恢复续跑会被该门禁拦截),保留 span。 + userCtx := context.WithValue(saveCtx, "user", &beans.User{Id: uint64(exec.UserId), UserName: exec.Creator, TenantId: exec.TenantId}) nodeGroupId := uuid.NewString() // 置运行中与图执行的节点组标识须一致 // 条件重置(原子防与用户断点续跑双跑):仅当仍可恢复(status=3 或 status=1 心跳陈旧)时才抢到重置权 staleBeforeMs := time.Now().UnixMilli() - int64(heartbeatStaleAfter/time.Millisecond) @@ -203,7 +173,7 @@ func recoverExecution(parentCtx context.Context, execId int64, attach *execAttac // 心跳连续失败达陈旧阈值时回调 execCancel 中止执行(心跳与执行同生共死)。 // 从函数顶部登记的 topCtx 派生执行 ctx:注入用户信息(供 insertHook 落库)+ 12h 执行超时上限。 // 关停时 cancelAllExecRuns 取消 topCancel → 本 ctx 随之取消,BuildExecution 中止后走下方错误分类落终态。 - execCtx, execCancel := context.WithTimeout(context.WithValue(topCtx, "user", &beans.User{UserName: exec.Creator, TenantId: exec.TenantId}), recoverExecTimeout) + execCtx, execCancel := context.WithTimeout(context.WithValue(topCtx, "user", &beans.User{Id: uint64(exec.UserId), UserName: exec.Creator, TenantId: exec.TenantId}), recoverExecTimeout) defer execCancel() stop := startHeartbeat(execCtx, execId, execCancel) defer stop() @@ -215,27 +185,24 @@ func recoverExecution(parentCtx context.Context, execId int64, attach *execAttac err = BuildExecution(progressCtx, false, exec.FlowId, execId, nodeGroupId, exec.SessionId, exec.RequestParams) if err != nil { // 用户显式取消(附着连接 workflow_cancel):永久取消,retryable=0,恢复扫描不再捞起, - // 杜绝"取消→恢复→再取消"循环;清 flow_async_task 孤儿缓存,落"用户已终止执行" + // 杜绝"取消→恢复→再取消"循环;产物(checkpoint/段/异步缓存)保留,落"用户已终止执行"。 + // 与 WS 路径一致:失败/取消不清,统一由"成功尾部 / 下一次 forceNewRun 起跑前"清理 if hub != nil && hub.UserCancelled() { retryable, retryCnt := 0, 0 - _ = flowDao.FlowAsyncTaskDao.DeleteByExecution(userCtx, execId) recordWorkflow(userCtx, execId, 0, context.Canceled, &retryable, &retryCnt) hub.Publish(&wsCommon.WsPushMsg{Type: "error", Message: errWorkflowTerminated}) return nil } // 续跑失败:恢复例程无用户,任意错误(含执行超时/租约丢失取消/模型/DB/网络/panic) // 一律 retryable=1 交下一轮扫描决定是否再恢复;重试耗尽才终局失败。 + // 产物(checkpoint/段/异步缓存)保留不在此清:重试耗尽后用户仍可手动同参数续跑 + // 复用已成功段/异步结果,换参数则由 forceNewRun 起跑前统一清理(见 BuildExecution) retryable, retryCnt := 1, exec.RetryCount+1 // 优雅关停导致的取消:换错误标记让 recordWorkflow 写"程序关停中断"(与 WS 路径一致), // 仍 retryable=1 下次启动扫描捞起续跑;非关停的取消(租约丢失/执行超时)保留原错误 if errors.Is(err, context.Canceled) && IsShuttingDown() { err = errInterruptedByShutdown } - // 终局清理(Task 5 用户裁定):重试耗尽 → 执行永久失败, - // 清理该 exec 残留的 flow_async_task 孤儿缓存,避免未来复用同一 execId 的运行误取到过期 done 结果 - if exec.RetryCount+1 >= execMaxRetryCount { - _ = flowDao.FlowAsyncTaskDao.DeleteByExecution(userCtx, execId) - } recordWorkflow(userCtx, execId, 0, err, &retryable, &retryCnt) // 终态广播(附着连接;scan 路径无连接则无人接收) if hub != nil { @@ -262,55 +229,3 @@ func recoverExecution(parentCtx context.Context, execId int64, attach *execAttac return } } - -// isRecoverable 判定可恢复(spec §3):僵尸运行中(status=1 且心跳陈旧)或可重试失败(status=3 且 retryable=1 且未耗尽)。 -// 与 ListRecoverable SQL 判定一致;last_heartbeat=0(老数据/默认)视为陈旧。 -func isRecoverable(exec *entity.ExecWorkflow, nowMs int64) bool { - if exec == nil || exec.Status == nil { - return false - } - status := *exec.Status - switch status { - case *flow.FlowExecutionStatusRunning.Code(): - return exec.LastHeartbeat < nowMs-int64(heartbeatStaleAfter/time.Millisecond) - case *flow.FlowExecutionStatusFailed.Code(): - return exec.Retryable == 1 && exec.RetryCount < execMaxRetryCount - default: - return false - } -} - -// startHeartbeat 后台心跳 goroutine:每 30s touch last_heartbeat,返回 stop 函数。 -// 覆盖正常执行与恢复执行,崩溃前最后一次心跳即崩溃近似时间戳(spec §4)。 -// onLeaseLost:心跳连续失败达到陈旧阈值(租约丢失)时回调——调用方应取消执行 ctx, -// 使心跳与执行同生共死,避免"心跳已过期但执行还活着"的窗口被其它节点恢复导致双跑。 -func startHeartbeat(ctx context.Context, execId int64, onLeaseLost func()) func() { - stopCh := make(chan struct{}) - done := make(chan struct{}) - go func() { - defer close(done) - ticker := time.NewTicker(heartbeatInterval) - defer ticker.Stop() - var consecutiveFail int - for { - select { - case <-ticker.C: - if err := sessionDao.ExecWorkflowDao.TouchHeartbeat(ctx, execId); err != nil { - g.Log().Warningf(ctx, "心跳落库失败 execId=%d: %v", execId, err) - consecutiveFail++ - if onLeaseLost != nil && consecutiveFail >= heartbeatMaxFail { - g.Log().Errorf(ctx, "心跳连续失败 %d 次,租约丢失,中止执行 execId=%d", consecutiveFail, execId) - onLeaseLost() - } - } else { - consecutiveFail = 0 - } - case <-stopCh: - return - case <-ctx.Done(): - return - } - } - }() - return func() { close(stopCh); <-done } -} diff --git a/workflow/service/flow/exec_shutdown.go b/workflow/service/flow/exec_shutdown.go deleted file mode 100644 index fc18b02..0000000 --- a/workflow/service/flow/exec_shutdown.go +++ /dev/null @@ -1,70 +0,0 @@ -package flow - -import ( - "context" - "sync" - "time" - - "github.com/gogf/gf/v2/frame/g" - "github.com/google/uuid" -) - -// 运行中执行跟踪:优雅关停时取消全部执行(含脱离连接的恢复执行)并等待落完终态再退出, -// 避免"程序停止 → exec 仍卡在 status=1"。 -// -// 背景:WS 执行(handleExecute)的 execCtx 派生自连接 closeCtx,Close() 取消连接即随之中止落终态; -// 但恢复执行(recoverExecution)的 ctx 经 context.WithoutCancel 脱离连接,程序关停时若不显式取消, -// 恢复中的 exec 不会落终态(仍 status=1),重启要等心跳陈旧 60s 才能再捞。这里登记所有运行中执行, -// SetShuttingDown 统一取消、main 等待全部落库后再退出。 -var ( - execRunMu sync.Mutex - execRuns = make(map[string]context.CancelFunc) -) - -// trackExecRun 登记一次运行中执行,返回 finish 在落完终态(recordWorkflow 后)调用。 -// 每次运行唯一 key(重试/断点续跑复用同一 execId 也不冲突)。 -func trackExecRun(cancel context.CancelFunc) (finish func()) { - key := uuid.NewString() - execRunMu.Lock() - execRuns[key] = cancel - execRunMu.Unlock() - return func() { - execRunMu.Lock() - delete(execRuns, key) - execRunMu.Unlock() - } -} - -// cancelAllExecRuns 优雅关停:取消所有运行中执行。恢复执行 ctx 经 WithoutCancel 脱离连接, -// 必须显式取消才能随 WS 执行一起落终态(status=3/retryable=1)。 -func cancelAllExecRuns() { - execRunMu.Lock() - cancels := make([]context.CancelFunc, 0, len(execRuns)) - for _, c := range execRuns { - cancels = append(cancels, c) - } - execRunMu.Unlock() - for _, c := range cancels { - c() - } -} - -// WaitExecRunsDrain 等待所有运行中执行落完终态(限时),供 main 优雅关停收尾后退出进程。 -// 关停后不再启动新执行(execute/reExecute/recoverExecution 顶部有 IsShuttingDown 守卫), -// 因此运行中集合只减不增,轮询安全。 -func WaitExecRunsDrain(timeout time.Duration) { - deadline := time.Now().Add(timeout) - for { - execRunMu.Lock() - n := len(execRuns) - execRunMu.Unlock() - if n == 0 { - return - } - if time.Now().After(deadline) { - g.Log().Warningf(context.Background(), "优雅关停等待执行落库超时,剩余 %d 个执行", n) - return - } - time.Sleep(100 * time.Millisecond) - } -} diff --git a/workflow/service/flow/flow_ws_exec.go b/workflow/service/flow/exec_ws.go similarity index 58% rename from workflow/service/flow/flow_ws_exec.go rename to workflow/service/flow/exec_ws.go index 6f9d404..14f3191 100644 --- a/workflow/service/flow/flow_ws_exec.go +++ b/workflow/service/flow/exec_ws.go @@ -12,12 +12,14 @@ import ( "encoding/json" "errors" "fmt" + "strings" "time" - "gitea.redpowerfuture.com/red-future/common/oss" + "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/net/ghttp" "github.com/gogf/gf/v2/os/glog" "github.com/gogf/gf/v2/util/gconv" "github.com/google/uuid" @@ -26,7 +28,7 @@ import ( // ====================== WebSocket 服务器 ====================== func init() { - // 工作流消息处理器注册在统一的 SessionWsService(见 ws_server.go)上: + // 工作流消息处理器注册在统一的 SessionWsService 上: // 首次连接仅升级,连接后按消息 type 路由,不再建连时区分普通对话/工作流 SessionWsService.OnMessage("workflow", handleExecute) SessionWsService.OnMessage("workflow_cancel", handleCancel) @@ -38,25 +40,6 @@ 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 -} - -// 进度上报由 exec_hub.go 的 execHub 实现(单执行事件中枢,可多连接订阅);wsProgressCtxKey/ProgressReporter/GetProgressReporter 保留。 - // ====================== 消息处理 ====================== // handleExecute 处理工作流执行(由 workerPool 异步调用,不阻塞读循环) @@ -70,9 +53,13 @@ func handleExecute(ctx context.Context, conn *wsCommon.WsConnection, payload int execCtx, execCancel := context.WithCancel(ctx) - // 替换旧 cancel:同一连接上重新执行时取消上一次遗留的执行 - if oldCancel := getExecCancel(conn); oldCancel != nil { - oldCancel() + // 同流程运行中再点"执行" = 附着看进度(此处不做预取消):下方 registerHubIfAbsent 复用同 + // session+flow 的运行中 hub,executeOrResume 返回 errExecAlreadyRunning,本连接附着其进度/取消, + // 不新建执行。原 getExecCancel 无条件预取消会把运行中的同流程执行一并杀掉,与附着语义冲突,故弃用。 + // 仅当本连接 meta execHub 指向的是另一流程(换流程执行)的运行中 hub 时终止它: + // 避免会话内串跑两个流程、旧流程在后台继续消耗计费(workflow_cancel 是显式取消入口)。 + if prev, ok := wsCommon.GetMetaT[*execHub](conn, "execHub"); ok && prev != nil && prev.flowId != execPayload.FlowId { + prev.CancelByUser() } // 异步执行工作流(直接 goroutine,不依赖上游 workerPool 二次排队) @@ -97,7 +84,6 @@ func handleExecute(ctx context.Context, conn *wsCommon.WsConnection, payload int start := time.Now() var execId int64 var execErr error - recorded := false owner := false // 本 goroutine 是否为执行持有者(附着/触发恢复时不持有) // 持有者收尾:落完终态广播后关闭 hub(退订全部连接/清 meta/注销)。 // 附着路径不置 owner,由运行中的持有方统一 Close;panic 路径在下方 recover defer 中置 owner 兜底。 @@ -111,16 +97,13 @@ func handleExecute(ctx context.Context, conn *wsCommon.WsConnection, payload int if r := recover(); r != nil { glog.Errorf(execCtx, "workflow panic: %v", r) execErr = fmt.Errorf("工作流异常: %v", r) - // panic 发生在 executeOrResume 内部(如节点 lambda panic)时, - // 多返回值赋值不会完成,此处 execId 可能为 0,需走兜底按会话+流程查最近"运行中"记录标记失败, - // 避免 exec_workflow 记录卡在 Running - if !recorded { - recordExecutionFailure(saveCtx, conn.SessionId, execPayload.FlowId, execId, execErr) - recorded = true - } + // 只置错误与所有权,不做终态落库/广播:recover 后本 goroutine 从 panic 处(execId/execErr + // 赋值可能未完成)继续执行,下方常规错误路径恰好覆盖此场景且只走一遍—— + // execId=0 → else-if 兜底 recordExecutionFailure(按会话+流程查最近 Running 标记失败); + // execId>0 → recordWorkflow;随后单次终态 Publish。若在此提前落库/广播,主路径会重复 + // 一次(前端连收两条错误、兜底标记被调两遍)。 // panic 必然发生在本 goroutine 持有的执行内(附着/恢复路径不跑 BuildExecution,不会 panic) owner = true - hub.Publish(&wsCommon.WsPushMsg{Type: "error", Message: "工作流异常", Error: fmt.Sprintf("%v", r)}) } }() @@ -134,7 +117,12 @@ func handleExecute(ctx context.Context, conn *wsCommon.WsConnection, payload int _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "error", Message: "工作流会话创建失败", Error: fmt.Sprintf("%v", e)}) } - _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "ack", Message: fmt.Sprintf("开始执行工作流(共 %d 个节点)", len(execPayload.FlowContent.Nodes))}) + // flowContent 经 WS 的 gconv 解析不跑 v:"required" 校验,缺省时按 0 节点提示,不 panic + nodeCount := 0 + if execPayload.FlowContent != nil { + nodeCount = len(execPayload.FlowContent.Nodes) + } + _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "ack", Message: fmt.Sprintf("开始执行工作流(共 %d 个节点)", nodeCount)}) execId, execErr = executeOrResume(progressCtx, conn, execPayload) if errors.Is(execErr, errExecAlreadyRunning) { @@ -176,22 +164,22 @@ func handleExecute(ctx context.Context, conn *wsCommon.WsConnection, payload int execErr = errInterruptedByShutdown case errors.Is(execErr, context.Canceled): retryable, retryCnt = intptr(0), intptr(0) + case errors.Is(execErr, errBillingGateBlocked): + // 计费门禁拦截:余额不足/钱包不可用/费率非法,终局失败不重试不恢复 + retryable, retryCnt = intptr(0), intptr(0) default: + // 程序报错且重试耗尽 → exec 永久失败:async/segment/checkpoint 一律**保留**, + // 供用户同参数再点续跑(reExecute)复用已产出、免重复调模型/免双扣; + // 换参数重跑走 execute→forceNewRun,BuildExecution 起跑前统一清空,不误用旧残留 retryable, retryCnt = intptr(1), intptr(retryCount) - if retryCount >= execMaxRetryCount { - // 终局清理(Task 5 裁定):重试耗尽,exec 永久失败,清 flow_async_task 孤儿缓存 - _ = flowDao.FlowAsyncTaskDao.DeleteByExecution(saveCtx, execId) - } } } if !g.IsEmpty(execId) { glog.Infof(saveCtx, "工作流执行完成,execId: %v", execId) recordWorkflow(saveCtx, execId, time.Since(start), execErr, retryable, retryCnt) - recorded = true } else if execErr != nil { // 查询/创建执行记录失败(拿不到 execId)时,兜底把该会话+工作流最近一条"运行中"记录标记为失败 recordExecutionFailure(saveCtx, conn.SessionId, execPayload.FlowId, execId, execErr) - recorded = true } if execErr != nil { // 终态广播(发起连接 + 附着订阅者) @@ -209,125 +197,13 @@ func handleExecute(ctx context.Context, conn *wsCommon.WsConnection, payload int }() } -// recordExecutionFailure 记录一次失败状态。有 execId 直接更新该记录;拿不到 execId -// (executeOrResume 在创建记录后、返回前 panic,或查询/创建执行记录失败)时, -// 兜底按会话+工作流查最近一条仍处于"运行中"的记录标记为失败,避免前端已报错但记录卡在 Running。 -// 最近记录已是成功/失败状态则不处理(可能是上一次执行的结果,不应误改)。 -func recordExecutionFailure(ctx context.Context, sessionId string, flowId int64, execId int64, runErr error) { - // 兜底失败路径同样携带 retryable 分类(与 handleExecute 一致): - // 用户取消=0;其余(程序关停中断/程序报错)可重试=1,保证任意 status=3 写库都带分类 - var retryable, retryCnt *int - if runErr != nil { - if errors.Is(runErr, context.Canceled) && !IsShuttingDown() { - retryable, retryCnt = intptr(0), intptr(0) - } else { - retryable, retryCnt = intptr(1), intptr(0) - } - } - if !g.IsEmpty(execId) { - recordWorkflow(ctx, execId, 0, runErr, retryable, retryCnt) - return - } - lastExec, err := sessionDao.ExecWorkflowDao.GetLatestBySessionAndFlow(ctx, sessionId, flowId) - if err != nil || lastExec == nil { - glog.Errorf(ctx, "兜底标记失败状态失败: sessionId=%s flowId=%d err=%v", sessionId, flowId, err) - return - } - if lastExec.Status == nil || *lastExec.Status != *flow.FlowExecutionStatusRunning.Code() { - return - } - recordWorkflow(ctx, lastExec.Id, 0, runErr, retryable, retryCnt) -} - -// intptr 取 int 值指针(recordWorkflow 的 retryable/retryCount 参数:nil=不改动) -func intptr(v int) *int { return &v } - -// recordWorkflow 把一次工作流执行写入 exec_workflow/exec_workflow_result:运行记录 + 输出文件结果。 -// retryable/retryCount 非 nil 时随终态同语句原子落库,避免"先 UpdateRetry 再 Update"两步写部分生效 -// 导致 retryable 与 status/error_message 不一致(关停中断场景曾出现 retryable=0 但 message=程序关停中断) -func recordWorkflow(ctx context.Context, id int64, duration time.Duration, runErr error, retryable, retryCount *int) { - // 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 - switch { - case errors.Is(runErr, context.Canceled): - errorMessage = errWorkflowTerminated - case errors.Is(runErr, errInterruptedByShutdown): - errorMessage = "程序关停中断" - default: - errorMessage = "工作流执行失败" - errorDetail = runErr.Error() - } - } - data := map[string]any{ - entity.ExecWorkflowCol.Status: *status.Code(), - } - if d := int64(duration.Seconds()); d != 0 { - data[entity.ExecWorkflowCol.Duration] = d - } - if errorMessage != "" { - data[entity.ExecWorkflowCol.ErrorMessage] = errorMessage - } - if errorDetail != "" { - data[entity.ExecWorkflowCol.Error] = errorDetail - } - if retryable != nil { - data[entity.ExecWorkflowCol.Retryable] = *retryable - data[entity.ExecWorkflowCol.RetryCount] = *retryCount - } - if err := sessionDao.ExecWorkflowDao.UpdateMap(ctx, id, data); err != nil { - glog.Errorf(ctx, "exec_workflow 终态落库失败 execId=%d: %v", id, err) - return - } - // 执行成功:重新执行复用了同一条记录,需显式清空,避免上一次失败的报错残留 - 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, _ := oss.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 是"并发触发仲裁"的用户侧决策点:同一条 exec 的多个触发方 +// (自动重试 / 手动续跑 / 恢复扫描 / 用户点击)同时发生时,谁真正跑由条件重置 +// (ResetRunning/ResetRunningIfRecoverable,DB 行级原子)定夺,输家一律附着观察或 +// 放弃(errExecAlreadyRunning),绝不双跑。点击落点分情形(status=1 心跳新鲜→附着 / +// 陈旧→触发恢复 / status=3 同参→手动续跑 / 其余→execute)见根目录 +// 《工作流执行并发仲裁设计.md》§1/§4。 +// // executeOrResume 决策工作流执行方式: // - 同会话+同工作流的最近一次执行失败,且本次传递参数与上次一致 → 断点续跑(reExecute,复用原执行记录,从失败断点继续) // - 其余情况(上次成功 / 上次参数与本次不同 / 无历史记录 / 查询出错)→ 全新执行(execute) @@ -341,7 +217,8 @@ func executeOrResume(ctx context.Context, conn *wsCommon.WsConnection, req *sess if *lastExec.Status == *flow.FlowExecutionStatusRunning.Code() { // status=1:可能正在跑(本节点或其它节点后台恢复)或僵尸遗留。 // 心跳新鲜 → 真在跑,不新建避免双跑;心跳陈旧 → 僵尸,触发后台恢复。 - // 都不新建执行,恢复在后台完成,完成后状态自然收敛(spec §9)。 + // 都不新建执行:恢复在后台完成,完成后状态自然收敛——输家不写终态、由持有方收敛 + // (仲裁语义见《工作流执行并发仲裁设计.md》§1/§4) nowMs := time.Now().UnixMilli() hub := getProgressHub(ctx) owned := hub != nil && hub.Owned() @@ -396,9 +273,18 @@ func execute(ctx context.Context, conn *wsCommon.WsConnection, execId int64, sta return 0, context.Canceled } var nodeGroupId = uuid.NewString() - if g.IsEmpty(execId) { - glog.Infof(ctx, "工作流全新执行execute,无历史记录") - execId, err = sessionDao.ExecWorkflowDao.Insert(ctx, &sessionDto.CreateWorkflowReq{ + // 记录发起执行用户的数字 ID:崩溃恢复续跑无 WS/HTTP 用户,需按 exec.user_id 补全合成用户 + // 的 Id,外发 model-gateway/modelCall 的 X-User-Info 才能过单次调用最低余额门禁 + // (creator 仅存 userName,推不回数字 id)。取不到用户不阻塞执行(openBillingOrder 后续会拦); + // user_id=0 仅影响此类记录自身的恢复续跑。 + var execUserId int64 + if u, e := utils.GetUserInfo(ctx); e == nil && u != nil { + execUserId = int64(u.Id) + } + // 全新执行与复用旧 ID 新建两条路径共用同一插入逻辑,收敛为 createExec + createExec := func() (int64, error) { + execId, err := sessionDao.ExecWorkflowDao.Insert(ctx, &sessionDto.CreateWorkflowReq{ + UserId: execUserId, SessionId: conn.SessionId, FlowId: req.FlowId, NodeGroupId: nodeGroupId, @@ -406,61 +292,42 @@ func execute(ctx context.Context, conn *wsCommon.WsConnection, execId int64, sta RequestParams: req.FlowContent, LastHeartbeat: time.Now().UnixMilli(), }) - if err != nil || g.IsEmpty(execId) { + if err == nil && g.IsEmpty(execId) { + err = fmt.Errorf("创建执行记录返回空ID") + } + if err != nil { glog.Errorf(ctx, "工作流执行记录创建失败: %v", err) + return 0, err + } + return execId, nil + } + // 复用失败记录重跑:仅当上次执行为失败状态(executeOrResume 传入的 lastExec.Status)时重置复用; + // 上次成功 / 无历史 → 一律新建执行记录(createExec)。 + // FlowExecutionStatus 是 *int8 别名,Code() 返回包级指针,直接 == 是地址比较恒为 false, + // 需解引用按值比较,否则复用失败记录时不会重置为 Running、也不更新 RequestParams + if execId > 0 && status != nil && *status == *flow.FlowExecutionStatusFailed.Code() { + var reset bool + reset, err = sessionDao.ExecWorkflowDao.ResetRunning(ctx, execId, nodeGroupId, *flow.FlowExecutionStatusFailed.Code()) + if err != nil { + return + } + if !reset { + // 已被其它路径(恢复例程/并发触发)抢先重置为运行中:放弃本次执行,状态由持有方收敛 + return execId, errExecAlreadyRunning + } + // 复用失败记录时参数可能已变:把新参数落库,供后续 reExecute/恢复例程按记录参数续跑 + _, err = sessionDao.ExecWorkflowDao.Update(ctx, &sessionDto.UpdateWorkflowReq{Id: execId, RequestParams: req.FlowContent}) + if err != nil { return } } else { - // FlowExecutionStatus 是 *int8 别名,Code() 返回包级指针,直接 == 是地址比较恒为 false, - // 需解引用按值比较,否则复用失败记录时不会重置为 Running、也不更新 RequestParams - if status != nil && *status == *flow.FlowExecutionStatusFailed.Code() { - glog.Infof(ctx, "工作流断点续跑execute,execId: %v", execId) - var reset bool - reset, err = sessionDao.ExecWorkflowDao.ResetRunning(ctx, execId, nodeGroupId, *flow.FlowExecutionStatusFailed.Code()) - if err != nil { - return - } - if !reset { - // 已被其它路径(恢复例程/并发触发)抢先重置为运行中:放弃本次执行,状态由持有方收敛 - return execId, errExecAlreadyRunning - } - // 复用失败记录时参数可能已变:把新参数落库,供后续 reExecute/恢复例程按记录参数续跑 - _, err = sessionDao.ExecWorkflowDao.Update(ctx, &sessionDto.UpdateWorkflowReq{Id: execId, RequestParams: req.FlowContent}) - if err != nil { - return - } - } else { - glog.Infof(ctx, "工作流全新执行execute,lastExec: %v", execId) - execId, err = sessionDao.ExecWorkflowDao.Insert(ctx, &sessionDto.CreateWorkflowReq{ - SessionId: conn.SessionId, - FlowId: req.FlowId, - NodeGroupId: nodeGroupId, - Status: flow.FlowExecutionStatusRunning.Code(), - RequestParams: req.FlowContent, - LastHeartbeat: time.Now().UnixMilli(), - }) - if err != nil || g.IsEmpty(execId) { - glog.Errorf(ctx, "工作流执行记录创建失败: %v", err) - return - } + execId, err = createExec() + if err != nil { + return } } - _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "round_start", Message: "运行开始", Data: map[string]interface{}{ - "id": execId, - }}) - // WS 路径心跳与 BuildExecution 已共享同一 ctx(连接取消/DB 故障同生共死),无需租约丢失回调 - stop := startHeartbeat(ctx, execId, nil) - defer stop() - // 接管候选 hub:本 goroutine 确认真正启动执行前标记 owned,供并发附着连接识别运行中持有者 - if h := getProgressHub(ctx); h != nil && !h.Owned() { - h.MarkOwned() - } - err = BuildExecution(ctx, true, req.FlowId, execId, nodeGroupId, conn.SessionId, req.FlowContent) - if err != nil { - return - } - return execId, nil + return launchExecution(ctx, conn, execId, req.FlowId, nodeGroupId, conn.SessionId, req.FlowContent, true) } // reExecute 重新执行工作流。 @@ -485,18 +352,37 @@ func reExecute(ctx context.Context, execWorkflowId int64, prevStatus int8) (id i // 状态已被其它路径(恢复例程/并发触发)抢先重置为运行中:放弃续跑 return flowInfo.Id, errExecAlreadyRunning } + return launchExecution(ctx, nil, flowInfo.Id, flowInfo.FlowId, nodeGroupId, flowInfo.SessionId, flowInfo.RequestParams, false) +} + +// launchExecution execute 与 reExecute 共用的启动尾部:计费建单 →(可选)推送 round_start → +// 心跳 → hub 接管 → BuildExecution。 +// conn 非 nil(WS 路径)时在计费通过后推送 round_start;reExecute 无连接传 nil 不推送。 +// 返回语义与调用方原始约定一致:计费门禁失败返回 (execId, err)(终局失败), +// BuildExecution 失败返回 (0, err),成功返回 (execId, nil)。 +func launchExecution(ctx context.Context, conn *wsCommon.WsConnection, execId int64, flowId int64, nodeGroupId string, sessionId string, flowContent *entity.FlowInfo, forceNewRun bool) (id int64, err error) { + // 工作流计费:建计费单(门禁:余额>=min_balance,钱包须存在)。 + // 业务错误(余额不足/钱包不可用/费率非法)→ errBillingGateBlocked 终局失败,不重试不恢复; + // 续跑复用 execId,原计费单仍 CREATED 则幂等沿用、原单已终态则开新单 + if err := openBillingOrder(ctx, execId, flowContent); err != nil { + return execId, err + } + if conn != nil { + _ = writeJSON(conn, &wsCommon.WsPushMsg{Type: "round_start", Message: "运行开始", Data: map[string]interface{}{ + "id": execId, + }}) + } // WS 路径心跳与 BuildExecution 已共享同一 ctx(连接取消/DB 故障同生共死),无需租约丢失回调 - stop := startHeartbeat(ctx, flowInfo.Id, nil) + stop := startHeartbeat(ctx, execId, nil) defer stop() - // 接管候选 hub:确认续跑启动前标记 owned(与 execute 一致) + // 接管候选 hub:确认真正启动执行前标记 owned,供并发附着连接识别运行中持有者 if h := getProgressHub(ctx); h != nil && !h.Owned() { h.MarkOwned() } - err = BuildExecution(ctx, false, flowInfo.FlowId, flowInfo.Id, nodeGroupId, flowInfo.SessionId, flowInfo.RequestParams) - if err != nil { - return + if err = BuildExecution(ctx, forceNewRun, flowId, execId, nodeGroupId, sessionId, flowContent); err != nil { + return 0, err } - return flowInfo.Id, nil + return execId, nil } func BuildExecution(ctx context.Context, forceNewRun bool, flowId, executionId int64, nodeGroupId string, sessionId string, flowContent *entity.FlowInfo) (err error) { @@ -513,14 +399,7 @@ func BuildExecution(ctx context.Context, forceNewRun bool, flowId, executionId i // ========================================================================= // 构建 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 - } + configMap := buildConfigMap(flowContent, nodeList) // ========================================================================= // 构建全局执行入参 @@ -534,20 +413,48 @@ func BuildExecution(ctx context.Context, forceNewRun bool, flowId, executionId i ForceNewRun: forceNewRun, } - var opts []compose.Option - opts = append(opts, compose.WithCheckPointID(gconv.String(executionId))) - if forceNewRun { - opts = append(opts, compose.WithForceNewRun()) - } - // 全新执行前清理该执行残留段结果:forceNewRun 复用同一条 exec 记录时可能留有旧参数生成的段, - // 不清则续跑会误复用。只按 execution_id 清理(放在图启动前,避免多视频节点互相误删) + // 全新执行前统一清理该执行残留(checkpoint + 段结果 + 异步任务缓存): + // forceNewRun 复用同一条 exec 记录时可能留有旧参数/旧运行的产物,若不清: + // - 旧 checkpoint 会让"本应全新跑"误断点续跑(WithForceNewRun 虽绕开 Eino 读, + // 但 DB 残留键会被后续复用拾取,且成功尾部 Delete 也是清此键); + // - 旧异步 done 结果会在 asyncCallAction 里被复用(key 含 execId+nodeId+segIdx, + // 参数已变时语义过期); + // - 旧段结果被 lambda 复用生成错参视频。 + // 统一收口:只按 execution_id 清理(放在图启动前,避免多视频节点互相误删)。失败/取消不清, + // 交给成功尾部(下方同三处)或下一次 forceNewRun 起跑前统一清。 if forceNewRun { + if err := flowDao.FlowCheckpointDao.Delete(ctx, gconv.String(executionId)); err != nil { + return fmt.Errorf("清理断点失败: %v", err) + } + if err := flowDao.FlowAsyncTaskDao.DeleteByExecution(ctx, executionId); err != nil { + return fmt.Errorf("清理异步任务缓存失败: %v", err) + } if err := flowDao.FlowSegmentResultDao.DeleteByExecution(ctx, executionId); err != nil { return fmt.Errorf("清理段结果失败: %v", err) } } - _, err = runGraph.Invoke(ctx, execInput, opts...) - if err != nil { + // 驱动循环:编译期已对每个业务节点注册 WithInterruptAfterNodes(graph_build.go),节点正常完成后 + // Eino 自动暂停并落 checkpoint。这里识别出"纯进度暂停"后同 checkpoint id 立即续跑,直至图完整跑完 + // (err==nil) 或遇到真正终态(节点失败 / 用户取消 / 非中断错误)。崩溃硬杀恢复 BuildExecution(false) + // 走同一循环:查得断点即从断点续跑,已完成的同步节点不再重跑/重复计费。 + // 判别器:节点失败中断 RerunNodes 恒非空(HandleFailedNodeExecution→compose.Interrupt);ctx 取消由 + // 下方 ctx.Err() 短路;异步模型调用是阻塞式(async.go),不产生空 RerunNodes 伪暂停。故 RerunNodes + // 为空 = 编译期纯进度暂停 → 续跑。详见《工作流节点断点续跑技术设计.md》。 + // WithForceNewRun 只允许出现在全新跑首轮(此时起跑前三清也已只做一次);续跑/暂停轮绝不能带, + // 否则把断点续跑打成"忽略断点从图头重跑"。 + first := forceNewRun + maxIter := len(flowContent.Nodes)*2 + 8 // 纯进度暂停每轮必推进 ≥1 节点, 正常轮数 ≤ 节点数+1; 超限即疑似死循环 + iter := 0 + for { + runOpts := []compose.Option{compose.WithCheckPointID(gconv.String(executionId))} + if first { + runOpts = append(runOpts, compose.WithForceNewRun()) + first = false + } + _, err = runGraph.Invoke(ctx, execInput, runOpts...) + if err == nil { + break // 图完整跑完 → 下方成功尾部三清 + } // 图执行被 ctx 取消(WS 断连/用户终止):返回 context.Canceled 语义,让 recordWorkflow // 记为"用户已终止执行"。此时 Eino 已把断点写入 checkpoint store(DbCheckPointStore 用 // WithoutCancel 落库),重新提交相同参数即可断点续跑。 @@ -555,22 +462,32 @@ func BuildExecution(ctx context.Context, forceNewRun bool, flowId, executionId i return fmt.Errorf("执行工作流失败: %w", ctxErr) } 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 !infoOk { + return fmt.Errorf("执行工作流失败: %v", err) + } + if len(info.RerunNodes) == 0 { + // 纯进度暂停(编译期 after-node checkpoint 已落库)→ 同 id 续跑下一段 + iter++ + if iter > maxIter { + return fmt.Errorf("执行工作流失败: 断点续跑超限(%d 次), 疑似死循环", maxIter) } - if !g.IsEmpty(errMsg) { - err = fmt.Errorf("%v个节点,%v", errNodeCount, errMsg) + continue + } + // 节点失败中断 → 终态失败 + var sb strings.Builder + 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) + fmt.Fprintf(&sb, "\n节点:%v, 失败原因:%v", valMap["node"], valMap["error"]) + } + } + if sb.Len() > 0 { + err = fmt.Errorf("%v个节点,%v", errNodeCount, strings.TrimPrefix(sb.String(), "\n")) } return fmt.Errorf("执行工作流失败: %v", err) } @@ -582,3 +499,33 @@ func BuildExecution(ctx context.Context, forceNewRun bool, flowId, executionId i _ = flowDao.FlowAsyncTaskDao.DeleteByExecution(ctx, executionId) return } + +// SessionWsService 会话 WebSocket 服务器:普通对话与工作流共用一条连接, +// 首次连接仅升级,后续按消息 type 路由到对话/工作流处理器 +// (对话处理器在 react_ws_exec.go 注册,工作流处理器在上方 init 注册)。 +var SessionWsService = wsCommon.NewWsServer( + wsCommon.WithConnKeyPrefix("ws:session:"), +) + +// WsConnect 控制器统一入口:升级 WebSocket(普通对话/工作流均由消息 type 区分,此处不区分) +func WsConnect(ctx context.Context, r *ghttp.Request, req *sessionDto.WebSocketConnectReq) error { + _, err := SessionWsService.Upgrade(ctx, r, req.SessionId) + return err +} + +// ensureSession 解析前端 sessionId 并确保会话存在:命中已存在会话则复用其 id,否则按 name 新建。 +// 普通对话(react_ws_exec.go)与工作流(exec_ws.go)共用。 +func ensureSession(ctx context.Context, sessionId string, name string) error { + exist, err := sessionDao.SessionDao.GetById(ctx, sessionId) + if err != nil { + return err + } + if exist != nil { + return nil + } + if r := []rune(name); len(r) > 128 { // session_name VARCHAR(128) + name = string(r[:128]) + } + _, err = sessionDao.SessionDao.Insert(ctx, &sessionDto.CreateSessionReq{SessionId: sessionId, SessionName: name}) + return err +} diff --git a/workflow/service/flow/flow_helper.go b/workflow/service/flow/flow_helper.go index 09ac4bf..38ea497 100644 --- a/workflow/service/flow/flow_helper.go +++ b/workflow/service/flow/flow_helper.go @@ -3,10 +3,7 @@ package flow import ( "ai-agent/workflow/model/entity" "net/url" - "path" "path/filepath" - "regexp" - "strconv" "strings" ) @@ -45,28 +42,8 @@ func FindEndNodes(startNodeId string, edges []entity.FlowEdge) []string { return res } -// ExtractFlowNodeFrom 从 FlowInfo 中提取节点列表,并自动补齐 DataMerge 节点的 InputSource +// ExtractFlowNodeFrom 从 FlowInfo 中提取节点列表(返回指针切片) func ExtractFlowNodeFrom(flowContent *entity.FlowInfo) []*entity.FlowNode { - // 构建每个节点的上游节点映射 - upstreamMap := make(map[string][]string) - for _, edge := range flowContent.Edges { - upstreamMap[edge.To] = append(upstreamMap[edge.To], edge.From) - } - - // 同时更新 flowContent.Nodes 中的 DataMerge 节点 - //for i := range flowContent.Nodes { - // n := &flowContent.Nodes[i] - // // 对于 DataMerge 节点,自动根据边关系填充 InputSource - // if n.NodeCode == node.NodeTypeDataMerge { - // n.InputSource = nil - // for _, fromId := range upstreamMap[n.Id] { - // n.InputSource = append(n.InputSource, entity.FlowNodeInputSource{ - // NodeId: fromId, - // }) - // } - // } - //} - var flowNodes []*entity.FlowNode for _, item := range flowContent.Nodes { flowNodes = append(flowNodes, &item) @@ -108,110 +85,3 @@ func GetFileTypeByPath(filePath string) string { return "" } } - -// GetUrlSuffix 获取URL文件后缀 -// rawUrl: 原始链接 -// withDot: true 返回 .mp4 false 返回 mp4 -func GetUrlSuffix(rawUrl string, withDot bool) string { - // 解析URL,剥离查询参数 - u, err := url.Parse(rawUrl) - if err != nil { - return "" - } - - // 提取路径部分 - filePath := u.Path - // 获取文件名 - fileName := path.Base(filePath) - if fileName == "" || !strings.Contains(fileName, ".") { - return "" - } - - // 截取后缀 - suffix := path.Ext(fileName) - if !withDot { - suffix = strings.TrimPrefix(suffix, ".") - } - return suffix -} - -// ExtractImageCount 修复:支持单引号/双引号 + 换行 + 空格 -func ExtractImageCount(content string) int { - // 🔥 关键:支持 class='image-count' (单引号) - re := regexp.MustCompile(`

]*>.*?(\d+).*?

`) - match := re.FindStringSubmatch(content) - if len(match) >= 2 { - num, err := strconv.Atoi(match[1]) - if err == nil { - return num - } - } - return 0 -} - -func ImageTagRegex(html string) string { - // 🔥 修复:支持单引号、双引号、空格、换行,100% 删除

- imageTagRegex := regexp.MustCompile(`

]*>[\s\S]*?

`) - return imageTagRegex.ReplaceAllString(html, "") -} - -// StripHtmlTags 去掉所有HTML标签,保留换行和文本结构,并删除配图标记行 -func StripHtmlTags(html string) string { - // 1. 替换块级标签为换行,保证排版 - blockTags := regexp.MustCompile(`]*>`) - text := blockTags.ReplaceAllString(html, "\n") - - // 2. 去掉所有剩余的 HTML 标签 - allTags := regexp.MustCompile(`<[^>]+>`) - text = allTags.ReplaceAllString(text, "") - - // 4. 清理多余空行(多个换行只保留一个) - text = regexp.MustCompile(`\n\s*\n`).ReplaceAllString(text, "\n") - - // 5. 只去掉首尾空白,中间换行保留 - text = strings.TrimSpace(text) - - return text -} - -// SplitMultiContents 拆分模型返回的多条文案(基于HTML标签分隔) -func SplitMultiContents(htmlContent string) []string { - var contents []string - // 正则匹配
包裹的内容 - re := regexp.MustCompile(`
([\s\S]*?)
`) - matches := re.FindAllStringSubmatch(htmlContent, -1) - for _, match := range matches { - if len(match) > 1 { - // 清理空内容 - trimmed := strings.TrimSpace(match[1]) - if trimmed != "" { - contents = append(contents, trimmed) - } - } - } - // 兜底:如果没有匹配到结构化内容,按换行/分隔符拆分 - if len(contents) == 0 { - contents = strings.Split(htmlContent, "===分隔符===") // 提示词中可新增此兜底规则 - } - return contents -} - -// GetAllImgSrcFromHtml 先把提取img src的工具方法放在外面 -func GetAllImgSrcFromHtml(html string) []string { - var imgSrcList []string - re := regexp.MustCompile(`]*src\s*=\s*["']([^"']+)["']`) - submatch := re.FindAllStringSubmatch(html, -1) - for _, match := range submatch { - if len(match) >= 2 { - imgSrcList = append(imgSrcList, match[1]) - } - } - return imgSrcList -} - -// ReplaceImgSrc 替换img src的方法 -func ReplaceImgSrc(html string, oldSrc string, newSrc string) string { - // 精准替换:找到 - re := regexp.MustCompile(`(]*src\s*=\s*["'])` + regexp.QuoteMeta(oldSrc) + `(["'])`) - return re.ReplaceAllString(html, `${1}`+newSrc+`${2}`) -} diff --git a/workflow/service/flow/flow_graph_builder.go b/workflow/service/flow/graph_build.go similarity index 65% rename from workflow/service/flow/flow_graph_builder.go rename to workflow/service/flow/graph_build.go index 651fbed..04bde00 100644 --- a/workflow/service/flow/flow_graph_builder.go +++ b/workflow/service/flow/graph_build.go @@ -14,9 +14,9 @@ import ( "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 类型合并 +// init 注册自定义合并函数:处理 *flowDto.FlowExecutionInput 类型合并。 +// 合并函数全局唯一且与图内容无关,放包级 init 注册一次(避免每次 BuildGraph 重注册全局状态)。 +func init() { compose.RegisterValuesMergeFunc(func(values []*flowDto.FlowExecutionInput) (*flowDto.FlowExecutionInput, error) { if len(values) == 0 { return nil, nil @@ -59,7 +59,10 @@ func BuildGraph(ctx context.Context, flowContent *entity.FlowInfo) ([]entity.Flo } return base, nil }) +} +// BuildGraph 根据 FlowInfo 构建完整的 Eino Graph 拓扑 +func BuildGraph(ctx context.Context, flowContent *entity.FlowInfo) ([]entity.FlowNode, *compose.Graph[any, any]) { graph := compose.NewGraph[any, any]( // 本地状态初始化 compose.WithGenLocalState(func(ctx context.Context) *flowDto.NodeExecutionState { @@ -68,12 +71,8 @@ func BuildGraph(ctx context.Context, flowContent *entity.FlowInfo) ([]entity.Flo ) // 注册所有节点 - nodeMap := make(map[string]entity.FlowNode) for _, item := range flowContent.Nodes { - nodeMap[item.Id] = item - //if item.NodeCode != node.NodeTypeJudge { registerNodeToGraph(graph, item) - //} } // 注册开始节点 @@ -102,71 +101,15 @@ func BuildGraph(ctx context.Context, flowContent *entity.FlowInfo) ([]entity.Flo } // 构建边关系 - 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) } } @@ -176,10 +119,40 @@ func BuildGraph(ctx context.Context, flowContent *entity.FlowInfo) ([]entity.Flo // 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)) + // BuildGraph 已把 summary(保存结果)节点追加进 flowContent.Nodes,此时是全部已注册节点的完整集合。 + // 方案: 每个业务节点正常完成后自动暂停并落 checkpoint(编译期 WithInterruptAfterNodes), + // 崩溃恢复(BuildExecution(false)续跑)即跳过已完成的同步节点, 不再重跑/重复计费。 + // 详见根目录《工作流节点断点续跑技术设计.md》。Start 型空载节点不为它落 cp; 只接 END 的节点 + // Eino 不落暂停(无下游续跑点), 列了也无副作用。 + interruptAfter := make([]string, 0, len(flowContent.Nodes)) + for _, n := range flowContent.Nodes { + if n.NodeCode == node.NodeTypeStart { + continue + } + interruptAfter = append(interruptAfter, n.Id) + } + compile, err := graph.Compile(ctx, + compose.WithGraphName("auto_build_workflow"), + compose.WithCheckPointStore(NewDbCheckPointStore()), + compose.WithNodeTriggerMode(compose.AllPredecessor), + compose.WithInterruptAfterNodes(interruptAfter), + ) return nodeList, compile, err } +// buildConfigMap 由 FlowInfo + 图节点列表构建 ConfigMap:先放流程配置节点,再放图中补充节点 +// (保存结果节点等),供节点执行时按 nodeId 查配置。nodeList 取自 BuildGraph 返回值。 +func buildConfigMap(flowContent *entity.FlowInfo, nodeList []entity.FlowNode) map[string]*entity.FlowNode { + configMap := make(map[string]*entity.FlowNode) + for _, cfg := range ExtractFlowNodeFrom(flowContent) { + configMap[cfg.Id] = cfg + } + for i := range nodeList { + configMap[nodeList[i].Id] = &nodeList[i] + } + return configMap +} + // registerNodeToGraph 将单个节点注册到图中(包含通用包装逻辑) func registerNodeToGraph(graph *compose.Graph[any, any], flowNode entity.FlowNode) { // 通用包装:全程入参都是 *FlowExecutionInput @@ -200,11 +173,7 @@ func registerNodeToGraph(graph *compose.Graph[any, any], flowNode entity.FlowNod // 上报节点执行进度(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)) + reporter.ReportStart(flowNode.Id, flowNodeDesc, nodeReportIndex(execInput, flowNode.Id, 1), len(execInput.ConfigMap)) } // 上传入参到OSS @@ -236,11 +205,7 @@ func registerNodeToGraph(graph *compose.Graph[any, any], flowNode entity.FlowNod // 上报节点执行进度(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)) + reporter.ReportComplete(flowNode.Id, flowNodeDesc, nodeReportIndex(execInput, flowNode.Id, 0), len(execInput.ConfigMap)) } // 返回整个 execInput,让下一个节点继续用 @@ -265,25 +230,18 @@ func registerNodeToGraph(graph *compose.Graph[any, any], flowNode entity.FlowNod _ = 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))) } } +// nodeReportIndex 计算节点在进度上报中的序号:节点已在已执行列表则取其位置, +// 否则按当前已执行数 + offset(start 上报时节点尚未入列 offset=1;complete 后已入列命中 IndexOf) +func nodeReportIndex(execInput *flowDto.FlowExecutionInput, nodeId string, offset int) int { + if idx := IndexOf(execInput.ExecutedNodes, nodeId); idx != -1 { + return idx + } + return len(execInput.ExecutedNodes) + offset +} + // IndexOf 返回元素第一次出现的下标,不存在返回 -1 func IndexOf(slice []flowDto.ExecutedNode, target string) int { for i, v := range slice { diff --git a/workflow/service/flow/lambda_core.go b/workflow/service/flow/lambda_core.go new file mode 100644 index 0000000..a32566c --- /dev/null +++ b/workflow/service/flow/lambda_core.go @@ -0,0 +1,252 @@ +package flow + +import ( + "ai-agent/gateway" + flowDao "ai-agent/workflow/dao/flow" + nodeDao "ai-agent/workflow/dao/node" + flowDto "ai-agent/workflow/model/dto/flow" + nodeDto "ai-agent/workflow/model/dto/node" + "ai-agent/workflow/model/entity" + "ai-agent/workflow/service/flow/processor/builtin/media" + "ai-agent/workflow/service/flow/values" + "context" + "fmt" + "sync" + + "github.com/gogf/gf/v2/util/gconv" +) + +// StartLambda 启动节点 +func StartLambda(ctx context.Context, input any) (any, error) { + return input, nil +} + +// FormLambda 表单调用节点 +func FormLambda(ctx context.Context, input any) (any, error) { + nodeInput, ok := input.(*flowDto.NodeExecutionInput) + if !ok { + return nil, fmt.Errorf("入参类型错误") + } + // 解析 valueSource 引用,填充表单节点输出配置(供下游引用) + for _, output := range nodeInput.Config.OutputConfig { + values.ProcessValueSourceRecursive(output, nodeInput.Global) + } + return nodeInput, nil +} + +// ModelLambda 模型调用节点 +func ModelLambda(ctx context.Context, input any) (any, error) { + nodeInput, ok := input.(*flowDto.NodeExecutionInput) + if !ok { + return nil, fmt.Errorf("入参类型错误") + } + + modelParams, err := values.BuildModelRequestBody(nodeInput.Config.ModelConfig.ModelRequestParamsPath, nodeInput.Global) + if err != nil { + return nil, err + } + + // 2. 前置工具:决定模型调用入参(单次/多次) + // 入参统一为扁平模型请求体(BuildModelRequestBody 输出,key 为点分路径)。 + // 分批处理器按默认上限拆分集合字段,其余前置工具(如 split_shots_pipeline)读取扁平参数。 + preToolParams := modelParams + paramsList, err := invokePreTool(ctx, nodeInput.Config.PreTool, preToolParams) + if err != nil { + return nil, err + } + + // 3. 逐批调用模型,汇总输出(保持请求顺序),累计 token/费用供节点记录落库 + var outputRes []map[string]any + var totalTokens int64 + var totalPrompt int64 + var totalCompletion int64 + var totalCost float64 + var totalDuration int64 + // 计价用生效模型 id:引用行由 model-gateway 解析为系统模型 id(ModelCallRes.ModelId,计价按系统模型); + // 未返回(model-gateway 旧版本)时回落节点配置的模型 id。media_type 供 per_token 命中媒体价。 + effModelID := nodeInput.Config.ModelConfig.ModelId + var effMediaType string + if len(paramsList) > 1 { + // 段级续跑仅在"多段 + 视频模型"启用;非视频分段(批量文本等)走原逻辑零影响。 + // 段身份 = 列表位置(0-based):paramsList 顺序即段序,concat 按列表顺序拼接; + // 位置互不重复且跨 reExecute 稳定(参数一致 → 段数/顺序不变)。不依赖 params 里的 + // segment_index——真实链路(上游 split_shots_pipeline 转写 → 下游 split_segment 按 + // __segment_fields 拆分,invokePreTool 剥离 __ 内部键)下 paramsList 只有模型参数。 + // segVideo=false 时走既有非段级合并路径(全量生成、不落库、不复用),全新执行行为不变, + // 后续自动 concat 判断(独立的 isVideoModel 调用)仍正常执行。 + segVideo := isVideoModel(ctx, nodeInput.Config.ModelConfig.ModelId) + + // 续跑(!ForceNewRun)时读取该节点已成功段;全新执行不查(BuildExecution 已清旧段),saved 为 nil → 全量重生成 + var saved map[int]entity.SegmentRef + if !nodeInput.Global.ForceNewRun && segVideo { + saved, err = flowDao.FlowSegmentResultDao.ListByNode(ctx, nodeInput.Global.ExecutionId, nodeInput.Config.Id) + if err != nil { + return nil, err + } + } + + idxList, needGen := planSegmentResume(paramsList, saved) + + results := make([][]map[string]any, len(paramsList)) + tokenRes := make([]*gateway.ModelCallRes, len(paramsList)) + errs := make([]error, len(paramsList)) + saveErrs := make([]error, len(paramsList)) + isInference := make([]bool, len(paramsList)) + var wg sync.WaitGroup + for i, params := range paramsList { + if !needGen[i] { + continue + } + wg.Add(1) + go func(i int, params map[string]any) { + defer wg.Done() + // 每段单次调用,不原地重试:段失败即走节点失败收口(HandleFailedNodeExecution → Interrupt), + // 下次 reExecute 由 planSegmentResume 复用已成功段、仅重生成失败段 + results[i], tokenRes[i], isInference[i], errs[i] = ModelCallResultLambda(ctx, nodeInput.Config.ModelConfig.ModelId, nodeInput.Global.SessionId, params, nodeInput.Config.Prompt, nodeInput.Global.ExecutionId, nodeInput.Config.Id, idxList[i]) + // 每段成功立即落库:该段刚成功即持久化,其他段仍在跑时已成功段也不丢; + // 后续段失败或进程崩溃(panic/OOM/kill)时,已完成段已在库中,reExecute 可直接复用 + if segVideo && errs[i] == nil { + for _, rec := range results[i] { + key := media.FindVideoKey(rec) + url := media.FindVideoURL(ctx, rec) + if key == "" || url == "" { + continue + } + if err := flowDao.FlowSegmentResultDao.Save(ctx, nodeInput.Global.ExecutionId, nodeInput.Config.Id, idxList[i], key, url); err != nil { + saveErrs[i] = err + } + } + } + }(i, params) + } + wg.Wait() + + // 仍有失败段或落库失败 → 节点失败(成功段已立即落库,供下次 reExecute 复用) + for i := range results { + if saveErrs[i] != nil { + return nil, saveErrs[i] + } + if needGen[i] && errs[i] != nil { + return nil, errs[i] + } + if needGen[i] && tokenRes[i] != nil { + totalTokens += tokenRes[i].TotalTokens + totalPrompt += tokenRes[i].PromptTokens + totalCompletion += tokenRes[i].CompletionTokens + totalCost += tokenRes[i].Cost + if isVideoModel(ctx, nodeInput.Config.ModelConfig.ModelId) { + totalDuration += tokenRes[i].Duration + } + if tokenRes[i].ModelId > 0 { + effModelID = tokenRes[i].ModelId + } + if effMediaType == "" { + effMediaType = tokenRes[i].MediaType + } + } + } + + if segVideo { + // 复用段 + 新生段按段序号升序合并,concat 按列表顺序拼接 → 顺序保证 + outputRes = mergeSegmentOutputs(idxList, needGen, results, saved) + } else { + if isInference[0] { + outputRes = mergeInferenceBatchResults(results) + } else { + for _, res := range results { + outputRes = append(outputRes, res...) + } + } + } + } else { + for _, params := range paramsList { + res, modelRes, _, err := ModelCallResultLambda(ctx, nodeInput.Config.ModelConfig.ModelId, nodeInput.Global.SessionId, params, nodeInput.Config.Prompt, nodeInput.Global.ExecutionId, nodeInput.Config.Id, flowDao.FlowAsyncSegSentinel) + if err != nil { + return nil, err + } + if modelRes != nil { + totalTokens += modelRes.TotalTokens + totalPrompt += modelRes.PromptTokens + totalCompletion += modelRes.CompletionTokens + totalCost += modelRes.Cost + if isVideoModel(ctx, nodeInput.Config.ModelConfig.ModelId) { + totalDuration += modelRes.Duration + } + if modelRes.ModelId > 0 { + effModelID = modelRes.ModelId + } + if effMediaType == "" { + effMediaType = modelRes.MediaType + } + } + outputRes = append(outputRes, res...) + } + } + + // 3.5 把本次节点消耗的 token/费用/生成视频时长写入节点执行记录,供汇总节点聚合到 exec_workflow + // model_id 供 per_token 结算按模型聚合 token;total_duration 供 per_item/per_second 按生成视频总时长计费; + // per_char 模型把输出字数映射到 completion_tokens 传输,随 token 拆分一并落库 + if nodeInput.NodeExecutionId > 0 && (totalTokens > 0 || totalCost > 0 || totalDuration > 0) { + if _, err = nodeDao.NodeExecutionDao.Update(ctx, &nodeDto.UpdateNodeExecutionReq{ + Id: nodeInput.NodeExecutionId, + TokenInfo: []map[string]any{{ + // model_id 写字符串:token_info 为 JSONB,int64 落库成 JSON 数字,读回是 float64, + // billing 侧按 model 聚合时 (string) 断言会失败导致 per_token 永远记 0。 + // 与 ModelItem.ModelId 的 json:"modelId,string" 约定一致,字符串精确回环(雪花id>2^53 无损)。 + // effModelID 为解析后的系统模型 id(引用行),per_token 结算按此查价。 + "model_id": gconv.String(effModelID), + "prompt_tokens": totalPrompt, + "completion_tokens": totalCompletion, + "media_type": effMediaType, + "total_tokens": totalTokens, + "total_fee": totalCost, + "total_duration": totalDuration, + }}, + }); err != nil { + return nil, fmt.Errorf("节点:%v 写入token信息失败: %v", nodeInput.Config.Name, err) + } + } + + // 4.5 视频模型节点返回多个视频时,自动调用视频合成工具(concat_videos)合并为单条; + // 已显式配置 concat_videos 后置工具时跳过,避免重复合并 + if nodeInput.Config.PostTool != media.ProcessorName && len(outputRes) > 1 && isVideoModel(ctx, nodeInput.Config.ModelConfig.ModelId) { + outputRes, err = invokePostTool(ctx, media.ProcessorName, outputRes, map[string]any{"callback_url": "callback_url", "upload": true}) + if err != nil { + return nil, err + } + } else { + // 4. 后置工具:加工模型输出(透传原始请求参数,供后置工具读取合并配置等) + outputRes, err = invokePostTool(ctx, nodeInput.Config.PostTool, outputRes, modelParams) + if err != nil { + return nil, err + } + } + nodeInput.Config.OutputResult = outputRes + return nodeInput, nil +} + +// mergeInferenceBatchResults 推理模型分批结果拼接为单条输出记录: +// 各批结果按批序对同名 key 的值做字符串拼接("拼到一个字段"),最终返回单条 {key:值} 记录。 +// 非字符串值(如结构/数组字段)取最后一份,避免误拼接。 +func mergeInferenceBatchResults(results [][]map[string]any) []map[string]any { + merged := make(map[string]any) + for _, res := range results { + for _, record := range res { + for key, val := range record { + prev, has := merged[key] + if !has { + merged[key] = val + continue + } + sPrev, pOK := prev.(string) + sVal, vOK := val.(string) + if pOK && vOK { + merged[key] = sPrev + "\n" + sVal + continue + } + merged[key] = val + } + } + } + return []map[string]any{merged} +} diff --git a/workflow/service/flow/lambda_http.go b/workflow/service/flow/lambda_http.go new file mode 100644 index 0000000..1f7a0dd --- /dev/null +++ b/workflow/service/flow/lambda_http.go @@ -0,0 +1,21 @@ +package flow + +import ( + flowDto "ai-agent/workflow/model/dto/flow" + "context" + "fmt" +) + +// HttpLambda 构建HTTP(S)接口 +func HttpLambda(ctx context.Context, input any) (any, error) { + nodeInput, ok := input.(*flowDto.NodeExecutionInput) + if !ok { + return nil, fmt.Errorf("入参类型错误") + } + outputRes, err := HttpCallResultLambda(ctx, nodeInput) + if err != nil { + return nil, err + } + nodeInput.Config.OutputResult = outputRes + return nodeInput, nil +} diff --git a/workflow/service/flow/lambda_node.go b/workflow/service/flow/lambda_node.go deleted file mode 100644 index 583e0c9..0000000 --- a/workflow/service/flow/lambda_node.go +++ /dev/null @@ -1,742 +0,0 @@ -package flow - -import ( - "ai-agent/gateway" - "ai-agent/workflow/consts/flow" - "ai-agent/workflow/consts/model" - "ai-agent/workflow/consts/node" - "ai-agent/workflow/consts/public" - flowDao "ai-agent/workflow/dao/flow" - nodeDao "ai-agent/workflow/dao/node" - sessionDao "ai-agent/workflow/dao/session" - flowDto "ai-agent/workflow/model/dto/flow" - nodeDto "ai-agent/workflow/model/dto/node" - sessionDto "ai-agent/workflow/model/dto/session" - "ai-agent/workflow/model/entity" - "ai-agent/workflow/service/flow/processor" - "ai-agent/workflow/service/flow/processor/builtin/media" - "context" - "encoding/base64" - "encoding/json" - "fmt" - "strings" - "sync" - - "gitea.redpowerfuture.com/red-future/common/db/gfdb" - "gitea.redpowerfuture.com/red-future/common/oss" - "github.com/cloudwego/eino-examples/compose/batch/batch" - "github.com/cloudwego/eino/compose" - "github.com/gogf/gf/v2/database/gdb" - "github.com/gogf/gf/v2/frame/g" - "github.com/gogf/gf/v2/util/gconv" - "github.com/google/uuid" -) - -// StartLambda 启动节点 -func StartLambda(ctx context.Context, input any) (any, error) { - return input, nil -} - -// FormLambda 表单调用节点 -func FormLambda(ctx context.Context, input any) (any, error) { - nodeInput, ok := input.(*flowDto.NodeExecutionInput) - if !ok { - return nil, fmt.Errorf("入参类型错误") - } - // 解析 valueSource 引用,填充表单节点输出配置(供下游引用) - for _, output := range nodeInput.Config.OutputConfig { - ProcessValueSourceRecursive(output, nodeInput.Global) - } - return nodeInput, nil -} - -// ModelLambda 模型调用节点 -func ModelLambda(ctx context.Context, input any) (any, error) { - nodeInput, ok := input.(*flowDto.NodeExecutionInput) - if !ok { - return nil, fmt.Errorf("入参类型错误") - } - - modelParams, err := BuildModelRequestBody(nodeInput.Config.ModelConfig.ModelRequestParamsPath, nodeInput.Global) - if err != nil { - return nil, err - } - - // 2. 前置工具:决定模型调用入参(单次/多次) - // 入参统一为扁平模型请求体(BuildModelRequestBody 输出,key 为点分路径)。 - // 分批处理器按默认上限拆分集合字段,其余前置工具(如 split_shots_pipeline)读取扁平参数。 - preToolParams := modelParams - paramsList, err := invokePreTool(ctx, nodeInput.Config.PreTool, preToolParams) - if err != nil { - return nil, err - } - - // 3. 逐批调用模型,汇总输出(保持请求顺序),累计 token/费用供节点记录落库 - var outputRes []map[string]any - var totalTokens int64 - var totalCost float64 - if len(paramsList) > 1 { - // 段级续跑仅在"多段 + 视频模型"启用;非视频分段(批量文本等)走原逻辑零影响。 - // 段身份 = 列表位置(0-based):paramsList 顺序即段序,concat 按列表顺序拼接; - // 位置互不重复且跨 reExecute 稳定(参数一致 → 段数/顺序不变)。不依赖 params 里的 - // segment_index——真实链路(上游 split_shots_pipeline 转写 → 下游 split_segment 按 - // __segment_fields 拆分,invokePreTool 剥离 __ 内部键)下 paramsList 只有模型参数。 - // segVideo=false 时走既有非段级合并路径(全量生成、不落库、不复用),全新执行行为不变, - // 后续自动 concat 判断(独立的 isVideoModel 调用)仍正常执行。 - segVideo := isVideoModel(ctx, nodeInput.Config.ModelConfig.ModelId) - - // 续跑(!ForceNewRun)时读取该节点已成功段;全新执行不查(BuildExecution 已清旧段),saved 为 nil → 全量重生成 - var saved map[int]entity.SegmentRef - if !nodeInput.Global.ForceNewRun && segVideo { - saved, err = flowDao.FlowSegmentResultDao.ListByNode(ctx, nodeInput.Global.ExecutionId, nodeInput.Config.Id) - if err != nil { - return nil, err - } - } - - idxList, needGen := planSegmentResume(paramsList, saved) - - results := make([][]map[string]any, len(paramsList)) - tokenRes := make([]*gateway.ModelCallRes, len(paramsList)) - errs := make([]error, len(paramsList)) - saveErrs := make([]error, len(paramsList)) - isInference := make([]bool, len(paramsList)) - var wg sync.WaitGroup - for i, params := range paramsList { - if !needGen[i] { - continue - } - wg.Add(1) - go func(i int, params map[string]any) { - defer wg.Done() - // 视频段每段失败自动重试 1 次(共 2 次尝试);非视频保持单次调用 - for attempt := 0; attempt < segmentGenerateMaxAttempts; attempt++ { - results[i], tokenRes[i], isInference[i], errs[i] = ModelCallResultLambda(ctx, nodeInput.Config.ModelConfig.ModelId, nodeInput.Global.SessionId, params, nodeInput.Config.Prompt, nodeInput.Global.ExecutionId, nodeInput.Config.Id, idxList[i]) - if errs[i] == nil || !segVideo { - break - } - } - // 每段成功立即落库:该段刚成功即持久化,其他段仍在跑/重试时已成功段也不丢; - // 后续段失败或进程崩溃(panic/OOM/kill)时,已完成段已在库中,reExecute 可直接复用 - if segVideo && errs[i] == nil { - for _, rec := range results[i] { - key := media.FindVideoKey(rec) - url := media.FindVideoURL(ctx, rec) - if key == "" || url == "" { - continue - } - if err := flowDao.FlowSegmentResultDao.Save(ctx, nodeInput.Global.ExecutionId, nodeInput.Config.Id, idxList[i], key, url); err != nil { - saveErrs[i] = err - } - } - } - }(i, params) - } - wg.Wait() - - // 仍有失败段或落库失败 → 节点失败(成功段已立即落库,供下次 reExecute 复用) - for i := range results { - if saveErrs[i] != nil { - return nil, saveErrs[i] - } - if needGen[i] && errs[i] != nil { - return nil, errs[i] - } - if needGen[i] && tokenRes[i] != nil { - totalTokens += tokenRes[i].TotalTokens - totalCost += tokenRes[i].Cost - } - } - - if segVideo { - // 复用段 + 新生段按段序号升序合并,concat 按列表顺序拼接 → 顺序保证 - outputRes = mergeSegmentOutputs(idxList, needGen, results, saved) - } else { - if isInference[0] { - outputRes = mergeInferenceBatchResults(results) - } else { - for _, res := range results { - outputRes = append(outputRes, res...) - } - } - } - } else { - for _, params := range paramsList { - res, modelRes, _, err := ModelCallResultLambda(ctx, nodeInput.Config.ModelConfig.ModelId, nodeInput.Global.SessionId, params, nodeInput.Config.Prompt, nodeInput.Global.ExecutionId, nodeInput.Config.Id, flowDao.FlowAsyncSegSentinel) - if err != nil { - return nil, err - } - if modelRes != nil { - totalTokens += modelRes.TotalTokens - totalCost += modelRes.Cost - } - outputRes = append(outputRes, res...) - } - } - - // 3.5 把本次节点消耗的 token/费用写入节点执行记录,供汇总节点聚合到 exec_workflow - if nodeInput.NodeExecutionId > 0 && (totalTokens > 0 || totalCost > 0) { - if _, err = nodeDao.NodeExecutionDao.Update(ctx, &nodeDto.UpdateNodeExecutionReq{ - Id: nodeInput.NodeExecutionId, - TokenInfo: []map[string]any{{ - "total_tokens": totalTokens, - "total_fee": totalCost, - }}, - }); err != nil { - return nil, fmt.Errorf("节点:%v 写入token信息失败: %v", nodeInput.Config.Name, err) - } - } - - // 4.5 视频模型节点返回多个视频时,自动调用视频合成工具(concat_videos)合并为单条; - // 已显式配置 concat_videos 后置工具时跳过,避免重复合并 - if nodeInput.Config.PostTool != media.ProcessorName && len(outputRes) > 1 && isVideoModel(ctx, nodeInput.Config.ModelConfig.ModelId) { - g.Log().Debugf(ctx, "modelId1:%v ,outputRes: %v", nodeInput.Config.ModelConfig.ModelId, outputRes) - outputRes, err = invokePostTool(ctx, media.ProcessorName, outputRes, map[string]any{"callback_url": "callback_url", "upload": true}) - g.Log().Debugf(ctx, "modelId2:%v ,outputRes: %v", nodeInput.Config.ModelConfig.ModelId, outputRes) - if err != nil { - return nil, err - } - } else { - // 4. 后置工具:加工模型输出(透传原始请求参数,供后置工具读取合并配置等) - outputRes, err = invokePostTool(ctx, nodeInput.Config.PostTool, outputRes, modelParams) - if err != nil { - return nil, err - } - } - g.Log().Debugf(ctx, "modelId3:%v ,outputRes: %v", nodeInput.Config.ModelConfig.ModelId, outputRes) - nodeInput.Config.OutputResult = outputRes - return nodeInput, nil -} - -// isVideoModel 判断模型是否为视频模型(模型类型 TypeVideo=600),用于视频节点多视频自动合成判断 -func isVideoModel(ctx context.Context, modelId int64) bool { - modelInfo, err := gateway.GetModelInfoById(ctx, &gateway.GetModelInfoByIdReq{ModelId: modelId}) - if err != nil { - g.Log().Warningf(ctx, "查询模型配置失败,跳过自动视频合成 modelId=%d err=%v", modelId, err) - return false - } - return modelInfo.ModelManage.ModelType != nil && *modelInfo.ModelManage.ModelType == model.TypeVideo -} - -// mergeInferenceBatchResults 推理模型分批结果拼接为单条输出记录: -// 各批结果按批序对同名 key 的值做字符串拼接("拼到一个字段"),最终返回单条 {key:值} 记录。 -// 非字符串值(如结构/数组字段)取最后一份,避免误拼接。 -func mergeInferenceBatchResults(results [][]map[string]any) []map[string]any { - merged := make(map[string]any) - for _, res := range results { - for _, record := range res { - for key, val := range record { - prev, has := merged[key] - if !has { - merged[key] = val - continue - } - sPrev, pOK := prev.(string) - sVal, vOK := val.(string) - if pOK && vOK { - merged[key] = sPrev + "\n" + sVal - continue - } - merged[key] = val - } - } - } - return []map[string]any{merged} -} - -// invokePreTool 执行前置处理器,把模型请求参数转换为模型调用入参列表。 -// 前置处理器契约:入参即模型请求参数本体;返回值: -// - map[string]any 一次模型调用,入参为返回值 -// - []map[string]any 多次模型调用,逐个入参请求 -// - nil 视为异常,节点失败(不允许静默跳过模型调用) -func invokePreTool(ctx context.Context, processorName string, modelParams map[string]any) (paramsList []map[string]any, err error) { - if processorName == "" { - return []map[string]any{stripInternalKeys(modelParams)}, nil - } - data, err := processor.Call(ctx, processorName, modelParams) - if err != nil { - return nil, fmt.Errorf("执行前置处理器[%s]失败: %v", processorName, err) - } - switch v := data.(type) { - case nil: - return nil, fmt.Errorf("前置处理器[%s]返回空", processorName) - case map[string]any: - return []map[string]any{stripInternalKeys(v)}, nil - case []map[string]any: - list := make([]map[string]any, 0, len(v)) - for _, m := range v { - list = append(list, stripInternalKeys(m)) - } - return list, nil - default: - return nil, fmt.Errorf("前置处理器[%s]返回类型不支持: %T", processorName, data) - } -} - -// stripInternalKeys 剥离 __ 前缀的内部键(如 __segment_fields/__produced), -// 模型网关做参数严格校验(CheckParams strictUnknown)会拒绝未知字段,内部标记不得随请求体下发。 -func stripInternalKeys(params map[string]any) map[string]any { - if params == nil { - return params - } - for k := range params { - if strings.HasPrefix(k, "__") { - delete(params, k) - } - } - return params -} - -// invokePostTool 执行后置处理器,加工模型调用结果。 -// 后置处理器契约:入参 {"output": 模型输出结果列表, "request": 原始模型请求参数}(列表须包成对象传入);返回值: -// - []map[string]any 替换模型输出 -// - map[string]any 替换为单条输出 -// - nil 保留原输出 -func invokePostTool(ctx context.Context, processorName string, outputRes []map[string]any, requestParams map[string]any) ([]map[string]any, error) { - if processorName == "" { - return outputRes, nil - } - data, err := processor.Call(ctx, processorName, map[string]any{"output": outputRes, "request": requestParams}) - if err != nil { - return nil, fmt.Errorf("执行后置处理器[%s]失败: %v", processorName, err) - } - switch v := data.(type) { - case nil: - return outputRes, nil - case []map[string]any: - return v, nil - case map[string]any: - return []map[string]any{v}, nil - default: - return nil, fmt.Errorf("后置处理器[%s]返回类型不支持: %T", processorName, data) - } -} - -func SubFlowLambda(ctx context.Context, input any) (any, error) { - // 1. 类型断言(和其他节点保持一致的入参结构) - nodeExecInput, ok := input.(*flowDto.NodeExecutionInput) - if !ok { - return nil, fmt.Errorf("子流程节点入参类型错误,期望*flowDto.NodeExecutionInput,实际%T", input) - } - // 2. 解析子流程配置 - subFlowConfig := nodeExecInput.Config.SubConfig - if subFlowConfig == nil { - return nil, fmt.Errorf("子流程节点缺少配置") - } - getRes, err := FlowUserService.Get(ctx, &flowDto.GetFlowUserReq{ - Id: subFlowConfig.WorkflowId, - }) - if err != nil { - return nil, err - } - // 3. 引入参数解析:把首页表单值/上游引用值/静态默认值写入子流程开始节点 outputConfig。 - // 须在 BuildGraph / ExtractFlowNodeFrom 之前执行,batchInputs 深拷贝的才是注入后的开始节点。 - injectSubFlowFields(nodeExecInput.Global, getRes.FlowContent, subFlowConfig.Fields) - // 4. 并发数:从主流程开始节点 outputConfig 的 maxConcurrency 字段读取(前端把子流程节点生成次数表单字段聚合到主流程开始节点),读不到再用子流程节点配置兜底 - maxConcurrency := mainFlowMaxConcurrency(nodeExecInput.Global, subFlowConfig.MaxConcurrency) - // 4. 编译子流程Graph(复用现有 BuildGraphFromFlowContent 逻辑) - nodeList, subGraph := BuildGraph(ctx, getRes.FlowContent) - // 4. 构建子流程Workflow(绑定START/END,和示例对齐) - innerWorkflow := compose.NewWorkflow[*flowDto.FlowExecutionInput, *flowDto.FlowExecutionInput]() - // 挂载子图节点并绑定全局START - innerWorkflow.AddGraphNode("sub_flow_graph", subGraph).AddInput(compose.START) - // 绑定子图输出到全局END - innerWorkflow.End().AddInput("sub_flow_graph") - // 生成次数(批量条数):maxConcurrency<=0 时按 1 次兜底 - batchCount := maxConcurrency - if batchCount <= 0 { - batchCount = 1 - } - // 5. 构建BatchNode(批量执行子流程,复用示例逻辑) - batchNode := batch.NewBatchNode(&batch.NodeConfig[*flowDto.FlowExecutionInput, *flowDto.FlowExecutionInput]{ - Name: fmt.Sprintf("sub_flow_batch_%s", nodeExecInput.Config.Id), - InnerTask: innerWorkflow, - MaxConcurrency: batchCount, - }) - - // 6. 提取批量输入:按生成次数生成 N 份(每份独立克隆 ConfigMap,避免并发执行时节点输出写串) - nodeInputParams := ExtractFlowNodeFrom(getRes.FlowContent) - configMap := make(map[string]*entity.FlowNode) - for _, cfg := range nodeInputParams { - configMap[cfg.Id] = cfg - } - for _, i := range nodeList { - configMap[i.Id] = &i - } - batchInputs := make([]*flowDto.FlowExecutionInput, 0, batchCount) - for j := 0; j < batchCount; j++ { - batchInputs = append(batchInputs, &flowDto.FlowExecutionInput{ - NodeGroupId: nodeExecInput.Global.NodeGroupId, - ExecutionId: nodeExecInput.Global.ExecutionId, - FlowId: nodeExecInput.Global.FlowId, - ConfigMap: cloneConfigMap(configMap), - SessionId: nodeExecInput.Global.SessionId, - }) - } - // 7. 执行批量子流程 - batchOutput, err := batchNode.Invoke(ctx, batchInputs) - if err != nil { - return nil, fmt.Errorf("执行子流程BatchNode失败: %v", err) - } - // 8. 展平每份子流程执行的节点输出,写回当前节点 OutputResult 供下游引用 - var outputRes []map[string]any - for _, single := range batchOutput { - if single == nil { - continue - } - outputRes = append(outputRes, collectFlowNodeResults(single)...) - } - g.Log().Info(ctx, fmt.Sprintf("子流程执行完成,共 %d 次,输出 %d 条", batchCount, len(outputRes))) - nodeExecInput.Config.OutputResult = outputRes - return nodeExecInput, nil -} - -// injectSubFlowFields 将子流程节点引入参数(subConfig.Fields)解析后写入子流程开始节点 -// outputConfig,使子流程启动时能读到首页表单值/上游引用值/静态默认值。 -// 每个字段的取值优先级:valueSource 引用解析成功 → field.value → field.defaultValue; -// 匹配键为 field(前端约定以 field 为主,不兼容 path)。 -func injectSubFlowFields(global *flowDto.FlowExecutionInput, subFlowContent *entity.FlowInfo, fields []map[string]any) { - if global == nil || subFlowContent == nil || len(fields) == 0 { - return - } - startNode := subFlowStartNode(subFlowContent) - if startNode == nil { - return - } - byField := make(map[string]map[string]any, len(startNode.OutputConfig)) - for _, output := range startNode.OutputConfig { - byField[gconv.String(output["field"])] = output - } - for _, field := range fields { - entry := byField[gconv.String(field["field"])] - if entry == nil { - continue - } - value := field["value"] - if vs, has := field["valueSource"]; has && vs != nil { - if vsNodeId, vsField := firstValueSource(vs); vsNodeId != "" && vsField != "" { - if v, _, ok := resolveValueSource(global, vsNodeId, vsField); ok { - value = v - } - } - } - if value == nil { - value = field["defaultValue"] - } - if value != nil { - entry["value"] = value - } - } -} - -// firstValueSource 从 valueSource 提取第一个引用源 (nodeId, field)。 -// 前端统一发送数组 [{nodeId, field}](见 serializeSubFlowConfig),旧 DSL 可能是单对象 -// {nodeId, fieldName},两种形态都兼容;子流程字段与引用源一一对应,只取第一个。 -func firstValueSource(vs any) (nodeId, field string) { - if vs == nil { - return - } - switch v := vs.(type) { - case []any: - if len(v) > 0 { - return firstValueSource(v[0]) - } - return - case []map[string]any: - if len(v) > 0 { - return firstValueSource(v[0]) - } - return - } - m := gconv.Map(vs) - nodeId = gconv.String(m["nodeId"]) - field = gconv.String(m["fieldName"]) - if field == "" { - field = gconv.String(m["field"]) - } - return -} - -// subFlowStartNode 返回工作流开始节点 -func subFlowStartNode(content *entity.FlowInfo) *entity.FlowNode { - if content == nil { - return nil - } - for i := range content.Nodes { - if content.Nodes[i].Id == content.StartNodeId { - return &content.Nodes[i] - } - } - return nil -} - -// mainFlowMaxConcurrency 取子流程批量执行并发数:从主流程开始节点 outputConfig -// 的 maxConcurrency 字段读取(前端把子流程节点的生成次数表单字段聚合到主流程开始节点), -// 读不到再用子流程节点配置的兜底值。 -func mainFlowMaxConcurrency(global *flowDto.FlowExecutionInput, fallback int) int { - if global == nil { - return fallback - } - for _, n := range global.ConfigMap { - if n == nil || n.NodeCode != node.NodeTypeStart { - continue - } - for _, output := range n.OutputConfig { - if gconv.String(output["field"]) != "maxConcurrency" && gconv.String(output["path"]) != "maxConcurrency" { - continue - } - if v := gconv.Int(output["value"]); v > 0 { - return v - } - } - return fallback - } - return fallback -} - -// cloneConfigMap 深拷贝 ConfigMap,保证各批次子流程并发执行时节点输出互不串扰。 -// 浅拷贝会共享 *entity.FlowNode,并发写 OutputResult 产生竞态。 -func cloneConfigMap(src map[string]*entity.FlowNode) map[string]*entity.FlowNode { - dst := make(map[string]*entity.FlowNode, len(src)) - for k, v := range src { - data, err := json.Marshal(v) - if err != nil { - dst[k] = v - continue - } - n := new(entity.FlowNode) - if err = json.Unmarshal(data, n); err != nil { - dst[k] = v - continue - } - dst[k] = n - } - return dst -} - -// collectFlowNodeResults 收集一次子流程执行中所有已执行节点的输出,展平成 {字段:值} 列表 -func collectFlowNodeResults(execInput *flowDto.FlowExecutionInput) []map[string]any { - var res []map[string]any - for _, executed := range execInput.ExecutedNodes { - if nodeConfig := execInput.ConfigMap[executed.NodeId]; nodeConfig != nil { - res = append(res, nodeConfig.OutputResult...) - } - } - return res -} - -// HttpLambda 构建HTTP(S)接口 -func HttpLambda(ctx context.Context, input any) (any, error) { - nodeInput, ok := input.(*flowDto.NodeExecutionInput) - if !ok { - return nil, fmt.Errorf("入参类型错误") - } - outputRes, err := HttpCallResultLambda(ctx, nodeInput) - if err != nil { - return nil, err - } - nodeInput.Config.OutputResult = outputRes - return nodeInput, nil -} - -func DataMergeLambda(ctx context.Context, input any) (res any, err error) { - nodeInput, ok := input.(*flowDto.NodeExecutionInput) - if !ok { - return nil, fmt.Errorf("参数合并入参类型错误") - } - return nodeInput, nil -} - -func SummaryLambda(ctx context.Context, input any) (any, error) { - execInput, ok := input.(*flowDto.NodeExecutionInput) - if !ok { - return nil, fmt.Errorf("汇总节点入参类型错误,实际是 %T", input) - } - - // 聚合所有已执行节点中需入库的文件结果(两层规则) - summaryResult := collectSaveFileResults(ctx, execInput.Global) - - // 把汇总结果存入当前节点的输出 - g.Log().Info(ctx, fmt.Sprintf("结果汇总完成,汇总数据:%+v", summaryResult)) - - err := gfdb.DB(ctx, public.DbNameBlackDeacon).Transaction(ctx, func(ctx context.Context, tx gdb.TX) error { - res, _, err := nodeDao.NodeExecutionDao.ListByFlowExecutionId(ctx, &nodeDto.ListNodeExecutionByFlowReq{ - NodeGroupId: execInput.Global.NodeGroupId, - }, entity.NodeExecutionCol.TokenInfo) - if err != nil { - return err - } - var totalTokens int - var totalFee float64 - for _, item := range res { - for _, itemToken := range item.TokenInfo { - m := gconv.Map(itemToken) - totalTokens += gconv.Int(m["total_tokens"]) - totalFee += gconv.Float64(m["total_fee"]) - } - } - _, err = sessionDao.ExecWorkflowDao.Update(ctx, &sessionDto.UpdateWorkflowReq{ - Id: execInput.Global.ExecutionId, - Status: flow.FlowExecutionStatusSuccess.Code(), - TotalTokens: totalTokens, - TotalFee: totalFee, - }) - if err != nil { - return err - } - if len(summaryResult) > 0 { - _, err = sessionDao.ExecWorkflowResultDao.BatchInsert(ctx, summaryResult) - if err != nil { - return err - } - } - return nil - }) - - return execInput, err -} - -// collectSaveFileResults 按两层规则收集需入库的文件结果: -// 第一层:节点须开启"保存文件"(IsSaveFile); -// 第二层:key 取自节点 OutputResult 的各字段,命中 ModelResponseBodyMapping 才入库; -// 原始响应体 key(respBody)恒入库(不要求映射声明);HTTP 节点产出以 http_file_url:{key} -// 标记的字段(IsSaveFile 时由 HttpCallResultLambda 生成)恒入库(无模型响应映射可查)。 -// 结果值为 http(s) URL 或 MinIO 对象裸路径直接使用;非路径值(base64 图片/文本)先上传 OSS 换取 URL, -// 文本内容以 .inc 扩展名存储。 -func collectSaveFileResults(ctx context.Context, execInput *flowDto.FlowExecutionInput) []*sessionDto.CreateWorkflowResultReq { - if execInput == nil { - return nil - } - var summaryResult []*sessionDto.CreateWorkflowResultReq - for _, executedNode := range execInput.ExecutedNodes { - nodeConfig := execInput.ConfigMap[executedNode.NodeId] - if nodeConfig == nil || len(nodeConfig.OutputResult) == 0 || !nodeConfig.IsSaveFile { - continue - } - // 第二层:key 取自节点 OutputResult 的各字段, - // 命中 ModelResponseBodyMapping 才入库;respBody 与 HTTP 节点 http_file_url:{key} 标记恒入库 - saveKeys := nodeConfig.ModelConfig.ModelResponseBodyMapping - for _, respBody := range nodeConfig.OutputResult { - for key, val := range gconv.Map(respBody) { - isHTTPFile := strings.HasPrefix(key, "http_file_url:") - if !isHTTPFile { - if _, ok := saveKeys[key]; !ok && key != "respBody" { - continue - } - } - fileUrl, err := resolveSaveFileResult(ctx, val) - if err != nil { - g.Log().Warningf(ctx, "collectSaveFileResults 上传结果文件失败 key=%s err=%v", key, err) - continue - } - summaryResult = append(summaryResult, &sessionDto.CreateWorkflowResultReq{ - SessionId: execInput.SessionId, - FlowId: execInput.FlowId, - ExecId: execInput.ExecutionId, - ResultFileUrl: fileUrl, - }) - } - } - } - return summaryResult -} - -// resolveSaveFileResult 解析结果值为可入库的 URL: -// - 已是 http(s) URL 或 MinIO 对象裸路径 → 直接返回 -// - 非路径(base64 图片/文本)→ 上传 OSS 换取 URL -func resolveSaveFileResult(ctx context.Context, val any) (string, error) { - isPath, path, fileBytes, ext := resolveFileContent(val) - if isPath { - return path, nil - } - if ext == "" { - ext = ".png" - } - fileUrl, err := gateway.Upload(ctx, fmt.Sprintf("workflow_result_%s%s", uuid.NewString(), ext), fileBytes) - if err != nil { - return "", err - } - return fileUrl, nil -} - -// resolveFileContent 判断结果值形态: -// - 已是 URL 路径(http/https 开头)→ 直接使用 -// - data URI(data:;base64,)→ 解码为字节,扩展名按 mime 推断 -// - 纯 base64(可解码且长度足以认为是编码数据)→ 解码为字节,默认 .png -// - 其余(文本)→ 以 .inc 扩展名上传原文 -func resolveFileContent(val any) (isPath bool, path string, fileBytes []byte, ext string) { - s := gconv.String(val) - if isFileURL(s) { - return true, s, nil, "" - } - // MinIO 对象裸路径(无 http 前缀,模型网关转存 OSS 后返回) - if oss.IsOSSPath(s) { - return true, s, nil, "" - } - // data URI:data:;base64, - if b, mime, ok := parseDataURI(s); ok { - return false, "", b, extOfMime(mime) - } - // 纯 base64:可解码且长度足够,视为编码后的文件内容 - trimmed := strings.TrimSpace(s) - if len(trimmed) >= 64 { - if b, err := base64.StdEncoding.DecodeString(trimmed); err == nil && len(b) > 0 { - return false, "", b, ".png" - } - } - // 文本:以 .inc 存储 - return false, "", []byte(s), ".inc" -} - -// isFileURL 判断字符串是否已是对外可访问的 URL 路径(http/https 开头) -func isFileURL(s string) bool { - lower := strings.ToLower(s) - return strings.HasPrefix(lower, "http://") || strings.HasPrefix(lower, "https://") -} - -// extOfMime 按 MIME 类型推断文件扩展名 -func extOfMime(mime string) string { - switch strings.ToLower(strings.TrimSpace(mime)) { - case "image/png", "png": - return ".png" - case "image/jpeg", "image/jpg", "jpeg", "jpg": - return ".jpg" - case "image/webp": - return ".webp" - case "image/gif": - return ".gif" - case "audio/mpeg", "audio/mp3", "mp3": - return ".mp3" - case "audio/wav", "wav": - return ".wav" - case "video/mp4", "mp4": - return ".mp4" - case "application/json", "json": - return ".json" - default: - return "" - } -} - -// parseDataURI 解析 data URI:data:;base64,,返回解码字节与 mime -func parseDataURI(s string) ([]byte, string, bool) { - const prefix = "data:" - if !strings.HasPrefix(s, prefix) { - return nil, "", false - } - rest := s[len(prefix):] - comma := strings.Index(rest, ",") - if comma < 0 { - return nil, "", false - } - mime := rest[:comma] - if semicolon := strings.Index(mime, ";"); semicolon >= 0 { - mime = mime[:semicolon] - } - payload := strings.TrimPrefix(rest[comma+1:], "base64,") - b, err := base64.StdEncoding.DecodeString(payload) - if err != nil { - return nil, "", false - } - return b, mime, true -} diff --git a/workflow/service/flow/lambda_savefile.go b/workflow/service/flow/lambda_savefile.go new file mode 100644 index 0000000..a97a2bb --- /dev/null +++ b/workflow/service/flow/lambda_savefile.go @@ -0,0 +1,113 @@ +package flow + +import ( + "ai-agent/gateway" + "context" + "encoding/base64" + "fmt" + "strings" + + "gitea.redpowerfuture.com/red-future/common/oss" + "github.com/gogf/gf/v2/util/gconv" + "github.com/google/uuid" +) + +// resolveSaveFileResult 解析结果值为可入库的 URL: +// - 已是 http(s) URL 或 MinIO 对象裸路径 → 直接返回 +// - 非路径(base64 图片/文本)→ 上传 OSS 换取 URL +func resolveSaveFileResult(ctx context.Context, val any) (string, error) { + isPath, path, fileBytes, ext := resolveFileContent(val) + if isPath { + return path, nil + } + if ext == "" { + ext = ".png" + } + fileUrl, err := gateway.Upload(ctx, fmt.Sprintf("workflow_result_%s%s", uuid.NewString(), ext), fileBytes) + if err != nil { + return "", err + } + return fileUrl, nil +} + +// resolveFileContent 判断结果值形态: +// - 已是 URL 路径(http/https 开头)→ 直接使用 +// - data URI(data:;base64,)→ 解码为字节,扩展名按 mime 推断 +// - 纯 base64(可解码且长度足以认为是编码数据)→ 解码为字节,默认 .png +// - 其余(文本)→ 以 .inc 扩展名上传原文 +func resolveFileContent(val any) (isPath bool, path string, fileBytes []byte, ext string) { + s := gconv.String(val) + if isFileURL(s) { + return true, s, nil, "" + } + // MinIO 对象裸路径(无 http 前缀,模型网关转存 OSS 后返回) + if oss.IsOSSPath(s) { + return true, s, nil, "" + } + // data URI:data:;base64, + if b, mime, ok := parseDataURI(s); ok { + return false, "", b, extOfMime(mime) + } + // 纯 base64:可解码且长度足够,视为编码后的文件内容 + trimmed := strings.TrimSpace(s) + if len(trimmed) >= 64 { + if b, err := base64.StdEncoding.DecodeString(trimmed); err == nil && len(b) > 0 { + return false, "", b, ".png" + } + } + // 文本:以 .inc 存储 + return false, "", []byte(s), ".inc" +} + +// isFileURL 判断字符串是否已是对外可访问的 URL 路径(http/https 开头) +func isFileURL(s string) bool { + lower := strings.ToLower(s) + return strings.HasPrefix(lower, "http://") || strings.HasPrefix(lower, "https://") +} + +// extOfMime 按 MIME 类型推断文件扩展名 +func extOfMime(mime string) string { + switch strings.ToLower(strings.TrimSpace(mime)) { + case "image/png", "png": + return ".png" + case "image/jpeg", "image/jpg", "jpeg", "jpg": + return ".jpg" + case "image/webp": + return ".webp" + case "image/gif": + return ".gif" + case "audio/mpeg", "audio/mp3", "mp3": + return ".mp3" + case "audio/wav", "wav": + return ".wav" + case "video/mp4", "mp4": + return ".mp4" + case "application/json", "json": + return ".json" + default: + return "" + } +} + +// parseDataURI 解析 data URI:data:;base64,,返回解码字节与 mime +func parseDataURI(s string) ([]byte, string, bool) { + const prefix = "data:" + if !strings.HasPrefix(s, prefix) { + return nil, "", false + } + rest := s[len(prefix):] + comma := strings.Index(rest, ",") + if comma < 0 { + return nil, "", false + } + mime := rest[:comma] + if semicolon := strings.Index(mime, ";"); semicolon >= 0 { + mime = mime[:semicolon] + } + payload := strings.TrimPrefix(rest[comma+1:], "base64,") + b, err := base64.StdEncoding.DecodeString(payload) + if err != nil { + return nil, "", false + } + return b, mime, true +} diff --git a/workflow/service/flow/lambda_script_transcribe.go b/workflow/service/flow/lambda_script_transcribe.go index 52dcbc5..ea433b9 100644 --- a/workflow/service/flow/lambda_script_transcribe.go +++ b/workflow/service/flow/lambda_script_transcribe.go @@ -4,6 +4,7 @@ import ( "ai-agent/workflow/consts/node" "ai-agent/workflow/service/flow/processor" "ai-agent/workflow/service/flow/processor/builtin/split_shots_pipeline" + "ai-agent/workflow/service/flow/values" "context" "encoding/json" "fmt" @@ -81,7 +82,7 @@ func ScriptTranscribeLambda(ctx context.Context, input any) (any, error) { } } - modelParams, err := BuildModelRequestBody(nodeInput.Config.ModelConfig.ModelRequestParamsPath, nodeInput.Global) + modelParams, err := values.BuildModelRequestBody(nodeInput.Config.ModelConfig.ModelRequestParamsPath, nodeInput.Global) if err != nil { return nil, err } diff --git a/workflow/service/flow/lambda_segment_resume.go b/workflow/service/flow/lambda_segment_resume.go index 92a6fff..70f6616 100644 --- a/workflow/service/flow/lambda_segment_resume.go +++ b/workflow/service/flow/lambda_segment_resume.go @@ -6,9 +6,6 @@ import ( "ai-agent/workflow/model/entity" ) -// segmentGenerateMaxAttempts 视频段生成最大尝试次数(失败自动重试 1 次,共 2 次尝试),参数化可调 -const segmentGenerateMaxAttempts = 2 - // planSegmentResume 段级续跑决策:段身份取列表位置(0-based,paramsList 顺序即段序), // 把各段映射到"是否需重新生成"。savedMap 为该节点已成功段(段序号 → {key,url}); // 段在表中缺失或地址为空则需重新生成。返回值与 paramsList 对齐。 diff --git a/workflow/service/flow/lambda_subflow.go b/workflow/service/flow/lambda_subflow.go new file mode 100644 index 0000000..debd14c --- /dev/null +++ b/workflow/service/flow/lambda_subflow.go @@ -0,0 +1,219 @@ +package flow + +import ( + "ai-agent/workflow/consts/node" + flowDto "ai-agent/workflow/model/dto/flow" + "ai-agent/workflow/model/entity" + "ai-agent/workflow/service/flow/values" + "context" + "encoding/json" + "fmt" + + "github.com/cloudwego/eino-examples/compose/batch/batch" + "github.com/cloudwego/eino/compose" + "github.com/gogf/gf/v2/frame/g" + "github.com/gogf/gf/v2/util/gconv" +) + +func SubFlowLambda(ctx context.Context, input any) (any, error) { + // 1. 类型断言(和其他节点保持一致的入参结构) + nodeExecInput, ok := input.(*flowDto.NodeExecutionInput) + if !ok { + return nil, fmt.Errorf("子流程节点入参类型错误,期望*flowDto.NodeExecutionInput,实际%T", input) + } + // 2. 解析子流程配置 + subFlowConfig := nodeExecInput.Config.SubConfig + if subFlowConfig == nil { + return nil, fmt.Errorf("子流程节点缺少配置") + } + getRes, err := FlowUserService.Get(ctx, &flowDto.GetFlowUserReq{ + Id: subFlowConfig.WorkflowId, + }) + if err != nil { + return nil, err + } + // 3. 引入参数解析:把首页表单值/上游引用值/静态默认值写入子流程开始节点 outputConfig。 + // 须在 BuildGraph / ExtractFlowNodeFrom 之前执行,batchInputs 深拷贝的才是注入后的开始节点。 + injectSubFlowFields(nodeExecInput.Global, getRes.FlowContent, subFlowConfig.Fields) + // 4. 并发数:从主流程开始节点 outputConfig 的 maxConcurrency 字段读取(前端把子流程节点生成次数表单字段聚合到主流程开始节点),读不到再用子流程节点配置兜底 + maxConcurrency := mainFlowMaxConcurrency(nodeExecInput.Global, subFlowConfig.MaxConcurrency) + // 4. 编译子流程Graph(复用现有 BuildGraphFromFlowContent 逻辑) + nodeList, subGraph := BuildGraph(ctx, getRes.FlowContent) + // 4. 构建子流程Workflow(绑定START/END,和示例对齐) + innerWorkflow := compose.NewWorkflow[*flowDto.FlowExecutionInput, *flowDto.FlowExecutionInput]() + // 挂载子图节点并绑定全局START + innerWorkflow.AddGraphNode("sub_flow_graph", subGraph).AddInput(compose.START) + // 绑定子图输出到全局END + innerWorkflow.End().AddInput("sub_flow_graph") + // 生成次数(批量条数):maxConcurrency<=0 时按 1 次兜底 + batchCount := maxConcurrency + if batchCount <= 0 { + batchCount = 1 + } + // 5. 构建BatchNode(批量执行子流程,复用示例逻辑) + batchNode := batch.NewBatchNode(&batch.NodeConfig[*flowDto.FlowExecutionInput, *flowDto.FlowExecutionInput]{ + Name: fmt.Sprintf("sub_flow_batch_%s", nodeExecInput.Config.Id), + InnerTask: innerWorkflow, + MaxConcurrency: batchCount, + }) + + // 6. 提取批量输入:按生成次数生成 N 份(每份独立克隆 ConfigMap,避免并发执行时节点输出写串) + configMap := buildConfigMap(getRes.FlowContent, nodeList) + batchInputs := make([]*flowDto.FlowExecutionInput, 0, batchCount) + for j := 0; j < batchCount; j++ { + batchInputs = append(batchInputs, &flowDto.FlowExecutionInput{ + NodeGroupId: nodeExecInput.Global.NodeGroupId, + ExecutionId: nodeExecInput.Global.ExecutionId, + FlowId: nodeExecInput.Global.FlowId, + ConfigMap: cloneConfigMap(configMap), + SessionId: nodeExecInput.Global.SessionId, + }) + } + // 7. 执行批量子流程 + batchOutput, err := batchNode.Invoke(ctx, batchInputs) + if err != nil { + return nil, fmt.Errorf("执行子流程BatchNode失败: %v", err) + } + // 8. 展平每份子流程执行的节点输出,写回当前节点 OutputResult 供下游引用 + var outputRes []map[string]any + for _, single := range batchOutput { + if single == nil { + continue + } + outputRes = append(outputRes, collectFlowNodeResults(single)...) + } + g.Log().Info(ctx, fmt.Sprintf("子流程执行完成,共 %d 次,输出 %d 条", batchCount, len(outputRes))) + nodeExecInput.Config.OutputResult = outputRes + return nodeExecInput, nil +} + +// injectSubFlowFields 将子流程节点引入参数(subConfig.Fields)解析后写入子流程开始节点 +// outputConfig,使子流程启动时能读到首页表单值/上游引用值/静态默认值。 +// 每个字段的取值优先级:valueSource 引用解析成功 → field.value → field.defaultValue; +// 匹配键为 field(前端约定以 field 为主,不兼容 path)。 +func injectSubFlowFields(global *flowDto.FlowExecutionInput, subFlowContent *entity.FlowInfo, fields []map[string]any) { + if global == nil || subFlowContent == nil || len(fields) == 0 { + return + } + startNode := subFlowStartNode(subFlowContent) + if startNode == nil { + return + } + byField := make(map[string]map[string]any, len(startNode.OutputConfig)) + for _, output := range startNode.OutputConfig { + byField[gconv.String(output["field"])] = output + } + for _, field := range fields { + entry := byField[gconv.String(field["field"])] + if entry == nil { + continue + } + value := field["value"] + if vs, has := field["valueSource"]; has && vs != nil { + if vsNodeId, vsField := firstValueSource(vs); vsNodeId != "" && vsField != "" { + if v, _, ok := values.ResolveValueSource(global, vsNodeId, vsField); ok { + value = v + } + } + } + if value == nil { + value = field["defaultValue"] + } + if value != nil { + entry["value"] = value + } + } +} + +// firstValueSource 从 valueSource 提取第一个引用源 (nodeId, field)。 +// 前端契约统一数组 [{nodeId, field}],旧 DSL 可能是单对象 {nodeId, field},两种形态都兼容; +// 子流程字段与引用源一一对应,只取第一个。 +func firstValueSource(vs any) (nodeId, field string) { + if vs == nil { + return + } + switch v := vs.(type) { + case []any: + if len(v) > 0 { + return firstValueSource(v[0]) + } + return + case []map[string]any: + if len(v) > 0 { + return firstValueSource(v[0]) + } + return + } + m := gconv.Map(vs) + nodeId = gconv.String(m["nodeId"]) + field = gconv.String(m["field"]) + return +} + +// subFlowStartNode 返回工作流开始节点 +func subFlowStartNode(content *entity.FlowInfo) *entity.FlowNode { + if content == nil { + return nil + } + for i := range content.Nodes { + if content.Nodes[i].Id == content.StartNodeId { + return &content.Nodes[i] + } + } + return nil +} + +// mainFlowMaxConcurrency 取子流程批量执行并发数:从主流程开始节点 outputConfig +// 的 maxConcurrency 字段读取(前端把子流程节点的生成次数表单字段聚合到主流程开始节点), +// 读不到再用子流程节点配置的兜底值。 +func mainFlowMaxConcurrency(global *flowDto.FlowExecutionInput, fallback int) int { + if global == nil { + return fallback + } + for _, n := range global.ConfigMap { + if n == nil || n.NodeCode != node.NodeTypeStart { + continue + } + for _, output := range n.OutputConfig { + if gconv.String(output["field"]) != "maxConcurrency" && gconv.String(output["path"]) != "maxConcurrency" { + continue + } + if v := gconv.Int(output["value"]); v > 0 { + return v + } + } + return fallback + } + return fallback +} + +// cloneConfigMap 深拷贝 ConfigMap,保证各批次子流程并发执行时节点输出互不串扰。 +// 浅拷贝会共享 *entity.FlowNode,并发写 OutputResult 产生竞态。 +func cloneConfigMap(src map[string]*entity.FlowNode) map[string]*entity.FlowNode { + dst := make(map[string]*entity.FlowNode, len(src)) + for k, v := range src { + data, err := json.Marshal(v) + if err != nil { + dst[k] = v + continue + } + n := new(entity.FlowNode) + if err = json.Unmarshal(data, n); err != nil { + dst[k] = v + continue + } + dst[k] = n + } + return dst +} + +// collectFlowNodeResults 收集一次子流程执行中所有已执行节点的输出,展平成 {字段:值} 列表 +func collectFlowNodeResults(execInput *flowDto.FlowExecutionInput) []map[string]any { + var res []map[string]any + for _, executed := range execInput.ExecutedNodes { + if nodeConfig := execInput.ConfigMap[executed.NodeId]; nodeConfig != nil { + res = append(res, nodeConfig.OutputResult...) + } + } + return res +} diff --git a/workflow/service/flow/lambda_summary.go b/workflow/service/flow/lambda_summary.go new file mode 100644 index 0000000..fdaae14 --- /dev/null +++ b/workflow/service/flow/lambda_summary.go @@ -0,0 +1,122 @@ +package flow + +import ( + "ai-agent/workflow/consts/flow" + "ai-agent/workflow/consts/public" + nodeDao "ai-agent/workflow/dao/node" + sessionDao "ai-agent/workflow/dao/session" + flowDto "ai-agent/workflow/model/dto/flow" + nodeDto "ai-agent/workflow/model/dto/node" + sessionDto "ai-agent/workflow/model/dto/session" + "ai-agent/workflow/model/entity" + "context" + "fmt" + "strings" + + "gitea.redpowerfuture.com/red-future/common/db/gfdb" + "github.com/gogf/gf/v2/database/gdb" + "github.com/gogf/gf/v2/frame/g" + "github.com/gogf/gf/v2/util/gconv" +) + +func DataMergeLambda(ctx context.Context, input any) (res any, err error) { + nodeInput, ok := input.(*flowDto.NodeExecutionInput) + if !ok { + return nil, fmt.Errorf("参数合并入参类型错误") + } + return nodeInput, nil +} + +func SummaryLambda(ctx context.Context, input any) (any, error) { + execInput, ok := input.(*flowDto.NodeExecutionInput) + if !ok { + return nil, fmt.Errorf("汇总节点入参类型错误,实际是 %T", input) + } + + // 聚合所有已执行节点中需入库的文件结果(两层规则) + summaryResult := collectSaveFileResults(ctx, execInput.Global) + + // 把汇总结果存入当前节点的输出 + g.Log().Info(ctx, fmt.Sprintf("结果汇总完成,汇总数据:%+v", summaryResult)) + + err := gfdb.DB(ctx, public.DbNameBlackDeacon).Transaction(ctx, func(ctx context.Context, tx gdb.TX) error { + res, _, err := nodeDao.NodeExecutionDao.ListByFlowExecutionId(ctx, &nodeDto.ListNodeExecutionByFlowReq{ + NodeGroupId: execInput.Global.NodeGroupId, + }, entity.NodeExecutionCol.TokenInfo) + if err != nil { + return err + } + var totalTokens int + var totalFee float64 + for _, item := range res { + for _, itemToken := range item.TokenInfo { + m := gconv.Map(itemToken) + totalTokens += gconv.Int(m["total_tokens"]) + totalFee += gconv.Float64(m["total_fee"]) + } + } + _, err = sessionDao.ExecWorkflowDao.Update(ctx, &sessionDto.UpdateWorkflowReq{ + Id: execInput.Global.ExecutionId, + Status: flow.FlowExecutionStatusSuccess.Code(), + TotalTokens: totalTokens, + TotalFee: totalFee, + }) + if err != nil { + return err + } + if len(summaryResult) > 0 { + _, err = sessionDao.ExecWorkflowResultDao.BatchInsert(ctx, summaryResult) + if err != nil { + return err + } + } + return nil + }) + + return execInput, err +} + +// collectSaveFileResults 按两层规则收集需入库的文件结果: +// 第一层:节点须开启"保存文件"(IsSaveFile); +// 第二层:key 取自节点 OutputResult 的各字段,命中 ModelResponseBodyMapping 才入库; +// 原始响应体 key(respBody)恒入库(不要求映射声明);HTTP 节点产出以 http_file_url:{key} +// 标记的字段(IsSaveFile 时由 HttpCallResultLambda 生成)恒入库(无模型响应映射可查)。 +// 结果值为 http(s) URL 或 MinIO 对象裸路径直接使用;非路径值(base64 图片/文本)先上传 OSS 换取 URL, +// 文本内容以 .inc 扩展名存储。 +func collectSaveFileResults(ctx context.Context, execInput *flowDto.FlowExecutionInput) []*sessionDto.CreateWorkflowResultReq { + if execInput == nil { + return nil + } + var summaryResult []*sessionDto.CreateWorkflowResultReq + for _, executedNode := range execInput.ExecutedNodes { + nodeConfig := execInput.ConfigMap[executedNode.NodeId] + if nodeConfig == nil || len(nodeConfig.OutputResult) == 0 || !nodeConfig.IsSaveFile { + continue + } + // 第二层:key 取自节点 OutputResult 的各字段, + // 命中 ModelResponseBodyMapping 才入库;respBody 与 HTTP 节点 http_file_url:{key} 标记恒入库 + saveKeys := nodeConfig.ModelConfig.ModelResponseBodyMapping + for _, respBody := range nodeConfig.OutputResult { + for key, val := range gconv.Map(respBody) { + isHTTPFile := strings.HasPrefix(key, "http_file_url:") + if !isHTTPFile { + if _, ok := saveKeys[key]; !ok && key != "respBody" { + continue + } + } + fileUrl, err := resolveSaveFileResult(ctx, val) + if err != nil { + g.Log().Warningf(ctx, "collectSaveFileResults 上传结果文件失败 key=%s err=%v", key, err) + continue + } + summaryResult = append(summaryResult, &sessionDto.CreateWorkflowResultReq{ + SessionId: execInput.SessionId, + FlowId: execInput.FlowId, + ExecId: execInput.ExecutionId, + ResultFileUrl: fileUrl, + }) + } + } + } + return summaryResult +} diff --git a/workflow/service/flow/lambda_tool.go b/workflow/service/flow/lambda_tool.go new file mode 100644 index 0000000..0b26218 --- /dev/null +++ b/workflow/service/flow/lambda_tool.go @@ -0,0 +1,90 @@ +package flow + +import ( + "ai-agent/gateway" + "ai-agent/workflow/consts/model" + "ai-agent/workflow/service/flow/processor" + "context" + "fmt" + "strings" + + "github.com/gogf/gf/v2/frame/g" +) + +// isVideoModel 判断模型是否为视频模型(模型类型 TypeVideo=600),用于视频节点多视频自动合成判断 +func isVideoModel(ctx context.Context, modelId int64) bool { + modelInfo, err := gateway.GetModelInfoById(ctx, &gateway.GetModelInfoByIdReq{ModelId: modelId}) + if err != nil { + g.Log().Warningf(ctx, "查询模型配置失败,跳过自动视频合成 modelId=%d err=%v", modelId, err) + return false + } + return modelInfo.ModelManage.ModelType != nil && *modelInfo.ModelManage.ModelType == model.TypeVideo +} + +// invokePreTool 执行前置处理器,把模型请求参数转换为模型调用入参列表。 +// 前置处理器契约:入参即模型请求参数本体;返回值: +// - map[string]any 一次模型调用,入参为返回值 +// - []map[string]any 多次模型调用,逐个入参请求 +// - nil 视为异常,节点失败(不允许静默跳过模型调用) +func invokePreTool(ctx context.Context, processorName string, modelParams map[string]any) (paramsList []map[string]any, err error) { + if processorName == "" { + return []map[string]any{stripInternalKeys(modelParams)}, nil + } + data, err := processor.Call(ctx, processorName, modelParams) + if err != nil { + return nil, fmt.Errorf("执行前置处理器[%s]失败: %v", processorName, err) + } + switch v := data.(type) { + case nil: + return nil, fmt.Errorf("前置处理器[%s]返回空", processorName) + case map[string]any: + return []map[string]any{stripInternalKeys(v)}, nil + case []map[string]any: + list := make([]map[string]any, 0, len(v)) + for _, m := range v { + list = append(list, stripInternalKeys(m)) + } + return list, nil + default: + return nil, fmt.Errorf("前置处理器[%s]返回类型不支持: %T", processorName, data) + } +} + +// stripInternalKeys 剥离 __ 前缀的内部键(如 __segment_fields/__produced), +// 模型网关做参数严格校验(CheckParams strictUnknown)会拒绝未知字段,内部标记不得随请求体下发。 +func stripInternalKeys(params map[string]any) map[string]any { + if params == nil { + return params + } + for k := range params { + if strings.HasPrefix(k, "__") { + delete(params, k) + } + } + return params +} + +// invokePostTool 执行后置处理器,加工模型调用结果。 +// 后置处理器契约:入参 {"output": 模型输出结果列表, "request": 原始模型请求参数}(列表须包成对象传入);返回值: +// - []map[string]any 替换模型输出 +// - map[string]any 替换为单条输出 +// - nil 保留原输出 +func invokePostTool(ctx context.Context, processorName string, outputRes []map[string]any, requestParams map[string]any) ([]map[string]any, error) { + if processorName == "" { + return outputRes, nil + } + data, err := processor.Call(ctx, processorName, map[string]any{"output": outputRes, "request": requestParams}) + if err != nil { + return nil, fmt.Errorf("执行后置处理器[%s]失败: %v", processorName, err) + } + switch v := data.(type) { + case nil: + return outputRes, nil + case []map[string]any: + return v, nil + case map[string]any: + return []map[string]any{v}, nil + default: + return nil, fmt.Errorf("后置处理器[%s]返回类型不支持: %T", processorName, data) + } +} diff --git a/workflow/service/flow/lambda_value_source.go b/workflow/service/flow/lambda_value_source.go deleted file mode 100644 index e0f90c9..0000000 --- a/workflow/service/flow/lambda_value_source.go +++ /dev/null @@ -1,703 +0,0 @@ -package flow - -import ( - "ai-agent/workflow/consts/node" - flowDto "ai-agent/workflow/model/dto/flow" - "ai-agent/workflow/model/entity" - "context" - "encoding/json" - "reflect" - "regexp" - "strconv" - "strings" - - "github.com/gogf/gf/v2/frame/g" - "github.com/gogf/gf/v2/os/glog" - "github.com/gogf/gf/v2/util/gconv" - "github.com/tidwall/gjson" -) - -var ( - // 匹配 [数字] - regNumIndex = regexp.MustCompile(`\[\d+\]`) - // 匹配 .attrs - regAttrs = regexp.MustCompile(`\.attrs`) - // 匹配带捕获组的数组下标,转扁平点分路径用 - arrayIndexPath = regexp.MustCompile(`\[(\d+)\]`) -) - -// CleanFieldPath 清理字段路径:移除 .attrs、数字下标转为 [*] -// 示例:usage.attrs.total_tokens → usage.total_tokens -// 示例:choices.attrs[0].attrs.message.attrs.content → choices[*].message.content -func CleanFieldPath(path string) string { - //// 1. 替换 [数字] 为 [*] - //s := regNumIndex.ReplaceAllString(path, `.#`) - //// 2. 移除所有 .attrs - //s = regAttrs.ReplaceAllString(s, "") - index := CleanFieldPathReplaceNumIndex(path) - attrs := CleanFieldPathRemoveAttrs(index) - return attrs -} - -func CleanFieldPathReplaceNumIndex(path string) string { - // 1. 替换 [数字] 为 [*] - s := regNumIndex.ReplaceAllString(path, `.#`) - return s -} - -func CleanFieldPathRemoveAttrs(path string) string { - // 2. 移除所有 .attrs - s := regAttrs.ReplaceAllString(path, "") - return s -} - -// UnwrapSchemaWrapper 递归剥掉 json-schema-editor 输出的 {type, value/attrs} 包裹层, -// 只保留干净的 key/value 嵌套结构。 -// 示例: -// -// {"a": {"type":"string","value":"hi"}} → {"a": "hi"} -// {"b": {"type":"object","attrs":{"c":1}}} → {"b": {"c": 1}} -// {"arr": {"type":"array","attrs":[{"type":"number","value":1}]}} → {"arr": [1]} -func UnwrapSchemaWrapper(v any) any { - switch val := v.(type) { - case map[string]any: - // 识别包裹节点:{type: "", value/attrs: <实际值>, ...} - if t, ok := val["type"].(string); ok && isSchemaEditorType(t) { - dataKey := "value" - if t == "object" || t == "array" { - dataKey = "attrs" - } - if raw, has := val[dataKey]; has { - return UnwrapSchemaWrapper(raw) - } - } - res := make(map[string]any, len(val)) - for k, child := range val { - res[k] = UnwrapSchemaWrapper(child) - } - return res - case []any: - res := make([]any, len(val)) - for i, item := range val { - res[i] = UnwrapSchemaWrapper(item) - } - return res - default: - return val - } -} - -// isSchemaEditorType 是否为 json-schema-editor 的 6 种类型标识 -func isSchemaEditorType(t string) bool { - switch t { - case "string", "number", "boolean", "null", "object", "array": - return true - } - return false -} - -// MapResultByTemplate 按 template 定义的结构,从 source 中拷贝对应字段的值。 -// 只保留 template 里出现的字段:对象字段按同名字段递归拷贝,数组字段按模板元素结构逐元素过滤,标量字段直接拷贝 source 的值。 -func MapResultByTemplate(template map[string]any, source map[string]any) map[string]any { - result := make(map[string]any, len(template)) - for key, tmplVal := range template { - srcVal, ok := source[key] - if !ok { - continue - } - if tmplMap, isMap := tmplVal.(map[string]any); isMap { - if srcMap, isMap := srcVal.(map[string]any); isMap { - result[key] = MapResultByTemplate(tmplMap, srcMap) - } - continue - } - if tmplArr, isArr := tmplVal.([]any); isArr { - result[key] = mapTemplateArray(tmplArr, srcVal) - continue - } - result[key] = srcVal - } - return result -} - -// mapTemplateArray 按模板数组的元素结构映射 source 数组: -// 模板首元素为对象时,逐元素按 MapResultByTemplate 过滤只保留模板字段; -// 模板数组为空或首元素非对象(无法确定元素结构)时,原样拷贝 source 数组。 -func mapTemplateArray(tmplArr []any, srcVal any) any { - srcList, ok := srcVal.([]any) - if !ok || len(tmplArr) == 0 { - return srcVal - } - elemTmpl, ok := tmplArr[0].(map[string]any) - if !ok { - return srcVal - } - result := make([]any, 0, len(srcList)) - for _, srcElem := range srcList { - if srcMap, isMap := srcElem.(map[string]any); isMap { - result = append(result, MapResultByTemplate(elemTmpl, srcMap)) - } else { - result = append(result, srcElem) - } - } - return result -} - -// ProcessValueSourceRecursive 递归遍历map,同级同时存在value和valueSource则把value设置为"AA" -func ProcessValueSourceRecursive(rawParams map[string]interface{}, globalParams *flowDto.FlowExecutionInput) { - walkMap(rawParams, globalParams) -} - -// resolveValueSource 解析 valueSource {nodeId, fieldName} 引用的实际值。 -// 返回 (value, refsName, ok);ok=false 表示引用节点不存在或引用值仍为空。 -// - 开始/表单节点:OutputConfig 平铺条目按 field == fieldName 匹配(前端约定以 field 为主, -// 不兼容 path),直接读 entry 的 value / refsName -// - scriptTranscribe 节点:OutputResult 是各段扁平请求参数,按段序收集字段为数组(段位留 nil) -// - 其他节点:读 OutputResult 中 fieldName 路径对应的值 -func resolveValueSource(global *flowDto.FlowExecutionInput, nodeId, field string) (value any, refsName any, ok bool) { - if global == nil || global.ConfigMap == nil { - return nil, nil, false - } - nodeConfig := global.ConfigMap[nodeId] - if nodeConfig == nil { - return nil, nil, false - } - switch nodeConfig.NodeCode { - case node.NodeTypeStart, node.NodeTypeForm: - for _, output := range nodeConfig.OutputConfig { - if gconv.String(output["field"]) != field { - continue - } - if !g.IsEmpty(output["value"]) { - return output["value"], output["refsName"], true - } - } - case node.NodeTypeScriptTranscribe: - // 脚本转写节点 OutputResult 是各段扁平请求参数(split_shots_pipeline 产出,key 为字面量 - // prompt/duration/seed 等),按段序读取 output[field] 收集为数组,供分段模型节点整体引用。 - // 每段都占一位(字段缺失/为空留 nil),保证数组与段序对齐,供 split_segment 按段取值。 - var list []any - for _, output := range nodeConfig.OutputResult { - list = append(list, output[field]) - } - for _, v := range list { - if !g.IsEmpty(v) { - return list, "", true - } - } - default: - // templates 是模型节点在前端配置的静态输出模板,不在 OutputResult 中,需单独取 - if field == "templates" { - if !g.IsEmpty(nodeConfig.Templates) { - return nodeConfig.Templates, "", true - } - return nil, nil, false - } - for _, output := range nodeConfig.OutputResult { - // 模型节点输出记录是单 key 的字面量扁平 key(如 "choices.attrs[0].attrs.delta.attrs.content"), - // gjson 会把 . 和 [0] 当结构路径解析,无法命中字面量 key,故先按字面量 key 直接取值; - // 未命中再回退 gjson 路径查询(兼容真正嵌套的输出结构)。 - if v, has := output[field]; has { - value = v - } else { - value = gjson.Get(gconv.String(output), field).Value() - } - if !g.IsEmpty(value) { - return value, gjson.Get(gconv.String(output), CleanFieldPath("refsName")).Value(), true - } - } - } - return nil, nil, false -} - -// walkMap 递归处理map/数组 -func walkMap(data interface{}, globalParams *flowDto.FlowExecutionInput) { - switch v := data.(type) { - case map[string]interface{}: - // 有 valueSource:解析引用节点值 - if valueSource, hasSource := v["valueSource"]; hasSource { - sources := new([]entity.ValueSource) - gconv.Structs(valueSource, sources) - - // 多个引用源:把各源解析出的值拼成 "label: value"(无 label 只拼值),逗号分隔 - if len(*sources) > 1 { - parts := make([]string, 0, len(*sources)) - var refsName any - for _, src := range *sources { - value, rn, ok := resolveValueSource(globalParams, src.NodeId, src.Field) - if !ok || schemaValueEmpty(value) { - continue - } - // 引用非模型节点(开始/表单/HTTP/脚本转写等)时,值按当前字段声明的 type 做类型化转换; - // 模型节点值由模型网关处理,复制时不需要转换 - if !isModelSourceNode(globalParams, src.NodeId) { - value = assignBySchemaType(v, value) - } - text := toPlainString(value) - if src.Label != "" { - text = src.Label + ": " + text - } - parts = append(parts, text) - if !g.IsEmpty(rn) && g.IsEmpty(refsName) { - refsName = rn - } - } - if len(parts) > 0 { - v["value"] = strings.Join(parts, ", ") - if !g.IsEmpty(refsName) { - v["refsName"] = refsName - } - return - } - } else if len(*sources) == 1 { - // 单个引用源:保持旧行为,值按原样赋值(非模型节点按声明 type 转换) - src := (*sources)[0] - value, refsName, ok := resolveValueSource(globalParams, src.NodeId, src.Field) - if ok && !isModelSourceNode(globalParams, src.NodeId) { - value = assignBySchemaType(v, value) - } - if ok && !schemaValueEmpty(value) { - v["value"] = value - if !g.IsEmpty(refsName) { - v["refsName"] = refsName - } - return - } - } - } - // 统一兜底:无 valueSource(或解析失败/值为空)时,value 为空或 0 则取 defaultValue - if defaultValue, hasDefault := v["defaultValue"]; hasDefault && isEmptyForFallback(v["value"]) && !schemaValueEmpty(defaultValue) { - v["value"] = defaultValue - } - // 递归遍历所有子元素 - for _, child := range v { - walkMap(child, globalParams) - } - case []interface{}: - // 数组遍历 - for _, item := range v { - walkMap(item, globalParams) - } - } -} - -// isEmptyForFallback 兜底场景判空:除 schemaValueEmpty 规则外,数字 0 也视为未填写, -// 便于配置了 defaultValue 的字段在值为 0 时用默认值兜底。 -// 覆盖 json.Number(gconv 反序列化数字的运行时类型)与字符串 "0"/"0.0"。 -func isEmptyForFallback(v interface{}) bool { - if schemaValueEmpty(v) { - return true - } - switch val := v.(type) { - case float32: - return val == 0 - case float64: - return val == 0 - case int: - return val == 0 - case int8: - return val == 0 - case int16: - return val == 0 - case int32: - return val == 0 - case int64: - return val == 0 - case uint: - return val == 0 - case uint8: - return val == 0 - case uint16: - return val == 0 - case uint32: - return val == 0 - case uint64: - return val == 0 - case json.Number: - if f, err := val.Float64(); err == nil { - return f == 0 - } - case string: - if f, err := strconv.ParseFloat(val, 64); err == nil { - return f == 0 - } - } - return false -} - -// isModelSourceNode 判断引用源节点是否为模型节点(值由模型网关处理,复制时不转换) -func isModelSourceNode(global *flowDto.FlowExecutionInput, nodeId string) bool { - if global == nil || global.ConfigMap == nil { - return false - } - nodeConfig := global.ConfigMap[nodeId] - return nodeConfig != nil && nodeConfig.NodeCode == node.NodeTypeModel -} - -// assignBySchemaType 按当前字段声明的 schema 类型把值类型化: -// string 遇数组/对象转 JSON 字符串;number/boolean 解析字符串;object/array 解析 JSON 字符串;其余原样返回 -func assignBySchemaType(node map[string]interface{}, value any) any { - t, _ := node["type"].(string) - return assignByType(t, value) -} - -// assignByType 按字段声明的 type 把值类型化;walkMap 的 schema 节点与 parseMap 的模型参数共用 -func assignByType(t string, value any) any { - switch t { - case "string": - return toSchemaString(value) - case "number": - return toSchemaNumber(value) - case "boolean": - return toSchemaBool(value) - case "object", "array": - return toSchemaStruct(value) - default: - return value - } -} - -// toSchemaString 转 string:字符串原样,数组元素拼成字符串(单元素取元素本身,多元素逗号连接),对象序列化为 JSON 字符串 -func toSchemaString(v any) any { - switch val := v.(type) { - case []interface{}: - parts := make([]string, 0, len(val)) - for _, item := range val { - parts = append(parts, toPlainString(item)) - } - return strings.Join(parts, ",") - case map[string]interface{}: - if b, err := json.Marshal(val); err == nil { - return string(b) - } - } - return v -} - -// toPlainString 把数组元素转成不带括号的纯字符串: -// 数组([]any / 类型化切片)逐元素取纯字符串,单元素取元素本身,多元素逗号连接; -// 对象序列化为 JSON 字符串;其余原样字符串化。 -func toPlainString(v any) string { - if s, ok := v.(string); ok { - return s - } - switch val := v.(type) { - case []interface{}: - parts := make([]string, 0, len(val)) - for _, item := range val { - parts = append(parts, toPlainString(item)) - } - return strings.Join(parts, ",") - case map[string]interface{}: - if b, err := json.Marshal(val); err == nil { - return string(b) - } - } - rv := reflect.ValueOf(v) - if rv.IsValid() && (rv.Kind() == reflect.Slice || rv.Kind() == reflect.Array) { - parts := make([]string, 0, rv.Len()) - for i := 0; i < rv.Len(); i++ { - parts = append(parts, toPlainString(rv.Index(i).Interface())) - } - return strings.Join(parts, ",") - } - if b, err := json.Marshal(v); err == nil { - return string(b) - } - return gconv.String(v) -} - -// toSchemaNumber 转 number:数字原样,字符串尝试解析为 float64,失败原样返回 -func toSchemaNumber(v any) any { - if s, ok := v.(string); ok { - if f, err := strconv.ParseFloat(s, 64); err == nil { - return f - } - } - return v -} - -// toSchemaBool 转 boolean:布尔原样,字符串尝试解析为 bool,失败原样返回 -func toSchemaBool(v any) any { - if s, ok := v.(string); ok { - if b, err := strconv.ParseBool(s); err == nil { - return b - } - } - return v -} - -// toSchemaStruct 转 object/array:合法 JSON 字符串解析为结构化数据,否则原样返回 -func toSchemaStruct(v any) any { - s, ok := v.(string) - if !ok { - return v - } - if !json.Valid([]byte(s)) { - return v - } - var out any - if err := json.Unmarshal([]byte(s), &out); err != nil { - return v - } - return out -} - -// CleanEmptyModelParams 剔除模型请求参数中 value 为空的字段; -// 数组/枚举(attrs / enumValues)元素整体为空时移除整个元素。0/false 视为有效值。 -func CleanEmptyModelParams(params map[string]interface{}) { - cleanSchemaMap(params) -} - -// cleanSchemaMap 递归清理普通 map:包装节点按 schema 语义清理,空字段删除 -func cleanSchemaMap(m map[string]interface{}) { - for key, val := range m { - switch v := val.(type) { - case map[string]interface{}: - if isSchemaWrapperNode(v) { - cleanSchemaWrapper(v) - if isSchemaNodeEmpty(v) { - delete(m, key) - } - } else { - cleanSchemaMap(v) - } - case []interface{}: - m[key] = cleanSchemaSlice(v) - } - } -} - -// cleanSchemaWrapper 清理单个 {type,...} 包装节点:递归 value / attrs / enumValues 容器 -func cleanSchemaWrapper(node map[string]interface{}) { - if mv, ok := node["value"].(map[string]interface{}); ok { - cleanSchemaMap(mv) - } - if lv, ok := node["value"].([]interface{}); ok { - node["value"] = cleanSchemaSlice(lv) - } - if attrs, ok := node["attrs"].(map[string]interface{}); ok { - cleanSchemaMap(attrs) - } - if attrs, ok := node["attrs"].([]interface{}); ok { - node["attrs"] = cleanSchemaSlice(attrs) - } - if evs, ok := node["enumValues"].([]interface{}); ok { - node["enumValues"] = cleanSchemaSlice(evs) - } -} - -// cleanSchemaSlice 清理数组/枚举元素,元素为包装节点且整体为空时移除 -func cleanSchemaSlice(list []interface{}) []interface{} { - i := 0 - for i < len(list) { - if item, ok := list[i].(map[string]interface{}); ok { - if isSchemaWrapperNode(item) { - cleanSchemaWrapper(item) - if isSchemaNodeEmpty(item) { - list = append(list[:i], list[i+1:]...) - continue - } - } else { - cleanSchemaMap(item) - } - } - i++ - } - return list -} - -// isSchemaWrapperNode 是否为 {type: } 包装节点 -func isSchemaWrapperNode(m map[string]interface{}) bool { - t, ok := m["type"].(string) - return ok && isSchemaEditorType(t) -} - -// isSchemaNodeEmpty 判断 schema 节点是否已无有效内容: -// 标量看 value(0/false 有效);object/array 看 value/attrs/enumValues 容器是否都为空 -func isSchemaNodeEmpty(node map[string]interface{}) bool { - t, _ := node["type"].(string) - switch t { - case "object": - return schemaContainerEmpty(node, "value") && schemaContainerEmpty(node, "attrs") - case "array": - return schemaContainerEmpty(node, "value") && schemaContainerEmpty(node, "attrs") && schemaContainerEmpty(node, "enumValues") - default: - return schemaValueEmpty(node["value"]) - } -} - -// schemaContainerEmpty 容器(value/attrs/enumValues)是否为空 -func schemaContainerEmpty(node map[string]interface{}, key string) bool { - switch v := node[key].(type) { - case []interface{}: - return len(v) == 0 - case map[string]interface{}: - return len(v) == 0 - default: - return v == nil - } -} - -// schemaValueEmpty 值是否为空;0/false 视为有效值不剔除 -func schemaValueEmpty(v interface{}) bool { - switch val := v.(type) { - case nil: - return true - case string: - return val == "" - case []interface{}: - return len(val) == 0 - case map[string]interface{}: - return len(val) == 0 - default: - return false - } -} - -// BuildModelRequestBody 从参数定义 + 全局执行上下文构建最终嵌套 JSON 请求体。 -func BuildModelRequestBody(params []entity.FlowModelParams, globalParams *flowDto.FlowExecutionInput) (map[string]interface{}, error) { - // 1. 解析引用、过滤空值 - resolved := parseMap(params, globalParams) - - // 2. 转扁平路径映射 - flat := toFlatMap(resolved) - - // 3. 引用了脚本转写节点的字段打内部标记 __segment_fields(逗号分隔的扁平路径), - // 供前置处理器 split_segment 按段拆批;__ 前缀内部键由 invokePreTool 统一剥离,不传给模型网关 - if seg := segmentFields(resolved, globalParams); len(seg) > 0 { - flat["__segment_fields"] = strings.Join(seg, ",") - } - - return flat, nil -} - -// segmentFields 收集引用了脚本转写节点的字段扁平路径(分段字段)。 -// 脚本转写节点按段产出一份扁平参数列表,下游模型节点单源引用其字段时, -// 值按段序聚合成数组(见 resolveValueSource 的 scriptTranscribe 分支),需随批拆分。 -func segmentFields(resolved []entity.FlowModelParams, globalParams *flowDto.FlowExecutionInput) []string { - if globalParams == nil || globalParams.ConfigMap == nil { - return nil - } - var fields []string - for _, p := range resolved { - if g.IsEmpty(p.Path) || len(p.ValueSource) != 1 { - continue - } - src := p.ValueSource[0] - if nodeConfig := globalParams.ConfigMap[src.NodeId]; nodeConfig != nil && - nodeConfig.NodeCode == node.NodeTypeScriptTranscribe { - fields = append(fields, flatPath(p.Path)) - } - } - return fields -} - -// parseMap 解析模型请求参数 -func parseMap(data []entity.FlowModelParams, globalParams *flowDto.FlowExecutionInput) []entity.FlowModelParams { - newData := make([]entity.FlowModelParams, 0, len(data)) - for _, item := range data { - var d entity.FlowModelParams - d.Path = item.Path - d.Type = item.Type - - // 无引用源:直接取静态值 - if g.IsEmpty(item.ValueSource) { - if isParamEmpty(item.Value) { - continue - } - d = item - newData = append(newData, d) - continue - } - - // 有引用源:单源保持旧行为;多源把各源解析出的值拼成 "label: value"(无 label 只拼值),逗号分隔 - var value, refsName any - if len(item.ValueSource) > 1 { - parts := make([]string, 0, len(item.ValueSource)) - for _, src := range item.ValueSource { - v, rn, ok := resolveValueSource(globalParams, src.NodeId, src.Field) - if !ok || isParamEmpty(v) { - continue - } - // 模型解析时引用其他节点的值不做类型化转换; - // 多引用源需拼接为字符串,用 toPlainString 渲染(数组取元素去括号) - text := toPlainString(v) - if src.Label != "" { - text = src.Label + ": " + text - } - parts = append(parts, text) - if !g.IsEmpty(rn) && g.IsEmpty(refsName) { - refsName = rn - } - } - if len(parts) == 0 { - // 解析失败不静默,留日志便于排查引用丢失 - glog.Debugf(context.Background(), - "resolve value source failed, nodeId=%+v path=%s", - item.ValueSource, item.Path) - continue - } - d.Value = strings.Join(parts, ", ") - d.RefsName = gconv.String(refsName) - d.ValueSource = item.ValueSource - newData = append(newData, d) - continue - } - - // 单个引用源:解析取非空值;模型解析时引用其他节点的数组值不做类型化转换,整体传给模型 - src := item.ValueSource[0] - value, refsName, ok := resolveValueSource(globalParams, src.NodeId, src.Field) - if !ok || isParamEmpty(value) { - // 解析失败不静默,留日志便于排查引用丢失 - glog.Debugf(context.Background(), - "resolve value source failed, nodeId=%+v path=%s", - item.ValueSource, item.Path) - continue - } - d.Value = value - d.RefsName = gconv.String(refsName) - d.ValueSource = item.ValueSource - newData = append(newData, d) - } - return newData -} - -// toFlatMap 将解析后的参数列表转为 sjson 可用的扁平路径映射。 -// key 为 Path,value 为参数值;保留 RefsName 供上层追踪引用来源。 -func toFlatMap(params []entity.FlowModelParams) map[string]interface{} { - m := make(map[string]interface{}, len(params)) - for _, p := range params { - if g.IsEmpty(p.Path) { - continue - } - m[flatPath(p.Path)] = p.Value - } - return m -} - -// flatPath 把数组下标路径转扁平点分路径:a[0].b → a.0.b。 -func flatPath(path string) string { - return arrayIndexPath.ReplaceAllString(path, `.$1`) -} - -// isParamEmpty 判断参数值是否为"空"。 -// 仅 nil、空字符串、空切片/映射视为空;0、false 等零值是合法值,保留。 -func isParamEmpty(v interface{}) bool { - if v == nil { - return true - } - switch val := v.(type) { - case string: - return val == "" - case []byte: - return len(val) == 0 - case []interface{}: - return len(val) == 0 - case map[string]interface{}: - return len(val) == 0 - default: - // 数字、布尔、结构体等一律视为非空 - return false - } -} diff --git a/workflow/service/flow/lambda_node_util.go b/workflow/service/flow/model_call.go similarity index 55% rename from workflow/service/flow/lambda_node_util.go rename to workflow/service/flow/model_call.go index 8d0fc3b..27e756e 100644 --- a/workflow/service/flow/lambda_node_util.go +++ b/workflow/service/flow/model_call.go @@ -5,12 +5,10 @@ import ( "ai-agent/workflow/consts/model" "ai-agent/workflow/consts/node" flowDto "ai-agent/workflow/model/dto/flow" + "ai-agent/workflow/service/flow/values" "context" "fmt" - "regexp" "strings" - "sync" - "unicode/utf8" commonHttp "gitea.redpowerfuture.com/red-future/common/http" "gitea.redpowerfuture.com/red-future/common/oss" @@ -21,47 +19,6 @@ import ( "github.com/google/uuid" ) -// 全局等待任务回调的工具 -var ( - asyncMu sync.Mutex - asyncTasks = make(map[string]chan any) -) - -// Wait 阻塞等待回调结果 -// 调用后会一直卡住,直到 Notify 唤醒 或 超时/取消 -func Wait(ctx context.Context, taskId string) (any, error) { - asyncMu.Lock() - ch := make(chan any, 1) - asyncTasks[taskId] = ch - asyncMu.Unlock() - - defer close(ch) - for { - select { - case result := <-ch: - return result, nil - case <-ctx.Done(): - asyncMu.Lock() - delete(asyncTasks, taskId) - asyncMu.Unlock() - return nil, ctx.Err() - } - } -} - -// Notify 回调时调用,唤醒等待的任务 -func Notify(taskId string, result any) { - asyncMu.Lock() - defer asyncMu.Unlock() - - ch, exist := asyncTasks[taskId] - if !exist { - return - } - ch <- result - delete(asyncTasks, taskId) -} - // ModelCallResultLambda 调用模型并返回输出内容列表,同时回传本次调用的 token/费用(*gateway.ModelCallRes) // 与是否推理模型(供 ModelLambda 决定分批结果是否拼接),供调用方(ModelLambda)累计写入节点执行记录 // token_info,最后由汇总节点聚合到 exec_workflow。 @@ -126,7 +83,7 @@ func HttpCallResultLambda(ctx context.Context, nodeInput *flowDto.NodeExecutionI body = gconv.Map(item.Value) case "response": // 先剥掉 {type, value/attrs} 包裹层,得到干净的输出结构模板 - responseMapping = gconv.Map(UnwrapSchemaWrapper(gconv.Map(item.Value))) + responseMapping = gconv.Map(values.UnwrapSchemaWrapper(gconv.Map(item.Value))) case "responseType": responseType = gconv.String(item.Value) if responseType == "callback" { @@ -143,38 +100,18 @@ func HttpCallResultLambda(ctx context.Context, nodeInput *flowDto.NodeExecutionI } if headers == nil { - headers = make(map[string]string) - if r := g.RequestFromCtx(ctx); r != nil { - for k, v := range r.Request.Header { - if len(v) > 0 { - headers[k] = v[0] - } - } - } - // 后台恢复续跑无 HTTP 请求时(ctx 携带合成 user、无 token),补充 X-User-Info 供下游内部服务 - // GetUserInfo 识别租户;正常执行 ctx 无 user(仅 request 有 Authorization),不会注入, - // 避免把内部身份透传给任意外部 URL。与 gateway/model.go requestHeaders 的恢复兜底保持一致 - if headers["X-User-Info"] == "" { - if u := ctx.Value("user"); u != nil { - headers["X-User-Info"] = gconv.String(u) - } - } + headers = utils.HeadersFromCtx(ctx) } - g.Log().Debugf(ctx, "httpCallResultLambda: body: %v", body) - // 构建请求参数 - ProcessValueSourceRecursive(body, nodeInput.Global) + values.ProcessValueSourceRecursive(body, nodeInput.Global) // 递归剥掉 {type, value/attrs} 包裹层,只保留 key/value - wrapper := UnwrapSchemaWrapper(body) + wrapper := values.UnwrapSchemaWrapper(body) newBody := gconv.Map(wrapper) // body 值若为 MinIO 裸路径(模型网关转存 OSS 后返回,无 http 前缀), // 补上前缀供目标 HTTP 服务直接下载文件 addFilePathPrefix(ctx, url, newBody) - // 打印入参 - g.Log().Debugf(ctx, "httpCallResultLambda: newBody: %v", newBody) - // 1. 自己生成唯一 taskId(不用前端给) taskId := "my_task_" + uuid.New().String() // 自己生成唯一ID if responseType == "callback" { @@ -205,7 +142,7 @@ func HttpCallResultLambda(ctx context.Context, nodeInput *flowDto.NodeExecutionI if responseType == "sync" { httpResultJson := gconv.String(rawHttpResult) // 按 responseMapping 定义的结构,从 http 返回结果中拷贝对应字段 - finalResult = MapResultByTemplate(responseMapping, rawHttpResult) + finalResult = values.MapResultByTemplate(responseMapping, rawHttpResult) e = httpResultJson } if responseType == "callback" { @@ -221,7 +158,7 @@ func HttpCallResultLambda(ctx context.Context, nodeInput *flowDto.NodeExecutionI bodyStr := request.GetBodyString() // 按 responseMapping 定义的结构,从回调结果中拷贝对应字段 - finalResult = MapResultByTemplate(responseMapping, gconv.Map(bodyStr)) + finalResult = values.MapResultByTemplate(responseMapping, gconv.Map(bodyStr)) e = bodyStr } if responseType == "pull" { @@ -320,161 +257,3 @@ func prependFilePathPrefix(prefix string, v any) any { return val } } - -// punctRe 切分/剥离用的中文标点(含顿号、) -var punctRe = regexp.MustCompile(`[,。;!?、]`) - -// BuildSubtitles 核心工具:单个sentence生成多条subtitle -func BuildSubtitles(sents *[]flowDto.Sentence) ([]flowDto.Subtitle, error) { - var subtitles []flowDto.Subtitle - - for _, sent := range *sents { - // 1. 先按标点把文本拆成多个片段 - segList := splitTextByPunct(sent.Text) - if len(segList) == 0 { - continue - } - - // 去标点后得到纯净片段(纯空白/纯标点片段跳过) - var cleans []string - for _, seg := range segList { - c := strings.TrimSpace(cleanPunct(seg)) - if c != "" { - cleans = append(cleans, c) - } - } - if len(cleans) == 0 || len(sent.Words) == 0 { - continue - } - - // 2. 词级文本与句子文本一致时,按词精确对齐取首尾词时间(最准) - if spans, ok := alignAllSegments(sent.Words, cleans); ok { - for i, span := range spans { - subtitles = append(subtitles, flowDto.Subtitle{ - Start: sent.Words[span[0]].StartTime, - End: sent.Words[span[1]].EndTime, - Text: cleans[i], - }) - } - continue - } - - // 3. ASR 词级转写与句子文本不一致时(如 血→谑、数字写法不一), - // 整句回退为按片段字符占比分配时间,避免整句被吞成一条字幕 - segWords := allocWordsByProportion(sent.Words, cleans) - for i, ws := range segWords { - if len(ws) == 0 { - continue - } - subtitles = append(subtitles, flowDto.Subtitle{ - Start: ws[0].StartTime, - End: ws[len(ws)-1].EndTime, - Text: cleans[i], - }) - } - } - - return subtitles, nil -} - -// splitTextByPunct 按中文标点分割句子,同时保留标点在分段内 -// 例如:"这个叫高血压调理方,注意是根源调理不是临时缓解," -// 会变成:["这个叫高血压调理方,", "注意是根源调理不是临时缓解,"] -func splitTextByPunct(raw string) []string { - // 匹配中文标点并保留在文本中,按标点位置切分 - indexes := punctRe.FindAllStringIndex(raw, -1) - if len(indexes) == 0 { - return []string{raw} - } - - var res []string - prev := 0 - for _, idx := range indexes { - end := idx[1] // 标点的结束位置 - seg := raw[prev:end] - res = append(res, seg) - prev = end - } - // 处理最后一段没有标点的文本 - if prev < len(raw) { - res = append(res, raw[prev:]) - } - return res -} - -// cleanPunct 去掉中文标点,得到纯净文本 -func cleanPunct(raw string) string { - return punctRe.ReplaceAllString(raw, "") -} - -// alignAllSegments 按顺序把各纯净片段与词级文本逐字符对齐(允许个别字符不一致)。 -// 全部片段对齐成功且词被完整覆盖时返回各片段对应的词区间,否则 ok=false, -// 由调用方回退到时间占比分配。 -func alignAllSegments(words []flowDto.Word, cleans []string) ([][2]int, bool) { - spans := make([][2]int, len(cleans)) - wordIdx := 0 - for i, seg := range cleans { - start := wordIdx - segRunes := []rune(seg) - s := 0 - for wordIdx < len(words) && s < len(segRunes) { - for _, r := range []rune(words[wordIdx].Word) { - if s < len(segRunes) && r == segRunes[s] { - s++ - } - } - wordIdx++ - } - // 片段文本没被完整匹配,或该片段没吃到任何词 → 无法精确对齐 - if s < len(segRunes) || start == wordIdx { - return nil, false - } - spans[i] = [2]int{start, wordIdx - 1} - } - // 有剩余词未被任何片段覆盖,说明对齐失败,避免吞掉剩余时间 - if wordIdx < len(words) { - return nil, false - } - return spans, true -} - -// allocWordsByProportion 按纯净片段字符占比把整句时间区间切成段,再按时间中点把 -// 每个 word 归属到所属片段(对词级转写与句子文本不一致的情况兜底)。 -func allocWordsByProportion(words []flowDto.Word, cleans []string) [][]flowDto.Word { - runes := make([]int, len(cleans)) - totalChars := 0 - for i, c := range cleans { - runes[i] = utf8.RuneCountInString(c) - totalChars += runes[i] - } - - sentStart := words[0].StartTime - sentEnd := words[len(words)-1].EndTime - duration := sentEnd - sentStart - if duration < 0 { - duration = 0 - } - - bounds := make([]float64, len(cleans)+1) - bounds[0] = sentStart - accum := 0.0 - for i := range cleans { - if totalChars > 0 { - accum += float64(runes[i]) / float64(totalChars) - } - bounds[i+1] = sentStart + accum*duration - } - - segWords := make([][]flowDto.Word, len(cleans)) - for _, w := range words { - mid := (w.StartTime + w.EndTime) / 2 - idx := 0 - for b := 0; b < len(bounds)-1; b++ { - if mid >= bounds[b+1] { - idx = b + 1 - } - } - segWords[idx] = append(segWords[idx], w) - } - return segWords -} diff --git a/workflow/service/flow/processor/builtin/media/media.go b/workflow/service/flow/processor/builtin/media/media.go index f54e1be..6981053 100644 --- a/workflow/service/flow/processor/builtin/media/media.go +++ b/workflow/service/flow/processor/builtin/media/media.go @@ -9,6 +9,7 @@ import ( commonHttp "gitea.redpowerfuture.com/red-future/common/http" "gitea.redpowerfuture.com/red-future/common/oss" + "gitea.redpowerfuture.com/red-future/common/utils" "github.com/gogf/gf/v2/frame/g" "github.com/gogf/gf/v2/util/gconv" ) @@ -158,7 +159,7 @@ func SubmitMergeAsync(ctx context.Context, videoURLs, audioURLs []string, upload func submitMediaTask(ctx context.Context, kind TaskKind, req *mergeSubmitReq) (string, error) { path := "media/video/" + string(kind) + "/async" res := new(mergeSubmitRes) - if err := commonHttp.Post(ctx, path, requestHeaders(ctx), res, req); err != nil { + if err := commonHttp.Post(ctx, path, utils.HeadersFromCtx(ctx), res, req); err != nil { return "", fmt.Errorf("提交%s任务失败: %v", kind, err) } if res.TaskID == "" { @@ -199,31 +200,12 @@ func WaitMediaTask(ctx context.Context, kind TaskKind, taskID string, timeout ti func getMediaTask(ctx context.Context, kind TaskKind, taskID string) (*MergeTask, error) { path := "media/video/" + string(kind) + "/task/" + taskID res := new(MergeTask) - if err := commonHttp.Get(ctx, path, requestHeaders(ctx), res); err != nil { + if err := commonHttp.Get(ctx, path, utils.HeadersFromCtx(ctx), res); err != nil { return nil, err } return res, nil } -// requestHeaders 透传当前请求头(含 Authorization / X-User-Info),供内部服务鉴权使用; -// 后台恢复续跑无 HTTP 请求时(ctx 携带合成 user),补充 X-User-Info 供下游 GetUserInfo 识别租户 -func requestHeaders(ctx context.Context) map[string]string { - headers := make(map[string]string) - if r := g.RequestFromCtx(ctx); r != nil { - for k, v := range r.Request.Header { - if len(v) > 0 { - headers[k] = v[0] - } - } - } - if headers["X-User-Info"] == "" { - if u := ctx.Value("user"); u != nil { - headers["X-User-Info"] = gconv.String(u) - } - } - return headers -} - // segmentResult 一段视频生成的产出(原 ai-agent/video/plan.SegmentResult,plan 包已并入处理器树)。 type segmentResult struct { SegmentIndex int `json:"segment_index"` diff --git a/workflow/service/flow/react_ws_exec.go b/workflow/service/flow/react_ws_exec.go index a1ba55d..67fa326 100644 --- a/workflow/service/flow/react_ws_exec.go +++ b/workflow/service/flow/react_ws_exec.go @@ -21,7 +21,7 @@ import ( ) func init() { - // 普通对话消息处理器(会话服务器 SessionWsService 见 ws_server.go) + // 普通对话消息处理器(会话服务器 SessionWsService 在 exec_ws.go 定义) SessionWsService.OnMessage("agent", handleToolAgent) SessionWsService.OnMessage("agent_cancel", handleToolAgentCancel) } diff --git a/workflow/service/flow/subtitle.go b/workflow/service/flow/subtitle.go new file mode 100644 index 0000000..8f8b9f4 --- /dev/null +++ b/workflow/service/flow/subtitle.go @@ -0,0 +1,166 @@ +package flow + +import ( + flowDto "ai-agent/workflow/model/dto/flow" + "regexp" + "strings" + "unicode/utf8" +) + +// punctRe 切分/剥离用的中文标点(含顿号、) +var punctRe = regexp.MustCompile(`[,。;!?、]`) + +// BuildSubtitles 核心工具:单个sentence生成多条subtitle +func BuildSubtitles(sents *[]flowDto.Sentence) ([]flowDto.Subtitle, error) { + var subtitles []flowDto.Subtitle + + for _, sent := range *sents { + // 1. 先按标点把文本拆成多个片段 + segList := splitTextByPunct(sent.Text) + if len(segList) == 0 { + continue + } + + // 去标点后得到纯净片段(纯空白/纯标点片段跳过) + var cleans []string + for _, seg := range segList { + c := strings.TrimSpace(cleanPunct(seg)) + if c != "" { + cleans = append(cleans, c) + } + } + if len(cleans) == 0 || len(sent.Words) == 0 { + continue + } + + // 2. 词级文本与句子文本一致时,按词精确对齐取首尾词时间(最准) + if spans, ok := alignAllSegments(sent.Words, cleans); ok { + for i, span := range spans { + subtitles = append(subtitles, flowDto.Subtitle{ + Start: sent.Words[span[0]].StartTime, + End: sent.Words[span[1]].EndTime, + Text: cleans[i], + }) + } + continue + } + + // 3. ASR 词级转写与句子文本不一致时(如 血→谑、数字写法不一), + // 整句回退为按片段字符占比分配时间,避免整句被吞成一条字幕 + segWords := allocWordsByProportion(sent.Words, cleans) + for i, ws := range segWords { + if len(ws) == 0 { + continue + } + subtitles = append(subtitles, flowDto.Subtitle{ + Start: ws[0].StartTime, + End: ws[len(ws)-1].EndTime, + Text: cleans[i], + }) + } + } + + return subtitles, nil +} + +// splitTextByPunct 按中文标点分割句子,同时保留标点在分段内 +// 例如:"这个叫高血压调理方,注意是根源调理不是临时缓解," +// 会变成:["这个叫高血压调理方,", "注意是根源调理不是临时缓解,"] +func splitTextByPunct(raw string) []string { + // 匹配中文标点并保留在文本中,按标点位置切分 + indexes := punctRe.FindAllStringIndex(raw, -1) + if len(indexes) == 0 { + return []string{raw} + } + + var res []string + prev := 0 + for _, idx := range indexes { + end := idx[1] // 标点的结束位置 + seg := raw[prev:end] + res = append(res, seg) + prev = end + } + // 处理最后一段没有标点的文本 + if prev < len(raw) { + res = append(res, raw[prev:]) + } + return res +} + +// cleanPunct 去掉中文标点,得到纯净文本 +func cleanPunct(raw string) string { + return punctRe.ReplaceAllString(raw, "") +} + +// alignAllSegments 按顺序把各纯净片段与词级文本逐字符对齐(允许个别字符不一致)。 +// 全部片段对齐成功且词被完整覆盖时返回各片段对应的词区间,否则 ok=false, +// 由调用方回退到时间占比分配。 +func alignAllSegments(words []flowDto.Word, cleans []string) ([][2]int, bool) { + spans := make([][2]int, len(cleans)) + wordIdx := 0 + for i, seg := range cleans { + start := wordIdx + segRunes := []rune(seg) + s := 0 + for wordIdx < len(words) && s < len(segRunes) { + for _, r := range []rune(words[wordIdx].Word) { + if s < len(segRunes) && r == segRunes[s] { + s++ + } + } + wordIdx++ + } + // 片段文本没被完整匹配,或该片段没吃到任何词 → 无法精确对齐 + if s < len(segRunes) || start == wordIdx { + return nil, false + } + spans[i] = [2]int{start, wordIdx - 1} + } + // 有剩余词未被任何片段覆盖,说明对齐失败,避免吞掉剩余时间 + if wordIdx < len(words) { + return nil, false + } + return spans, true +} + +// allocWordsByProportion 按纯净片段字符占比把整句时间区间切成段,再按时间中点把 +// 每个 word 归属到所属片段(对词级转写与句子文本不一致的情况兜底)。 +func allocWordsByProportion(words []flowDto.Word, cleans []string) [][]flowDto.Word { + runes := make([]int, len(cleans)) + totalChars := 0 + for i, c := range cleans { + runes[i] = utf8.RuneCountInString(c) + totalChars += runes[i] + } + + sentStart := words[0].StartTime + sentEnd := words[len(words)-1].EndTime + duration := sentEnd - sentStart + if duration < 0 { + duration = 0 + } + + bounds := make([]float64, len(cleans)+1) + bounds[0] = sentStart + accum := 0.0 + for i := range cleans { + if totalChars > 0 { + accum += float64(runes[i]) / float64(totalChars) + } + bounds[i+1] = sentStart + accum*duration + } + + segWords := make([][]flowDto.Word, len(cleans)) + for _, w := range words { + mid := (w.StartTime + w.EndTime) / 2 + idx := 0 + for b := 0; b < len(bounds)-1; b++ { + if mid >= bounds[b+1] { + idx = b + 1 + } + } + segWords[idx] = append(segWords[idx], w) + } + return segWords +} diff --git a/workflow/service/flow/values/value_request.go b/workflow/service/flow/values/value_request.go new file mode 100644 index 0000000..4b5df2c --- /dev/null +++ b/workflow/service/flow/values/value_request.go @@ -0,0 +1,140 @@ +package values + +import ( + "ai-agent/workflow/consts/node" + flowDto "ai-agent/workflow/model/dto/flow" + "ai-agent/workflow/model/entity" + "context" + "strings" + + "github.com/gogf/gf/v2/frame/g" + "github.com/gogf/gf/v2/os/glog" + "github.com/gogf/gf/v2/util/gconv" +) + +// BuildModelRequestBody 从参数定义 + 全局执行上下文构建最终嵌套 JSON 请求体。 +func BuildModelRequestBody(params []entity.FlowModelParams, globalParams *flowDto.FlowExecutionInput) (map[string]interface{}, error) { + // 1. 解析引用、过滤空值 + resolved := parseMap(params, globalParams) + + // 2. 转扁平路径映射 + flat := toFlatMap(resolved) + + // 3. 引用了脚本转写节点的字段打内部标记 __segment_fields(逗号分隔的扁平路径), + // 供前置处理器 split_segment 按段拆批;__ 前缀内部键由 invokePreTool 统一剥离,不传给模型网关 + if seg := segmentFields(resolved, globalParams); len(seg) > 0 { + flat["__segment_fields"] = strings.Join(seg, ",") + } + + return flat, nil +} + +// segmentFields 收集引用了脚本转写节点的字段扁平路径(分段字段)。 +// 脚本转写节点按段产出一份扁平参数列表,下游模型节点单源引用其字段时, +// 值按段序聚合成数组(见 ResolveValueSource 的 scriptTranscribe 分支),需随批拆分。 +func segmentFields(resolved []entity.FlowModelParams, globalParams *flowDto.FlowExecutionInput) []string { + if globalParams == nil || globalParams.ConfigMap == nil { + return nil + } + var fields []string + for _, p := range resolved { + if g.IsEmpty(p.Path) || len(p.ValueSource) != 1 { + continue + } + src := p.ValueSource[0] + if nodeConfig := globalParams.ConfigMap[src.NodeId]; nodeConfig != nil && + nodeConfig.NodeCode == node.NodeTypeScriptTranscribe { + fields = append(fields, flatPath(p.Path)) + } + } + return fields +} + +// parseMap 解析模型请求参数 +func parseMap(data []entity.FlowModelParams, globalParams *flowDto.FlowExecutionInput) []entity.FlowModelParams { + newData := make([]entity.FlowModelParams, 0, len(data)) + for _, item := range data { + var d entity.FlowModelParams + d.Path = item.Path + d.Type = item.Type + + // 无引用源:直接取静态值 + if g.IsEmpty(item.ValueSource) { + if isParamEmpty(item.Value) { + continue + } + d = item + newData = append(newData, d) + continue + } + + // 有引用源:单源解析取非空值(模型解析引用其他节点的数组值不做类型化转换,整体传给模型); + // 多源把各源解析出的值拼成 "label: value"(无 label 只拼值),逗号分隔 + var value, refsName any + if len(item.ValueSource) > 1 { + text, rn, ok := joinValueSources(globalParams, item.ValueSource, isParamEmpty, nil) + if !ok { + parseMapLogFail(item) + continue + } + value, refsName = text, rn + } else { + src := item.ValueSource[0] + var ok bool + value, refsName, ok = ResolveValueSource(globalParams, src.NodeId, src.Field) + if !ok || isParamEmpty(value) { + parseMapLogFail(item) + continue + } + } + d.Value = value + d.RefsName = gconv.String(refsName) + d.ValueSource = item.ValueSource + newData = append(newData, d) + } + return newData +} + +// parseMapLogFail 引用解析失败留日志(单源/多源共用,不静默,便于排查引用丢失) +func parseMapLogFail(item entity.FlowModelParams) { + glog.Debugf(context.Background(), "resolve value source failed, nodeId=%+v path=%s", item.ValueSource, item.Path) +} + +// toFlatMap 将解析后的参数列表转为 sjson 可用的扁平路径映射。 +// key 为 Path,value 为参数值;保留 RefsName 供上层追踪引用来源。 +func toFlatMap(params []entity.FlowModelParams) map[string]interface{} { + m := make(map[string]interface{}, len(params)) + for _, p := range params { + if g.IsEmpty(p.Path) { + continue + } + m[flatPath(p.Path)] = p.Value + } + return m +} + +// flatPath 把数组下标路径转扁平点分路径:a[0].b → a.0.b。 +func flatPath(path string) string { + return arrayIndexPath.ReplaceAllString(path, `.$1`) +} + +// isParamEmpty 判断参数值是否为"空"。 +// 仅 nil、空字符串、空切片/映射视为空;0、false 等零值是合法值,保留。 +func isParamEmpty(v interface{}) bool { + if v == nil { + return true + } + switch val := v.(type) { + case string: + return val == "" + case []byte: + return len(val) == 0 + case []interface{}: + return len(val) == 0 + case map[string]interface{}: + return len(val) == 0 + default: + // 数字、布尔、结构体等一律视为非空 + return false + } +} diff --git a/workflow/service/flow/values/value_resolve.go b/workflow/service/flow/values/value_resolve.go new file mode 100644 index 0000000..a32a509 --- /dev/null +++ b/workflow/service/flow/values/value_resolve.go @@ -0,0 +1,254 @@ +package values + +import ( + "ai-agent/workflow/consts/node" + flowDto "ai-agent/workflow/model/dto/flow" + "ai-agent/workflow/model/entity" + "encoding/json" + "regexp" + "strconv" + "strings" + + "github.com/gogf/gf/v2/frame/g" + "github.com/gogf/gf/v2/util/gconv" + "github.com/tidwall/gjson" +) + +var ( + // 匹配 [数字] + regNumIndex = regexp.MustCompile(`\[\d+\]`) + // 匹配 .attrs + regAttrs = regexp.MustCompile(`\.attrs`) + // 匹配带捕获组的数组下标,转扁平点分路径用 + arrayIndexPath = regexp.MustCompile(`\[(\d+)\]`) +) + +// CleanFieldPath 清理字段路径:移除 .attrs、数字下标转为 .#(gjson 数组通配符) +// 示例:usage.attrs.total_tokens → usage.total_tokens +// 示例:choices.attrs[0].attrs.message.attrs.content → choices.#.message.content +func CleanFieldPath(path string) string { + index := CleanFieldPathReplaceNumIndex(path) + attrs := CleanFieldPathRemoveAttrs(index) + return attrs +} + +func CleanFieldPathReplaceNumIndex(path string) string { + // 1. 替换 [数字] 为 [*] + s := regNumIndex.ReplaceAllString(path, `.#`) + return s +} + +func CleanFieldPathRemoveAttrs(path string) string { + // 2. 移除所有 .attrs + s := regAttrs.ReplaceAllString(path, "") + return s +} + +// ProcessValueSourceRecursive 递归遍历map,同级同时存在value和valueSource则把value设置为"AA" +func ProcessValueSourceRecursive(rawParams map[string]interface{}, globalParams *flowDto.FlowExecutionInput) { + walkMap(rawParams, globalParams) +} + +// ResolveValueSource 解析 valueSource {nodeId, field} 引用的实际值。 +// 返回 (value, refsName, ok);ok=false 表示引用节点不存在或引用值仍为空。 +// - 开始/表单节点:OutputConfig 平铺条目按 field == 引用字段匹配(前端约定以 field 为主, +// 不兼容 path),直接读 entry 的 value / refsName +// - scriptTranscribe 节点:OutputResult 是各段扁平请求参数,按段序收集字段为数组(段位留 nil) +// - 其他节点:读 OutputResult 中引用字段 field 路径对应的值 +func ResolveValueSource(global *flowDto.FlowExecutionInput, nodeId, field string) (value any, refsName any, ok bool) { + if global == nil || global.ConfigMap == nil { + return nil, nil, false + } + nodeConfig := global.ConfigMap[nodeId] + if nodeConfig == nil { + return nil, nil, false + } + switch nodeConfig.NodeCode { + case node.NodeTypeStart, node.NodeTypeForm: + for _, output := range nodeConfig.OutputConfig { + if gconv.String(output["field"]) != field { + continue + } + if !g.IsEmpty(output["value"]) { + return output["value"], output["refsName"], true + } + } + case node.NodeTypeScriptTranscribe: + // 脚本转写节点 OutputResult 是各段扁平请求参数(split_shots_pipeline 产出,key 为字面量 + // prompt/duration/seed 等),按段序读取 output[field] 收集为数组,供分段模型节点整体引用。 + // 每段都占一位(字段缺失/为空留 nil),保证数组与段序对齐,供 split_segment 按段取值。 + var list []any + for _, output := range nodeConfig.OutputResult { + list = append(list, output[field]) + } + for _, v := range list { + if !g.IsEmpty(v) { + return list, "", true + } + } + default: + // templates 是模型节点在前端配置的静态输出模板,不在 OutputResult 中,需单独取 + if field == "templates" { + if !g.IsEmpty(nodeConfig.Templates) { + return nodeConfig.Templates, "", true + } + return nil, nil, false + } + for _, output := range nodeConfig.OutputResult { + // 模型节点输出记录是单 key 的字面量扁平 key(如 "choices.attrs[0].attrs.delta.attrs.content"), + // gjson 会把 . 和 [0] 当结构路径解析,无法命中字面量 key,故先按字面量 key 直接取值; + // 未命中再回退 gjson 路径查询(兼容真正嵌套的输出结构)。 + if v, has := output[field]; has { + value = v + } else { + value = gjson.Get(gconv.String(output), field).Value() + } + if !g.IsEmpty(value) { + return value, gjson.Get(gconv.String(output), CleanFieldPath("refsName")).Value(), true + } + } + } + return nil, nil, false +} + +// walkMap 递归处理map/数组 +func walkMap(data interface{}, globalParams *flowDto.FlowExecutionInput) { + switch v := data.(type) { + case map[string]interface{}: + // 有 valueSource:解析引用节点值 + if valueSource, hasSource := v["valueSource"]; hasSource { + sources := new([]entity.ValueSource) + gconv.Structs(valueSource, sources) + + // 多个引用源:把各源解析出的值拼成 "label: value"(无 label 只拼值),逗号分隔 + if len(*sources) > 1 { + text, refsName, ok := joinValueSources(globalParams, *sources, schemaValueEmpty, + func(src entity.ValueSource, value any) any { + // 引用非模型节点(开始/表单/HTTP/脚本转写等)时,值按当前字段声明的 type 做类型化转换; + // 模型节点值由模型网关处理,复制时不需要转换 + if !isModelSourceNode(globalParams, src.NodeId) { + return assignBySchemaType(v, value) + } + return value + }) + if ok { + v["value"] = text + if !g.IsEmpty(refsName) { + v["refsName"] = refsName + } + return + } + } else if len(*sources) == 1 { + // 单个引用源:保持旧行为,值按原样赋值(非模型节点按声明 type 转换) + src := (*sources)[0] + value, refsName, ok := ResolveValueSource(globalParams, src.NodeId, src.Field) + if ok && !isModelSourceNode(globalParams, src.NodeId) { + value = assignBySchemaType(v, value) + } + if ok && !schemaValueEmpty(value) { + v["value"] = value + if !g.IsEmpty(refsName) { + v["refsName"] = refsName + } + return + } + } + } + // 统一兜底:无 valueSource(或解析失败/值为空)时,value 为空或 0 则取 defaultValue + if defaultValue, hasDefault := v["defaultValue"]; hasDefault && isEmptyForFallback(v["value"]) && !schemaValueEmpty(defaultValue) { + v["value"] = defaultValue + } + // 递归遍历所有子元素 + for _, child := range v { + walkMap(child, globalParams) + } + case []interface{}: + // 数组遍历 + for _, item := range v { + walkMap(item, globalParams) + } + } +} + +// isEmptyForFallback 兜底场景判空:除 schemaValueEmpty 规则外,数字 0 也视为未填写, +// 便于配置了 defaultValue 的字段在值为 0 时用默认值兜底。 +// 覆盖 json.Number(gconv 反序列化数字的运行时类型)与字符串 "0"/"0.0"。 +func isEmptyForFallback(v interface{}) bool { + if schemaValueEmpty(v) { + return true + } + switch val := v.(type) { + case float32: + return val == 0 + case float64: + return val == 0 + case int: + return val == 0 + case int8: + return val == 0 + case int16: + return val == 0 + case int32: + return val == 0 + case int64: + return val == 0 + case uint: + return val == 0 + case uint8: + return val == 0 + case uint16: + return val == 0 + case uint32: + return val == 0 + case uint64: + return val == 0 + case json.Number: + if f, err := val.Float64(); err == nil { + return f == 0 + } + case string: + if f, err := strconv.ParseFloat(val, 64); err == nil { + return f == 0 + } + } + return false +} + +// isModelSourceNode 判断引用源节点是否为模型节点(值由模型网关处理,复制时不转换) +func isModelSourceNode(global *flowDto.FlowExecutionInput, nodeId string) bool { + if global == nil || global.ConfigMap == nil { + return false + } + nodeConfig := global.ConfigMap[nodeId] + return nodeConfig != nil && nodeConfig.NodeCode == node.NodeTypeModel +} + +// joinValueSources 拼接多个 valueSource 的解析值为 "label: value"(无 label 只拼值),逗号分隔。 +// 返回拼接文本与第一个非空 refsName;所有源都为空/解析失败时 ok=false。 +// isEmpty 为各场景的空值判断(walkMap 用 schemaValueEmpty,模型请求解析用 isParamEmpty); +// transform 对每个非空解析值做转换(walkMap 按声明 type 类型化、模型源不转换;模型请求场景传 nil)。 +func joinValueSources(global *flowDto.FlowExecutionInput, sources []entity.ValueSource, + isEmpty func(any) bool, transform func(src entity.ValueSource, value any) any) (text string, refsName any, ok bool) { + var parts []string + for _, src := range sources { + value, rn, ok := ResolveValueSource(global, src.NodeId, src.Field) + if !ok || isEmpty(value) { + continue + } + if transform != nil { + value = transform(src, value) + } + s := toPlainString(value) + if src.Label != "" { + s = src.Label + ": " + s + } + parts = append(parts, s) + if !g.IsEmpty(rn) && g.IsEmpty(refsName) { + refsName = rn + } + } + if len(parts) == 0 { + return "", nil, false + } + return strings.Join(parts, ", "), refsName, true +} diff --git a/workflow/service/flow/values/value_schema.go b/workflow/service/flow/values/value_schema.go new file mode 100644 index 0000000..3c21bfd --- /dev/null +++ b/workflow/service/flow/values/value_schema.go @@ -0,0 +1,227 @@ +package values + +import ( + "encoding/json" + "reflect" + "strconv" + "strings" + + "github.com/gogf/gf/v2/util/gconv" +) + +// UnwrapSchemaWrapper 递归剥掉 json-schema-editor 输出的 {type, value/attrs} 包裹层, +// 只保留干净的 key/value 嵌套结构。 +// 示例: +// +// {"a": {"type":"string","value":"hi"}} → {"a": "hi"} +// {"b": {"type":"object","attrs":{"c":1}}} → {"b": {"c": 1}} +// {"arr": {"type":"array","attrs":[{"type":"number","value":1}]}} → {"arr": [1]} +func UnwrapSchemaWrapper(v any) any { + switch val := v.(type) { + case map[string]any: + // 识别包裹节点:{type: "", value/attrs: <实际值>, ...} + if t, ok := val["type"].(string); ok && isSchemaEditorType(t) { + dataKey := "value" + if t == "object" || t == "array" { + dataKey = "attrs" + } + if raw, has := val[dataKey]; has { + return UnwrapSchemaWrapper(raw) + } + } + res := make(map[string]any, len(val)) + for k, child := range val { + res[k] = UnwrapSchemaWrapper(child) + } + return res + case []any: + res := make([]any, len(val)) + for i, item := range val { + res[i] = UnwrapSchemaWrapper(item) + } + return res + default: + return val + } +} + +// isSchemaEditorType 是否为 json-schema-editor 的 6 种类型标识 +func isSchemaEditorType(t string) bool { + switch t { + case "string", "number", "boolean", "null", "object", "array": + return true + } + return false +} + +// MapResultByTemplate 按 template 定义的结构,从 source 中拷贝对应字段的值。 +// 只保留 template 里出现的字段:对象字段按同名字段递归拷贝,数组字段按模板元素结构逐元素过滤,标量字段直接拷贝 source 的值。 +func MapResultByTemplate(template map[string]any, source map[string]any) map[string]any { + result := make(map[string]any, len(template)) + for key, tmplVal := range template { + srcVal, ok := source[key] + if !ok { + continue + } + if tmplMap, isMap := tmplVal.(map[string]any); isMap { + if srcMap, isMap := srcVal.(map[string]any); isMap { + result[key] = MapResultByTemplate(tmplMap, srcMap) + } + continue + } + if tmplArr, isArr := tmplVal.([]any); isArr { + result[key] = mapTemplateArray(tmplArr, srcVal) + continue + } + result[key] = srcVal + } + return result +} + +// mapTemplateArray 按模板数组的元素结构映射 source 数组: +// 模板首元素为对象时,逐元素按 MapResultByTemplate 过滤只保留模板字段; +// 模板数组为空或首元素非对象(无法确定元素结构)时,原样拷贝 source 数组。 +func mapTemplateArray(tmplArr []any, srcVal any) any { + srcList, ok := srcVal.([]any) + if !ok || len(tmplArr) == 0 { + return srcVal + } + elemTmpl, ok := tmplArr[0].(map[string]any) + if !ok { + return srcVal + } + result := make([]any, 0, len(srcList)) + for _, srcElem := range srcList { + if srcMap, isMap := srcElem.(map[string]any); isMap { + result = append(result, MapResultByTemplate(elemTmpl, srcMap)) + } else { + result = append(result, srcElem) + } + } + return result +} + +// assignBySchemaType 按当前字段声明的 schema 类型把值类型化: +// string 遇数组/对象转 JSON 字符串;number/boolean 解析字符串;object/array 解析 JSON 字符串;其余原样返回 +func assignBySchemaType(node map[string]interface{}, value any) any { + t, _ := node["type"].(string) + return assignByType(t, value) +} + +// assignByType 按字段声明的 type 把值类型化;walkMap 的 schema 节点与 parseMap 的模型参数共用 +func assignByType(t string, value any) any { + switch t { + case "string": + return toSchemaString(value) + case "number": + return toSchemaNumber(value) + case "boolean": + return toSchemaBool(value) + case "object", "array": + return toSchemaStruct(value) + default: + return value + } +} + +// toSchemaString 转 string:字符串原样,数组元素拼成字符串(单元素取元素本身,多元素逗号连接),对象序列化为 JSON 字符串 +func toSchemaString(v any) any { + switch val := v.(type) { + case []interface{}: + parts := make([]string, 0, len(val)) + for _, item := range val { + parts = append(parts, toPlainString(item)) + } + return strings.Join(parts, ",") + case map[string]interface{}: + if b, err := json.Marshal(val); err == nil { + return string(b) + } + } + return v +} + +// toPlainString 把数组元素转成不带括号的纯字符串: +// 数组([]any / 类型化切片)逐元素取纯字符串,单元素取元素本身,多元素逗号连接; +// 对象序列化为 JSON 字符串;其余原样字符串化。 +func toPlainString(v any) string { + if s, ok := v.(string); ok { + return s + } + switch val := v.(type) { + case []interface{}: + parts := make([]string, 0, len(val)) + for _, item := range val { + parts = append(parts, toPlainString(item)) + } + return strings.Join(parts, ",") + case map[string]interface{}: + if b, err := json.Marshal(val); err == nil { + return string(b) + } + } + rv := reflect.ValueOf(v) + if rv.IsValid() && (rv.Kind() == reflect.Slice || rv.Kind() == reflect.Array) { + parts := make([]string, 0, rv.Len()) + for i := 0; i < rv.Len(); i++ { + parts = append(parts, toPlainString(rv.Index(i).Interface())) + } + return strings.Join(parts, ",") + } + if b, err := json.Marshal(v); err == nil { + return string(b) + } + return gconv.String(v) +} + +// toSchemaNumber 转 number:数字原样,字符串尝试解析为 float64,失败原样返回 +func toSchemaNumber(v any) any { + if s, ok := v.(string); ok { + if f, err := strconv.ParseFloat(s, 64); err == nil { + return f + } + } + return v +} + +// toSchemaBool 转 boolean:布尔原样,字符串尝试解析为 bool,失败原样返回 +func toSchemaBool(v any) any { + if s, ok := v.(string); ok { + if b, err := strconv.ParseBool(s); err == nil { + return b + } + } + return v +} + +// toSchemaStruct 转 object/array:合法 JSON 字符串解析为结构化数据,否则原样返回 +func toSchemaStruct(v any) any { + s, ok := v.(string) + if !ok { + return v + } + if !json.Valid([]byte(s)) { + return v + } + var out any + if err := json.Unmarshal([]byte(s), &out); err != nil { + return v + } + return out +} + +// schemaValueEmpty 值是否为空;0/false 视为有效值不剔除 +func schemaValueEmpty(v interface{}) bool { + switch val := v.(type) { + case nil: + return true + case string: + return val == "" + case []interface{}: + return len(val) == 0 + case map[string]interface{}: + return len(val) == 0 + default: + return false + } +} diff --git a/workflow/service/flow/ws_server.go b/workflow/service/flow/ws_server.go deleted file mode 100644 index 0c22140..0000000 --- a/workflow/service/flow/ws_server.go +++ /dev/null @@ -1,41 +0,0 @@ -package flow - -import ( - "context" - - sessionDao "ai-agent/workflow/dao/session" - sessionDto "ai-agent/workflow/model/dto/session" - - wsCommon "gitea.redpowerfuture.com/red-future/common/websocket" - "github.com/gogf/gf/v2/net/ghttp" -) - -// SessionWsService 会话 WebSocket 服务器:普通对话与工作流共用一条连接, -// 首次连接仅升级,后续按消息 type 路由到对话/工作流处理器 -// (对话处理器在 react_ws_exec.go 注册,工作流处理器在 flow_ws_exec.go 注册)。 -var SessionWsService = wsCommon.NewWsServer( - wsCommon.WithConnKeyPrefix("ws:session:"), -) - -// WsConnect 控制器统一入口:升级 WebSocket(普通对话/工作流均由消息 type 区分,此处不区分) -func WsConnect(ctx context.Context, r *ghttp.Request, req *sessionDto.WebSocketConnectReq) error { - _, err := SessionWsService.Upgrade(ctx, r, req.SessionId) - return err -} - -// ensureSession 解析前端 sessionId 并确保会话存在:命中已存在会话则复用其 id,否则按 name 新建。 -// 普通对话(react_ws_exec.go)与工作流(flow_ws_exec.go)共用。 -func ensureSession(ctx context.Context, sessionId string, name string) error { - exist, err := sessionDao.SessionDao.GetById(ctx, sessionId) - if err != nil { - return err - } - if exist != nil { - return nil - } - if r := []rune(name); len(r) > 128 { // session_name VARCHAR(128) - name = string(r[:128]) - } - _, err = sessionDao.SessionDao.Insert(ctx, &sessionDto.CreateSessionReq{SessionId: sessionId, SessionName: name}) - return err -} diff --git a/workflow/service/session/session_service.go b/workflow/service/session/session_service.go index 8df79ca..9b83079 100644 --- a/workflow/service/session/session_service.go +++ b/workflow/service/session/session_service.go @@ -211,6 +211,7 @@ func workflowExecVO(w *entity.ExecWorkflow, resultFileUrl string) *sessionDto.VO ResultFileUrl: resultFileUrl, TotalTokens: w.TotalTokens, TotalFee: w.TotalFee, + ActualAmount: w.ActualAmount, ErrorMsg: w.ErrorMessage, Error: w.Error, CreatedAt: w.CreatedAt, diff --git a/workflow/service/util_service.go b/workflow/service/util_service.go index 531b514..8824817 100644 --- a/workflow/service/util_service.go +++ b/workflow/service/util_service.go @@ -4,7 +4,7 @@ import ( "context" commonHttp "gitea.redpowerfuture.com/red-future/common/http" - "github.com/gogf/gf/v2/frame/g" + "gitea.redpowerfuture.com/red-future/common/utils" ) var UtilService = &utilService{} @@ -13,16 +13,8 @@ type utilService struct{} // IsAdmin 调用admin-go服务检查是否是管理员 func (s *utilService) IsAdmin(ctx context.Context) (res bool, err error) { - headers := make(map[string]string) - if r := g.RequestFromCtx(ctx); r != nil { - for k, v := range r.Request.Header { - if len(v) > 0 { - headers[k] = v[0] - } - } - } var r = make(map[string]bool) - if err = commonHttp.Get(ctx, "admin-go/api/v1/system/user/checkIsSuperAdmin", headers, &r); err != nil { + if err = commonHttp.Get(ctx, "admin-go/api/v1/system/user/checkIsSuperAdmin", utils.HeadersFromCtx(ctx), &r); err != nil { return false, err } return r["isSuperAdmin"], err