446 lines
15 KiB
Go
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(¤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))
|
|
}
|
|
}
|