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 }