Files
model-gateway/service/model_task_end_service.go
T
19904408334 cb9d04648c fix: 确保任务结束释放 Redis 锁
将 Redis 锁的删除从成功路径移动到 defer,保证任务无论成功失败都会释放锁,避免失败时残留锁导致任务在 TTL 内无法被重新获取。
2026-08-22 08:26:43 +08:00

396 lines
13 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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,
})
// 任务处理结束(无论成功失败)都释放 Redis 锁,避免失败路径残留锁、
// 在锁 TTL1200s)内阻塞任务被其他 worker 重新获取
defer func() {
if _, delErr := g.Redis().Del(asyncCtx, "model_video_task:"+gconv.String(item.Id)); delErr != nil {
g.Log().Errorf(asyncCtx, "清理任务锁失败: %v", delErr)
}
}()
// 按 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)
// 组装查询请求体:POST 查询接口需要 body(从 RequestBodyMapping 出发,替换 {…} 占位符为任务 ID)
reqBody := buildAsyncTaskBody(modelInfo.AsyncTaskMapping.RequestBodyMapping, item.TaskId)
// 发起HTTP请求
modelRespBody, err := ModelHttpNormalRequest(
asyncCtx,
url,
modelInfo.AsyncTaskMapping.RequestHeadMapping,
modelInfo.AsyncTaskMapping.HttpMethod, reqBody,
)
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)
}
// 发布消息
if err = TaskMsgPublish(asyncCtx, item.MsgTopic, docMsg); err != nil {
g.Log().Errorf(asyncCtx, "模型消息发布失败: %v", err)
}
})
}
// 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
}
}
// 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
}
//ossRes.FileAddressPrefix + ossRes.FileURL, true
return 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
}