Add backend capacity manager
This commit is contained in:
@@ -523,7 +523,7 @@ unlimited by omission, and no runtime call is scheduled yet.
|
|||||||
|
|
||||||
## Stage 2 — Engine-Local Capacity Manager
|
## Stage 2 — Engine-Local Capacity Manager
|
||||||
|
|
||||||
**Status:** Pending.
|
**Status:** Complete.
|
||||||
|
|
||||||
### Goal
|
### Goal
|
||||||
|
|
||||||
|
|||||||
40
internal/capacity/client.go
Normal file
40
internal/capacity/client.go
Normal file
@@ -0,0 +1,40 @@
|
|||||||
|
package capacity
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||||
|
"gitea.maximumdirect.net/eric/promptkit/internal/llm"
|
||||||
|
)
|
||||||
|
|
||||||
|
type client struct {
|
||||||
|
manager *Manager
|
||||||
|
next llm.Client
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewClient wraps next with configured active-generation limits. A nil manager
|
||||||
|
// leaves next unchanged.
|
||||||
|
func NewClient(manager *Manager, next llm.Client) llm.Client {
|
||||||
|
if manager == nil {
|
||||||
|
return next
|
||||||
|
}
|
||||||
|
return &client{
|
||||||
|
manager: manager,
|
||||||
|
next: next,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *client) Generate(
|
||||||
|
ctx context.Context,
|
||||||
|
req domain.GenerateRequest,
|
||||||
|
) (*domain.GenerateResponse, error) {
|
||||||
|
pool := c.manager.getPool(req.Target.BackendID)
|
||||||
|
if pool == nil {
|
||||||
|
return c.next.Generate(ctx, req)
|
||||||
|
}
|
||||||
|
if err := pool.acquire(ctx); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer pool.releaseActive()
|
||||||
|
return c.next.Generate(ctx, req)
|
||||||
|
}
|
||||||
517
internal/capacity/client_test.go
Normal file
517
internal/capacity/client_test.go
Normal file
@@ -0,0 +1,517 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
160
internal/capacity/manager.go
Normal file
160
internal/capacity/manager.go
Normal file
@@ -0,0 +1,160 @@
|
|||||||
|
// Package capacity coordinates engine-local run admission and model-generation
|
||||||
|
// concurrency for configured backends.
|
||||||
|
package capacity
|
||||||
|
|
||||||
|
import (
|
||||||
|
"container/list"
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ErrCapacityExceeded identifies an admission rejected because a backend's
|
||||||
|
// configured run capacity is full.
|
||||||
|
var ErrCapacityExceeded = errors.New("backend capacity exceeded")
|
||||||
|
|
||||||
|
// Manager owns independent backend capacity pools with immutable limits.
|
||||||
|
type Manager struct {
|
||||||
|
pools map[string]*pool
|
||||||
|
}
|
||||||
|
|
||||||
|
type pool struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
concurrencyLimit int
|
||||||
|
totalCapacity int
|
||||||
|
admitted int
|
||||||
|
active int
|
||||||
|
waiters list.List
|
||||||
|
}
|
||||||
|
|
||||||
|
type waiter struct {
|
||||||
|
ready chan struct{}
|
||||||
|
element *list.Element
|
||||||
|
granted bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewManager constructs independent pools from normalized backend policies.
|
||||||
|
func NewManager(policies map[string]domain.BackendCapacityPolicy) (*Manager, error) {
|
||||||
|
manager := &Manager{
|
||||||
|
pools: make(map[string]*pool, len(policies)),
|
||||||
|
}
|
||||||
|
maxInt := int(^uint(0) >> 1)
|
||||||
|
for id, policy := range policies {
|
||||||
|
if strings.TrimSpace(id) == "" {
|
||||||
|
return nil, errors.New("backend capacity policy ID must not be blank")
|
||||||
|
}
|
||||||
|
if policy.ConcurrencyLimit <= 0 {
|
||||||
|
return nil, fmt.Errorf(
|
||||||
|
"backend %q concurrency limit must be positive",
|
||||||
|
id,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if policy.QueueCapacity < 0 {
|
||||||
|
return nil, fmt.Errorf(
|
||||||
|
"backend %q queue capacity must not be negative",
|
||||||
|
id,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if policy.QueueCapacity > maxInt-policy.ConcurrencyLimit {
|
||||||
|
return nil, fmt.Errorf("backend %q total capacity overflows int", id)
|
||||||
|
}
|
||||||
|
manager.pools[id] = &pool{
|
||||||
|
concurrencyLimit: policy.ConcurrencyLimit,
|
||||||
|
totalCapacity: policy.ConcurrencyLimit + policy.QueueCapacity,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return manager, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Admit immediately reserves one configured backend run slot. Backends without
|
||||||
|
// a configured pool are unlimited.
|
||||||
|
func (m *Manager) Admit(ctx context.Context, backendID string) (func(), error) {
|
||||||
|
pool := m.getPool(backendID)
|
||||||
|
if pool == nil {
|
||||||
|
return releaseNothing, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
pool.mu.Lock()
|
||||||
|
defer pool.mu.Unlock()
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if pool.admitted >= pool.totalCapacity {
|
||||||
|
return nil, ErrCapacityExceeded
|
||||||
|
}
|
||||||
|
pool.admitted++
|
||||||
|
|
||||||
|
var once sync.Once
|
||||||
|
return func() {
|
||||||
|
once.Do(func() {
|
||||||
|
pool.mu.Lock()
|
||||||
|
pool.admitted--
|
||||||
|
pool.mu.Unlock()
|
||||||
|
})
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func releaseNothing() {}
|
||||||
|
|
||||||
|
func (m *Manager) getPool(backendID string) *pool {
|
||||||
|
if m == nil || backendID == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return m.pools[backendID]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *pool) acquire(ctx context.Context) error {
|
||||||
|
p.mu.Lock()
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
p.mu.Unlock()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if p.active < p.concurrencyLimit && p.waiters.Len() == 0 {
|
||||||
|
p.active++
|
||||||
|
p.mu.Unlock()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
waiter := &waiter{ready: make(chan struct{})}
|
||||||
|
waiter.element = p.waiters.PushBack(waiter)
|
||||||
|
p.mu.Unlock()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-waiter.ready:
|
||||||
|
return nil
|
||||||
|
case <-ctx.Done():
|
||||||
|
p.mu.Lock()
|
||||||
|
if !waiter.granted {
|
||||||
|
p.waiters.Remove(waiter.element)
|
||||||
|
waiter.element = nil
|
||||||
|
p.mu.Unlock()
|
||||||
|
return ctx.Err()
|
||||||
|
}
|
||||||
|
p.mu.Unlock()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *pool) releaseActive() {
|
||||||
|
var ready chan struct{}
|
||||||
|
|
||||||
|
p.mu.Lock()
|
||||||
|
if element := p.waiters.Front(); element != nil {
|
||||||
|
waiter := element.Value.(*waiter)
|
||||||
|
p.waiters.Remove(element)
|
||||||
|
waiter.element = nil
|
||||||
|
waiter.granted = true
|
||||||
|
ready = waiter.ready
|
||||||
|
} else {
|
||||||
|
p.active--
|
||||||
|
}
|
||||||
|
p.mu.Unlock()
|
||||||
|
|
||||||
|
if ready != nil {
|
||||||
|
close(ready)
|
||||||
|
}
|
||||||
|
}
|
||||||
158
internal/capacity/manager_test.go
Normal file
158
internal/capacity/manager_test.go
Normal file
@@ -0,0 +1,158 @@
|
|||||||
|
package capacity
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNewManagerRejectsInvalidPolicies(t *testing.T) {
|
||||||
|
maxInt := int(^uint(0) >> 1)
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
id string
|
||||||
|
policy domain.BackendCapacityPolicy
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "blank ID",
|
||||||
|
id: " \t ",
|
||||||
|
policy: domain.BackendCapacityPolicy{ConcurrencyLimit: 1},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "zero concurrency",
|
||||||
|
id: "backend",
|
||||||
|
policy: domain.BackendCapacityPolicy{},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "negative concurrency",
|
||||||
|
id: "backend",
|
||||||
|
policy: domain.BackendCapacityPolicy{ConcurrencyLimit: -1},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "negative queue",
|
||||||
|
id: "backend",
|
||||||
|
policy: domain.BackendCapacityPolicy{
|
||||||
|
ConcurrencyLimit: 1,
|
||||||
|
QueueCapacity: -1,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "total overflow",
|
||||||
|
id: "backend",
|
||||||
|
policy: domain.BackendCapacityPolicy{
|
||||||
|
ConcurrencyLimit: maxInt,
|
||||||
|
QueueCapacity: 1,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
_, err := NewManager(map[string]domain.BackendCapacityPolicy{
|
||||||
|
tc.id: tc.policy,
|
||||||
|
})
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected invalid policy error")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestManagerAdmissionIsBoundedAndReleaseIsIdempotent(t *testing.T) {
|
||||||
|
policies := map[string]domain.BackendCapacityPolicy{
|
||||||
|
"limited": {
|
||||||
|
ConcurrencyLimit: 2,
|
||||||
|
QueueCapacity: 1,
|
||||||
|
},
|
||||||
|
"independent": {
|
||||||
|
ConcurrencyLimit: 1,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
manager, err := NewManager(policies)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("construct manager: %v", err)
|
||||||
|
}
|
||||||
|
policies["limited"] = domain.BackendCapacityPolicy{
|
||||||
|
ConcurrencyLimit: 100,
|
||||||
|
QueueCapacity: 100,
|
||||||
|
}
|
||||||
|
|
||||||
|
releases := make([]func(), 0, 3)
|
||||||
|
for range 3 {
|
||||||
|
release, err := manager.Admit(context.Background(), "limited")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("admit within configured capacity: %v", err)
|
||||||
|
}
|
||||||
|
releases = append(releases, release)
|
||||||
|
}
|
||||||
|
if release, err := manager.Admit(context.Background(), "limited"); release != nil ||
|
||||||
|
!errors.Is(err, ErrCapacityExceeded) {
|
||||||
|
t.Fatalf("admission beyond capacity=(release=%t, err=%v), want ErrCapacityExceeded",
|
||||||
|
release != nil, err)
|
||||||
|
}
|
||||||
|
independentRelease, err := manager.Admit(context.Background(), "independent")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("admit independent backend while first is full: %v", err)
|
||||||
|
}
|
||||||
|
independentRelease()
|
||||||
|
|
||||||
|
releases[0]()
|
||||||
|
releases[0]()
|
||||||
|
replacement, err := manager.Admit(context.Background(), "limited")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("admit after release: %v", err)
|
||||||
|
}
|
||||||
|
replacement()
|
||||||
|
releases[1]()
|
||||||
|
releases[2]()
|
||||||
|
|
||||||
|
pool := manager.pools["limited"]
|
||||||
|
pool.mu.Lock()
|
||||||
|
admitted := pool.admitted
|
||||||
|
pool.mu.Unlock()
|
||||||
|
if admitted != 0 {
|
||||||
|
t.Fatalf("admitted runs after releases=%d, want 0", admitted)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestManagerAdmissionHonorsContextAndUnlimitedBackends(t *testing.T) {
|
||||||
|
manager, err := NewManager(map[string]domain.BackendCapacityPolicy{
|
||||||
|
"limited": {ConcurrencyLimit: 1},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("construct manager: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
cancel()
|
||||||
|
if release, err := manager.Admit(ctx, "limited"); release != nil ||
|
||||||
|
!errors.Is(err, context.Canceled) {
|
||||||
|
t.Fatalf("canceled limited admission=(release=%t, err=%v), want context cancellation",
|
||||||
|
release != nil, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var nilManager *Manager
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
manager *Manager
|
||||||
|
backendID string
|
||||||
|
}{
|
||||||
|
{name: "nil manager", manager: nilManager, backendID: "limited"},
|
||||||
|
{name: "blank ID", manager: manager},
|
||||||
|
{name: "unknown ID", manager: manager, backendID: "unknown"},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
release, err := tc.manager.Admit(ctx, tc.backendID)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unlimited admission: %v", err)
|
||||||
|
}
|
||||||
|
if release == nil {
|
||||||
|
t.Fatal("unlimited admission returned nil release")
|
||||||
|
}
|
||||||
|
release()
|
||||||
|
release()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user