Files
chorus/portal/worker/worker_test.go
T

254 lines
9.6 KiB
Go

package worker
import (
"bytes"
"context"
"errors"
"image"
"image/color"
"image/png"
"io"
"testing"
"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"
)
type fakeQueue struct {
claim *queue.Claim
inputs []model.GenerationInput
begun []model.Attempt
finished []model.Attempt
succeeded []model.GenerationOutput
failedCode string
staleSucceed bool
}
func (q *fakeQueue) ClaimNext(context.Context, string, time.Duration) (*queue.Claim, error) {
claim := q.claim
q.claim = nil
return claim, nil
}
func (q *fakeQueue) BeginProviderAttempt(_ context.Context, _ uint64, _ string, memberID, providerModelID uint64) (model.Attempt, bool, error) {
started := time.Now()
attempt := model.Attempt{Type: "provider", RouteMemberID: memberID, ProviderModelID: providerModelID, ProviderOrdinal: uint32(len(q.begun) + 1), StartedAt: &started}
q.begun = append(q.begun, attempt)
return attempt, true, nil
}
func (q *fakeQueue) FinishProviderAttempt(_ context.Context, _ uint64, _ string, attempt model.Attempt) (bool, error) {
q.finished = append(q.finished, attempt)
return true, nil
}
func (q *fakeQueue) Succeed(_ context.Context, _ uint64, _ string, outputs []model.GenerationOutput, _ model.Attempt) (bool, error) {
if q.staleSucceed {
return false, nil
}
q.succeeded = append(q.succeeded, outputs...)
return true, nil
}
func (q *fakeQueue) Fail(_ context.Context, _ uint64, _ string, code, _ string, _ model.Attempt) (bool, error) {
q.failedCode = code
return true, nil
}
func (q *fakeQueue) Inputs(context.Context, uint64) ([]model.GenerationInput, error) {
return q.inputs, nil
}
type fakeController struct{ stopped bool }
func (c *fakeController) StopClaims() { c.stopped = true }
type fakeRuntime struct {
states []router.MemberState
unavailable map[uint64]bool
reservations []router.ReservationRequest
observations []router.CircuitObservation
}
func (r *fakeRuntime) MemberStates(context.Context, model.Capability, []router.MemberSnapshot) ([]router.MemberState, error) {
return append([]router.MemberState(nil), r.states...), nil
}
func (r *fakeRuntime) Reserve(_ context.Context, request router.ReservationRequest) (router.Reservation, error) {
r.reservations = append(r.reservations, request)
if r.unavailable[request.RoutePoolMemberID] {
return router.Reservation{}, router.ErrMemberUnavailable
}
return router.Reservation{RoutePoolMemberID: request.RoutePoolMemberID, State: router.CircuitClosed}, nil
}
func (r *fakeRuntime) Record(_ context.Context, _ router.Reservation, observation router.CircuitObservation) (bool, error) {
r.observations = append(r.observations, observation)
return true, nil
}
type fakeCatalog struct{ selections map[uint64]Selection }
func (c fakeCatalog) Resolve(_ context.Context, member router.MemberSnapshot, _ model.Capability) (Selection, error) {
selection, ok := c.selections[member.RoutePoolMemberID]
if !ok {
return Selection{}, ErrNoProvider
}
return selection, nil
}
type fakeFactory struct {
clients []provider.Client
index int
}
func (f *fakeFactory) New(Selection) (provider.Client, error) {
if f.index >= len(f.clients) {
return nil, errors.New("no fake provider")
}
client := f.clients[f.index]
f.index++
return client, nil
}
type fakeProvider struct {
outputs []provider.Output
err error
}
func (p fakeProvider) Generate(context.Context, provider.Request) ([]provider.Output, error) {
return p.outputs, p.err
}
type fakeStore struct {
deleted []string
}
func (s *fakeStore) Open(context.Context, string) (io.ReadCloser, corestorage.Object, error) {
return nil, corestorage.Object{}, errors.New("unexpected input read")
}
func (s *fakeStore) PutImage(_ context.Context, request ImageRequest) (ImageObjects, error) {
return ImageObjects{
Original: corestorage.Object{Key: request.Key, OwnerID: request.OwnerID, GenerationID: request.GenerationID, ContentType: request.ContentType, Size: 1},
Thumbnail: corestorage.Object{Key: request.ThumbnailKey, OwnerID: request.OwnerID, GenerationID: request.GenerationID, ContentType: "image/png", Size: 1},
}, nil
}
func (s *fakeStore) Delete(_ context.Context, key string) error {
s.deleted = append(s.deleted, key)
return nil
}
type firstRandom struct{}
func (firstRandom) Uint64n(uint64) uint64 { return 0 }
func TestWorkerRetriesOnlyRetryableProviderFailure(t *testing.T) {
q := newQueue(t, model.KindText, model.CapabilityText, 1)
runtime := &fakeRuntime{states: activeStates(1, 2)}
factory := &fakeFactory{clients: []provider.Client{
fakeProvider{err: &provider.Error{Code: provider.CodeRateLimited, Class: provider.FailureRateLimited}},
fakeProvider{outputs: []provider.Output{{Kind: model.KindText, Text: "done"}}},
}}
w := newWorker(t, q, runtime, fakeCatalog{selections: selections()}, factory, &fakeStore{})
worked, err := w.ProcessOne(context.Background())
if err != nil || !worked || len(q.begun) != 2 || len(q.finished) != 2 || len(q.succeeded) != 1 || q.failedCode != "" {
t.Fatalf("worked=%v err=%v begun=%#v finished=%#v succeeded=%#v failed=%s", worked, err, q.begun, q.finished, q.succeeded, q.failedCode)
}
if len(runtime.observations) != 2 || !runtime.observations[0].Retryable || !runtime.observations[1].Succeeded {
t.Fatalf("runtime observations=%#v", runtime.observations)
}
}
func TestWorkerStopsOnUnauthorizedAndOpensCircuit(t *testing.T) {
q := newQueue(t, model.KindText, model.CapabilityText, 1)
runtime := &fakeRuntime{states: activeStates(1, 2)}
factory := &fakeFactory{clients: []provider.Client{fakeProvider{err: &provider.Error{Code: provider.CodeUnauthorized, Class: provider.FailureUnauthorized}}}}
w := newWorker(t, q, runtime, fakeCatalog{selections: selections()}, factory, &fakeStore{})
_, err := w.ProcessOne(context.Background())
if err != nil || len(q.begun) != 1 || q.failedCode != string(provider.CodeUnauthorized) {
t.Fatalf("err=%v begun=%#v failed=%s", err, q.begun, q.failedCode)
}
if len(runtime.observations) != 1 || !runtime.observations[0].OpenImmediately || runtime.observations[0].Retryable {
t.Fatalf("runtime observations=%#v", runtime.observations)
}
}
func TestWorkerCleansStagedOutputAfterStaleSucceed(t *testing.T) {
q := newQueue(t, model.KindImage, model.CapabilityImageGenerate, 0)
q.staleSucceed = true
runtime := &fakeRuntime{states: activeStates(1)}
store := &fakeStore{}
w := newWorker(t, q, runtime, fakeCatalog{selections: map[uint64]Selection{1: {ProviderModelID: 101, APIType: model.APIImages, ModelID: "image"}}}, &fakeFactory{clients: []provider.Client{fakeProvider{outputs: []provider.Output{{Kind: model.KindImage, Content: []byte("image"), ContentType: "image/png"}}}}}, store)
_, err := w.ProcessOne(context.Background())
if err != nil || len(q.succeeded) != 0 || len(store.deleted) != 2 {
t.Fatalf("err=%v succeeded=%#v deleted=%#v", err, q.succeeded, store.deleted)
}
}
func TestWorkerFailsWhenRouteSnapshotIsMissing(t *testing.T) {
q := &fakeQueue{claim: &queue.Claim{Generation: model.Generation{ID: 10, UserID: 1, Kind: model.KindText, Attempts: []byte("[]")}, LeaseToken: "token"}}
w := newWorker(t, q, &fakeRuntime{}, fakeCatalog{}, &fakeFactory{}, &fakeStore{})
_, err := w.ProcessOne(context.Background())
if err != nil || q.failedCode != string(provider.CodeUnknown) {
t.Fatalf("err=%v failed=%s", err, q.failedCode)
}
}
func newQueue(t *testing.T, kind model.GenerationKind, capability model.Capability, maxFailover uint16) *fakeQueue {
t.Helper()
snapshot := router.RouteSnapshot{
Capability: capability, RoutePoolID: 1, RoutePoolVersion: 1, PromptTemplateID: 1,
PromptTemplateKey: "test", PromptTemplateVersion: 1, MaxFailover: maxFailover,
Members: []router.MemberSnapshot{
{RoutePoolMemberID: 1, ProviderModelID: 101, Weight: 1, FailureThreshold: 2, OpenSeconds: 60, HalfOpenMax: 1},
{RoutePoolMemberID: 2, ProviderModelID: 102, Weight: 1, FailureThreshold: 2, OpenSeconds: 60, HalfOpenMax: 1},
},
}
if maxFailover == 0 {
snapshot.Members = snapshot.Members[:1]
}
encoded, err := snapshot.Encode()
if err != nil {
t.Fatal(err)
}
return &fakeQueue{claim: &queue.Claim{Generation: model.Generation{
ID: 10, UserID: 1, Kind: kind, RenderedPrompt: "prompt", RouteSnapshot: encoded,
Attempts: []byte("[]"), AttemptCount: 1,
}, LeaseToken: "token"}}
}
func activeStates(memberIDs ...uint64) []router.MemberState {
states := make([]router.MemberState, 0, len(memberIDs))
for _, memberID := range memberIDs {
states = append(states, router.MemberState{RoutePoolMemberID: memberID, Enabled: true, ProviderEnabled: true, ModelEnabled: true, SupportsCapability: true, CircuitState: router.CircuitClosed, HalfOpenMax: 1})
}
return states
}
func selections() map[uint64]Selection {
return map[uint64]Selection{
1: {ProviderModelID: 101, APIType: model.APIChat, ModelID: "first"},
2: {ProviderModelID: 102, APIType: model.APIChat, ModelID: "second"},
}
}
func newWorker(t *testing.T, q Queue, runtime router.RuntimeRepository, catalog Catalog, factory ClientFactory, store ImageStore) *Worker {
t.Helper()
w, err := New(Config{Owner: "test", LeaseDuration: time.Second, PollInterval: time.Millisecond, Random: firstRandom{}}, q, &fakeController{}, catalog, factory, runtime, store)
if err != nil {
t.Fatal(err)
}
return w
}
func testPNG() []byte {
img := image.NewRGBA(image.Rect(0, 0, 4, 4))
for y := 0; y < 4; y++ {
for x := 0; x < 4; x++ {
img.Set(x, y, color.RGBA{R: 200, G: 80, B: 40, A: 255})
}
}
var output bytes.Buffer
if err := png.Encode(&output, img); err != nil {
panic(err)
}
return output.Bytes()
}