From d2c4051dd0f8dc3b4d6e601f5a49e7d788f8583e Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Wed, 29 Jul 2026 21:11:08 +0000 Subject: [PATCH] Wire backend capacity into engine runtime --- capacity_contract_test.go | 466 ++++++++++++++++++++++++++++++++ engine.go | 41 ++- errors.go | 3 + internal/usecase/runner_test.go | 214 +++++++++++++++ types.go | 11 +- 5 files changed, 720 insertions(+), 15 deletions(-) create mode 100644 capacity_contract_test.go diff --git a/capacity_contract_test.go b/capacity_contract_test.go new file mode 100644 index 0000000..b4f0bea --- /dev/null +++ b/capacity_contract_test.go @@ -0,0 +1,466 @@ +package promptkit_test + +import ( + "context" + "errors" + "runtime" + "strings" + "sync" + "testing" + "time" + + "gitea.maximumdirect.net/eric/promptkit" +) + +func TestEngineLimitsInjectedClientConcurrency(t *testing.T) { + release := make(chan struct{}) + client := newCapacityGateClient(release, 8) + engine := newBackendCapacityEngine(t, client, 2, capacityInt(4), nil) + + results := make(chan capacityRunResult, 6) + for i := 0; i < 6; i++ { + go runCapacityRequest(engine, context.Background(), promptkit.RunRequest{ + PromptID: "prompt", + Execution: &promptkit.ExecutionTargetOverride{ + Endpoint: "http://request.example/v1", + }, + }, results) + } + + first := awaitCapacityRequest(t, client.started) + second := awaitCapacityRequest(t, client.started) + if first.Target.BackendID != "limited" || second.Target.BackendID != "limited" { + t.Fatalf("endpoint override changed backend pool: first=%q second=%q", + first.Target.BackendID, second.Target.BackendID) + } + if active, peak, _ := client.snapshot(); active != 2 || peak != 2 { + t.Fatalf("client concurrency before release=(active=%d peak=%d), want 2", active, peak) + } + + close(release) + for i := 0; i < 6; i++ { + outcome := awaitCapacityRun(t, results) + if outcome.err != nil || outcome.result == nil { + t.Fatalf("run outcome=(%+v, %v), want success", outcome.result, outcome.err) + } + } + if _, peak, calls := client.snapshot(); peak > 2 || calls != 6 { + t.Fatalf("client observations=(peak=%d calls=%d), want peak <= 2 and 6 calls", peak, calls) + } +} + +func TestEngineRejectsRunBeforeCompletionWhenAdmissionIsFull(t *testing.T) { + artifactRelease := make(chan struct{}) + reader := &capacityArtifactReader{ + entered: make(chan struct{}, 2), + release: artifactRelease, + } + client := newCapacityGateClient(closedCapacityChannel(), 2) + engine := newBackendCapacityEngine(t, client, 1, capacityInt(0), reader) + firstResult := make(chan capacityRunResult, 1) + go runCapacityRequest(engine, context.Background(), capacityInputRequest("http://first.example/v1"), firstResult) + + awaitCapacitySignal(t, reader.entered, "first artifact read") + + result, err := engine.Run(context.Background(), capacityInputRequest("http://second.example/v1")) + if result != nil { + t.Fatalf("capacity rejection returned partial result: %+v", result) + } + if !errors.Is(err, promptkit.ErrCapacityExceeded) { + t.Fatalf("capacity rejection=%v, want ErrCapacityExceeded", err) + } + if errors.Is(err, promptkit.ErrInvalidRequest) || errors.Is(err, promptkit.ErrLLMGenerate) { + t.Fatalf("capacity rejection had an unrelated category: %v", err) + } + if calls := reader.callCount(); calls != 1 { + t.Fatalf("artifact calls=%d, want only the admitted run", calls) + } + if _, _, calls := client.snapshot(); calls != 0 { + t.Fatalf("client calls=%d before admitted run was released, want 0", calls) + } + + close(artifactRelease) + outcome := awaitCapacityRun(t, firstResult) + if outcome.err != nil || outcome.result == nil { + t.Fatalf("first run outcome=(%+v, %v), want success", outcome.result, outcome.err) + } +} + +func TestBackendCapacityIsIndependentBetweenEngines(t *testing.T) { + firstRelease := make(chan struct{}) + firstClient := newCapacityGateClient(firstRelease, 1) + firstEngine := newBackendCapacityEngine(t, firstClient, 1, capacityInt(0), nil) + secondClient := newCapacityGateClient(closedCapacityChannel(), 1) + secondEngine := newBackendCapacityEngine(t, secondClient, 1, capacityInt(0), nil) + + firstResult := make(chan capacityRunResult, 1) + go runCapacityRequest(firstEngine, context.Background(), promptkit.RunRequest{PromptID: "prompt"}, firstResult) + awaitCapacityRequest(t, firstClient.started) + + result, err := secondEngine.Run(context.Background(), promptkit.RunRequest{PromptID: "prompt"}) + if err != nil || result == nil { + t.Fatalf("second engine run=(%+v, %v), want independent success", result, err) + } + if _, _, calls := secondClient.snapshot(); calls != 1 { + t.Fatalf("second engine client calls=%d, want 1", calls) + } + + close(firstRelease) + outcome := awaitCapacityRun(t, firstResult) + if outcome.err != nil || outcome.result == nil { + t.Fatalf("first engine run=(%+v, %v), want success", outcome.result, outcome.err) + } +} + +func TestUnlimitedBackendsRetainInjectedClientConcurrency(t *testing.T) { + tests := []struct { + name string + configure func(*testing.T, promptkit.LLMClient) *promptkit.Engine + }{ + { + name: "custom backend", + configure: func(t *testing.T, client promptkit.LLMClient) *promptkit.Engine { + return newBackendCapacityEngine(t, client, 0, nil, nil) + }, + }, + { + name: "endpoint-only profile", + configure: func(t *testing.T, client promptkit.LLMClient) *promptkit.Engine { + engine, err := promptkit.NewEngine(promptkit.Config{}, + promptkit.WithPromptFS(contractPromptFS("prompt", "profile", "message"), "."), + promptkit.WithProfiles(promptkit.Profile{ + ID: "profile", Endpoint: "http://endpoint.example/v1", Model: "model", + }), + promptkit.WithLLMClient(client), + ) + if err != nil { + t.Fatalf("construct endpoint-only engine: %v", err) + } + return engine + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + release := make(chan struct{}) + client := newCapacityGateClient(release, 2) + engine := tc.configure(t, client) + results := make(chan capacityRunResult, 2) + for i := 0; i < 2; i++ { + go runCapacityRequest( + engine, + context.Background(), + promptkit.RunRequest{PromptID: "prompt"}, + results, + ) + } + + awaitCapacityRequest(t, client.started) + awaitCapacityRequest(t, client.started) + if active, peak, _ := client.snapshot(); active != 2 || peak != 2 { + t.Fatalf("unlimited concurrency=(active=%d peak=%d), want 2", active, peak) + } + close(release) + for i := 0; i < 2; i++ { + outcome := awaitCapacityRun(t, results) + if outcome.err != nil || outcome.result == nil { + t.Fatalf("run outcome=(%+v, %v), want success", outcome.result, outcome.err) + } + } + }) + } +} + +func TestCancelWhileWaitingForGenerationCapacityPreservesErrorAndCapacity(t *testing.T) { + release := make(chan struct{}) + client := newCapacityGateClient(release, 3) + engine := newBackendCapacityEngine(t, client, 1, capacityInt(1), nil) + + firstResult := make(chan capacityRunResult, 1) + go runCapacityRequest(engine, context.Background(), promptkit.RunRequest{PromptID: "prompt"}, firstResult) + awaitCapacityRequest(t, client.started) + + waitingContext := newObservedCancelContext() + secondResult := make(chan capacityRunResult, 1) + go runCapacityRequest(engine, waitingContext, promptkit.RunRequest{PromptID: "prompt"}, secondResult) + awaitCapacitySignal(t, waitingContext.waiting, "generation-capacity wait") + waitingContext.cancel() + + outcome := awaitCapacityRun(t, secondResult) + if outcome.result != nil { + t.Fatalf("canceled capacity wait returned partial result: %+v", outcome.result) + } + if !errors.Is(outcome.err, context.Canceled) || !errors.Is(outcome.err, promptkit.ErrLLMGenerate) { + t.Fatalf("canceled capacity wait=%v, want context.Canceled and ErrLLMGenerate", outcome.err) + } + if _, _, calls := client.snapshot(); calls != 1 { + t.Fatalf("client calls=%d after canceled waiter, want only active call", calls) + } + + close(release) + first := awaitCapacityRun(t, firstResult) + if first.err != nil || first.result == nil { + t.Fatalf("first run=(%+v, %v), want success", first.result, first.err) + } + result, err := engine.Run(context.Background(), promptkit.RunRequest{PromptID: "prompt"}) + if err != nil || result == nil { + t.Fatalf("run after cancellation=(%+v, %v), want reusable capacity", result, err) + } + if _, _, calls := client.snapshot(); calls != 2 { + t.Fatalf("client calls=%d after reuse, want 2", calls) + } +} + +func TestCapacityExceededSentinelContract(t *testing.T) { + if promptkit.ErrCapacityExceeded == nil { + t.Fatal("ErrCapacityExceeded is nil") + } + for _, unrelated := range []error{ + promptkit.ErrInvalidConfig, + promptkit.ErrInvalidRequest, + promptkit.ErrLLMGenerate, + promptkit.ErrValidation, + } { + if errors.Is(promptkit.ErrCapacityExceeded, unrelated) || + errors.Is(unrelated, promptkit.ErrCapacityExceeded) { + t.Fatalf("ErrCapacityExceeded aliases unrelated sentinel %v", unrelated) + } + } +} + +type capacityRunResult struct { + result *promptkit.RunResult + err error +} + +func runCapacityRequest( + engine *promptkit.Engine, + ctx context.Context, + request promptkit.RunRequest, + results chan<- capacityRunResult, +) { + result, err := engine.Run(ctx, request) + results <- capacityRunResult{result: result, err: err} +} + +func newBackendCapacityEngine( + t *testing.T, + client promptkit.LLMClient, + limit int, + queueCapacity *int, + reader promptkit.ArtifactReader, +) *promptkit.Engine { + t.Helper() + promptFS := contractPromptFS("prompt", "profile", "message") + if reader != nil { + promptFS = contractInputPromptFS() + } + options := []promptkit.Option{ + promptkit.WithPromptFS(promptFS, "."), + promptkit.WithBackend(promptkit.Backend{ + ID: "limited", + Endpoint: "http://backend.example/v1", + ConcurrencyLimit: limit, + QueueCapacity: queueCapacity, + }), + promptkit.WithProfiles(promptkit.Profile{ + ID: "profile", BackendID: "limited", Model: "model", + }), + promptkit.WithLLMClient(client), + } + if reader != nil { + options = append(options, promptkit.WithArtifactReader(reader)) + } + engine, err := promptkit.NewEngine(promptkit.Config{}, options...) + if err != nil { + t.Fatalf("construct capacity engine: %v", err) + } + return engine +} + +func capacityInputRequest(endpoint string) promptkit.RunRequest { + return promptkit.RunRequest{ + PromptID: "input-prompt", + Inputs: map[string]promptkit.ArtifactRef{ + "input": promptkit.Inline("input"), + }, + Execution: &promptkit.ExecutionTargetOverride{Endpoint: endpoint}, + } +} + +type capacityGateClient struct { + mu sync.Mutex + active int + peak int + calls int + started chan promptkit.GenerateRequest + release <-chan struct{} +} + +func newCapacityGateClient(release <-chan struct{}, buffer int) *capacityGateClient { + return &capacityGateClient{ + started: make(chan promptkit.GenerateRequest, buffer), + release: release, + } +} + +func (c *capacityGateClient) Generate( + ctx context.Context, + request promptkit.GenerateRequest, +) (*promptkit.GenerateResponse, error) { + c.mu.Lock() + c.calls++ + c.active++ + if c.active > c.peak { + c.peak = c.active + } + c.mu.Unlock() + defer func() { + c.mu.Lock() + c.active-- + c.mu.Unlock() + }() + + c.started <- request + select { + case <-c.release: + return &promptkit.GenerateResponse{Content: "ok"}, nil + case <-ctx.Done(): + return nil, ctx.Err() + } +} + +func (c *capacityGateClient) snapshot() (active, peak, calls int) { + c.mu.Lock() + defer c.mu.Unlock() + return c.active, c.peak, c.calls +} + +type capacityArtifactReader struct { + mu sync.Mutex + calls int + entered chan struct{} + release <-chan struct{} +} + +func (r *capacityArtifactReader) Read( + ctx context.Context, + _ promptkit.ArtifactRef, +) (*promptkit.Artifact, error) { + r.mu.Lock() + r.calls++ + r.mu.Unlock() + r.entered <- struct{}{} + select { + case <-r.release: + return &promptkit.Artifact{Body: []byte("input")}, nil + case <-ctx.Done(): + return nil, ctx.Err() + } +} + +func (r *capacityArtifactReader) callCount() int { + r.mu.Lock() + defer r.mu.Unlock() + return r.calls +} + +type observedCancelContext struct { + done chan struct{} + waiting chan struct{} + once sync.Once +} + +func newObservedCancelContext() *observedCancelContext { + return &observedCancelContext{ + done: make(chan struct{}), + waiting: make(chan struct{}, 1), + } +} + +func (c *observedCancelContext) Deadline() (time.Time, bool) { + return time.Time{}, false +} + +func (c *observedCancelContext) Done() <-chan struct{} { + var callers [8]uintptr + frames := runtime.CallersFrames(callers[:runtime.Callers(2, callers[:])]) + for { + frame, more := frames.Next() + if strings.Contains(frame.Function, "internal/capacity.(*pool).acquire") { + select { + case c.waiting <- struct{}{}: + default: + } + break + } + if !more { + break + } + } + return c.done +} + +func (c *observedCancelContext) Err() error { + select { + case <-c.done: + return context.Canceled + default: + return nil + } +} + +func (c *observedCancelContext) Value(any) any { + return nil +} + +func (c *observedCancelContext) cancel() { + c.once.Do(func() { + close(c.done) + }) +} + +func awaitCapacityRequest( + t *testing.T, + requests <-chan promptkit.GenerateRequest, +) promptkit.GenerateRequest { + t.Helper() + select { + case request := <-requests: + return request + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for client invocation") + return promptkit.GenerateRequest{} + } +} + +func awaitCapacityRun(t *testing.T, results <-chan capacityRunResult) capacityRunResult { + t.Helper() + select { + case result := <-results: + return result + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for Run") + return capacityRunResult{} + } +} + +func awaitCapacitySignal(t *testing.T, signal <-chan struct{}, name string) { + t.Helper() + select { + case <-signal: + case <-time.After(5 * time.Second): + t.Fatalf("timed out waiting for %s", name) + } +} + +func capacityInt(value int) *int { + return &value +} + +func closedCapacityChannel() <-chan struct{} { + channel := make(chan struct{}) + close(channel) + return channel +} diff --git a/engine.go b/engine.go index 63a94db..71089ca 100644 --- a/engine.go +++ b/engine.go @@ -13,6 +13,7 @@ import ( artifactadapter "gitea.maximumdirect.net/eric/promptkit/internal/artifact" "gitea.maximumdirect.net/eric/promptkit/internal/backend" + "gitea.maximumdirect.net/eric/promptkit/internal/capacity" "gitea.maximumdirect.net/eric/promptkit/internal/defaults" "gitea.maximumdirect.net/eric/promptkit/internal/domain" "gitea.maximumdirect.net/eric/promptkit/internal/llm" @@ -62,6 +63,10 @@ var ( // ErrPromptRender identifies a failure to render prompt messages or the // session ID from the resolved inputs and variables. ErrPromptRender = errors.New("failed to render prompt") + // ErrCapacityExceeded identifies a Run rejected because the selected backend + // already admitted ConcurrencyLimit + QueueCapacity calls. It is not an + // invalid request, provider rate limit, or ErrLLMGenerate. + ErrCapacityExceeded = errors.New("backend capacity exceeded") // ErrLLMGenerate identifies a model-client failure or a nil successful // response. Errors returned by an injected LLMClient remain available // through errors.Is. @@ -75,7 +80,9 @@ var ( // Engine prepares and runs Promptkit prompt requests. // // An Engine is safe for concurrent calls to [Engine.Prepare] and [Engine.Run]. -// Injected collaborators may consequently be invoked concurrently. +// Each Engine owns independent backend-capacity pools that coordinate Run +// admission and model generation. Injected collaborators may still be invoked +// concurrently across different backend pools or for unlimited backends. type Engine struct { runner *usecase.Runner } @@ -141,8 +148,10 @@ type engineOptions struct { // WithLLMClient replaces the built-in model client used by [Engine.Run]. // -// A nil client makes NewEngine fail with ErrInvalidConfig. The client may be -// called concurrently and is not used by [Engine.Prepare]. +// A nil client makes NewEngine fail with ErrInvalidConfig. The Engine schedules +// Generate calls according to the selected backend's capacity policy, but the +// client may still be called concurrently across different backend pools or for +// unlimited backends. The client is not used by [Engine.Prepare]. func WithLLMClient(client LLMClient) Option { return optionFunc(func(options *engineOptions) error { if client == nil { @@ -313,8 +322,9 @@ func WithSchemaFile(path string) Option { // or Run needs them. // // NewEngine returns an error matching ErrInvalidConfig for invalid -// configuration or options. It does not perform model requests or require -// credentials. +// configuration, options, or backend-capacity policies. Each constructed +// Engine has independent backend-capacity pools. Construction does not perform +// model requests or require credentials. func NewEngine(cfg Config, opts ...Option) (*Engine, error) { var options engineOptions for _, opt := range opts { @@ -347,6 +357,11 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) { return nil, fmt.Errorf("%w: failed to construct backend registry: %v", ErrInvalidConfig, err) } + capacityManager, err := capacity.NewManager(backendRegistry.CapacityPolicies()) + if err != nil { + return nil, fmt.Errorf("%w: failed to construct backend capacity manager: %v", ErrInvalidConfig, err) + } + validator := options.validator if !options.validatorSource { schemaDir := cfg.SchemaDir @@ -367,6 +382,7 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) { return nil, fmt.Errorf("%w: %v", ErrInvalidConfig, err) } } + llmClient = capacity.NewClient(capacityManager, llmClient) artifacts := options.artifactReader if !options.artifactSource { @@ -382,7 +398,7 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) { prompt.NewGoRenderer(), llmClient, validator, - nil, + capacityManager, ), }, nil } @@ -451,11 +467,14 @@ func (e *Engine) Prepare(ctx context.Context, req RunRequest) (*PreparedRun, err // single-pass even when OutputContract.RepairAttempts is positive. // // Run can return every error category documented by [Engine.Prepare], plus -// ErrLLMGenerate. Errors from injected clients remain available through -// errors.Is. Cancellation is passed through the active collaborator and is -// reported in the applicable operation category; no general errors.Is -// relationship to ctx.Err is promised. A nil Engine returns ErrInvalidConfig. -// Run returns no partial result on error. +// ErrCapacityExceeded and ErrLLMGenerate. ErrCapacityExceeded identifies +// rejection before artifacts, schemas, rendering, or model generation because +// the selected backend's admission capacity is full; it does not match +// ErrInvalidRequest or ErrLLMGenerate. Errors from injected clients remain +// available through errors.Is. Cancellation while waiting for model-generation +// capacity matches both ErrLLMGenerate and the context error. Cancellation +// otherwise follows the active collaborator's documented behavior. A nil +// Engine returns ErrInvalidConfig. Run returns no partial result on error. func (e *Engine) Run(ctx context.Context, req RunRequest) (*RunResult, error) { if e == nil || e.runner == nil { return nil, fmt.Errorf("%w: engine is nil", ErrInvalidConfig) diff --git a/errors.go b/errors.go index 763b59c..9957eed 100644 --- a/errors.go +++ b/errors.go @@ -4,6 +4,7 @@ import ( "errors" "fmt" + "gitea.maximumdirect.net/eric/promptkit/internal/capacity" "gitea.maximumdirect.net/eric/promptkit/internal/profile" "gitea.maximumdirect.net/eric/promptkit/internal/promptdef" "gitea.maximumdirect.net/eric/promptkit/internal/usecase" @@ -38,6 +39,8 @@ func publicErrorFor(err error) error { return ErrProfileLoad case errors.Is(err, usecase.ErrAPIKeyEnvMissing): return errors.Join(ErrInvalidRequest, ErrAPIKeyEnvMissing) + case errors.Is(err, capacity.ErrCapacityExceeded): + return ErrCapacityExceeded case errors.Is(err, usecase.ErrArtifactLoad): return ErrArtifactLoad case errors.Is(err, usecase.ErrPromptRender): diff --git a/internal/usecase/runner_test.go b/internal/usecase/runner_test.go index ba94e7c..59e8a9b 100644 --- a/internal/usecase/runner_test.go +++ b/internal/usecase/runner_test.go @@ -10,7 +10,9 @@ import ( "reflect" "regexp" "strings" + "sync" "testing" + "time" "gitea.maximumdirect.net/eric/promptkit/internal/capacity" "gitea.maximumdirect.net/eric/promptkit/internal/defaults" @@ -39,6 +41,52 @@ type fakeBackendResolver struct { backends map[string]domain.Backend } +type staticPromptRepo struct { + def *domain.PromptDefinition +} + +func (r staticPromptRepo) GetPromptDefinition( + context.Context, + string, + string, +) (*domain.PromptDefinition, error) { + return r.def, nil +} + +type staticExecutionProfileRepo struct { + profile *domain.ExecutionProfile +} + +func (r staticExecutionProfileRepo) GetProfile( + context.Context, + string, +) (*domain.ExecutionProfile, error) { + value := *r.profile + return &value, nil +} + +type staticArtifactReader struct{} + +func (staticArtifactReader) Read( + context.Context, + domain.ArtifactRef, +) (*domain.Artifact, error) { + return &domain.Artifact{Body: []byte("artifact"), Hash: hashString("artifact")}, nil +} + +type staticRenderer struct{} + +func (staticRenderer) Render( + context.Context, + *domain.PromptDefinition, + map[string]*domain.Artifact, + map[string]string, +) (*domain.RenderedPrompt, error) { + return &domain.RenderedPrompt{ + Messages: []domain.RenderedMessage{{Role: "user", Content: "hello"}}, + }, nil +} + func (f fakeBackendResolver) GetBackend(id string) (domain.Backend, error) { value, ok := f.backends[id] if !ok { @@ -166,6 +214,81 @@ type fakeRunAdmitter struct { releaseCalls int } +type recordingRunAdmitter struct { + mu sync.Mutex + next RunAdmitter + backendIDs []string +} + +func (a *recordingRunAdmitter) Admit(ctx context.Context, backendID string) (func(), error) { + a.mu.Lock() + a.backendIDs = append(a.backendIDs, backendID) + a.mu.Unlock() + return a.next.Admit(ctx, backendID) +} + +func (a *recordingRunAdmitter) admittedBackendIDs() []string { + a.mu.Lock() + defer a.mu.Unlock() + return append([]string(nil), a.backendIDs...) +} + +type controlledRepairLLM struct { + mu sync.Mutex + active int + peak int + calls int + repairCalls int + events chan controlledGeneration + backendIDs []string +} + +type controlledGeneration struct { + release chan struct{} +} + +func (c *controlledRepairLLM) Generate( + ctx context.Context, + req domain.GenerateRequest, +) (*domain.GenerateResponse, error) { + isRepair := len(req.Prompt.Messages) > 0 && + strings.HasPrefix(req.Prompt.Messages[0].Content, "You repair invalid JSON") + c.mu.Lock() + c.calls++ + c.active++ + if c.active > c.peak { + c.peak = c.active + } + if isRepair { + c.repairCalls++ + } + c.backendIDs = append(c.backendIDs, req.Target.BackendID) + c.mu.Unlock() + defer func() { + c.mu.Lock() + c.active-- + c.mu.Unlock() + }() + + event := controlledGeneration{release: make(chan struct{})} + c.events <- event + select { + case <-event.release: + case <-ctx.Done(): + return nil, ctx.Err() + } + if isRepair { + return &domain.GenerateResponse{Content: `{}`}, nil + } + return &domain.GenerateResponse{Content: `{"broken":`}, nil +} + +func (c *controlledRepairLLM) snapshot() (peak, calls, repairCalls int, backendIDs []string) { + c.mu.Lock() + defer c.mu.Unlock() + return c.peak, c.calls, c.repairCalls, append([]string(nil), c.backendIDs...) +} + func (f *fakeRunAdmitter) Admit(_ context.Context, backendID string) (func(), error) { f.backendIDs = append(f.backendIDs, backendID) if f.err != nil { @@ -1957,6 +2080,97 @@ func TestRunnerRunStructuredRepairRemainsBoundedAndUsesEffectiveModelSettings(t } } +func TestRunnerSchedulesInitialAndRepairGenerationThroughOneBackendPool(t *testing.T) { + manager, err := capacity.NewManager(map[string]domain.BackendCapacityPolicy{ + "custom": { + ConcurrencyLimit: 1, + QueueCapacity: 1, + }, + }) + if err != nil { + t.Fatalf("construct capacity manager: %v", err) + } + baseClient := &controlledRepairLLM{ + events: make(chan controlledGeneration, 4), + } + scheduledClient := capacity.NewClient(manager, baseClient) + admitter := &recordingRunAdmitter{next: manager} + runner := NewRunnerWithRepairer( + staticPromptRepo{def: promptDef(domain.FormatJSON, domain.ValidationJSON, 1)}, + staticExecutionProfileRepo{profile: &domain.ExecutionProfile{ + ID: "exec", BackendID: "custom", Model: "model", + }}, + fakeBackendResolver{backends: map[string]domain.Backend{ + "custom": {ID: "custom", Endpoint: "http://backend.example/v1"}, + }}, + staticArtifactReader{}, + staticRenderer{}, + scheduledClient, + validate.NewStandardValidator("."), + NewDefaultOutputRepairer(scheduledClient), + admitter, + ) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + type runOutcome struct { + result *domain.RunResult + err error + } + outcomes := make(chan runOutcome, 2) + request := domain.RunRequest{ + PromptID: "p", + ProfileID: "exec", + Inputs: singleInputRef(), + } + for i := 0; i < 2; i++ { + go func() { + result, runErr := runner.Run(ctx, request) + outcomes <- runOutcome{result: result, err: runErr} + }() + } + + for i := 0; i < 4; i++ { + select { + case event := <-baseClient.events: + close(event.release) + case <-ctx.Done(): + t.Fatalf("timed out waiting for generation %d: %v", i+1, ctx.Err()) + } + } + for i := 0; i < 2; i++ { + select { + case outcome := <-outcomes: + if outcome.err != nil || outcome.result == nil { + t.Fatalf("run outcome=(%+v, %v), want success", outcome.result, outcome.err) + } + if outcome.result.Validation.Status != domain.ValidationPassed || + outcome.result.Validation.RepairAttempts != 1 { + t.Fatalf("unexpected repaired result: %+v", outcome.result.Validation) + } + case <-ctx.Done(): + t.Fatalf("timed out waiting for repaired run: %v", ctx.Err()) + } + } + + admitted := admitter.admittedBackendIDs() + if !reflect.DeepEqual(admitted, []string{"custom", "custom"}) { + t.Fatalf("admission backend IDs=%#v, want one admission per run", admitted) + } + peak, calls, repairCalls, backendIDs := baseClient.snapshot() + if peak != 1 || calls != 4 || repairCalls != 2 { + t.Fatalf( + "generation observations=(peak=%d calls=%d repairs=%d), want (1, 4, 2)", + peak, + calls, + repairCalls, + ) + } + if !reflect.DeepEqual(backendIDs, []string{"custom", "custom", "custom", "custom"}) { + t.Fatalf("generation backend IDs=%#v, want custom for initial and repair calls", backendIDs) + } +} + func TestRunnerRunRepairCarriesEffectiveSessionID(t *testing.T) { llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: `{"broken":`}} runner := NewRunnerWithRepairer( diff --git a/types.go b/types.go index c33f734..f26dc8d 100644 --- a/types.go +++ b/types.go @@ -561,10 +561,13 @@ type StructuredOutputJSONSpec struct { // LLMClient executes rendered prompts for [Engine.Run]. // -// Generate may be called concurrently. It must honor context cancellation to -// make Run responsive to cancellation. The request and all nested maps, -// slices, and pointers are client-owned copies and may be mutated or retained -// without affecting engine state. +// Generate is scheduled according to the resolved backend's capacity policy. +// It may still be called concurrently for different backend pools or unlimited +// backends. Cancellation while waiting for capacity can prevent Generate from +// being called. Once invoked, it must honor context cancellation to make Run +// responsive to cancellation. The request and all nested maps, slices, and +// pointers are client-owned copies and may be mutated or retained without +// affecting engine state. // // Generate receives rendered messages and may receive a direct API key. A // client must protect those values and any raw output in its logging, storage,