From 861da355d840dc6c1c1892f83df5393c8a005b6b Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Wed, 29 Jul 2026 21:01:37 +0000 Subject: [PATCH] Add early backend run admission --- docs/roadmap/implementation.md | 2 +- engine.go | 1 + internal/usecase/runner.go | 122 +++++++++-- internal/usecase/runner_test.go | 372 +++++++++++++++++++++++++++----- 4 files changed, 425 insertions(+), 72 deletions(-) diff --git a/docs/roadmap/implementation.md b/docs/roadmap/implementation.md index d558fc9..dcb0b55 100644 --- a/docs/roadmap/implementation.md +++ b/docs/roadmap/implementation.md @@ -586,7 +586,7 @@ are independent, and the wrapper is transparent apart from waiting. ## Stage 3 — Shared Preparation And Early Run Admission -**Status:** Pending. +**Status:** Complete. ### Goal diff --git a/engine.go b/engine.go index 4f6148c..63a94db 100644 --- a/engine.go +++ b/engine.go @@ -382,6 +382,7 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) { prompt.NewGoRenderer(), llmClient, validator, + nil, ), }, nil } diff --git a/internal/usecase/runner.go b/internal/usecase/runner.go index 607d9dd..23738a5 100644 --- a/internal/usecase/runner.go +++ b/internal/usecase/runner.go @@ -14,6 +14,7 @@ import ( "unicode" "gitea.maximumdirect.net/eric/promptkit/internal/artifact" + "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" @@ -46,6 +47,7 @@ type Runner struct { llm llm.Client validator validate.Validator repairer OutputRepairer + admitter RunAdmitter } // BackendResolver resolves one normalized backend ID. @@ -53,6 +55,22 @@ type BackendResolver interface { GetBackend(string) (domain.Backend, error) } +// RunAdmitter reserves capacity for one resolved backend run. +type RunAdmitter interface { + Admit(context.Context, string) (func(), error) +} + +type preparationState struct { + definition *domain.PromptDefinition + directSessionID string + promptDefinitionHash string + selectedProfileID string + effectiveModel domain.ExecutionTarget + targetPresence domain.ExecutionTargetPresence + effectiveContract domain.OutputContract + start time.Time +} + func NewRunner( promptDefs promptdef.Repository, profiles profile.Repository, @@ -61,8 +79,19 @@ func NewRunner( renderer prompt.Renderer, llmClient llm.Client, validator validate.Validator, + admitter RunAdmitter, ) *Runner { - return NewRunnerWithRepairer(promptDefs, profiles, backends, artifacts, renderer, llmClient, validator, nil) + return NewRunnerWithRepairer( + promptDefs, + profiles, + backends, + artifacts, + renderer, + llmClient, + validator, + nil, + admitter, + ) } func NewRunnerWithRepairer( @@ -74,6 +103,7 @@ func NewRunnerWithRepairer( llmClient llm.Client, validator validate.Validator, repairer OutputRepairer, + admitter RunAdmitter, ) *Runner { return &Runner{ promptDefs: promptDefs, @@ -84,6 +114,7 @@ func NewRunnerWithRepairer( llm: llmClient, validator: validator, repairer: repairer, + admitter: admitter, } } @@ -95,7 +126,27 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes start := time.Now().UTC() - prepared, err := r.Prepare(ctx, req) + state, err := r.resolvePreparation(ctx, req, time.Now().UTC()) + if err != nil { + return nil, err + } + + if r.admitter != nil { + release, admitErr := r.admitter.Admit(ctx, state.effectiveModel.BackendID) + if admitErr != nil { + if errors.Is(admitErr, capacity.ErrCapacityExceeded) { + return nil, fmt.Errorf( + "backend %q admission: %w", + state.effectiveModel.BackendID, + admitErr, + ) + } + return nil, admitErr + } + defer release() + } + + prepared, err := r.completePreparation(ctx, req, state) if err != nil { return nil, err } @@ -177,6 +228,18 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes } func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.PreparedRun, error) { + state, err := r.resolvePreparation(ctx, req, time.Now().UTC()) + if err != nil { + return nil, err + } + return r.completePreparation(ctx, req, state) +} + +func (r *Runner) resolvePreparation( + ctx context.Context, + req domain.RunRequest, + start time.Time, +) (*preparationState, error) { if strings.TrimSpace(req.PromptID) == "" { return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidRequest) } @@ -185,8 +248,6 @@ func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.Pr return nil, fmt.Errorf("%w: session_id: %v", ErrInvalidRequest, err) } - start := time.Now().UTC() - def, err := r.promptDefs.GetPromptDefinition(ctx, req.PromptID, req.PromptVersion) if err != nil { return nil, fmt.Errorf("%w: %w", ErrPromptLoad, err) @@ -238,7 +299,28 @@ func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.Pr } effectiveContract := resolveOutputContract(def, req.Validation) - structuredOutput, err := r.resolveStructuredOutput(ctx, def, effectiveContract) + return &preparationState{ + definition: def, + directSessionID: directSessionID, + promptDefinitionHash: promptDefinitionHash, + selectedProfileID: selectedProfileID, + effectiveModel: effectiveModel, + targetPresence: targetPresence, + effectiveContract: effectiveContract, + start: start, + }, nil +} + +func (r *Runner) completePreparation( + ctx context.Context, + req domain.RunRequest, + state *preparationState, +) (*domain.PreparedRun, error) { + structuredOutput, err := r.resolveStructuredOutput( + ctx, + state.definition, + state.effectiveContract, + ) if err != nil { return nil, err } @@ -257,9 +339,9 @@ func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.Pr inputHashes[name] = art.Hash } - definitionToRender := def - if directSessionID != "" { - definitionCopy := *def + definitionToRender := state.definition + if state.directSessionID != "" { + definitionCopy := *state.definition definitionCopy.SessionID = "" definitionToRender = &definitionCopy } @@ -267,28 +349,28 @@ func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.Pr if err != nil { return nil, fmt.Errorf("%w: %w", ErrPromptRender, err) } - if directSessionID != "" { - renderedPrompt.SessionID = directSessionID + if state.directSessionID != "" { + renderedPrompt.SessionID = state.directSessionID } end := time.Now().UTC() return &domain.PreparedRun{ - PromptID: def.ID, - PromptVersion: def.Version, - PromptHash: promptDefinitionHash, - SelectedProfileID: selectedProfileID, - SelectedBackendID: effectiveModel.BackendID, - EffectiveModelParams: effectiveModel, - TargetPresence: targetPresence, - OutputContract: effectiveContract, + PromptID: state.definition.ID, + PromptVersion: state.definition.Version, + PromptHash: state.promptDefinitionHash, + SelectedProfileID: state.selectedProfileID, + SelectedBackendID: state.effectiveModel.BackendID, + EffectiveModelParams: state.effectiveModel, + TargetPresence: state.targetPresence, + OutputContract: state.effectiveContract, StructuredOutput: structuredOutput, InputHashes: inputHashes, SessionID: renderedPrompt.SessionID, RenderedPromptHash: hashRenderedPrompt(*renderedPrompt), Messages: renderedPrompt.Messages, - StartTime: start, + StartTime: state.start, EndTime: end, - DurationMS: end.Sub(start).Milliseconds(), + DurationMS: end.Sub(state.start).Milliseconds(), }, nil } diff --git a/internal/usecase/runner_test.go b/internal/usecase/runner_test.go index bf58335..ba94e7c 100644 --- a/internal/usecase/runner_test.go +++ b/internal/usecase/runner_test.go @@ -12,6 +12,7 @@ import ( "strings" "testing" + "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" @@ -70,9 +71,11 @@ func (f *fakePromptRepo) GetPromptDefinition(ctx context.Context, id string, ver type fakeArtifactReader struct { artifactsByURI map[string]*domain.Artifact errByURI map[string]error + calls int } func (f *fakeArtifactReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domain.Artifact, error) { + f.calls++ if err, ok := f.errByURI[ref.URI]; ok { return nil, err } @@ -86,9 +89,11 @@ func (f *fakeArtifactReader) Read(ctx context.Context, ref domain.ArtifactRef) ( type fakeRenderer struct { rendered *domain.RenderedPrompt err error + calls int } func (f *fakeRenderer) Render(ctx context.Context, def *domain.PromptDefinition, inputs map[string]*domain.Artifact, vars map[string]string) (*domain.RenderedPrompt, error) { + f.calls++ if f.err != nil { return nil, f.err } @@ -125,9 +130,11 @@ type fakeValidator struct { schemaErr error schemaLoadPath string schemaLoads int + validateCalls int } func (f *fakeValidator) Validate(ctx context.Context, artifact *domain.Artifact, contract domain.OutputContract) (domain.ValidationResult, error) { + f.validateCalls++ if f.err != nil { return domain.ValidationResult{}, f.err } @@ -153,6 +160,22 @@ type fakeRepairer struct { reqs []RepairRequest } +type fakeRunAdmitter struct { + backendIDs []string + err error + releaseCalls int +} + +func (f *fakeRunAdmitter) Admit(_ context.Context, backendID string) (func(), error) { + f.backendIDs = append(f.backendIDs, backendID) + if f.err != nil { + return nil, f.err + } + return func() { + f.releaseCalls++ + }, nil +} + func (f *fakeRepairer) Repair(ctx context.Context, req RepairRequest) (*domain.GenerateResponse, error) { f.calls++ f.reqs = append(f.reqs, req) @@ -179,7 +202,7 @@ func TestRunnerPrepareWithExplicitProfileSelection(t *testing.T) { renderer := &fakeRenderer{rendered: &domain.RenderedPrompt{SessionID: "session-123", Messages: []domain.RenderedMessage{{Role: "system", Content: "sys"}, {Role: "user", Content: "usr"}}}} llmClient := &fakeLLM{forbid: true} - runner := NewRunner(promptRepo, execRepo, nil, reader, renderer, llmClient, nil) + runner := NewRunner(promptRepo, execRepo, nil, reader, renderer, llmClient, nil, nil) prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", PromptVersion: "1", @@ -234,8 +257,8 @@ func TestRunnerDirectSessionResolution(t *testing.T) { defaultArtifactReader(), prompt.NewGoRenderer(), &fakeLLM{forbid: true}, - nil, - ) + nil, nil) + req := domain.RunRequest{ PromptID: "p", ProfileID: "exec", @@ -281,8 +304,7 @@ func TestRunnerDirectSessionResolution(t *testing.T) { defaultArtifactReader(), prompt.NewGoRenderer(), &fakeLLM{forbid: true}, - nil, - ) + nil, nil) prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", @@ -312,8 +334,7 @@ func TestRunnerDirectSessionResolution(t *testing.T) { defaultArtifactReader(), prompt.NewGoRenderer(), &fakeLLM{forbid: true}, - nil, - ) + nil, nil) prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", @@ -340,8 +361,7 @@ func TestRunnerDirectSessionResolution(t *testing.T) { defaultArtifactReader(), defaultRenderer(), llmClient, - nil, - ) + nil, nil) _, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -408,7 +428,7 @@ func TestRunnerPrepareSelectedProfileDoesNotExistFails(t *testing.T) { } func TestRunnerPreparePromptLoadFailure(t *testing.T) { - runner := NewRunner(&fakePromptRepo{err: errors.New("boom")}, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{}, nil) + runner := NewRunner(&fakePromptRepo{err: errors.New("boom")}, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{}, nil, nil) _, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p"}) if !errors.Is(err, ErrPromptLoad) { t.Fatalf("expected ErrPromptLoad, got %v", err) @@ -432,7 +452,7 @@ func TestRunnerPrepareRuntimeOverrideBeatsSelectedProfileValue(t *testing.T) { ServiceTier: "priority", }, }} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil, nil) prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", @@ -535,7 +555,7 @@ func TestRunnerPrepareRequestNumericOverridePresence(t *testing.T) { defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, - nil) + nil, nil) prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", @@ -582,7 +602,7 @@ func TestRunnerPrepareInvalidRequestNumericOverridesFail(t *testing.T) { defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, - nil) + nil, nil) _, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", @@ -609,7 +629,7 @@ func TestRunnerPrepareSelectedProfileBeatsBuiltInDefault(t *testing.T) { ServiceTier: "priority", }, }} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil, nil) prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", @@ -651,7 +671,7 @@ func TestRunnerPrepareFileBackedPromptBodiesRenderCorrectly(t *testing.T) { reader, prompt.NewGoRenderer(), llmClient, - nil) + nil, nil) prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "valid-file-backed", @@ -682,7 +702,7 @@ func TestRunnerPrepareRequiredInputMissingFails(t *testing.T) { defaultArtifactReader(), prompt.NewGoRenderer(), &fakeLLM{forbid: true}, - nil) + nil, nil) _, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", @@ -709,7 +729,7 @@ func TestRunnerPrepareUnknownTemplateInputReferenceFails(t *testing.T) { defaultArtifactReader(), prompt.NewGoRenderer(), &fakeLLM{forbid: true}, - nil) + nil, nil) _, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", @@ -732,7 +752,7 @@ func TestRunnerPrepareAPIKeyEnvNameIncludedButNotResolvedValue(t *testing.T) { execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: envName}, }} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil, nil) prepared, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) if err != nil { @@ -765,7 +785,7 @@ func TestRunnerPrepareJSONSchemaBuildsStructuredOutputSpec(t *testing.T) { defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, - validator) + validator, nil) prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", @@ -809,7 +829,7 @@ func TestRunnerPrepareJSONSchemaSchemaLoadFailureReturnsValidationError(t *testi defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, - validator) + validator, nil) _, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", @@ -839,7 +859,7 @@ func TestRunnerRunJSONSchemaSchemaLoadFailureFailsBeforeLLM(t *testing.T) { defaultArtifactReader(), defaultRenderer(), llmClient, - validator) + validator, nil) _, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -989,7 +1009,7 @@ func TestRunnerRunSuccessful(t *testing.T) { renderer := &fakeRenderer{rendered: &domain.RenderedPrompt{SessionID: "session-123", Messages: []domain.RenderedMessage{{Role: "system", Content: "sys"}, {Role: "user", Content: "usr"}}}} llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "# recap", Usage: domain.TokenUsage{TotalTokens: 7}}} - runner := NewRunner(promptRepo, execRepo, nil, reader, renderer, llmClient, nil) + runner := NewRunner(promptRepo, execRepo, nil, reader, renderer, llmClient, nil, nil) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", PromptVersion: "1", @@ -1071,7 +1091,7 @@ func TestRunnerRunPassesExtraParamsToGenerateRequestTarget(t *testing.T) { }, }} llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil, nil) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1097,7 +1117,7 @@ func TestRunnerRunAndPrepareResolveSameProfileAndEffectiveSettings(t *testing.T) }} renderer := &fakeRenderer{rendered: &domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "system", Content: "sys"}, {Role: "user", Content: "usr"}}}} llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "# recap"}} - runner := NewRunner(promptRepo, execRepo, nil, reader, renderer, llmClient, nil) + runner := NewRunner(promptRepo, execRepo, nil, reader, renderer, llmClient, nil, nil) req := domain.RunRequest{ PromptID: "p", @@ -1135,6 +1155,253 @@ func TestRunnerRunAndPrepareResolveSameProfileAndEffectiveSettings(t *testing.T) } } +func TestRunnerAdmissionUsesResolvedBackendIdentity(t *testing.T) { + t.Run("selected backend survives endpoint override", func(t *testing.T) { + admitter := &fakeRunAdmitter{} + runner := NewRunner( + &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}, + &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "exec": { + ID: "exec", + BackendID: "custom", + Model: "model", + }, + }}, + fakeBackendResolver{backends: map[string]domain.Backend{ + "custom": {ID: "custom", Endpoint: "http://backend.example/v1"}, + }}, + defaultArtifactReader(), + defaultRenderer(), + &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, + nil, + admitter, + ) + + result, err := runner.Run(context.Background(), domain.RunRequest{ + PromptID: "p", + ProfileID: "exec", + Inputs: singleInputRef(), + Execution: &domain.ExecutionTargetOverride{ + Endpoint: "http://override.example/v1", + }, + }) + if err != nil { + t.Fatalf("run: %v", err) + } + if !reflect.DeepEqual(admitter.backendIDs, []string{"custom"}) { + t.Fatalf("admitted backend IDs=%#v, want custom", admitter.backendIDs) + } + if result.SelectedBackendID != "custom" || + result.Endpoint != "http://override.example/v1" { + t.Fatalf("unexpected routed result: %+v", result) + } + }) + + t.Run("endpoint-only preparation remains unrestricted", func(t *testing.T) { + admitter := &fakeRunAdmitter{} + runner := NewRunner( + &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}, + &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "exec": defaultExecutionProfile(), + }}, + nil, + defaultArtifactReader(), + defaultRenderer(), + &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, + nil, + admitter, + ) + request := domain.RunRequest{ + PromptID: "p", + ProfileID: "exec", + Inputs: singleInputRef(), + } + + if _, err := runner.Prepare(context.Background(), request); err != nil { + t.Fatalf("prepare: %v", err) + } + if len(admitter.backendIDs) != 0 { + t.Fatalf("prepare called admission with %#v", admitter.backendIDs) + } + if _, err := runner.Run(context.Background(), request); err != nil { + t.Fatalf("run: %v", err) + } + if !reflect.DeepEqual(admitter.backendIDs, []string{""}) { + t.Fatalf("admitted backend IDs=%#v, want blank ID", admitter.backendIDs) + } + }) +} + +func TestRunnerAdmissionFailureSkipsCompletionCollaborators(t *testing.T) { + tests := []struct { + name string + admissionError error + wantBackendContext bool + }{ + { + name: "capacity exhausted", + admissionError: capacity.ErrCapacityExceeded, + wantBackendContext: true, + }, + { + name: "context canceled", + admissionError: context.Canceled, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + def := promptDef(domain.FormatJSON, domain.ValidationJSONSchema, 1) + def.Validation.SchemaPath = "schema.json" + reader := defaultArtifactReader() + renderer := defaultRenderer() + llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: `{}`}} + validator := &fakeValidator{} + repairer := &fakeRepairer{ + responses: []*domain.GenerateResponse{{Content: `{}`}}, + } + admitter := &fakeRunAdmitter{err: tc.admissionError} + runner := NewRunnerWithRepairer( + &fakePromptRepo{def: def}, + &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "exec": {ID: "exec", BackendID: "custom", Model: "model"}, + }}, + fakeBackendResolver{backends: map[string]domain.Backend{ + "custom": {ID: "custom", Endpoint: "http://backend.example/v1"}, + }}, + reader, + renderer, + llmClient, + validator, + repairer, + admitter, + ) + + result, err := runner.Run(context.Background(), domain.RunRequest{ + PromptID: "p", + ProfileID: "exec", + Inputs: singleInputRef(), + }) + if result != nil { + t.Fatalf("admission failure returned partial result: %+v", result) + } + if !errors.Is(err, tc.admissionError) { + t.Fatalf("admission error=%v, want identity %v", err, tc.admissionError) + } + if errors.Is(err, ErrInvalidRequest) || errors.Is(err, ErrLLMGenerate) { + t.Fatalf("admission error was recategorized: %v", err) + } + if tc.wantBackendContext && !strings.Contains(err.Error(), "custom") { + t.Fatalf("capacity error lacks backend context: %v", err) + } + if !reflect.DeepEqual(admitter.backendIDs, []string{"custom"}) { + t.Fatalf("admitted backend IDs=%#v, want custom", admitter.backendIDs) + } + if admitter.releaseCalls != 0 || + validator.schemaLoads != 0 || + validator.validateCalls != 0 || + reader.calls != 0 || + renderer.calls != 0 || + llmClient.calls != 0 || + repairer.calls != 0 { + t.Fatalf( + "later collaborators invoked: releases=%d schema=%d validate=%d artifacts=%d render=%d llm=%d repair=%d", + admitter.releaseCalls, + validator.schemaLoads, + validator.validateCalls, + reader.calls, + renderer.calls, + llmClient.calls, + repairer.calls, + ) + } + }) + } +} + +func TestRunnerReleasesAdmissionAcrossRunOutcomes(t *testing.T) { + artifactFailure := errors.New("artifact failed") + generationFailure := errors.New("generation failed") + validationFailure := errors.New("validation failed") + tests := []struct { + name string + artifactError error + generationErr error + validationErr error + wantError error + }{ + {name: "success"}, + { + name: "completion failure", + artifactError: artifactFailure, + wantError: ErrArtifactLoad, + }, + { + name: "generation failure", + generationErr: generationFailure, + wantError: ErrLLMGenerate, + }, + { + name: "validation failure", + validationErr: validationFailure, + wantError: ErrValidation, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + reader := defaultArtifactReader() + if tc.artifactError != nil { + reader.errByURI = map[string]error{"a://ok": tc.artifactError} + } + llmClient := &fakeLLM{ + resp: &domain.GenerateResponse{Content: "ok"}, + err: tc.generationErr, + } + validator := &fakeValidator{ + result: domain.ValidationResult{ + Status: domain.ValidationPassed, + Mode: domain.ValidationBasic, + IsValid: true, + }, + err: tc.validationErr, + } + admitter := &fakeRunAdmitter{} + runner := NewRunner( + &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationBasic, 0)}, + &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "exec": defaultExecutionProfile(), + }}, + nil, + reader, + defaultRenderer(), + llmClient, + validator, + admitter, + ) + + result, err := runner.Run(context.Background(), domain.RunRequest{ + PromptID: "p", + ProfileID: "exec", + Inputs: singleInputRef(), + }) + if tc.wantError == nil { + if err != nil || result == nil { + t.Fatalf("successful run=(%+v, %v)", result, err) + } + } else { + if result != nil || !errors.Is(err, tc.wantError) { + t.Fatalf("failed run=(%+v, %v), want %v", result, err, tc.wantError) + } + } + if len(admitter.backendIDs) != 1 || admitter.releaseCalls != 1 { + t.Fatalf("admission calls=%#v releases=%d, want one each", + admitter.backendIDs, admitter.releaseCalls) + } + }) + } +} + func TestRunnerRunExplicitProfileIDIsUsed(t *testing.T) { promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} promptRepo.def.DefaultProfile = "default-prof" @@ -1227,7 +1494,7 @@ func TestRunnerRunExplicitRuntimeOverrideBeatsSelectedProfileValue(t *testing.T) }, }} llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil, nil) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1271,7 +1538,7 @@ func TestRunnerRunSelectedProfileBeatsBuiltInDefault(t *testing.T) { }, }} llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil, nil) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1298,7 +1565,7 @@ func TestRunnerRunBuiltInDefaultsUsedWhenProfileOmitsOptionalFields(t *testing.T "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model"}, }} llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil, nil) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1328,7 +1595,7 @@ func TestRunnerRunAPIKeyEnvResolvesFromEnvironment(t *testing.T) { execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: "PROMPTKIT_TEST_API_KEY"}, }} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil, nil) res, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) if err != nil { @@ -1344,7 +1611,7 @@ func TestRunnerRunAPIKeyEnvMissingEnvironmentValueFailsClearly(t *testing.T) { execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: "PROMPTKIT_MISSING_KEY"}, }} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil, nil) _, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) if !errors.Is(err, ErrInvalidRequest) { @@ -1365,7 +1632,7 @@ func TestRunnerRunDirectAPIKeyBypassesMissingEnvAndReachesLLM(t *testing.T) { "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: "PROMPTKIT_MISSING_KEY"}, }} llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil, nil) _, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1389,7 +1656,7 @@ func TestRunnerPrepareAPIKeyRequiredFailsWithoutDirectKey(t *testing.T) { execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyRequired: true}, }} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil, nil) _, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1408,7 +1675,7 @@ func TestRunnerRunAPIKeyRequiredSucceedsWithDirectKey(t *testing.T) { "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyRequired: true}, }} llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil, nil) _, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1435,7 +1702,7 @@ func TestRunnerRunRuntimeAPIKeyEnvOverrideWorks(t *testing.T) { execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model"}, }} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil, nil) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1460,7 +1727,7 @@ func TestRunnerRunRuntimeAPIKeyEnvOverrideBeatsProfile(t *testing.T) { execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: profileEnv}, }} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil, nil) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1484,7 +1751,7 @@ func TestRunnerRunAPIKeyValueNotPresentInMetadata(t *testing.T) { execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: envName}, }} - runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil) + runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil, nil) res, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) if err != nil { @@ -1500,7 +1767,7 @@ func TestRunnerRunAPIKeyValueNotPresentInMetadata(t *testing.T) { } func TestRunnerRunPromptLoadFailure(t *testing.T) { - runner := NewRunner(&fakePromptRepo{err: errors.New("boom")}, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{}, nil) + runner := NewRunner(&fakePromptRepo{err: errors.New("boom")}, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{}, nil, nil) _, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p"}) if !errors.Is(err, ErrPromptLoad) { t.Fatalf("expected ErrPromptLoad, got %v", err) @@ -1518,7 +1785,7 @@ func TestRunnerRunArtifactLoadFailure(t *testing.T) { &fakeArtifactReader{errByURI: map[string]error{"a://bad": errors.New("read failed")}}, &fakeRenderer{rendered: &domain.RenderedPrompt{}}, &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, - nil) + nil, nil) _, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1538,7 +1805,7 @@ func TestRunnerRunPromptRenderFailure(t *testing.T) { defaultArtifactReader(), &fakeRenderer{err: errors.New("render failed")}, &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, - nil) + nil, nil) _, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1558,7 +1825,7 @@ func TestRunnerRunLLMFailure(t *testing.T) { defaultArtifactReader(), defaultRenderer(), &fakeLLM{err: errors.New("llm failed")}, - nil) + nil, nil) _, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1578,7 +1845,7 @@ func TestRunnerRunCancellationPreservesGenerationCategory(t *testing.T) { defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ignored"}}, - nil) + nil, nil) ctx, cancel := context.WithCancel(context.Background()) cancel() @@ -1604,7 +1871,7 @@ func TestRunnerRunLLMInvalidRequestMapsToUsecaseInvalidRequest(t *testing.T) { defaultArtifactReader(), defaultRenderer(), &fakeLLM{err: llm.ErrInvalidRequest}, - nil) + nil, nil) _, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1628,7 +1895,7 @@ func TestRunnerRunValidationStillWorks(t *testing.T) { defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "raw output"}}, - validator) + validator, nil) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1659,7 +1926,7 @@ func TestRunnerRunStructuredRepairRemainsBoundedAndUsesEffectiveModelSettings(t defaultRenderer(), llmClient, validate.NewStandardValidator("."), - repairer) + repairer, nil) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1702,8 +1969,7 @@ func TestRunnerRunRepairCarriesEffectiveSessionID(t *testing.T) { defaultRenderer(), llmClient, validate.NewStandardValidator("."), - NewDefaultOutputRepairer(llmClient), - ) + NewDefaultOutputRepairer(llmClient), nil) result, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -1757,7 +2023,7 @@ func TestRunnerRunJSONSchemaRepairCarriesStructuredOutputSpec(t *testing.T) { defaultRenderer(), llmClient, validator, - repairer) + repairer, nil) _, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", @@ -2114,7 +2380,8 @@ func TestRunnerPrepareBackendResolutionAndCredentialPrecedence(t *testing.T) { t.Run("unknown backend is a profile load failure", func(t *testing.T) { runner := NewRunner(promptRepo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ "exec": {ID: "exec", BackendID: "unknown", Model: "model"}, - }}, resolver, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil) + }}, resolver, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil, nil) + _, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) if !errors.Is(err, ErrProfileLoad) { t.Fatalf("expected ErrProfileLoad, got %v", err) @@ -2124,7 +2391,8 @@ func TestRunnerPrepareBackendResolutionAndCredentialPrecedence(t *testing.T) { t.Run("nil resolver is a profile load failure", func(t *testing.T) { runner := NewRunner(promptRepo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ "exec": {ID: "exec", BackendID: "custom", Model: "model"}, - }}, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil) + }}, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil, nil) + _, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) if !errors.Is(err, ErrProfileLoad) { t.Fatalf("expected ErrProfileLoad, got %v", err) @@ -2135,7 +2403,8 @@ func TestRunnerPrepareBackendResolutionAndCredentialPrecedence(t *testing.T) { t.Setenv("REQUEST_KEY", "request-secret") runner := NewRunner(promptRepo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ "exec": {ID: "exec", BackendID: "custom", Model: "model", APIKeyEnv: "PROFILE_KEY"}, - }}, resolver, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil) + }}, resolver, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil, nil) + prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ PromptID: "p", ProfileID: "exec", Inputs: singleInputRef(), Execution: &domain.ExecutionTargetOverride{APIKeyEnv: "REQUEST_KEY"}, @@ -2152,7 +2421,8 @@ func TestRunnerPrepareBackendResolutionAndCredentialPrecedence(t *testing.T) { t.Setenv("BACKEND_KEY", "backend-secret") runner := NewRunner(promptRepo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ "exec": {ID: "exec", BackendID: "custom", Model: "model", APIKeyRequired: true}, - }}, resolver, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil) + }}, resolver, defaultArtifactReader(), defaultRenderer(), &fakeLLM{forbid: true}, nil, nil) + _, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) if !errors.Is(err, ErrAPIKeyRequired) { t.Fatalf("expected ErrAPIKeyRequired, got %v", err) @@ -2203,6 +2473,6 @@ func newMinimalRunner(promptRepo *fakePromptRepo, execRepo *fakeExecutionProfile defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, - nil) + nil, nil) }