worker 的终态判定改用随 BeginProviderAttempt 推进的真实尝试计数, 不再使用取任务时的内存快照,避免"已尝试过但无候选可换"被误报为 route_unavailable。 queue 的 Fail() 不再用 JSON_MERGE_PATCH 把终态码写进已结束的尾条 尝试,已完成的尝试记录保持不可变;只有尚未结束的尾条尝试(租约、 快照解码等非 provider 阶段失败)才补写终态信息。 Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
455 lines
15 KiB
Go
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
|
|
}
|