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(`?(div|p|h1|h2|h3|h4|h5|h6|li|ul|ol|br|tr|td|th)[^>]*>`)
- 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