Files
chorus/portal/worker/worker.go
T

407 lines
14 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"
)
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)
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 {
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
}
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 || queueRepository == nil || controller == nil || catalog == nil || factory == nil || runtime == nil || storage == nil {
return nil, ErrInvalidConfig
}
if config.Random == nil {
config.Random = router.CryptoRandom{}
}
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)
if claim.Generation.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 claim.Generation.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 - claim.Generation.ProviderAttemptCount)
if remaining < len(candidates) {
candidates = candidates[:remaining]
}
for _, member := range candidates {
reservation, reserveErr := w.runtime.Reserve(ctx, router.ReservationRequest{
RoutePoolMemberID: member.RoutePoolMemberID, Capability: snapshot.Capability,
Owner: w.config.Owner, LeaseDuration: w.config.LeaseDuration,
})
if errors.Is(reserveErr, router.ErrMemberUnavailable) {
continue
}
if reserveErr != nil {
return w.fail(ctx, claim, provider.CodeUnknown)
}
selection, resolveErr := w.catalog.Resolve(ctx, member, snapshot.Capability)
if resolveErr != nil {
continue
}
attempt, begun, beginErr := w.queue.BeginProviderAttempt(ctx, claim.Generation.ID, claim.LeaseToken, member.RoutePoolMemberID, selection.ProviderModelID)
if beginErr != nil {
return beginErr
}
if !begun {
return nil
}
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 claim.Generation.ProviderAttemptCount > 0 {
return w.fail(ctx, claim, provider.ErrorCode(router.CodeFailoverExhausted))
}
return w.fail(ctx, claim, provider.ErrorCode(router.CodeRouteUnavailable))
}
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
}