package capacity import ( "context" "errors" "reflect" "runtime" "sync" "sync/atomic" "testing" "time" "gitea.maximumdirect.net/eric/promptkit/internal/domain" "gitea.maximumdirect.net/eric/promptkit/internal/llm" ) type generateResult struct { response *domain.GenerateResponse err error } type clientFunc func( context.Context, domain.GenerateRequest, ) (*domain.GenerateResponse, error) func (f clientFunc) Generate( ctx context.Context, req domain.GenerateRequest, ) (*domain.GenerateResponse, error) { return f(ctx, req) } type blockingClient struct { mu sync.Mutex active int peak int calls map[string]int started chan string releases map[string]chan struct{} } func newBlockingClient(releases map[string]chan struct{}) *blockingClient { return &blockingClient{ calls: make(map[string]int), started: make(chan string, 64), releases: releases, } } func (c *blockingClient) Generate( ctx context.Context, req domain.GenerateRequest, ) (*domain.GenerateResponse, error) { id := req.Prompt.SessionID c.mu.Lock() c.active++ if c.active > c.peak { c.peak = c.active } c.calls[id]++ c.mu.Unlock() defer func() { c.mu.Lock() c.active-- c.mu.Unlock() }() c.started <- id if release := c.releases[id]; release != nil { select { case <-release: case <-ctx.Done(): return nil, ctx.Err() } } return &domain.GenerateResponse{Content: id}, nil } func (c *blockingClient) callCount(id string) int { c.mu.Lock() defer c.mu.Unlock() return c.calls[id] } func (c *blockingClient) peakConcurrency() int { c.mu.Lock() defer c.mu.Unlock() return c.peak } func generateAsync( client llm.Client, ctx context.Context, backendID string, id string, ) <-chan generateResult { result := make(chan generateResult, 1) go func() { response, err := client.Generate(ctx, domain.GenerateRequest{ Prompt: domain.RenderedPrompt{SessionID: id}, Target: domain.ExecutionTarget{BackendID: backendID}, }) result <- generateResult{response: response, err: err} }() return result } func waitForWaiterCount(t *testing.T, manager *Manager, backendID string, want int) { t.Helper() pool := manager.pools[backendID] deadline := time.Now().Add(2 * time.Second) for { pool.mu.Lock() got := pool.waiters.Len() pool.mu.Unlock() if got == want { return } if time.Now().After(deadline) { t.Fatalf("waiter count=%d, want %d", got, want) } runtime.Gosched() } } func receiveStarted(t *testing.T, started <-chan string) string { t.Helper() select { case id := <-started: return id case <-time.After(2 * time.Second): t.Fatal("timed out waiting for wrapped client invocation") return "" } } func receiveResult(t *testing.T, result <-chan generateResult) generateResult { t.Helper() select { case got := <-result: return got case <-time.After(2 * time.Second): t.Fatal("timed out waiting for generation result") return generateResult{} } } func newTestManager(t *testing.T, policies map[string]domain.BackendCapacityPolicy) *Manager { t.Helper() manager, err := NewManager(policies) if err != nil { t.Fatalf("construct manager: %v", err) } return manager } func TestClientLimitsPeakConcurrencyAndServesWaitersFIFO(t *testing.T) { manager := newTestManager(t, map[string]domain.BackendCapacityPolicy{ "limited": {ConcurrencyLimit: 1}, }) firstRelease := make(chan struct{}) secondRelease := make(chan struct{}) thirdRelease := make(chan struct{}) next := newBlockingClient(map[string]chan struct{}{ "first": firstRelease, "second": secondRelease, "third": thirdRelease, }) client := NewClient(manager, next) first := generateAsync(client, context.Background(), "limited", "first") if got := receiveStarted(t, next.started); got != "first" { t.Fatalf("first invocation=%q, want first", got) } second := generateAsync(client, context.Background(), "limited", "second") waitForWaiterCount(t, manager, "limited", 1) third := generateAsync(client, context.Background(), "limited", "third") waitForWaiterCount(t, manager, "limited", 2) close(firstRelease) if got := receiveResult(t, first); got.err != nil { t.Fatalf("first generation: %v", got.err) } if got := receiveStarted(t, next.started); got != "second" { t.Fatalf("second invocation=%q, want second", got) } close(secondRelease) if got := receiveResult(t, second); got.err != nil { t.Fatalf("second generation: %v", got.err) } if got := receiveStarted(t, next.started); got != "third" { t.Fatalf("third invocation=%q, want third", got) } close(thirdRelease) if got := receiveResult(t, third); got.err != nil { t.Fatalf("third generation: %v", got.err) } if peak := next.peakConcurrency(); peak != 1 { t.Fatalf("peak concurrency=%d, want 1", peak) } } func TestClientPeakConcurrencyDoesNotExceedConfiguredLimit(t *testing.T) { const limit = 2 manager := newTestManager(t, map[string]domain.BackendCapacityPolicy{ "limited": {ConcurrencyLimit: limit}, }) gate := make(chan struct{}) releases := make(map[string]chan struct{}) for i := range 5 { releases[string(rune('a'+i))] = gate } next := newBlockingClient(releases) client := NewClient(manager, next) results := make([]<-chan generateResult, 0, len(releases)) for id := range releases { results = append(results, generateAsync(client, context.Background(), "limited", id)) } for range limit { receiveStarted(t, next.started) } waitForWaiterCount(t, manager, "limited", len(releases)-limit) close(gate) for _, result := range results { if got := receiveResult(t, result); got.err != nil { t.Fatalf("generation: %v", got.err) } } if peak := next.peakConcurrency(); peak != limit { t.Fatalf("peak concurrency=%d, want %d", peak, limit) } } func TestClientRemovesCanceledWaiters(t *testing.T) { tests := []struct { name string cancelID string wantOrder []string }{ {name: "first waiter", cancelID: "one", wantOrder: []string{"two", "three"}}, {name: "middle waiter", cancelID: "two", wantOrder: []string{"one", "three"}}, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { manager := newTestManager(t, map[string]domain.BackendCapacityPolicy{ "limited": {ConcurrencyLimit: 1}, }) holderRelease := make(chan struct{}) releases := map[string]chan struct{}{ "holder": holderRelease, "one": make(chan struct{}), "two": make(chan struct{}), "three": make(chan struct{}), } next := newBlockingClient(releases) client := NewClient(manager, next) holder := generateAsync(client, context.Background(), "limited", "holder") if got := receiveStarted(t, next.started); got != "holder" { t.Fatalf("initial invocation=%q, want holder", got) } contexts := make(map[string]context.Context) cancels := make(map[string]context.CancelFunc) results := make(map[string]<-chan generateResult) for _, id := range []string{"one", "two", "three"} { contexts[id], cancels[id] = context.WithCancel(context.Background()) results[id] = generateAsync(client, contexts[id], "limited", id) waitForWaiterCount(t, manager, "limited", len(results)) } cancels[tc.cancelID]() if got := receiveResult(t, results[tc.cancelID]); !errors.Is(got.err, context.Canceled) { t.Fatalf("canceled waiter error=%v, want context.Canceled", got.err) } waitForWaiterCount(t, manager, "limited", 2) close(holderRelease) if got := receiveResult(t, holder); got.err != nil { t.Fatalf("holder generation: %v", got.err) } for _, id := range tc.wantOrder { if got := receiveStarted(t, next.started); got != id { t.Fatalf("next invocation=%q, want %q", got, id) } close(releases[id]) if got := receiveResult(t, results[id]); got.err != nil { t.Fatalf("%s generation: %v", id, got.err) } } if calls := next.callCount(tc.cancelID); calls != 0 { t.Fatalf("canceled waiter invoked wrapped client %d times", calls) } for _, cancel := range cancels { cancel() } }) } } func TestClientGrantCancellationRaceDoesNotLeakPermit(t *testing.T) { const iterations = 200 for i := range iterations { manager := newTestManager(t, map[string]domain.BackendCapacityPolicy{ "limited": {ConcurrencyLimit: 1}, }) holderRelease := make(chan struct{}) var waiterCalls atomic.Int64 next := clientFunc(func( _ context.Context, req domain.GenerateRequest, ) (*domain.GenerateResponse, error) { if req.Prompt.SessionID == "holder" { <-holderRelease } else if req.Prompt.SessionID == "waiter" { waiterCalls.Add(1) } return &domain.GenerateResponse{Content: req.Prompt.SessionID}, nil }) client := NewClient(manager, next) holder := generateAsync(client, context.Background(), "limited", "holder") waitForActiveCount(t, manager, "limited", 1) ctx, cancel := context.WithCancel(context.Background()) waiterResult := generateAsync(client, ctx, "limited", "waiter") waitForWaiterCount(t, manager, "limited", 1) start := make(chan struct{}) var race sync.WaitGroup race.Add(2) go func() { defer race.Done() <-start cancel() }() go func() { defer race.Done() <-start close(holderRelease) }() close(start) race.Wait() if got := receiveResult(t, holder); got.err != nil { t.Fatalf("iteration %d holder generation: %v", i, got.err) } got := receiveResult(t, waiterResult) switch calls := waiterCalls.Load(); { case calls == 0 && errors.Is(got.err, context.Canceled): case calls == 1 && got.err == nil: default: t.Fatalf("iteration %d waiter calls=%d error=%v", i, calls, got.err) } probe := generateAsync(client, context.Background(), "limited", "probe") if got := receiveResult(t, probe); got.err != nil { t.Fatalf("iteration %d probe generation: %v", i, got.err) } waitForActiveCount(t, manager, "limited", 0) waitForWaiterCount(t, manager, "limited", 0) } } func waitForActiveCount(t *testing.T, manager *Manager, backendID string, want int) { t.Helper() pool := manager.pools[backendID] deadline := time.Now().Add(2 * time.Second) for { pool.mu.Lock() got := pool.active pool.mu.Unlock() if got == want { return } if time.Now().After(deadline) { t.Fatalf("active count=%d, want %d", got, want) } runtime.Gosched() } } func TestClientUsesIndependentPoolsAndUnlimitedFastPaths(t *testing.T) { manager := newTestManager(t, map[string]domain.BackendCapacityPolicy{ "alpha": {ConcurrencyLimit: 1}, "beta": {ConcurrencyLimit: 1}, }) alphaRelease := make(chan struct{}) betaRelease := make(chan struct{}) next := newBlockingClient(map[string]chan struct{}{ "alpha": alphaRelease, "beta": betaRelease, }) client := NewClient(manager, next) alpha := generateAsync(client, context.Background(), "alpha", "alpha") beta := generateAsync(client, context.Background(), "beta", "beta") started := map[string]bool{ receiveStarted(t, next.started): true, receiveStarted(t, next.started): true, } if !started["alpha"] || !started["beta"] { t.Fatalf("independent pools did not both start: %#v", started) } close(alphaRelease) close(betaRelease) if got := receiveResult(t, alpha); got.err != nil { t.Fatalf("alpha generation: %v", got.err) } if got := receiveResult(t, beta); got.err != nil { t.Fatalf("beta generation: %v", got.err) } for _, backendID := range []string{"", "unknown"} { response, err := client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{SessionID: backendID}, Target: domain.ExecutionTarget{BackendID: backendID}, }) if err != nil || response == nil { t.Fatalf("unlimited backend %q response=(%#v, %v)", backendID, response, err) } } if got := NewClient(nil, next); got != next { t.Fatal("nil manager did not return the wrapped client unchanged") } } func TestClientPreservesRequestsResponsesAndErrors(t *testing.T) { manager := newTestManager(t, map[string]domain.BackendCapacityPolicy{ "limited": {ConcurrencyLimit: 1}, }) request := domain.GenerateRequest{ Prompt: domain.RenderedPrompt{ SessionID: "session", Messages: []domain.RenderedMessage{ {Role: "user", Content: "content"}, }, }, Target: domain.ExecutionTarget{ BackendID: "limited", Model: "model", ExtraParams: map[string]any{"key": "value"}, }, } response := &domain.GenerateResponse{ Content: "output", Usage: domain.TokenUsage{TotalTokens: 7}, } collaboratorErr := errors.New("collaborator failure") tests := []struct { name string response *domain.GenerateResponse err error }{ {name: "successful response", response: response}, {name: "nil response"}, {name: "collaborator error", response: response, err: collaboratorErr}, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { var captured domain.GenerateRequest next := clientFunc(func( _ context.Context, req domain.GenerateRequest, ) (*domain.GenerateResponse, error) { captured = req return tc.response, tc.err }) gotResponse, gotErr := NewClient(manager, next).Generate(context.Background(), request) if !reflect.DeepEqual(captured, request) { t.Fatalf("request changed: %#v", captured) } if gotResponse != tc.response || gotErr != tc.err { t.Fatalf("response=(%p, %v), want (%p, %v)", gotResponse, gotErr, tc.response, tc.err) } }) } } func TestClientReleasesPermitDuringPanicUnwinding(t *testing.T) { manager := newTestManager(t, map[string]domain.BackendCapacityPolicy{ "limited": {ConcurrencyLimit: 1}, }) var calls atomic.Int64 next := clientFunc(func( _ context.Context, _ domain.GenerateRequest, ) (*domain.GenerateResponse, error) { if calls.Add(1) == 1 { panic("test panic") } return &domain.GenerateResponse{Content: "recovered"}, nil }) client := NewClient(manager, next) request := domain.GenerateRequest{ Target: domain.ExecutionTarget{BackendID: "limited"}, } func() { defer func() { if recover() == nil { t.Fatal("expected wrapped client panic") } }() _, _ = client.Generate(context.Background(), request) }() response, err := client.Generate(context.Background(), request) if err != nil || response == nil || response.Content != "recovered" { t.Fatalf("generation after panic=(%#v, %v)", response, err) } }