Files
model-gateway/service/model_task_end_service.go
T
19904408334 76c55fbb73 feat: 新增业务字段路径读写工具
新增 TakeBusinessFields、WriteBusinessFields、SetByPath 与 GetByPath 等工具,支持按映射路径写入请求体与解析响应,并更新相关依赖。
2026-08-18 10:11:09 +08:00

355 lines
11 KiB
Go

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
}