Files
NianAIGC/backend/internal/jobs/worker.go

230 lines
6.6 KiB
Go

package jobs
import (
"context"
"crypto/rand"
"encoding/json"
"fmt"
"time"
)
type Processor interface {
Advance(context.Context, Job) (Job, error)
}
type TerminalRefund interface {
Refund(context.Context, Job, string) (Job, error)
}
type UsageRecorder interface {
Record(context.Context, Job) error
}
type WebhookDelivery interface {
Deliver(context.Context, Job) (WebhookResult, error)
}
type WebhookResult struct {
Attempts int
LastStatus any
}
type WorkerConfig struct {
BatchSize int
LockTimeoutSeconds int
PollInterval time.Duration
RetryBase time.Duration
RetryMaximum time.Duration
}
type Worker struct {
store Store
processor Processor
refunds TerminalRefund
usage UsageRecorder
webhooks WebhookDelivery
config WorkerConfig
now func() time.Time
leaseOwner func(string) (string, error)
}
func NewWorker(store Store, processor Processor, refunds TerminalRefund, usage UsageRecorder, webhooks WebhookDelivery, config WorkerConfig, now func() time.Time) *Worker {
if now == nil {
now = time.Now
}
if config.BatchSize < 1 {
config.BatchSize = 3
}
if config.BatchSize > 20 {
config.BatchSize = 20
}
if config.LockTimeoutSeconds <= 0 {
config.LockTimeoutSeconds = 300
}
if config.PollInterval <= 0 {
config.PollInterval = 5 * time.Second
}
return &Worker{store: store, processor: processor, refunds: refunds, usage: usage, webhooks: webhooks, config: config, now: now, leaseOwner: newLeaseOwner}
}
type TickResult struct {
WorkerID string `json:"workerId"`
Claimed int `json:"claimed"`
Jobs []TickJob `json:"jobs"`
}
type TickJob struct {
ID string `json:"id"`
Status Status `json:"status"`
Action string `json:"action"`
Error string `json:"error,omitempty"`
}
func (worker *Worker) Tick(ctx context.Context, workerID string) (TickResult, error) {
return worker.TickLimit(ctx, workerID, worker.config.BatchSize)
}
// TickLimit preserves the legacy internal HTTP tick's per-request bounded limit.
func (worker *Worker) TickLimit(ctx context.Context, workerID string, limit int) (TickResult, error) {
if limit <= 0 {
limit = worker.config.BatchSize
}
limit = max(1, min(limit, 20))
leaseOwner, err := worker.leaseOwner(workerID)
if err != nil {
return TickResult{}, fmt.Errorf("create worker lease owner: %w", err)
}
claimed, err := worker.store.ClaimJobs(ctx, leaseOwner, limit, worker.config.LockTimeoutSeconds)
if err != nil {
return TickResult{}, err
}
result := TickResult{WorkerID: workerID, Claimed: len(claimed), Jobs: make([]TickJob, 0, len(claimed))}
for _, job := range claimed {
advanced, advanceErr := worker.processor.Advance(ctx, job)
// Adapters must preserve the lease token on returned snapshots so every
// downstream settlement/output/refund write remains fenced.
if advanced.ID != "" && advanced.LockedBy == "" {
advanced.LockedBy, advanced.LockedAt = job.LockedBy, job.LockedAt
}
if advanceErr != nil {
status := StatusFailed
jobError := &JobError{Message: advanceErr.Error(), Retryable: true}
advanced, err = worker.store.UpdateJob(ctx, job.ID, workerPatch(job, Patch{Status: &status, Error: jobError}))
if err != nil {
return result, err
}
}
settled, action, err := worker.settle(ctx, advanced)
if err != nil {
return result, err
}
item := TickJob{ID: settled.ID, Status: settled.Status, Action: action}
if advanceErr != nil {
item.Action = "failed"
item.Error = advanceErr.Error()
}
result.Jobs = append(result.Jobs, item)
}
return result, nil
}
func (worker *Worker) settle(ctx context.Context, job Job) (Job, string, error) {
now := worker.now().UTC()
leasePatch := func(patch Patch) Patch {
patch.ExpectedStatuses = []Status{job.Status}
if job.LockedBy != "" {
leaseOwner := job.LockedBy
patch.ExpectedLockedBy = &leaseOwner
}
return patch
}
if job.Status == StatusFailed && job.Error != nil && job.Error.Retryable && job.Attempts < maxAttempts(job) {
attempts := job.Attempts + 1
scheduled := now.Add(RetryDelay(attempts, worker.config.RetryBase, worker.config.RetryMaximum))
status := StatusQueued
// A known provider task is the recovery handle. Retrying a poll must keep
// it; clearing it would turn a transient query failure into a duplicate
// provider submission.
returnPatch := leasePatch(Patch{Status: &status, Attempts: &attempts, ScheduledAt: &scheduled, ClearError: true, ClearLease: true})
retried, err := worker.store.UpdateJob(ctx, job.ID, returnPatch)
return retried, "retry_scheduled", err
}
if !job.Status.Terminal() {
scheduled := now.Add(worker.config.PollInterval)
released, err := worker.store.UpdateJob(ctx, job.ID, leasePatch(Patch{ScheduledAt: &scheduled, ClearLease: true}))
return released, "released", err
}
if job.Status != StatusSucceeded && worker.refunds != nil {
var err error
job, err = worker.refunds.Refund(ctx, job, terminalReason(job))
if err != nil {
return Job{}, "", err
}
}
if job.Status == StatusSucceeded && worker.usage != nil {
if err := worker.usage.Record(ctx, job); err != nil {
return Job{}, "", err
}
}
completed := now
finalPatch := Patch{CompletedAt: &completed, FinalizedAt: &completed, ClearLease: true}
if job.Status == StatusFailed {
attempts := job.Attempts + 1
finalPatch.Attempts = &attempts
}
if worker.webhooks != nil {
delivery, err := worker.webhooks.Deliver(ctx, job)
if err != nil {
return Job{}, "", err
}
if delivery.LastStatus != nil {
lastStatus, err := json.Marshal(delivery.LastStatus)
if err != nil {
return Job{}, "", fmt.Errorf("encode webhook last status: %w", err)
}
job, err = worker.store.UpdateJob(ctx, job.ID, leasePatch(Patch{WebhookAttempts: &delivery.Attempts, WebhookLastStatus: lastStatus, SetWebhookStatus: true}))
if err != nil {
return Job{}, "", err
}
}
}
job, err := worker.store.UpdateJob(ctx, job.ID, leasePatch(finalPatch))
if err != nil {
return Job{}, "", err
}
return job, "processed", nil
}
func newLeaseOwner(workerID string) (string, error) {
var nonce [16]byte
if _, err := rand.Read(nonce[:]); err != nil {
return "", err
}
return fmt.Sprintf("%s:%x", workerID, nonce), nil
}
func workerPatch(job Job, patch Patch) Patch {
patch.ExpectedStatuses = []Status{job.Status}
if job.LockedBy != "" {
worker := job.LockedBy
patch.ExpectedLockedBy = &worker
}
return patch
}
func maxAttempts(job Job) int {
if job.MaxAttempts > 0 {
return job.MaxAttempts
}
return 3
}
func terminalReason(job Job) string {
if job.Error != nil && job.Error.Message != "" {
return job.Error.Message
}
return "任务" + string(job.Status)
}