From 47b38f221c80293cdcea10bf86500f989c662244 Mon Sep 17 00:00:00 2001 From: WangLiZhao <1838393649@qq.com> Date: Mon, 22 Jun 2026 16:13:39 +0800 Subject: [PATCH] =?UTF-8?q?refactor(task):=20=E9=87=8D=E6=9E=84=E4=BB=BB?= =?UTF-8?q?=E5=8A=A1=E8=B0=83=E5=BA=A6=E7=B3=BB=E7=BB=9F,=E6=96=B0?= =?UTF-8?q?=E5=A2=9E=20JobTask=20=E5=AE=9A=E6=97=B6=E4=BB=BB=E5=8A=A1?= =?UTF-8?q?=E5=8A=9F=E8=83=BD,=E5=B9=B6=E4=BC=98=E5=8C=96=E9=85=8D?= =?UTF-8?q?=E7=BD=AE=E7=BB=93=E6=9E=84?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config.yml | 21 +++---- dao/model_gateway_task_dao.go | 17 ++++++ main.go | 18 +++--- model/dto/model_gateway_task_dto.go | 1 + model/entity/model_gateway_task.go | 3 + service/task/task_service.go | 94 ++++++++++++++++++++++++++++- service/task/worker.go | 17 +++--- 7 files changed, 140 insertions(+), 31 deletions(-) diff --git a/config.yml b/config.yml index 8da4bef..f2dd0a7 100644 --- a/config.yml +++ b/config.yml @@ -60,16 +60,11 @@ jaeger: addr: 192.168.3.30:4318 # 本地调试用:可选自动执行 worker/cleaner(默认关闭) -asynch: - queryPending: - enabled: false - intervalSeconds: 10 # 每10秒轮询一次 - limit: 10 # 每次查10条 - worker: - enabled: false - intervalSeconds: 5 - batchSize: 10 - goroutines: 1 - cleaner: - enabled: false - intervalSeconds: 30 +queryPending: + enabled: false + intervalSeconds: 10 # 每10秒轮询一次 + limit: 10 # 每次查10条 +jobTask: + intervalSeconds: 10 # 轮询间隔(秒) + batchSize: 10 # 每批处理条数 + poolSize: 5 # 协程池大小 \ No newline at end of file diff --git a/dao/model_gateway_task_dao.go b/dao/model_gateway_task_dao.go index 0031c4e..330138f 100644 --- a/dao/model_gateway_task_dao.go +++ b/dao/model_gateway_task_dao.go @@ -103,6 +103,23 @@ func (d *modelGatewayTaskDao) ListByTaskIDs(ctx context.Context, taskIDs []strin return } +// ListPending 查询待处理任务 +func (d *modelGatewayTaskDao) ListPending(ctx context.Context, limit int) (list []*entity.ModelGatewayTask, err error) { + if limit <= 0 { + limit = 10 + } + model := gfdb.DB(ctx, public.DbNameModelGateway).Model(ctx, public.TableNameTask). + Where(entity.ModelGatewayTaskCol.State, 0). + OrderAsc(entity.ModelGatewayTaskCol.CreatedAt). + Limit(limit) + r, err := model.All() + if err != nil { + return nil, err + } + err = r.Structs(&list) + return +} + // MarkDownloadedByID 标记已下载 func (d *modelGatewayTaskDao) MarkDownloadedByID(ctx context.Context, id int64) error { _, err := gfdb.DB(ctx, public.DbNameModelGateway).Model(ctx, public.TableNameTask). diff --git a/main.go b/main.go index d00e8fe..aabe876 100644 --- a/main.go +++ b/main.go @@ -17,6 +17,7 @@ import ( _ "github.com/gogf/gf/contrib/drivers/pgsql/v2" _ "github.com/gogf/gf/contrib/nosql/redis/v2" "github.com/gogf/gf/v2/frame/g" + "github.com/gogf/gf/v2/os/grpool" ) func main() { @@ -24,6 +25,11 @@ func main() { defer cancel() defer jaeger.ShutDown(ctx) + // 初始化全局协程池 + poolSize := g.Cfg().MustGet(ctx, "jobTask.poolSize", 5).Int() + task.JobPool = grpool.New(poolSize) + defer task.JobPool.Close() + // 注册路由 http.RouteRegister([]interface{}{ controller.ModelGatewayModels, @@ -31,25 +37,23 @@ func main() { controller.ModelGatewayLogsStat, }) - // 本地调试:可选自动触发 worker/cleaner(由配置文件控制) + // 本地调试:可选自动触发 worker/cleaner startAutoRunner(ctx) - // 监听退出信号,确保 Ctrl+C 能完整退出(停止 worker/cleaner 并关闭 gateway server) + // 监听退出信号 quit := make(chan os.Signal, 1) signal.Notify(quit, os.Interrupt, syscall.SIGTERM) <-quit - g.Log().Infof(ctx, "[main] 收到退出信号,开始优雅退出...") cancel() - // 关闭 gateway server(RouteRegister 内部是 go Httpserver.Run() 启动的) _ = http.Httpserver.Shutdown() } func startAutoRunner(ctx context.Context) { // queryPending - if g.Cfg().MustGet(ctx, "asynch.queryPending.enabled").Bool() { - interval := g.Cfg().MustGet(ctx, "asynch.queryPending.intervalSeconds", 10).Int() - limit := g.Cfg().MustGet(ctx, "asynch.queryPending.limit", 10).Int() + if g.Cfg().MustGet(ctx, "queryPending.enabled").Bool() { + interval := g.Cfg().MustGet(ctx, "queryPending.intervalSeconds", 10).Int() + limit := g.Cfg().MustGet(ctx, "queryPending.limit", 10).Int() ticker := time.NewTicker(time.Duration(interval) * time.Second) go func() { defer ticker.Stop() diff --git a/model/dto/model_gateway_task_dto.go b/model/dto/model_gateway_task_dto.go index e2e779a..a0a2606 100644 --- a/model/dto/model_gateway_task_dto.go +++ b/model/dto/model_gateway_task_dto.go @@ -22,6 +22,7 @@ type JobTaskReq struct { g.Meta `path:"/jobTask" method:"post" tags:"任务管理" summary:"定时任务" dc:"循环执行待处理任务,按间隔时间和批次大小处理"` Interval int `json:"interval" dc:"循环间隔(秒)"` BatchSize int `json:"batchSize" dc:"每批执行条数"` + PoolSize int `json:"poolSize" dc:"协程池大小"` } type JobTaskRes struct { diff --git a/model/entity/model_gateway_task.go b/model/entity/model_gateway_task.go index 986839a..b9c168e 100644 --- a/model/entity/model_gateway_task.go +++ b/model/entity/model_gateway_task.go @@ -16,6 +16,7 @@ type modelGatewayTaskCol struct { ExpendTokens string DurationSeconds string RetryCount string + BuildType string RequestPayload string EpicycleId string } @@ -32,6 +33,7 @@ var ModelGatewayTaskCol = modelGatewayTaskCol{ ExpendTokens: "expend_tokens", DurationSeconds: "duration_seconds", RetryCount: "retry_count", + BuildType: "build_type", RequestPayload: "request_payload", EpicycleId: "epicycle_id", } @@ -44,6 +46,7 @@ type ModelGatewayTask struct { BizName string `orm:"biz_name" json:"bizName"` CallbackURL string `orm:"callback_url" json:"callbackUrl"` State int `orm:"state" json:"state"` + BuildType int64 `orm:"build_type" json:"buildType"` ErrorMsg string `orm:"error_msg" json:"errorMsg"` ResultFile *ResultFile `orm:"result_file" json:"resultFile"` ExpendTokens int64 `orm:"expend_tokens" json:"expendTokens"` diff --git a/service/task/task_service.go b/service/task/task_service.go index 1a628e6..49e22b6 100644 --- a/service/task/task_service.go +++ b/service/task/task_service.go @@ -6,6 +6,7 @@ import ( "fmt" "model-gateway/common/util" "model-gateway/consts/public" + "sync" "time" "model-gateway/dao" @@ -16,6 +17,7 @@ import ( "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/os/grpool" "github.com/gogf/gf/v2/util/gconv" "github.com/google/uuid" ) @@ -71,6 +73,7 @@ func (s *taskService) Create(ctx context.Context, req *dto.CreateTaskReq) (res * CallbackURL: req.CallbackUrl, RequestPayload: req.RequestPayload, EpicycleId: req.EpicycleId, + BuildType: req.BuildType, } // 4) 插入任务记录 @@ -104,15 +107,102 @@ func (s *taskService) Create(ctx context.Context, req *dto.CreateTaskReq) (res * }) // 6) 异步执行任务 - go AsyncWorker.handleOne(util.AsyncCtx(ctx), task, model, req) + go AsyncWorker.handleOne(util.AsyncCtx(ctx), task, model) 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() + } - return nil, err + 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{ + 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(ctx, task, model) + return nil } // GetResult 获取任务结果 diff --git a/service/task/worker.go b/service/task/worker.go index f25bfab..5b190e7 100644 --- a/service/task/worker.go +++ b/service/task/worker.go @@ -6,7 +6,6 @@ import ( "encoding/json" "fmt" "io" - "model-gateway/model/dto" "net/http" "strings" "sync" @@ -31,7 +30,7 @@ type asyncWorker struct { } // handleOne 执行一次完整的任务 -func (w *asyncWorker) handleOne(ctx context.Context, task *entity.ModelGatewayTask, model *entity.ModelGatewayModel, req *dto.CreateTaskReq) { +func (w *asyncWorker) handleOne(ctx context.Context, task *entity.ModelGatewayTask, model *entity.ModelGatewayModel) { var ( body = task.RequestPayload maxRetry = model.RetryTimes @@ -59,12 +58,12 @@ func (w *asyncWorker) handleOne(ctx context.Context, task *entity.ModelGatewayTa result, err = util.ParseStreamResponse(rawBytes, model.StreamConfig) } case model.CallMode != nil && *model.CallMode == public.CallModeAsync: - result, err = w.callModel(ctx, task, model, body) + result, err = w.callModel(ctx, model, body) if err == nil { result, err = util.PullTaskResult(ctx, result, model.QueryConfig, model.HeadMsg) } default: - result, err = w.callModel(ctx, task, model, body) + result, err = w.callModel(ctx, model, body) } if err == nil { @@ -83,7 +82,7 @@ func (w *asyncWorker) handleOne(ctx context.Context, task *entity.ModelGatewayTa // ============================================ // 2) 解析校验 + 响应映射(可重试) // ============================================ - result, err = w.parseAndRetry(ctx, result, task, model, maxRetry, req) + result, err = w.parseAndRetry(ctx, result, task, model, maxRetry) if err != nil { w.failTask(ctx, task, startTime, err.Error()) return @@ -144,7 +143,7 @@ var asyncTaskChan = sync.Map{} // taskID → chan asyncResult func (w *asyncWorker) callModelAsync(ctx context.Context, task *entity.ModelGatewayTask, model *entity.ModelGatewayModel, body map[string]any) (map[string]any, error) { // 1. 提交异步任务 - body, err := w.callModel(ctx, task, model, body) + body, err := w.callModel(ctx, model, body) if err != nil { return nil, err } @@ -186,7 +185,7 @@ func NotifyAsyncResult(taskID string, result map[string]any, err error) { } // callModel 调用模型 + 提取文本结果 -func (w *asyncWorker) callModel(ctx context.Context, task *entity.ModelGatewayTask, model *entity.ModelGatewayModel, body map[string]any) (map[string]any, error) { +func (w *asyncWorker) callModel(ctx context.Context, model *entity.ModelGatewayModel, body map[string]any) (map[string]any, error) { data, err := InvokeModel(ctx, model, body) if err != nil { return nil, err @@ -206,7 +205,7 @@ func (w *asyncWorker) callModel(ctx context.Context, task *entity.ModelGatewayTa } // parseAndRetry 解析模型返回结果,并重试 -func (w *asyncWorker) parseAndRetry(ctx context.Context, body map[string]any, task *entity.ModelGatewayTask, model *entity.ModelGatewayModel, maxRetry int, req *dto.CreateTaskReq) (map[string]any, error) { +func (w *asyncWorker) parseAndRetry(ctx context.Context, body map[string]any, task *entity.ModelGatewayTask, model *entity.ModelGatewayModel, maxRetry int) (map[string]any, error) { var lastErr error for attempt := 0; attempt <= maxRetry; attempt++ { if attempt > 0 { @@ -234,7 +233,7 @@ func (w *asyncWorker) parseAndRetry(ctx context.Context, body map[string]any, ta // 3) 解析 + 校验 var parsed map[string]any - switch req.BuildType { + switch task.BuildType { case public.BuildTypePrompt, public.BuildTypeNode: parsed, err = util.ParseAndValidate(mapped, model) if err == nil {