518 lines
14 KiB
Go
518 lines
14 KiB
Go
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)
|
|
}
|
|
}
|