package service import ( "context" "encoding/json" "fmt" "model-gateway/consts/public" "model-gateway/dao" "model-gateway/model/dto" "model-gateway/model/entity" "model-gateway/service/httpclient" modelUtils "model-gateway/service/utils" "regexp" "time" "gitea.redpowerfuture.com/red-future/common/beans" "gitea.redpowerfuture.com/red-future/common/db/gfdb" "gitea.redpowerfuture.com/red-future/common/oss" "gitea.redpowerfuture.com/red-future/common/utils" gmq "github.com/bjang03/gmq/core/gmq" "github.com/bjang03/gmq/mq" "github.com/bjang03/gmq/types" "github.com/gogf/gf/v2/container/gvar" "github.com/gogf/gf/v2/database/gdb" "github.com/gogf/gf/v2/frame/g" "github.com/gogf/gf/v2/util/gconv" ) var ModelTaskEndService = &modelTaskEndService{} type modelTaskEndService struct{} // GetTaskStartList 获取待执行任务 func (s *modelTaskEndService) GetTaskStartList(ctx context.Context) (err error) { workerNum := g.Cfg().MustGet(ctx, "pool.workerNum", modelUtils.DefaultWorkerNum).Int64() redisKey := "model_video_task:" var ( pageNum = gconv.Int64(1) remain = workerNum ) // 字段列表 cols := []string{ entity.ModelTaskStartCol.Id, entity.ModelTaskStartCol.TaskId, entity.ModelTaskStartCol.ModelId, entity.ModelTaskStartCol.BizName, entity.ModelTaskStartCol.Creator, entity.ModelTaskStartCol.TenantId, entity.ModelTaskStartCol.MsgTopic, entity.ModelTaskStartCol.MediaType, } for remain > 0 { req := &dto.GetModelTaskStartListReq{ Page: &beans.Page{ PageNum: pageNum, PageSize: remain, // 每页只查当前需要的数量 }, } var list []entity.ModelTaskStart list, err = dao.ModelTaskStart.ListByLimitNotTenantId(ctx, req, cols...) if err != nil { return fmt.Errorf("查询任务失败: %w", err) } if len(list) == 0 { break } // 3. 组装锁key,批量查询Redis(性能最优) taskMap := make(map[string]*entity.ModelTaskStart, len(list)) lockKeys := make([]string, 0, len(list)) for _, item := range list { key := redisKey + gconv.String(item.Id) taskMap[key] = &item lockKeys = append(lockKeys, key) } var mGetRes map[string]*gvar.Var mGetRes, err = g.Redis().MGet(ctx, lockKeys...) if err != nil { return fmt.Errorf("批量查询锁状态失败: %w", err) } // 4. 提交异步处理:锁在 goroutine 内抢(WithLock 单次尝试),此处 MGet 只做快速预筛 for _, key := range lockKeys { val := gconv.String(mGetRes[key]) // 已被其他实例抢占,跳过(MGet 只做快速预筛;真正互斥靠 goroutine 内的原子抢锁) if val != "" { continue } err = s.handleSingleTask(ctx, taskMap[key], key) if err != nil { g.Log().Errorf(ctx, "提交任务失败: %v", err) } remain-- // 占用一个槽位 if remain <= 0 { break // 槽位已满,终止遍历 } } pageNum++ // 页码动态累加,不再写死2 } return nil } // taskLockTTL 任务锁 TTL(秒)。utils.WithLock 自动续期锁住整个任务处理,TTL 仅作崩溃兜底: // worker 崩溃后续期停止,TTL 过期后其他 worker 重新抢占。 const taskLockTTL = 1200 var urlParamReg = regexp.MustCompile(`\{.+?\}`) // handleSingleTask 提交异步处理:锁在 goroutine 内抢(utils.WithLock 自动续期 + 单次尝试)。 // 自动续期:锁持满整个任务处理,任务 >20min 不提前过期,避免其它 worker 重新抢到导致重复处理; // 单次尝试:锁被其它 worker 持有(任务已被别人处理)时立刻跳过——等待会拿着过期 item 在行删除后 // 重复扣费/重复回调。Submit 失败(池关闭)goroutine 不运行、从没抢锁,无锁泄漏路径。 func (s *modelTaskEndService) handleSingleTask(ctx context.Context, item *entity.ModelTaskStart, lockKey string) error { return modelUtils.Submit(ctx, func(ctx context.Context) { asyncCtx := context.WithoutCancel(ctx) // OSS 桶名依赖 ctx 中的用户(GetBucketName → tenantid-{tenantId}), // 响应临时路径转存 OSS 需要用户信息,故在任务体最前面注入 asyncCtx = context.WithValue(asyncCtx, "user", &beans.User{ UserName: item.Creator, TenantId: item.TenantId, }) ok, err := utils.WithLock(asyncCtx, lockKey, taskLockTTL, func(ctx context.Context) error { return s.processClaimedTask(ctx, item) }, 1) if err != nil || !ok { // 锁被其它实例持有或抢锁失败:跳过,任务行保留由持有方处理,下轮扫描不再命中 g.Log().Warningf(asyncCtx, "任务锁未抢占,跳过 taskId=%d: %v", item.Id, err) } }) } // processClaimedTask 抢到任务锁后的完整处理:轮询模型结果 → 终态落库 + 发布。 // 终态结果统一承载:成功/错误/解析失败任何路径都写进 docMsg.ErrorMsg 后走 finalize 落库+发布, // 避免早期直接 return 把任务丢弃——任务行不删、无结果落库、调用方永远收不到通知,只会被其他 worker 反复重捡。 func (s *modelTaskEndService) processClaimedTask(asyncCtx context.Context, item *entity.ModelTaskStart) error { startTime := time.Now() docMsg := new(dto.ModelMsg) docMsg.TaskID = item.Id var respObj map[string]any // 终态处理:删任务行 → 插结果行(含 ErrorMsg)→ NATS 发布结果给调用方 finalize := func() { err := gfdb.DB(asyncCtx, public.DbNameModelGateway).Transaction(asyncCtx, func(asyncCtx context.Context, tx gdb.TX) (err error) { // 删除视频任务 _, err = dao.ModelTaskStart.Delete(asyncCtx, &dto.DeleteModelTaskStartReq{ Id: item.Id, }) if err != nil { return err } // 保存视频任务结果 _, err = dao.ModelTaskEnd.Insert(asyncCtx, &dto.CreateModelTaskEndReq{ ModelId: item.ModelId, BizName: item.BizName, MsgTopic: item.MsgTopic, TaskId: item.TaskId, ResponseParams: docMsg.Content, OriginalResponseParams: respObj, DurationSeconds: int64(time.Since(startTime).Seconds()), PromptTokens: docMsg.PromptTokens, CompletionTokens: docMsg.CompletionTokens, TotalTokens: docMsg.TotalTokens, TotalCost: docMsg.Cost, ErrorMsg: docMsg.ErrorMsg, }) if err != nil { return err } return }) if err != nil { g.Log().Errorf(asyncCtx, "保存视频任务结果失败: %v", err) } // 发布消息 if err = TaskMsgPublish(asyncCtx, item.MsgTopic, docMsg); err != nil { g.Log().Errorf(asyncCtx, "模型消息发布失败: %v", err) } } // 按 modelId 现查模型配置(异步映射/token 映射/计费规则不随任务快照,任务完成时取当前配置) modelInfo, err := dao.ModelManage.GetNotTenantId(asyncCtx, &dto.GetModelManageReq{Id: item.ModelId}) if err != nil { g.Log().Errorf(asyncCtx, "查询模型配置失败: modelId=%d err=%v", item.ModelId, err) docMsg.ErrorMsg = fmt.Sprintf("查询模型配置失败: %v", err) finalize() return nil } if modelInfo == nil { g.Log().Errorf(asyncCtx, "模型配置不存在: modelId=%d", item.ModelId) docMsg.ErrorMsg = fmt.Sprintf("模型配置不存在: modelId=%d", item.ModelId) finalize() return nil } // 引用行 → 解析为系统模型配置+本人 apiKey(轮询/计价均用系统模型) modelInfo, err = resolveModelConfig(asyncCtx, modelInfo) if err != nil { g.Log().Errorf(asyncCtx, "模型配置解析失败: modelId=%d err=%v", item.ModelId, err) docMsg.ErrorMsg = fmt.Sprintf("模型配置解析失败: %v", err) finalize() return nil } // 连续轮询失败上限:瞬时抖动(HTTP 错/空响应/解析失败)先有限重试,超限按终态错误落库 const maxPollErrRetries = 3 pollErrCnt := 0 LOOP: // 替换URL占位符 url := urlParamReg.ReplaceAllString(modelInfo.AsyncTaskMapping.Url, item.TaskId) // 组装查询请求体:POST 查询接口需要 body(从 RequestBodyMapping 出发,替换 {…} 占位符为任务 ID) reqBody := buildAsyncTaskBody(modelInfo.AsyncTaskMapping.RequestBodyMapping, item.TaskId) // 发起HTTP请求 modelRespBody, err := httpclient.ModelHttpNormalRequest( asyncCtx, url, modelInfo.AsyncTaskMapping.RequestHeadMapping, modelInfo.AsyncTaskMapping.HttpMethod, reqBody, ) if err != nil { g.Log().Errorf(asyncCtx, "模型请求失败: %v", err) if pollErrCnt < maxPollErrRetries { pollErrCnt++ time.Sleep(10 * time.Second) goto LOOP } docMsg.ErrorMsg = fmt.Sprintf("模型请求失败: %v", err) finalize() return nil } if modelRespBody == nil { g.Log().Errorf(asyncCtx, "模型返回参数为空") if pollErrCnt < maxPollErrRetries { pollErrCnt++ time.Sleep(10 * time.Second) goto LOOP } docMsg.ErrorMsg = "模型返回参数为空" finalize() return nil } pollErrCnt = 0 // 请求成功一次即重置连续失败计数 // 异常响应识别:兼容 OpenAI 嵌套 error / 扁平 code 两种形态(与任务创建端一致),无错误返回空串 if _, docMsg.ErrorMsg, err = parseModelError(modelRespBody); err != nil { g.Log().Errorf(asyncCtx, "模型返回参数解析失败:%v", err) if pollErrCnt < maxPollErrRetries { pollErrCnt++ time.Sleep(10 * time.Second) goto LOOP } docMsg.ErrorMsg = fmt.Sprintf("模型返回参数解析失败: %v", err) finalize() return nil } // 统一字段路径(GetByPath)读取基于该对象 if err = json.Unmarshal(modelRespBody, &respObj); err != nil { g.Log().Errorf(asyncCtx, "模型返回参数解析失败:%v", err) if pollErrCnt < maxPollErrRetries { pollErrCnt++ time.Sleep(10 * time.Second) goto LOOP } docMsg.ErrorMsg = fmt.Sprintf("模型返回参数解析失败: %v", err) finalize() return nil } // 无错误时才组装成功内容(错误响应按终态处理,跳过成功解析/轮询) if docMsg.ErrorMsg == "" { // 组装业务返回内容 respBodyMap := modelUtils.CleanMapFieldPath(modelInfo.ResponseBodyMapping) content := make(map[string]any, len(respBodyMap)) for bizKey, jsonPath := range respBodyMap { content[bizKey] = oss.TempURLToOSS(asyncCtx, modelUtils.GetByPathValue(respObj, modelUtils.CleanFieldPath(jsonPath))) } docMsg.Content = content // 解析Token totalTokPath := modelUtils.CleanFieldPath(modelInfo.TokenMapping.TotalTokens) promptTokPath := modelUtils.CleanFieldPath(modelInfo.TokenMapping.PromptTokens) compTokPath := modelUtils.CleanFieldPath(modelInfo.TokenMapping.CompletionTokens) docMsg.TotalTokens = gconv.Int64(modelUtils.GetByPathValue(respObj, totalTokPath)) docMsg.PromptTokens = gconv.Int64(modelUtils.GetByPathValue(respObj, promptTokPath)) docMsg.CompletionTokens = gconv.Int64(modelUtils.GetByPathValue(respObj, compTokPath)) // 调 shop-user-trade 按用量算费(媒体类型取任务创建时的快照;subject=解析后的系统模型 id) docMsg.Cost = calcModelCost(asyncCtx, modelInfo.Id, buildModelUsage(docMsg.PromptTokens, docMsg.CompletionTokens, 0, item.MediaType, time.Since(startTime).Seconds())) // 判断任务状态,轮询等待 statusPath := modelUtils.CleanFieldPath(modelInfo.AsyncTaskMapping.TaskStatus) status := gconv.String(modelUtils.GetByPathValue(respObj, statusPath)) if status == modelInfo.AsyncTaskMapping.TaskStatusPending || status == modelInfo.AsyncTaskMapping.TaskStatusRunning { time.Sleep(10 * time.Second) goto LOOP } } // 成功或已识别出错误的终态统一落库+发布(内容组装完成/错误消息已写入 docMsg) finalize() return nil } // buildAsyncTaskBody 组装异步任务查询请求体:从 AsyncTaskMapping.RequestBodyMapping 出发, // 把 {…} 占位符(如 {taskId})替换为实际任务 ID,兼容 POST 查询接口需要请求体的场景。 // 未配置映射时返回 nil(GET 查询/无需 body 的场景)。 func buildAsyncTaskBody(mapping map[string]any, taskID string) map[string]any { if len(mapping) == 0 { return nil } out := make(map[string]any, len(mapping)) for k, v := range mapping { out[k] = replaceTaskPlaceholder(v, taskID) } return out } // replaceTaskPlaceholder 递归替换结构体中的 {…} 占位符为任务 ID func replaceTaskPlaceholder(v any, taskID string) any { switch val := v.(type) { case string: return urlParamReg.ReplaceAllString(val, taskID) case map[string]any: m := make(map[string]any, len(val)) for k, x := range val { m[k] = replaceTaskPlaceholder(x, taskID) } return m case []any: arr := make([]any, len(val)) for i, x := range val { arr[i] = replaceTaskPlaceholder(x, taskID) } return arr default: return v } } func TaskMsgPublish(ctx context.Context, topic string, data *dto.ModelMsg) (err error) { err = gmq.GetGmq(public.GmqMsgPluginsName).GmqPublish(ctx, &mq.NatsPubMessage{ PubMessage: types.PubMessage{ Topic: topic, Data: data, }, Durable: true, }) if err != nil { g.Log().Errorf(ctx, "[TaskMsgPublish] 发布消息失败 [Error]: %v", err) return fmt.Errorf("发布消息失败") } return }