package ws import ( "context" "encoding/json" "errors" "fmt" "log" "net/url" "strings" "sync" "time" "github.com/gin-gonic/gin" "github.com/google/uuid" "gorm.io/gorm" "everypublish/server/internal/config" "everypublish/server/internal/credentials" "everypublish/server/internal/models" "everypublish/server/internal/platform" "everypublish/server/internal/task" ) // Dispatcher 任务执行器(Web-only 本机 worker;审批/重试即时触发) type Dispatcher struct { DB *gorm.DB Hub *Hub Cfg *config.Config Registry *platform.Registry Credentials *credentials.Store Interval time.Duration runsMu sync.Mutex runs map[uint64]context.CancelFunc } type publishTarget struct { account models.Account adapter platform.Adapter } // NewDispatcher 构造下发器 func NewDispatcher(db *gorm.DB, hub *Hub, cfg *config.Config) *Dispatcher { return &Dispatcher{DB: db, Hub: hub, Cfg: cfg, Interval: time.Second, runs: make(map[uint64]context.CancelFunc)} } // Start 后台轮询:把到期且排队的任务交给当前执行模式 func (d *Dispatcher) Start(stop <-chan struct{}) { if d.Interval <= 0 { d.Interval = time.Second } ticker := time.NewTicker(d.Interval) defer ticker.Stop() for { select { case <-stop: return case <-ticker.C: d.poll() } } } func (d *Dispatcher) poll() { var rows []models.Task d.DB.Where("status = ?", task.Queued). Where("(schedule_at is null or schedule_at <= ?)", time.Now()). Order("priority desc, id asc").Limit(20).Find(&rows) for i := range rows { if d.TryDispatch(rows[i].ID) { log.Printf("[dispatch] task %d pushed", rows[i].ID) } } } // TryDispatch 立即执行单个任务(审批/重试后调用)。 func (d *Dispatcher) TryDispatch(taskID uint64) bool { var row models.Task if err := d.DB.First(&row, taskID).Error; err != nil || row.Status != task.Queued { return false } if row.ScheduleAt != nil && row.ScheduleAt.After(time.Now()) { return false } if d.Cfg != nil && d.Cfg.ExecutorMode != "mock" && d.Registry != nil { ids := decodeIDs(row.AccountIDs) targets := make([]publishTarget, 0, len(ids)) allRegistered := len(ids) > 0 hasRegistered := false for _, id := range ids { var account models.Account if d.DB.Where("id = ? and workspace_id = ?", id, row.WorkspaceID).First(&account).Error != nil { allRegistered = false continue } adapter := d.Registry.Get(account.Platform) if adapter == nil { allRegistered = false continue } hasRegistered = true targets = append(targets, publishTarget{account: account, adapter: adapter}) } if allRegistered { return d.dispatchPlatform(&row, targets) } if hasRegistered { return d.failUnsupportedPlatform(&row) } } return d.dispatchLocalMock(&row) } func (d *Dispatcher) failUnsupportedPlatform(row *models.Task) bool { next, ok := task.Next(row.Status, task.ActionDispatch) if !ok || d.DB.Model(row).Update("status", next).Error != nil { return false } row.Status = next d.recordEvent(row, string(task.ActionDispatch), "任务包含尚未接入真实适配器的平台") d.broadcastTask(*row) if next, ok = task.Next(row.Status, task.ActionFail); ok { message := "任务包含尚未接入真实适配器的平台,请拆分任务后重试" _ = d.DB.Model(row).Updates(map[string]interface{}{"status": next, "error_message": message}).Error row.Status = next row.ErrorMessage = message d.recordEvent(row, string(task.ActionFail), message) d.broadcastTask(*row) } return true } // dispatchPlatform runs a registered adapter on the server itself. A task is // marked dispatched before the goroutine starts so retries/cancel operations // observe the same state machine as the local mock executor. func (d *Dispatcher) dispatchPlatform(row *models.Task, targets []publishTarget) bool { next, ok := task.Next(row.Status, task.ActionDispatch) if !ok { return false } claim := d.DB.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", row.ID, row.WorkspaceID, task.Queued).Update("status", next) if claim.Error != nil || claim.RowsAffected != 1 { return false } row.Status = next d.recordEvent(row, string(task.ActionDispatch), "平台适配器已接收任务") d.broadcastTask(*row) ctx, cancel := context.WithCancel(context.Background()) d.runsMu.Lock() if d.runs == nil { d.runs = make(map[uint64]context.CancelFunc) } d.runs[row.ID] = cancel d.runsMu.Unlock() go d.runPlatform(ctx, row.ID, row.WorkspaceID, targets) return true } func (d *Dispatcher) runPlatform(ctx context.Context, taskID, workspaceID uint64, targets []publishTarget) { defer func() { d.runsMu.Lock() delete(d.runs, taskID) d.runsMu.Unlock() }() var row models.Task if err := d.DB.Where("id = ? and workspace_id = ?", taskID, workspaceID).First(&row).Error; err != nil || row.Status != task.Dispatched { return } if next, ok := task.Next(row.Status, task.ActionStart); ok { claim := d.DB.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", row.ID, row.WorkspaceID, task.Dispatched).Update("status", next) if claim.Error != nil || claim.RowsAffected != 1 { return } row.Status = next d.recordEvent(&row, string(task.ActionStart), "平台适配器开始处理") d.broadcastTask(row) } var err error results := make([]platform.PublishResult, 0, len(targets)) for _, target := range targets { input, inputErr := d.publishInput(row, target.account) if inputErr != nil { err = inputErr break } result, publishErr := target.adapter.Publish(ctx, input) if publishErr != nil { err = publishErr break } if strings.TrimSpace(result.URL) == "" { err = &platform.Error{Category: platform.PlatformChanged, Message: "平台未返回发布链接"} break } results = append(results, result) } if err == nil && len(results) == len(targets) { urlsList := make([]string, 0, len(results)) receiptsList := make([]string, 0, len(results)) for _, result := range results { urlsList = append(urlsList, result.URL) receiptsList = append(receiptsList, result.Receipt) } urls, _ := json.Marshal(urlsList) receipts, _ := json.Marshal(receiptsList) now := time.Now() tx := d.DB.Begin() committed := false if tx.Error == nil { result := tx.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", row.ID, row.WorkspaceID, task.Running).Updates(map[string]interface{}{ "status": task.Success, "published_urls": string(urls), "receipts": string(receipts), "error_message": "", "published_at": now, }) txErr := result.Error if txErr == nil && result.RowsAffected != 1 { txErr = context.Canceled } if txErr == nil { txErr = tx.Create(&models.TaskEvent{WorkspaceID: row.WorkspaceID, TaskID: row.ID, Type: string(task.ActionSuccess), Status: task.Success, Message: "平台发布完成"}).Error } if txErr == nil { txErr = tx.Commit().Error } if txErr != nil { tx.Rollback() } else { committed = true } } if committed { row.Status = task.Success row.PublishedURLs = string(urls) row.Receipts = string(receipts) row.PublishedAt = &now d.notifySuccess(row) d.broadcastTask(row) return } } message := "平台发布失败" if err != nil { message = safePlatformError(err) } var current models.Task if d.DB.Where("id = ? and workspace_id = ?", row.ID, row.WorkspaceID).First(¤t).Error == nil && current.Status == task.Cancelled { return } if next, ok := task.Next(row.Status, task.ActionFail); ok { updates := map[string]interface{}{"status": next, "error_message": message} if len(results) > 0 { partialURLs := make([]string, 0, len(results)) partialReceipts := make([]string, 0, len(results)) for _, result := range results { partialURLs = append(partialURLs, result.URL) partialReceipts = append(partialReceipts, result.Receipt) } partialURLJSON, _ := json.Marshal(partialURLs) partialReceiptJSON, _ := json.Marshal(partialReceipts) updates["published_urls"] = string(partialURLJSON) updates["receipts"] = string(partialReceiptJSON) message = "部分账号已发布;" + message updates["error_message"] = message } _ = d.DB.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", row.ID, row.WorkspaceID, task.Running).Updates(updates).Error row.Status = next row.ErrorMessage = message d.recordEvent(&row, string(task.ActionFail), message) d.broadcastTask(row) } } // Cancel interrupts a real adapter run. The task handler still performs the // persisted state transition so the database and worker agree on cancellation. func (d *Dispatcher) Cancel(taskID uint64) { d.runsMu.Lock() cancel := d.runs[taskID] d.runsMu.Unlock() if cancel != nil { cancel() } } func (d *Dispatcher) publishInput(row models.Task, account models.Account) (platform.PublishInput, error) { materialIDs := decodeIDs(row.MaterialIDs) if len(materialIDs) == 0 { return platform.PublishInput{}, &platform.Error{Category: platform.Validation, Message: "任务没有素材"} } var materials []models.Material if err := d.DB.Where("workspace_id = ? and id in ? and status = ?", row.WorkspaceID, materialIDs, "ready").Find(&materials).Error; err != nil { return platform.PublishInput{}, &platform.Error{Category: platform.Unknown, Message: "读取素材失败", Err: err} } if len(materials) != len(materialIDs) { return platform.PublishInput{}, &platform.Error{Category: platform.Validation, Message: "任务素材不存在或未就绪"} } urls := make([]string, 0, len(materials)) for _, material := range materials { token := models.TransferToken{MaterialID: material.ID, Token: uuid.NewString(), ExpiresAt: time.Now().Add(10 * time.Minute)} if err := d.DB.Create(&token).Error; err != nil { return platform.PublishInput{}, &platform.Error{Category: platform.Unknown, Message: "生成素材访问链接失败", Err: err} } base := "" if d.Cfg != nil { base = strings.TrimRight(d.Cfg.BaseURL, "/") } if base == "" { base = "http://127.0.0.1:8090" } urls = append(urls, base+"/api/v1/files/"+url.QueryEscape(token.Token)) } var tags []string if row.Tags != "" { _ = json.Unmarshal([]byte(row.Tags), &tags) } return platform.PublishInput{WorkspaceID: row.WorkspaceID, AccountID: account.ID, TaskID: row.ID, Title: row.Title, Content: row.Content, Tags: tags, MaterialURLs: urls, ScheduleAt: row.ScheduleAt}, nil } func (d *Dispatcher) notifySuccess(row models.Task) { if row.CreatedBy != 0 { _ = d.DB.Create(&models.Notification{WorkspaceID: row.WorkspaceID, UserID: row.CreatedBy, Kind: "task", Title: "任务执行完成", Content: row.Title}).Error } if d.Hub != nil { d.Hub.BroadcastWS(row.WorkspaceID, "notification.new", gin.H{"kind": "task", "title": "任务执行完成", "content": row.Title}) } } func decodeIDs(raw string) []uint64 { var ids []uint64 if err := json.Unmarshal([]byte(raw), &ids); err != nil { return nil } return ids } func safePlatformError(err error) string { var pe *platform.Error if errors.As(err, &pe) { return pe.Error() } return "平台发布失败" } // dispatchLocalMock is the Web-only V1 executor. It deliberately produces a // mock:// receipt so a local test can prove the state machine without claiming // that a real platform accepted the content. func (d *Dispatcher) dispatchLocalMock(row *models.Task) bool { next, ok := task.Next(row.Status, task.ActionDispatch) if !ok { return false } claim := d.DB.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", row.ID, row.WorkspaceID, task.Queued).Update("status", next) if claim.Error != nil || claim.RowsAffected != 1 { return false } row.Status = next d.recordEvent(row, string(task.ActionDispatch), "本机 Web 执行器已接收任务") d.broadcastTask(*row) go func(taskID uint64, workspaceID uint64) { time.Sleep(120 * time.Millisecond) var running models.Task if err := d.DB.Where("id = ? and workspace_id = ?", taskID, workspaceID).First(&running).Error; err != nil || running.Status != task.Dispatched { return } if next, ok := task.Next(running.Status, task.ActionStart); ok { claim := d.DB.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", running.ID, running.WorkspaceID, task.Dispatched).Update("status", next) if claim.Error != nil || claim.RowsAffected != 1 { return } running.Status = next d.recordEvent(&running, string(task.ActionStart), "本机 Web 执行器开始处理") d.broadcastTask(running) } time.Sleep(180 * time.Millisecond) if err := d.DB.Where("id = ? and workspace_id = ?", taskID, workspaceID).First(&running).Error; err != nil || running.Status != task.Running { return } next, ok := task.Next(running.Status, task.ActionSuccess) if !ok { return } urls, _ := json.Marshal([]string{fmt.Sprintf("mock://everypublish/task/%d", taskID)}) receipts, _ := json.Marshal([]string{"mock executor completed"}) now := time.Now() tx := d.DB.Begin() if tx.Error != nil { return } result := tx.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", running.ID, running.WorkspaceID, task.Running).Updates(map[string]interface{}{ "status": next, "published_urls": string(urls), "receipts": string(receipts), "error_message": "", "published_at": now, }) if result.Error != nil || result.RowsAffected != 1 { tx.Rollback() return } running.Status = next running.PublishedURLs = string(urls) running.Receipts = string(receipts) running.PublishedAt = &now if err := tx.Create(&models.TaskEvent{ WorkspaceID: running.WorkspaceID, TaskID: running.ID, Type: string(task.ActionSuccess), Status: running.Status, Message: "本机 mock 执行完成", }).Error; err != nil { tx.Rollback() return } if err := tx.Commit().Error; err != nil { return } if running.CreatedBy != 0 { _ = d.DB.Create(&models.Notification{ WorkspaceID: running.WorkspaceID, UserID: running.CreatedBy, Kind: "task", Title: "任务执行完成", Content: running.Title, }).Error } if d.Hub != nil { d.Hub.BroadcastWS(running.WorkspaceID, "notification.new", gin.H{"kind": "task", "title": "任务执行完成", "content": running.Title}) } d.broadcastTask(running) }(row.ID, row.WorkspaceID) return true } func (d *Dispatcher) recordEvent(row *models.Task, eventType, message string) { _ = d.DB.Create(&models.TaskEvent{ WorkspaceID: row.WorkspaceID, TaskID: row.ID, Type: eventType, Status: row.Status, Message: message, }).Error } func (d *Dispatcher) broadcastTask(row models.Task) { if d.Hub != nil { d.Hub.BroadcastWS(row.WorkspaceID, "task.status", taskEvent(row)) } }