377 lines
14 KiB
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
|
|
}
|