# 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
562 lines
16 KiB
Go
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
|
|
}
|