Files
chorus/portal/worker/worker.go
T
ilaandClaude Opus 5 52ebca2f62 fix: 修正生成失败错误码被终态码覆盖 (#63)
worker 的终态判定改用随 BeginProviderAttempt 推进的真实尝试计数,
不再使用取任务时的内存快照,避免"已尝试过但无候选可换"被误报为
route_unavailable。

queue 的 Fail() 不再用 JSON_MERGE_PATCH 把终态码写进已结束的尾条
尝试,已完成的尝试记录保持不可变;只有尚未结束的尾条尝试(租约、
快照解码等非 provider 阶段失败)才补写终态信息。

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-08-26 09:07:50 +08:00

455 lines
15 KiB
Go

package worker
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"strings"
"time"
"git.ilapage.cn/OPC/chorus/internal/core/model"
"git.ilapage.cn/OPC/chorus/internal/core/provider"
"git.ilapage.cn/OPC/chorus/internal/core/queue"
"git.ilapage.cn/OPC/chorus/internal/core/router"
corestorage "git.ilapage.cn/OPC/chorus/internal/core/storage"
"git.ilapage.cn/OPC/chorus/internal/platform/ratelimit"
)
var (
ErrInvalidConfig = errors.New("worker configuration is invalid")
ErrNoProvider = errors.New("routed provider model is unavailable")
)
type Queue interface {
ClaimNext(ctx context.Context, owner string, leaseDuration time.Duration) (*queue.Claim, error)
BeginProviderAttempt(ctx context.Context, generationID uint64, leaseToken string, routeMemberID, providerModelID uint64) (model.Attempt, bool, error)
FinishProviderAttempt(ctx context.Context, generationID uint64, leaseToken string, attempt model.Attempt) (bool, error)
Defer(ctx context.Context, generationID uint64, leaseToken string, availableAt time.Time) (bool, error)
Succeed(ctx context.Context, generationID uint64, leaseToken string, outputs []model.GenerationOutput, attempt model.Attempt) (bool, error)
Fail(ctx context.Context, generationID uint64, leaseToken string, code, message string, attempt model.Attempt) (bool, error)
Inputs(ctx context.Context, generationID uint64) ([]model.GenerationInput, error)
}
type ClaimController interface{ StopClaims() }
type Selection struct {
ProviderID uint64
ProviderModelID uint64
BaseURL string
AuthType provider.AuthType
APIKey string
ModelID string
APIType model.APIType
ExtraBody []byte
Timeout time.Duration
}
type Catalog interface {
Resolve(ctx context.Context, member router.MemberSnapshot, capability model.Capability) (Selection, error)
}
type ClientFactory interface {
New(selection Selection) (provider.Client, error)
}
type ImageStore interface {
Open(ctx context.Context, key string) (io.ReadCloser, corestorage.Object, error)
PutImage(ctx context.Context, request ImageRequest) (ImageObjects, error)
Delete(ctx context.Context, key string) error
}
type ImageRequest struct {
Key, ThumbnailKey string
OwnerID, GenerationID uint64
ContentType string
Source io.Reader
}
type ImageObjects struct{ Original, Thumbnail corestorage.Object }
type Config struct {
Owner string
LeaseDuration time.Duration
PollInterval time.Duration
Random router.RandomSource
RateLimiter *ratelimit.Limiter
ProviderPolicy ratelimit.Policy
Now func() time.Time
}
type Worker struct {
config Config
queue Queue
controller ClaimController
catalog Catalog
factory ClientFactory
runtime router.RuntimeRepository
storage ImageStore
}
func New(config Config, queueRepository Queue, controller ClaimController, catalog Catalog, factory ClientFactory, runtime router.RuntimeRepository, storage ImageStore) (*Worker, error) {
if strings.TrimSpace(config.Owner) == "" || config.LeaseDuration <= 0 || config.PollInterval <= 0 || config.RateLimiter == nil || !config.ProviderPolicy.Valid() || queueRepository == nil || controller == nil || catalog == nil || factory == nil || runtime == nil || storage == nil {
return nil, ErrInvalidConfig
}
if config.Random == nil {
config.Random = router.CryptoRandom{}
}
if config.Now == nil {
config.Now = func() time.Time { return time.Now().UTC() }
}
return &Worker{config: config, queue: queueRepository, controller: controller, catalog: catalog, factory: factory, runtime: runtime, storage: storage}, nil
}
func (w *Worker) Run(ctx context.Context) error {
stop := make(chan struct{})
go func() {
select {
case <-ctx.Done():
w.controller.StopClaims()
case <-stop:
}
}()
defer close(stop)
for {
if ctx.Err() != nil {
return nil
}
claim, err := w.queue.ClaimNext(ctx, w.config.Owner, w.config.LeaseDuration)
if errors.Is(err, queue.ErrClaimsStopped) || errors.Is(err, context.Canceled) {
return nil
}
if err != nil {
return fmt.Errorf("claim worker task: %w", err)
}
if claim == nil {
select {
case <-ctx.Done():
return nil
case <-time.After(w.config.PollInterval):
continue
}
}
workCtx, cancel := context.WithTimeout(context.Background(), w.config.LeaseDuration)
if err := w.process(workCtx, claim); err != nil {
cancel()
return fmt.Errorf("process worker task: %w", err)
}
cancel()
}
}
func (w *Worker) ProcessOne(ctx context.Context) (bool, error) {
claim, err := w.queue.ClaimNext(ctx, w.config.Owner, w.config.LeaseDuration)
if err != nil || claim == nil {
return false, err
}
return true, w.process(ctx, claim)
}
func (w *Worker) process(ctx context.Context, claim *queue.Claim) error {
snapshot, err := decodeSnapshot(claim.Generation.RouteSnapshot)
if err != nil {
return w.fail(ctx, claim, provider.CodeUnknown)
}
inputs, err := w.loadInputs(ctx, claim.Generation)
if err != nil {
return w.fail(ctx, claim, provider.CodeUnknown)
}
used := usedRouteMembers(claim.Generation.Attempts)
// providerAttemptCount 必须跟随 BeginProviderAttempt 的实际结果推进。
// claim.Generation.ProviderAttemptCount 是取任务时的快照,本次循环内不会更新,
// 用它判断终态会把"已尝试过但无候选可换"误报成"没有可用路由"。
providerAttemptCount := claim.Generation.ProviderAttemptCount
if providerAttemptCount >= uint32(snapshot.MaxFailover)+1 {
return w.fail(ctx, claim, provider.ErrorCode(router.CodeFailoverExhausted))
}
states, err := w.runtime.MemberStates(ctx, snapshot.Capability, snapshot.Members)
if err != nil {
return w.fail(ctx, claim, provider.CodeUnknown)
}
for index := range states {
if used[states[index].RoutePoolMemberID] {
states[index].Enabled = false
}
}
candidates, err := router.SelectCandidates(snapshot, states, w.config.Random)
if err != nil {
if errors.Is(err, router.ErrRouteUnavailable) {
if providerAttemptCount > 0 {
return w.fail(ctx, claim, provider.ErrorCode(router.CodeFailoverExhausted))
}
return w.fail(ctx, claim, provider.ErrorCode(router.CodeRouteUnavailable))
}
return w.fail(ctx, claim, provider.CodeUnknown)
}
remaining := int(uint32(snapshot.MaxFailover) + 1 - providerAttemptCount)
if remaining < len(candidates) {
candidates = candidates[:remaining]
}
var earliestRetry time.Time
otherUnavailable := false
for _, member := range candidates {
selection, resolveErr := w.catalog.Resolve(ctx, member, snapshot.Capability)
if resolveErr != nil || selection.ProviderID == 0 {
otherUnavailable = true
continue
}
decision := w.config.RateLimiter.Take(w.config.Now(), ratelimit.Bucket{
Key: fmt.Sprintf("provider:%d", selection.ProviderID), Policy: w.config.ProviderPolicy,
})
if !decision.Allowed {
if earliestRetry.IsZero() || decision.RetryAt.Before(earliestRetry) {
earliestRetry = decision.RetryAt
}
continue
}
rateReservation := decision.Reservation()
reservation, reserveErr := w.runtime.Reserve(ctx, router.ReservationRequest{
RoutePoolMemberID: member.RoutePoolMemberID, Capability: snapshot.Capability,
Owner: w.config.Owner, LeaseDuration: w.config.LeaseDuration,
})
if reserveErr != nil {
rateReservation.Cancel()
if errors.Is(reserveErr, router.ErrMemberUnavailable) {
otherUnavailable = true
continue
}
return w.fail(ctx, claim, provider.CodeUnknown)
}
attempt, begun, beginErr := w.queue.BeginProviderAttempt(ctx, claim.Generation.ID, claim.LeaseToken, member.RoutePoolMemberID, selection.ProviderModelID)
if beginErr != nil {
rateReservation.Cancel()
return beginErr
}
if !begun {
rateReservation.Cancel()
return nil
}
providerAttemptCount = attempt.ProviderOrdinal
client, factoryErr := w.factory.New(selection)
if factoryErr != nil {
finished, err := w.finishFailure(ctx, claim, reservation, attempt, provider.CodeUnknown, provider.FailureOther)
if err != nil {
return err
}
if !finished {
return nil
}
return w.fail(ctx, claim, provider.CodeUnknown)
}
requestCtx := ctx
var cancel context.CancelFunc
if selection.Timeout > 0 {
requestCtx, cancel = context.WithTimeout(ctx, selection.Timeout)
}
generated, generateErr := client.Generate(requestCtx, provider.Request{
Kind: claim.Generation.Kind, APIType: selection.APIType, ModelID: selection.ModelID,
RenderedPrompt: claim.Generation.RenderedPrompt, Inputs: inputs,
})
if cancel != nil {
cancel()
}
if generateErr != nil {
code, class := provider.CodeUnknown, provider.FailureOther
var providerErr *provider.Error
if errors.As(generateErr, &providerErr) {
code, class = providerErr.Code, providerErr.Class
}
finished, err := w.finishFailure(ctx, claim, reservation, attempt, code, class)
if err != nil {
return err
}
if !finished {
return nil
}
if router.ActionForFailure(class) == router.FailureTryNext {
continue
}
return w.fail(ctx, claim, code)
}
finished, err := w.finishSuccess(ctx, claim, reservation, attempt)
if err != nil {
return err
}
if !finished {
return nil
}
outputs, keys, saveErr := w.saveOutputs(ctx, claim.Generation, claim.LeaseToken, attempt.ProviderOrdinal, generated)
if saveErr != nil {
w.deleteKeys(keys)
return w.fail(ctx, claim, provider.CodeUnknown)
}
finalizeCtx, finalizeCancel := context.WithTimeout(context.Background(), 5*time.Second)
owned, succeedErr := w.queue.Succeed(finalizeCtx, claim.Generation.ID, claim.LeaseToken, outputs, model.Attempt{})
finalizeCancel()
if succeedErr != nil || !owned {
w.deleteKeys(keys)
return succeedErr
}
return nil
}
if !earliestRetry.IsZero() && !otherUnavailable {
return w.deferGeneration(claim, earliestRetry)
}
if providerAttemptCount > 0 {
return w.fail(ctx, claim, provider.ErrorCode(router.CodeFailoverExhausted))
}
return w.fail(ctx, claim, provider.ErrorCode(router.CodeRouteUnavailable))
}
func (w *Worker) deferGeneration(claim *queue.Claim, retryAt time.Time) error {
now := w.config.Now()
if !retryAt.After(now) {
retryAt = now.Add(w.config.PollInterval)
}
if maximum := now.Add(w.config.ProviderPolicy.Window); retryAt.After(maximum) {
retryAt = maximum
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_, err := w.queue.Defer(ctx, claim.Generation.ID, claim.LeaseToken, retryAt)
return err
}
func decodeSnapshot(encoded []byte) (router.RouteSnapshot, error) {
var snapshot router.RouteSnapshot
if len(encoded) == 0 || json.Unmarshal(encoded, &snapshot) != nil {
return router.RouteSnapshot{}, router.ErrInvalidSnapshot
}
if err := snapshot.Validate(); err != nil {
return router.RouteSnapshot{}, err
}
return snapshot, nil
}
func usedRouteMembers(encoded []byte) map[uint64]bool {
var attempts []model.Attempt
if json.Unmarshal(encoded, &attempts) != nil {
return map[uint64]bool{}
}
used := make(map[uint64]bool, len(attempts))
for _, attempt := range attempts {
if attempt.Type == "provider" && attempt.RouteMemberID != 0 {
used[attempt.RouteMemberID] = true
}
}
return used
}
func (w *Worker) finishFailure(ctx context.Context, claim *queue.Claim, reservation router.Reservation, attempt model.Attempt, code provider.ErrorCode, class provider.FailureClass) (bool, error) {
retryable := router.Retryable(class)
attempt.ErrorCode = string(code)
attempt.ErrorMessage = "upstream request failed"
attempt.LatencyMS = elapsedMilliseconds(attempt.StartedAt)
attempt.Retryable = &retryable
finished, err := w.queue.FinishProviderAttempt(ctx, claim.Generation.ID, claim.LeaseToken, attempt)
if err != nil {
return false, err
}
_, recordErr := w.runtime.Record(ctx, reservation, router.CircuitObservation{
Retryable: retryable, OpenImmediately: class == provider.FailureUnauthorized,
ErrorCode: string(code), ErrorMessage: "upstream request failed",
})
if recordErr != nil {
return false, recordErr
}
return finished, nil
}
func (w *Worker) finishSuccess(ctx context.Context, claim *queue.Claim, reservation router.Reservation, attempt model.Attempt) (bool, error) {
retryable := false
attempt.LatencyMS = elapsedMilliseconds(attempt.StartedAt)
attempt.Retryable = &retryable
finished, err := w.queue.FinishProviderAttempt(ctx, claim.Generation.ID, claim.LeaseToken, attempt)
if err != nil {
return false, err
}
_, recordErr := w.runtime.Record(ctx, reservation, router.CircuitObservation{Succeeded: true})
if recordErr != nil {
return false, recordErr
}
return finished, nil
}
func (w *Worker) loadInputs(ctx context.Context, generation model.Generation) ([]provider.Input, error) {
rows, err := w.queue.Inputs(ctx, generation.ID)
if err != nil {
return nil, err
}
inputs := make([]provider.Input, 0, len(rows))
for _, row := range rows {
reader, object, openErr := w.storage.Open(ctx, row.StorageKey)
if openErr != nil {
return nil, openErr
}
content, readErr := io.ReadAll(reader)
closeErr := reader.Close()
if readErr != nil {
return nil, readErr
}
if closeErr != nil {
return nil, closeErr
}
if object.OwnerID != generation.UserID || object.GenerationID != generation.ID {
return nil, fmt.Errorf("input storage ownership mismatch")
}
inputs = append(inputs, provider.Input{StorageKey: row.StorageKey, MIMEType: row.MIMEType, Role: row.Role, Position: row.Position, Content: content})
}
return inputs, nil
}
func (w *Worker) saveOutputs(ctx context.Context, generation model.Generation, leaseToken string, providerOrdinal uint32, generated []provider.Output) ([]model.GenerationOutput, []string, error) {
if len(generated) == 0 {
return nil, nil, fmt.Errorf("provider returned no outputs")
}
rows := make([]model.GenerationOutput, 0, len(generated))
keys := []string{}
for index, output := range generated {
if output.Text != "" {
text := output.Text
rows = append(rows, model.GenerationOutput{Kind: output.Kind, TextContent: &text})
continue
}
if output.Kind != model.KindImage || len(output.Content) == 0 {
return nil, keys, fmt.Errorf("provider returned invalid output")
}
key := fmt.Sprintf("outputs/%d/%d/%s/%d-%d", generation.UserID, generation.ID, leaseToken, providerOrdinal, index+1)
thumb := key + "-thumbnail"
objects, err := w.storage.PutImage(ctx, ImageRequest{Key: key, ThumbnailKey: thumb, OwnerID: generation.UserID, GenerationID: generation.ID, ContentType: output.ContentType, Source: bytes.NewReader(output.Content)})
if err != nil {
return nil, keys, err
}
keys = append(keys, objects.Original.Key, objects.Thumbnail.Key)
mimeType := objects.Original.ContentType
size := uint64(objects.Original.Size)
originalKey, thumbnailKey := objects.Original.Key, objects.Thumbnail.Key
rows = append(rows, model.GenerationOutput{Kind: output.Kind, StorageKey: &originalKey, ThumbnailStorageKey: &thumbnailKey, MIMEType: &mimeType, SizeBytes: &size})
}
return rows, keys, nil
}
func (w *Worker) fail(ctx context.Context, claim *queue.Claim, code provider.ErrorCode) error {
finalizeCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_, err := w.queue.Fail(finalizeCtx, claim.Generation.ID, claim.LeaseToken, string(code), "upstream request failed", model.Attempt{})
return err
}
func (w *Worker) deleteKeys(keys []string) {
for _, key := range keys {
_ = w.storage.Delete(context.Background(), key)
}
}
func elapsedMilliseconds(started *time.Time) int64 {
if started == nil {
return 1
}
elapsed := time.Since(*started).Milliseconds()
if elapsed < 1 {
return 1
}
return elapsed
}