Files
model-gateway/service/model_task_end_service.go

377 lines
14 KiB
Go

package service
import (
"context"
"encoding/json"
"fmt"
"model-gateway/consts/public"
"model-gateway/dao"
"model-gateway/model/domain"
"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 = modelUtils.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 // 请求成功一次即重置连续失败计数
// 异常响应识别:按模型 ErrorMessageMapping 解析,无错误返回空串
if _, docMsg.ErrorMsg, err = parseModelError(modelRespBody, modelInfo.ErrorMessageMapping); 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
// 解析 ResponseBusinessFieldMapping 字段
businessField := make(map[string]any, len(modelInfo.ResponseBusinessFieldMapping))
for key, value := range modelInfo.ResponseBusinessFieldMapping {
businessField[key] = modelUtils.GetByPathValue(respObj, modelUtils.CleanFieldPath(value))
}
businessFieldRes := new(domain.VideoFieldsRes)
err = gconv.Struct(businessField, businessFieldRes)
if err != nil {
docMsg.ErrorMsg = fmt.Sprintf("解析 ResponseBusinessFieldMapping 字段失败: %v", err)
finalize()
return nil
}
docMsg.Duration = businessFieldRes.Duration
// 解析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))
// 判断任务状态,轮询等待
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
}
// 调 shop-user-trade 按用量算费(媒体类型取任务创建时的快照;subject=解析后的系统模型 id)
docMsg.ModelId = modelInfo.Id // 引用行=系统模型 id,供 per_token 结算按系统模型计价
docMsg.MediaType = item.MediaType
docMsg.Cost = calcModelCost(asyncCtx, modelInfo.Id,
buildModelUsage(docMsg.PromptTokens, docMsg.CompletionTokens, 0, item.MediaType, docMsg.Duration))
}
// 成功或已识别出错误的终态统一落库+发布(内容组装完成/错误消息已写入 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
}