396 lines
13 KiB
Go
396 lines
13 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,
|
||
})
|
||
// 任务处理结束(无论成功失败)都释放 Redis 锁,避免失败路径残留锁、
|
||
// 在锁 TTL(1200s)内阻塞任务被其他 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
|
||
}
|