Files
model-gateway/service/task/task_service.go
T
WangLiZhao cc23697201 Merge branch 'dev未优化' into dev优化中
# Conflicts:
#	common/util/mapping.go
#	config.yml
#	model/entity/model_gateway_model.go
#	model/entity/model_gateway_task.go
#	service/gateway/gateway_http_service.go
#	service/task/task_service.go
#	service/task/worker.go
2026-07-03 09:52:41 +08:00

562 lines
16 KiB
Go

package task
import (
"context"
"errors"
"fmt"
"model-gateway/common/util"
"model-gateway/consts/public"
"time"
"model-gateway/dao"
"model-gateway/model/dto"
"model-gateway/model/entity"
"gitea.redpowerfuture.com/red-future/common/beans"
"gitea.redpowerfuture.com/red-future/common/utils"
"github.com/gogf/gf/v2/database/gdb"
"github.com/gogf/gf/v2/frame/g"
"github.com/gogf/gf/v2/util/gconv"
"github.com/google/uuid"
)
var ModelGatewayTask = &taskService{}
type taskService struct{}
// BuildMessages 构建消息(异步)
func (s *taskService) BuildMessages(ctx context.Context, req *dto.BuildMessagesReq) (*dto.BuildMessagesRes, error) {
user, err := utils.GetUserInfo(ctx)
if err != nil {
return nil, err
}
model, err := dao.ModelGatewayModels.Get(ctx, &entity.ModelGatewayModel{
SQLBaseDO: beans.SQLBaseDO{TenantId: user.TenantId, Creator: user.UserName},
ModelName: req.ModelName,
})
if err != nil || model == nil {
return nil, err
}
// 1) 创建构建记录
taskId := uuid.NewString()
record := &entity.ModelGatewayBuildRecord{
TaskID: taskId,
BuildType: req.BuildType,
ModelName: req.ModelName,
SkillName: req.SkillName,
SessionID: req.SessionId,
NodeID: req.NodeId,
RequestMessages: req.Messages,
CallbackURL: req.CallbackUrl,
Status: 0,
}
_, err = dao.ModelGatewayBuildRecord.Insert(ctx, record)
if err != nil {
return nil, err
}
// 2) 异步执行构建
go s.executeBuild(util.AsyncCtx(ctx), record, req, model)
return &dto.BuildMessagesRes{TaskId: taskId}, nil
}
// executeBuild 异步执行构建逻辑
func (s *taskService) executeBuild(ctx context.Context, record *entity.ModelGatewayBuildRecord, req *dto.BuildMessagesReq, model *entity.ModelGatewayModel) {
var (
startTime = time.Now()
result []map[string]any
err error
)
result, err = s.buildResult(ctx, req, model, record)
record.DurationSeconds = int(time.Since(startTime).Seconds())
if err != nil {
record.Status = 2
record.ErrorMsg = err.Error()
} else {
record.Status = 1
record.ResultMessages = result
}
_, _ = dao.ModelGatewayBuildRecord.Update(ctx, record)
gateway.CallbackBuildResult(ctx, record)
}
// buildResult 构建结果:推理模型手动拼接,视频模型调模型生成多轮
func (s *taskService) buildResult(ctx context.Context, req *dto.BuildMessagesReq, model *entity.ModelGatewayModel, record *entity.ModelGatewayBuildRecord) ([]map[string]any, error) {
messages := req.Messages
switch {
case model.ModelType == public.ModelTypeInference:
// 推理模型:拼接提示词 + 历史
systemPrompt := util.GetModelPrompt(ctx, model.ModelType)
skillContent := prompt.SkillMdContent(ctx, req.SkillName)
systemKey := util.GetRoleContentPath(model.Form, "system")
if systemKey != "" {
messages = util.MergePrompt(messages, systemKey, systemPrompt, skillContent, req.CustomPrompt)
}
history, _ := gateway.GetSessionHistory(ctx, req.NodeId, req.SessionId)
if len(history) > 0 {
messages = util.InjectHistory(messages, history, []string{"system", "history", "user"})
}
// 检查附件是否需要拆分多轮
rounds := util.SplitByAttachment(messages, model.Form)
if len(rounds) > 0 {
return rounds, nil
}
return []map[string]any{messages}, nil
case model.ModelType >= 600 && model.ModelType < 700:
// 视频模型:调推理模型生成多轮
chatModel, err := dao.ModelGatewayModels.Get(ctx, &entity.ModelGatewayModel{
SQLBaseDO: beans.SQLBaseDO{TenantId: model.TenantId, Creator: model.Creator},
IsChatModel: gconv.PtrInt(1),
})
if err != nil || chatModel == nil {
return nil, fmt.Errorf("未找到对话模型")
}
protocol, err := dao.ProviderProtocol.Get(ctx, &entity.ProviderProtocol{
ProviderName: chatModel.OperatorName,
Status: 1,
})
if err != nil || protocol == nil {
return nil, fmt.Errorf("未找到协议配置: %s", chatModel.OperatorName)
}
template := util.BuildTemplateFromForm(model.Form)
outputStruct := gjson.New(template).MustToJsonString()
durationForm, _ := util.GetFormByRole(model.Form, "duration")
totalDur := gconv.Int(gjson.New(req.Messages).Get(durationForm.Key).Val())
minDur := gconv.Int(durationForm.FieldConstraint.Min)
maxDur := gconv.Int(durationForm.FieldConstraint.Max)
systemPrompt := fmt.Sprintf(protocol.SystemPromptTemplate, outputStruct, totalDur, minDur, maxDur)
userContent := util.ExtractUserContent(req.Messages)
reqBody := util.BuildRequestBody(protocol.RequestTemplate, chatModel.ModelName, systemPrompt, userContent)
task := &entity.ModelGatewayTask{
ModelName: chatModel.ModelName,
TaskID: record.TaskID,
State: public.TaskStatusRunning,
BizName: "model-gateway",
RequestPayload: reqBody,
BuildType: req.BuildType,
}
id, err := dao.ModelGatewayTask.Insert(ctx, task)
if err != nil {
return nil, err
}
task.Id = id
rawData, err := AsyncWorker.callModel(chatModel, reqBody)
if err != nil {
task.State = public.TaskStatusFailed
task.ErrorMsg = err.Error()
_, _ = dao.ModelGatewayTask.Update(ctx, task)
return nil, err
}
mapped, err := util.MapResponsePayload(chatModel.ResponseMapping, rawData)
if err != nil {
return nil, err
}
if _, ok := mapped[entity.TotalTokens]; ok {
task.ExpendTokens = gconv.Int64(mapped[entity.TotalTokens])
}
var rounds []map[string]any
contentStr := gjson.New(mapped).Get(entity.ResponseBody).String()
if contentStr != "" {
if err = gjson.DecodeTo(contentStr, &rounds); err != nil {
task.State = public.TaskStatusFailed
task.ErrorMsg = err.Error()
_, _ = dao.ModelGatewayTask.Update(ctx, task)
return nil, err
}
}
oss, err := gateway.UploadByTask(ctx, gjson.New(rounds).MustToJson(), "json")
if err != nil {
task.State = public.TaskStatusFailed
task.ErrorMsg = err.Error()
_, _ = dao.ModelGatewayTask.Update(ctx, task)
return nil, err
}
task.State = public.TaskStatusSuccess
task.ResultFile = &entity.ResultFile{
OssFile: oss.FileAddressPrefix + oss.FileURL,
FileType: oss.FileFormat,
FileSize: int64(oss.FileSize),
}
_, _ = dao.ModelGatewayTask.Update(ctx, task)
return rounds, nil
default:
return nil, errors.New("不支持的模型类型")
}
}
// Create 创建任务
func (s *taskService) Create(ctx context.Context, req *dto.CreateTaskReq) (res *dto.CreateTaskRes, err error) {
taskID := uuid.NewString()
startAt := time.Now()
// 1) 获取用户信息
userInfo, err := utils.GetUserInfo(ctx)
if err != nil {
return nil, err
}
// 2) 检查模型配置
model, err := dao.ModelGatewayModels.Get(ctx, &entity.ModelGatewayModel{
SQLBaseDO: beans.SQLBaseDO{
TenantId: userInfo.TenantId,
Creator: userInfo.UserName,
},
ModelName: req.ModelName,
})
if err != nil {
return nil, err
}
if model == nil || (model.Enabled != nil && *model.Enabled != 1) {
return nil, errors.New("模型不存在或未启用")
}
// TODO: 排队控制暂时关闭,后续需要时取消注释
// limit := queue.GetRuntimeQueueLimit(ctx, req.ModelName, model.MaxConcurrency*2)
// if limit > 0 {
// ok, err := queue.AcquireQueueSlot(ctx, req.ModelName, taskID, limit, model.TimeoutSeconds)
// if err != nil {
// return nil, err
// }
// if !ok {
// return nil, errors.New("任务排队已满,请稍后再试")
// }
// }
// 3) 构建任务实体
task := &entity.ModelGatewayTask{
ModelName: model.ModelName,
TaskID: taskID,
State: public.TaskStatusRunning,
BizName: req.BizName,
CallbackURL: req.CallbackUrl,
RequestPayload: &entity.RequestPayload{
Body: req.RequestPayload,
Headers: util.ParseHeadMsgHeaders(model.HeadMsg),
},
EpicycleId: req.EpicycleId,
BuildModelName: req.BuildModelName,
}
// 4) 插入任务记录
id, err := dao.ModelGatewayTask.Insert(ctx, task)
if err != nil {
// TODO: 恢复排队逻辑后,此处需要回滚排队占位
// queue.ReleaseQueueSlot(ctx, req.ModelName, taskID)
return nil, err
}
task.Id = id
// 5) 记录操作日志(非关键路径,失败不影响主流程)
ip, ua := "", ""
if r := g.RequestFromCtx(ctx); r != nil {
ip = utils.GetLocalIP()
ua = r.UserAgent()
}
_, _ = dao.ModelGatewayLogsOp.Insert(ctx, &entity.ModelGatewayLogsOp{
IP: ip,
UserAgent: ua,
APIPath: "/task/createTask",
HttpMethod: "POST",
BizName: req.BizName,
ModelName: req.ModelName,
TaskID: taskID,
OpType: "createTask",
Success: 1,
CostMs: time.Since(startAt).Milliseconds(),
RequestPayload: task.RequestPayload,
ResponsePayload: gdb.Map{"taskId": taskID},
})
// 6) 模型计费
if len(model.BillingConfig) > 0 {
requestData := util.ExtractRequestBilling(ctx, model.BillingConfig, req.RequestPayload)
// 请求数据作为计费记录的基础字段,先存入数组
task.BillingData = append(task.BillingData, requestData)
_, _ = dao.ModelGatewayTask.Update(ctx, &entity.ModelGatewayTask{
SQLBaseDO: beans.SQLBaseDO{Id: task.Id},
BillingData: task.BillingData,
})
}
// 7) 异步执行任务
go AsyncWorker.handleOne(util.AsyncCtx(ctx), task, model, req)
return &dto.CreateTaskRes{TaskID: taskID}, nil
}
var JobPool *grpool.Pool
// JobTask 定时任务:循环执行待处理任务
func (s *taskService) JobTask(ctx context.Context, req *dto.JobTaskReq) (res *dto.JobTaskRes, err error) {
// 1) 参数默认值从配置取
if req.Interval <= 0 {
req.Interval = g.Cfg().MustGet(ctx, "jobTask.intervalSeconds", 5).Int()
}
if req.BatchSize <= 0 {
req.BatchSize = g.Cfg().MustGet(ctx, "jobTask.batchSize", 10).Int()
}
var (
totalProcessed int
successCount int
failCount int
mu sync.Mutex
wg sync.WaitGroup
)
// 2) 循环查询待处理任务
for {
select {
case <-ctx.Done():
wg.Wait()
return &dto.JobTaskRes{
TotalProcessed: totalProcessed,
SuccessCount: successCount,
FailCount: failCount,
}, nil
default:
}
// 3) 查询 state=0 的任务列表
tasks, err := dao.ModelGatewayTask.ListPending(ctx, req.BatchSize)
if err != nil {
g.Log().Warningf(ctx, "[定时任务] 查询任务失败: %v", err)
time.Sleep(time.Second)
continue
}
if len(tasks) == 0 {
time.Sleep(time.Duration(req.Interval) * time.Second)
continue
}
// 4) 提交到全局协程池执行
for _, task := range tasks {
wg.Add(1)
t := task
err = JobPool.Add(ctx, func(ctx context.Context) {
defer wg.Done()
mu.Lock()
totalProcessed++
mu.Unlock()
if execErr := s.executeTask(ctx, t); execErr != nil {
mu.Lock()
failCount++
mu.Unlock()
g.Log().Errorf(ctx, "[定时任务] 执行失败 taskId=%s err=%v", t.TaskID, execErr)
} else {
mu.Lock()
successCount++
mu.Unlock()
}
})
if err != nil {
return nil, err
}
}
}
}
// executeTask 执行单个任务
func (s *taskService) executeTask(ctx context.Context, task *entity.ModelGatewayTask) error {
// 1) 查询模型配置
model, err := dao.ModelGatewayModels.Get(ctx, &entity.ModelGatewayModel{
SQLBaseDO: beans.SQLBaseDO{
TenantId: task.TenantId,
Creator: task.Creator,
},
ModelName: task.ModelName,
})
if err != nil {
return fmt.Errorf("查询模型配置失败: %w", err)
}
if model == nil || (model.Enabled != nil && *model.Enabled != 1) {
return fmt.Errorf("模型不存在或未启用: %s", task.ModelName)
}
// 3) 调用 handleOne
AsyncWorker.handleOne(util.AsyncCtx(ctx), task, model)
return nil
}
// GetResult 获取任务结果
func (s *taskService) GetResult(ctx context.Context, taskID string) (res *dto.GetTaskResultRes, err error) {
t, err := dao.ModelGatewayTask.Get(ctx, &entity.ModelGatewayTask{
TaskID: taskID,
})
if err != nil {
return nil, err
}
if t == nil {
return nil, errors.New("任务不存在")
}
return &dto.GetTaskResultRes{
OssFile: t.ResultFile.OssFile,
State: t.State,
}, nil
}
// GetBatch 批量查询任务;将成功(state=2)的任务更新为已下载(state=4),并写入过期时间
func (s *taskService) GetBatch(ctx context.Context, req *dto.GetTaskBatchReq) (res *dto.GetTaskBatchRes, err error) {
if req == nil || len(req.TaskIDs) == 0 {
return &dto.GetTaskBatchRes{List: []dto.GetTaskBatchItem{}}, nil
}
// 1) 先查当前租户下的任务列表
list, err := dao.ModelGatewayTask.ListByTaskIDs(ctx, req.TaskIDs)
if err != nil {
return nil, err
}
// 2) 对成功(state=2)的任务:标记为已下载(state=4)
for _, t := range list {
if t == nil {
continue
}
if t.State != public.BuildTypeNode {
continue
}
_ = dao.ModelGatewayTask.MarkDownloadedByID(ctx, t.Id)
// 为了本次返回一致性,内存里也更新
t.State = public.TaskStatusDownloaded
}
// 3) 组装返回
items := make([]dto.GetTaskBatchItem, 0, len(list))
for _, t := range list {
if t == nil {
continue
}
items = append(items, dto.GetTaskBatchItem{
TaskID: t.TaskID,
State: t.State,
OssFile: t.ResultFile.OssFile,
TextResult: t.TextResult,
})
}
return &dto.GetTaskBatchRes{List: items}, nil
}
// List 获取任务列表
func (s *taskService) List(ctx context.Context, req *dto.ListTaskReq) (*dto.ListTaskRes, error) {
if req.PageNum <= 0 {
req.PageNum = 1
}
if req.PageSize <= 0 {
req.PageSize = 10
}
user, err := utils.GetUserInfo(ctx)
if err != nil {
return nil, err
}
list, total, err := dao.ModelGatewayTask.List(ctx, req.PageNum, req.PageSize, &entity.ModelGatewayTask{
SQLBaseDO: beans.SQLBaseDO{
Creator: user.UserName,
},
ModelName: req.ModelName,
BizName: req.BizName,
State: req.State,
TaskID: req.TaskID,
})
if err != nil {
return nil, err
}
return &dto.ListTaskRes{List: list, Total: total}, nil
}
// ModelTaskCallback 模型异步任务的回调通知
func (s *taskService) ModelTaskCallback(ctx context.Context, req *dto.ModelTaskCallbackReq) (*dto.ModelTaskCallbackRes, error) {
g.Log().Infof(ctx, "[模型回调] 收到通知 taskID=%s status=%s", req.TaskID, req.Status)
// 1. 查本地任务
task, err := dao.ModelGatewayTask.Get(ctx, &entity.ModelGatewayTask{
TaskID: req.TaskID,
})
if err != nil || task == nil {
return nil, fmt.Errorf("任务不存在: %s", req.TaskID)
}
// 2. 成功:取 video_url 和 usage
if req.Status == "succeeded" {
result := map[string]any{
"video_url": req.Content["video_url"],
"usage": req.Usage,
}
NotifyAsyncResult(req.TaskID, result, nil)
return &dto.ModelTaskCallbackRes{Success: true}, nil
}
// 3. 失败/过期
if req.Status == "failed" || req.Status == "expired" {
NotifyAsyncResult(req.TaskID, nil, fmt.Errorf(req.Status))
return &dto.ModelTaskCallbackRes{Success: true}, nil
}
return &dto.ModelTaskCallbackRes{Success: true}, nil
}
// QueryPendingTasks 批量轮询进行中的异步任务
func (s *taskService) QueryPendingTasks(ctx context.Context, req *dto.QueryPendingTasksReq) (*dto.QueryPendingTasksRes, error) {
limit := req.Limit
if limit <= 0 {
limit = g.Cfg().MustGet(ctx, "asynch.queryPending.limit", 10).Int()
}
// 1. 查 state=1(执行中)的异步任务
tasks, err := dao.ModelGatewayTask.GetPendingAsyncTasks(ctx, limit)
if err != nil {
return nil, err
}
// 2. 逐个查询
var results []dto.QueryTaskItem
for _, t := range tasks {
// 拿到模型配置
model, err := dao.ModelGatewayModels.GetByModelNameForTenant(ctx, t.TenantId, t.ModelName)
if err != nil || model == nil || model.QueryConfig == nil {
continue
}
result, err := util.PullTaskResult(ctx, nil, model.QueryConfig, model.HeadMsg)
if err != nil {
g.Log().Warningf(ctx, "[轮询] 查询失败 taskID=%s err=%v", t.TaskID, err)
continue
}
status := gconv.String(result["status"])
item := dto.QueryTaskItem{
TaskID: t.TaskID,
Status: status,
Content: result["content"].(map[string]any),
Usage: result["usage"].(map[string]any),
}
results = append(results, item)
// 如果任务完成,通知等待通道
if status == "succeeded" || status == "failed" || status == "expired" {
NotifyAsyncResult(t.TaskID, result["content"].(map[string]any), nil)
}
}
return &dto.QueryPendingTasksRes{
Total: len(results),
Results: results,
}, nil
}