Files
EveryPublish/server/internal/ws/dispatcher.go
T

446 lines
15 KiB
Go

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(&current).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))
}
}