254 lines
9.6 KiB
Go
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()
|
|
}
|