package service import ( "context" "encoding/json" "fmt" "io" "model-gateway/common/util" "model-gateway/consts/public" "model-gateway/dao" "model-gateway/model/dto" "model-gateway/model/entity" modelUtils "model-gateway/service/utils" "net/http" "regexp" "strings" "time" "gitea.redpowerfuture.com/red-future/common/beans" "gitea.redpowerfuture.com/red-future/common/db/gfdb" 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. 逐个原子抢锁,筛选可执行任务 for _, key := range lockKeys { val := gconv.String(mGetRes[key]) // 已被其他实例抢占,跳过 if val != "" { continue } tid := key[len(redisKey):] // SET NX EX 原子抢锁,防止并发竞争 err = g.Redis().SetEX(ctx, key, tid, 1200) if err != nil { return fmt.Errorf("抢占任务锁[%s]失败: %w", tid, err) } err = s.handleSingleTask(ctx, taskMap[key]) if err != nil { g.Log().Errorf(ctx, "处理任务失败: %v", err) } remain-- // 占用一个槽位 if remain <= 0 { break // 槽位已满,终止遍历 } } pageNum++ // 页码动态累加,不再写死2 } return nil } var urlParamReg = regexp.MustCompile(`\{.+?\}`) // handleSingleTask 处理单个视频任务(解耦原循环逻辑) func (s *modelTaskEndService) handleSingleTask(ctx context.Context, item *entity.ModelTaskStart) error { // 提交异步执行 return modelUtils.Submit(ctx, func(ctx context.Context) { startTime := time.Now() asyncCtx := context.WithoutCancel(ctx) // OSS 桶名依赖 ctx 中的用户(GetBucketName → tenantid-{tenantId}), // 响应临时路径转存 OSS 需要用户信息,故在任务体最前面注入 asyncCtx = context.WithValue(asyncCtx, "user", &beans.User{ UserName: item.Creator, TenantId: item.TenantId, }) // 按 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) return } if modelInfo == nil { g.Log().Errorf(asyncCtx, "模型配置不存在: modelId=%d", item.ModelId) return } LOOP: // 替换URL占位符 url := urlParamReg.ReplaceAllString(modelInfo.AsyncTaskMapping.Url, item.TaskId) // 发起HTTP请求 modelRespBody, err := ModelHttpNormalRequest( asyncCtx, url, modelInfo.AsyncTaskMapping.RequestHeadMapping, modelInfo.AsyncTaskMapping.HttpMethod, nil, ) if err != nil { g.Log().Errorf(asyncCtx, "模型请求失败: %v", err) return } if modelRespBody == nil { g.Log().Errorf(asyncCtx, "模型返回参数为空") return } // 解析错误响应 errMsg := new(dto.ModelErrorResp) if err = gconv.Struct(modelRespBody, errMsg); err != nil { g.Log().Errorf(asyncCtx, "模型返回参数解析失败:%v", err) return } // 统一字段路径(GetByPath)读取基于该对象 var respObj map[string]any if err = json.Unmarshal(modelRespBody, &respObj); err != nil { g.Log().Errorf(asyncCtx, "模型返回参数解析失败:%v", err) return } docMsg := new(dto.ModelMsg) docMsg.TaskID = item.Id if errMsg.Error.Code != "" { docMsg.ErrorMsg = errMsg.Error.Message } else { // 组装业务返回内容 respBodyMap := modelUtils.CleanMapFieldPath(modelInfo.ResponseBodyMapping) content := make(map[string]any, len(respBodyMap)) for bizKey, jsonPath := range respBodyMap { content[bizKey] = uploadTempURLToOSS(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)) // 按模型计费规则换算本次调用费用(未配置返回 0);媒体类型取任务创建时的快照 docMsg.Cost = calcCostWithMediaType(docMsg.PromptTokens, docMsg.CompletionTokens, 0, item.MediaType, modelInfo.PriceConfig) // 判断任务状态,轮询等待 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 } } // 按本次实际费用扣减租户余额(未产生费用不扣;异步任务无请求头,admin-go 租户接口无需鉴权可直接调用) if docMsg.Cost > 0 { if err := DeductBalance(asyncCtx, item.TenantId, docMsg.Cost); err != nil { g.Log().Errorf(asyncCtx, "[扣减余额] 异步任务扣费失败 taskId=%d cost=%.6f err=%v", item.Id, docMsg.Cost, err) } } 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) } // 删除redis视频任务 _, err = g.Redis().Del(asyncCtx, "model_video_task:"+gconv.String(item.Id)) if err != nil { return } // 发布消息 if err = TaskMsgPublish(asyncCtx, item.MsgTopic, docMsg); err != nil { g.Log().Errorf(asyncCtx, "模型消息发布失败: %v", err) } }) } // tempDownloadTimeout 临时路径下载超时 const tempDownloadTimeout = 5 * time.Minute // uploadTempURLToOSS 处理响应映射取值:模型返回的临时路径(http/https URL)会过期, // 需下载后转存 OSS,用 OSS 完整路径替换原值。 // - string 且以 http(s):// 开头 → 下载 → 转存 OSS → 返回 OSS 完整路径 // - []any → 逐元素处理,任一元素被替换则返回新切片 // - 其余类型 / 下载或上传失败 → 原样返回(失败仅记日志,不阻断任务) func uploadTempURLToOSS(ctx context.Context, value any) any { switch v := value.(type) { case string: if s, ok := uploadSingleURL(ctx, v); ok { return s } case []any: out := make([]any, len(v)) changed := false for i, e := range v { if s, isStr := e.(string); isStr { if ns, ok := uploadSingleURL(ctx, s); ok { out[i] = ns changed = true continue } } out[i] = e } if changed { return out } } return value } // uploadSingleURL 下载单个临时 URL 并转存 OSS;返回 OSS 完整路径 + 是否成功替换 func uploadSingleURL(ctx context.Context, rawURL string) (string, bool) { rawURL = strings.TrimSpace(rawURL) if !isHTTPURL(rawURL) { return rawURL, false } data, err := downloadTempURL(ctx, rawURL) if err != nil { g.Log().Errorf(ctx, "临时路径下载失败: url=%s err=%v", rawURL, err) return rawURL, false } ossRes, err := Upload(ctx, &dto.UploadFileBytesReq{ FileBytes: data, FileName: fmt.Sprintf("modelFile:%v%s", time.Now().UnixMilli(), extOfData(data)), }) if err != nil { g.Log().Errorf(ctx, "临时路径转存OSS失败: url=%s err=%v", rawURL, err) return rawURL, false } return ossRes.FileAddressPrefix + ossRes.FileURL, true } func isHTTPURL(s string) bool { return strings.HasPrefix(s, "http://") || strings.HasPrefix(s, "https://") } // extOfData 按下载内容嗅探文件后缀(不依赖 URL 路径,模型返回的临时路径可能无后缀) func extOfData(data []byte) string { _, ext := util.DetectFileType(data) if ext == "" || ext == ".octet-stream" { return ".bin" } return ext } // downloadTempURL 带超时下载 URL 内容 func downloadTempURL(ctx context.Context, rawURL string) ([]byte, error) { client := &http.Client{Timeout: tempDownloadTimeout} req, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil) if err != nil { return nil, err } resp, err := client.Do(req) if err != nil { return nil, err } defer func() { _ = resp.Body.Close() }() if resp.StatusCode < 200 || resp.StatusCode >= 300 { return nil, fmt.Errorf("HTTP状态码异常: %d", resp.StatusCode) } return io.ReadAll(resp.Body) } 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 }