Add early backend run admission

This commit is contained in:
2026-07-29 21:01:37 +00:00
parent a752f88166
commit 861da355d8
4 changed files with 425 additions and 72 deletions

View File

@@ -586,7 +586,7 @@ are independent, and the wrapper is transparent apart from waiting.
## Stage 3 — Shared Preparation And Early Run Admission ## Stage 3 — Shared Preparation And Early Run Admission
**Status:** Pending. **Status:** Complete.
### Goal ### Goal

View File

@@ -382,6 +382,7 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) {
prompt.NewGoRenderer(), prompt.NewGoRenderer(),
llmClient, llmClient,
validator, validator,
nil,
), ),
}, nil }, nil
} }

View File

@@ -14,6 +14,7 @@ import (
"unicode" "unicode"
"gitea.maximumdirect.net/eric/promptkit/internal/artifact" "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/defaults"
"gitea.maximumdirect.net/eric/promptkit/internal/domain" "gitea.maximumdirect.net/eric/promptkit/internal/domain"
"gitea.maximumdirect.net/eric/promptkit/internal/llm" "gitea.maximumdirect.net/eric/promptkit/internal/llm"
@@ -46,6 +47,7 @@ type Runner struct {
llm llm.Client llm llm.Client
validator validate.Validator validator validate.Validator
repairer OutputRepairer repairer OutputRepairer
admitter RunAdmitter
} }
// BackendResolver resolves one normalized backend ID. // BackendResolver resolves one normalized backend ID.
@@ -53,6 +55,22 @@ type BackendResolver interface {
GetBackend(string) (domain.Backend, error) 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( func NewRunner(
promptDefs promptdef.Repository, promptDefs promptdef.Repository,
profiles profile.Repository, profiles profile.Repository,
@@ -61,8 +79,19 @@ func NewRunner(
renderer prompt.Renderer, renderer prompt.Renderer,
llmClient llm.Client, llmClient llm.Client,
validator validate.Validator, validator validate.Validator,
admitter RunAdmitter,
) *Runner { ) *Runner {
return NewRunnerWithRepairer(promptDefs, profiles, backends, artifacts, renderer, llmClient, validator, nil) return NewRunnerWithRepairer(
promptDefs,
profiles,
backends,
artifacts,
renderer,
llmClient,
validator,
nil,
admitter,
)
} }
func NewRunnerWithRepairer( func NewRunnerWithRepairer(
@@ -74,6 +103,7 @@ func NewRunnerWithRepairer(
llmClient llm.Client, llmClient llm.Client,
validator validate.Validator, validator validate.Validator,
repairer OutputRepairer, repairer OutputRepairer,
admitter RunAdmitter,
) *Runner { ) *Runner {
return &Runner{ return &Runner{
promptDefs: promptDefs, promptDefs: promptDefs,
@@ -84,6 +114,7 @@ func NewRunnerWithRepairer(
llm: llmClient, llm: llmClient,
validator: validator, validator: validator,
repairer: repairer, repairer: repairer,
admitter: admitter,
} }
} }
@@ -95,7 +126,27 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
start := time.Now().UTC() 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 { if err != nil {
return nil, err 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) { 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) == "" { if strings.TrimSpace(req.PromptID) == "" {
return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidRequest) 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) return nil, fmt.Errorf("%w: session_id: %v", ErrInvalidRequest, err)
} }
start := time.Now().UTC()
def, err := r.promptDefs.GetPromptDefinition(ctx, req.PromptID, req.PromptVersion) def, err := r.promptDefs.GetPromptDefinition(ctx, req.PromptID, req.PromptVersion)
if err != nil { if err != nil {
return nil, fmt.Errorf("%w: %w", ErrPromptLoad, err) 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) 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 { if err != nil {
return nil, err return nil, err
} }
@@ -257,9 +339,9 @@ func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.Pr
inputHashes[name] = art.Hash inputHashes[name] = art.Hash
} }
definitionToRender := def definitionToRender := state.definition
if directSessionID != "" { if state.directSessionID != "" {
definitionCopy := *def definitionCopy := *state.definition
definitionCopy.SessionID = "" definitionCopy.SessionID = ""
definitionToRender = &definitionCopy definitionToRender = &definitionCopy
} }
@@ -267,28 +349,28 @@ func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.Pr
if err != nil { if err != nil {
return nil, fmt.Errorf("%w: %w", ErrPromptRender, err) return nil, fmt.Errorf("%w: %w", ErrPromptRender, err)
} }
if directSessionID != "" { if state.directSessionID != "" {
renderedPrompt.SessionID = directSessionID renderedPrompt.SessionID = state.directSessionID
} }
end := time.Now().UTC() end := time.Now().UTC()
return &domain.PreparedRun{ return &domain.PreparedRun{
PromptID: def.ID, PromptID: state.definition.ID,
PromptVersion: def.Version, PromptVersion: state.definition.Version,
PromptHash: promptDefinitionHash, PromptHash: state.promptDefinitionHash,
SelectedProfileID: selectedProfileID, SelectedProfileID: state.selectedProfileID,
SelectedBackendID: effectiveModel.BackendID, SelectedBackendID: state.effectiveModel.BackendID,
EffectiveModelParams: effectiveModel, EffectiveModelParams: state.effectiveModel,
TargetPresence: targetPresence, TargetPresence: state.targetPresence,
OutputContract: effectiveContract, OutputContract: state.effectiveContract,
StructuredOutput: structuredOutput, StructuredOutput: structuredOutput,
InputHashes: inputHashes, InputHashes: inputHashes,
SessionID: renderedPrompt.SessionID, SessionID: renderedPrompt.SessionID,
RenderedPromptHash: hashRenderedPrompt(*renderedPrompt), RenderedPromptHash: hashRenderedPrompt(*renderedPrompt),
Messages: renderedPrompt.Messages, Messages: renderedPrompt.Messages,
StartTime: start, StartTime: state.start,
EndTime: end, EndTime: end,
DurationMS: end.Sub(start).Milliseconds(), DurationMS: end.Sub(state.start).Milliseconds(),
}, nil }, nil
} }

View File

@@ -12,6 +12,7 @@ import (
"strings" "strings"
"testing" "testing"
"gitea.maximumdirect.net/eric/promptkit/internal/capacity"
"gitea.maximumdirect.net/eric/promptkit/internal/defaults" "gitea.maximumdirect.net/eric/promptkit/internal/defaults"
"gitea.maximumdirect.net/eric/promptkit/internal/domain" "gitea.maximumdirect.net/eric/promptkit/internal/domain"
"gitea.maximumdirect.net/eric/promptkit/internal/llm" "gitea.maximumdirect.net/eric/promptkit/internal/llm"
@@ -70,9 +71,11 @@ func (f *fakePromptRepo) GetPromptDefinition(ctx context.Context, id string, ver
type fakeArtifactReader struct { type fakeArtifactReader struct {
artifactsByURI map[string]*domain.Artifact artifactsByURI map[string]*domain.Artifact
errByURI map[string]error errByURI map[string]error
calls int
} }
func (f *fakeArtifactReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domain.Artifact, error) { func (f *fakeArtifactReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domain.Artifact, error) {
f.calls++
if err, ok := f.errByURI[ref.URI]; ok { if err, ok := f.errByURI[ref.URI]; ok {
return nil, err return nil, err
} }
@@ -86,9 +89,11 @@ func (f *fakeArtifactReader) Read(ctx context.Context, ref domain.ArtifactRef) (
type fakeRenderer struct { type fakeRenderer struct {
rendered *domain.RenderedPrompt rendered *domain.RenderedPrompt
err error 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) { 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 { if f.err != nil {
return nil, f.err return nil, f.err
} }
@@ -125,9 +130,11 @@ type fakeValidator struct {
schemaErr error schemaErr error
schemaLoadPath string schemaLoadPath string
schemaLoads int schemaLoads int
validateCalls int
} }
func (f *fakeValidator) Validate(ctx context.Context, artifact *domain.Artifact, contract domain.OutputContract) (domain.ValidationResult, error) { func (f *fakeValidator) Validate(ctx context.Context, artifact *domain.Artifact, contract domain.OutputContract) (domain.ValidationResult, error) {
f.validateCalls++
if f.err != nil { if f.err != nil {
return domain.ValidationResult{}, f.err return domain.ValidationResult{}, f.err
} }
@@ -153,6 +160,22 @@ type fakeRepairer struct {
reqs []RepairRequest 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) { func (f *fakeRepairer) Repair(ctx context.Context, req RepairRequest) (*domain.GenerateResponse, error) {
f.calls++ f.calls++
f.reqs = append(f.reqs, req) 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"}}}} renderer := &fakeRenderer{rendered: &domain.RenderedPrompt{SessionID: "session-123", Messages: []domain.RenderedMessage{{Role: "system", Content: "sys"}, {Role: "user", Content: "usr"}}}}
llmClient := &fakeLLM{forbid: true} 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{ prepared, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
PromptVersion: "1", PromptVersion: "1",
@@ -234,8 +257,8 @@ func TestRunnerDirectSessionResolution(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
prompt.NewGoRenderer(), prompt.NewGoRenderer(),
&fakeLLM{forbid: true}, &fakeLLM{forbid: true},
nil, nil, nil)
)
req := domain.RunRequest{ req := domain.RunRequest{
PromptID: "p", PromptID: "p",
ProfileID: "exec", ProfileID: "exec",
@@ -281,8 +304,7 @@ func TestRunnerDirectSessionResolution(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
prompt.NewGoRenderer(), prompt.NewGoRenderer(),
&fakeLLM{forbid: true}, &fakeLLM{forbid: true},
nil, nil, nil)
)
prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ prepared, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -312,8 +334,7 @@ func TestRunnerDirectSessionResolution(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
prompt.NewGoRenderer(), prompt.NewGoRenderer(),
&fakeLLM{forbid: true}, &fakeLLM{forbid: true},
nil, nil, nil)
)
prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ prepared, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -340,8 +361,7 @@ func TestRunnerDirectSessionResolution(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
defaultRenderer(), defaultRenderer(),
llmClient, llmClient,
nil, nil, nil)
)
_, err := runner.Run(context.Background(), domain.RunRequest{ _, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -408,7 +428,7 @@ func TestRunnerPrepareSelectedProfileDoesNotExistFails(t *testing.T) {
} }
func TestRunnerPreparePromptLoadFailure(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"}) _, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p"})
if !errors.Is(err, ErrPromptLoad) { if !errors.Is(err, ErrPromptLoad) {
t.Fatalf("expected ErrPromptLoad, got %v", err) t.Fatalf("expected ErrPromptLoad, got %v", err)
@@ -432,7 +452,7 @@ func TestRunnerPrepareRuntimeOverrideBeatsSelectedProfileValue(t *testing.T) {
ServiceTier: "priority", 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{ prepared, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -535,7 +555,7 @@ func TestRunnerPrepareRequestNumericOverridePresence(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
defaultRenderer(), defaultRenderer(),
&fakeLLM{forbid: true}, &fakeLLM{forbid: true},
nil) nil, nil)
prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ prepared, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -582,7 +602,7 @@ func TestRunnerPrepareInvalidRequestNumericOverridesFail(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
defaultRenderer(), defaultRenderer(),
&fakeLLM{forbid: true}, &fakeLLM{forbid: true},
nil) nil, nil)
_, err := runner.Prepare(context.Background(), domain.RunRequest{ _, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -609,7 +629,7 @@ func TestRunnerPrepareSelectedProfileBeatsBuiltInDefault(t *testing.T) {
ServiceTier: "priority", 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{ prepared, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -651,7 +671,7 @@ func TestRunnerPrepareFileBackedPromptBodiesRenderCorrectly(t *testing.T) {
reader, reader,
prompt.NewGoRenderer(), prompt.NewGoRenderer(),
llmClient, llmClient,
nil) nil, nil)
prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ prepared, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "valid-file-backed", PromptID: "valid-file-backed",
@@ -682,7 +702,7 @@ func TestRunnerPrepareRequiredInputMissingFails(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
prompt.NewGoRenderer(), prompt.NewGoRenderer(),
&fakeLLM{forbid: true}, &fakeLLM{forbid: true},
nil) nil, nil)
_, err := runner.Prepare(context.Background(), domain.RunRequest{ _, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -709,7 +729,7 @@ func TestRunnerPrepareUnknownTemplateInputReferenceFails(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
prompt.NewGoRenderer(), prompt.NewGoRenderer(),
&fakeLLM{forbid: true}, &fakeLLM{forbid: true},
nil) nil, nil)
_, err := runner.Prepare(context.Background(), domain.RunRequest{ _, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -732,7 +752,7 @@ func TestRunnerPrepareAPIKeyEnvNameIncludedButNotResolvedValue(t *testing.T) {
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: envName}, "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()}) prepared, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
if err != nil { if err != nil {
@@ -765,7 +785,7 @@ func TestRunnerPrepareJSONSchemaBuildsStructuredOutputSpec(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
defaultRenderer(), defaultRenderer(),
&fakeLLM{forbid: true}, &fakeLLM{forbid: true},
validator) validator, nil)
prepared, err := runner.Prepare(context.Background(), domain.RunRequest{ prepared, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -809,7 +829,7 @@ func TestRunnerPrepareJSONSchemaSchemaLoadFailureReturnsValidationError(t *testi
defaultArtifactReader(), defaultArtifactReader(),
defaultRenderer(), defaultRenderer(),
&fakeLLM{forbid: true}, &fakeLLM{forbid: true},
validator) validator, nil)
_, err := runner.Prepare(context.Background(), domain.RunRequest{ _, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -839,7 +859,7 @@ func TestRunnerRunJSONSchemaSchemaLoadFailureFailsBeforeLLM(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
defaultRenderer(), defaultRenderer(),
llmClient, llmClient,
validator) validator, nil)
_, err := runner.Run(context.Background(), domain.RunRequest{ _, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", 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"}}}} 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}}} 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{ res, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
PromptVersion: "1", PromptVersion: "1",
@@ -1071,7 +1091,7 @@ func TestRunnerRunPassesExtraParamsToGenerateRequestTarget(t *testing.T) {
}, },
}} }}
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} 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{ res, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", 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"}}}} renderer := &fakeRenderer{rendered: &domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "system", Content: "sys"}, {Role: "user", Content: "usr"}}}}
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "# recap"}} 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{ req := domain.RunRequest{
PromptID: "p", 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) { func TestRunnerRunExplicitProfileIDIsUsed(t *testing.T) {
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}
promptRepo.def.DefaultProfile = "default-prof" promptRepo.def.DefaultProfile = "default-prof"
@@ -1227,7 +1494,7 @@ func TestRunnerRunExplicitRuntimeOverrideBeatsSelectedProfileValue(t *testing.T)
}, },
}} }}
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} 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{ res, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1271,7 +1538,7 @@ func TestRunnerRunSelectedProfileBeatsBuiltInDefault(t *testing.T) {
}, },
}} }}
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} 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{ res, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1298,7 +1565,7 @@ func TestRunnerRunBuiltInDefaultsUsedWhenProfileOmitsOptionalFields(t *testing.T
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model"}, "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model"},
}} }}
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} 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{ res, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1328,7 +1595,7 @@ func TestRunnerRunAPIKeyEnvResolvesFromEnvironment(t *testing.T) {
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: "PROMPTKIT_TEST_API_KEY"}, "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()}) res, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
if err != nil { if err != nil {
@@ -1344,7 +1611,7 @@ func TestRunnerRunAPIKeyEnvMissingEnvironmentValueFailsClearly(t *testing.T) {
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: "PROMPTKIT_MISSING_KEY"}, "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()}) _, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
if !errors.Is(err, ErrInvalidRequest) { 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"}, "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: "PROMPTKIT_MISSING_KEY"},
}} }}
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} 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{ _, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1389,7 +1656,7 @@ func TestRunnerPrepareAPIKeyRequiredFailsWithoutDirectKey(t *testing.T) {
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyRequired: true}, "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{ _, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1408,7 +1675,7 @@ func TestRunnerRunAPIKeyRequiredSucceedsWithDirectKey(t *testing.T) {
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyRequired: true}, "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyRequired: true},
}} }}
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} 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{ _, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1435,7 +1702,7 @@ func TestRunnerRunRuntimeAPIKeyEnvOverrideWorks(t *testing.T) {
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model"}, "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{ res, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1460,7 +1727,7 @@ func TestRunnerRunRuntimeAPIKeyEnvOverrideBeatsProfile(t *testing.T) {
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: profileEnv}, "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{ res, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1484,7 +1751,7 @@ func TestRunnerRunAPIKeyValueNotPresentInMetadata(t *testing.T) {
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: envName}, "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()}) res, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
if err != nil { if err != nil {
@@ -1500,7 +1767,7 @@ func TestRunnerRunAPIKeyValueNotPresentInMetadata(t *testing.T) {
} }
func TestRunnerRunPromptLoadFailure(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"}) _, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p"})
if !errors.Is(err, ErrPromptLoad) { if !errors.Is(err, ErrPromptLoad) {
t.Fatalf("expected ErrPromptLoad, got %v", err) 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")}}, &fakeArtifactReader{errByURI: map[string]error{"a://bad": errors.New("read failed")}},
&fakeRenderer{rendered: &domain.RenderedPrompt{}}, &fakeRenderer{rendered: &domain.RenderedPrompt{}},
&fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}},
nil) nil, nil)
_, err := runner.Run(context.Background(), domain.RunRequest{ _, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1538,7 +1805,7 @@ func TestRunnerRunPromptRenderFailure(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
&fakeRenderer{err: errors.New("render failed")}, &fakeRenderer{err: errors.New("render failed")},
&fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}},
nil) nil, nil)
_, err := runner.Run(context.Background(), domain.RunRequest{ _, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1558,7 +1825,7 @@ func TestRunnerRunLLMFailure(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
defaultRenderer(), defaultRenderer(),
&fakeLLM{err: errors.New("llm failed")}, &fakeLLM{err: errors.New("llm failed")},
nil) nil, nil)
_, err := runner.Run(context.Background(), domain.RunRequest{ _, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1578,7 +1845,7 @@ func TestRunnerRunCancellationPreservesGenerationCategory(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
defaultRenderer(), defaultRenderer(),
&fakeLLM{resp: &domain.GenerateResponse{Content: "ignored"}}, &fakeLLM{resp: &domain.GenerateResponse{Content: "ignored"}},
nil) nil, nil)
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
cancel() cancel()
@@ -1604,7 +1871,7 @@ func TestRunnerRunLLMInvalidRequestMapsToUsecaseInvalidRequest(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
defaultRenderer(), defaultRenderer(),
&fakeLLM{err: llm.ErrInvalidRequest}, &fakeLLM{err: llm.ErrInvalidRequest},
nil) nil, nil)
_, err := runner.Run(context.Background(), domain.RunRequest{ _, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1628,7 +1895,7 @@ func TestRunnerRunValidationStillWorks(t *testing.T) {
defaultArtifactReader(), defaultArtifactReader(),
defaultRenderer(), defaultRenderer(),
&fakeLLM{resp: &domain.GenerateResponse{Content: "raw output"}}, &fakeLLM{resp: &domain.GenerateResponse{Content: "raw output"}},
validator) validator, nil)
res, err := runner.Run(context.Background(), domain.RunRequest{ res, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1659,7 +1926,7 @@ func TestRunnerRunStructuredRepairRemainsBoundedAndUsesEffectiveModelSettings(t
defaultRenderer(), defaultRenderer(),
llmClient, llmClient,
validate.NewStandardValidator("."), validate.NewStandardValidator("."),
repairer) repairer, nil)
res, err := runner.Run(context.Background(), domain.RunRequest{ res, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1702,8 +1969,7 @@ func TestRunnerRunRepairCarriesEffectiveSessionID(t *testing.T) {
defaultRenderer(), defaultRenderer(),
llmClient, llmClient,
validate.NewStandardValidator("."), validate.NewStandardValidator("."),
NewDefaultOutputRepairer(llmClient), NewDefaultOutputRepairer(llmClient), nil)
)
result, err := runner.Run(context.Background(), domain.RunRequest{ result, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -1757,7 +2023,7 @@ func TestRunnerRunJSONSchemaRepairCarriesStructuredOutputSpec(t *testing.T) {
defaultRenderer(), defaultRenderer(),
llmClient, llmClient,
validator, validator,
repairer) repairer, nil)
_, err := runner.Run(context.Background(), domain.RunRequest{ _, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
@@ -2114,7 +2380,8 @@ func TestRunnerPrepareBackendResolutionAndCredentialPrecedence(t *testing.T) {
t.Run("unknown backend is a profile load failure", func(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{ runner := NewRunner(promptRepo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", BackendID: "unknown", Model: "model"}, "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()}) _, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
if !errors.Is(err, ErrProfileLoad) { if !errors.Is(err, ErrProfileLoad) {
t.Fatalf("expected ErrProfileLoad, got %v", err) 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) { t.Run("nil resolver is a profile load failure", func(t *testing.T) {
runner := NewRunner(promptRepo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ runner := NewRunner(promptRepo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", BackendID: "custom", Model: "model"}, "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()}) _, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
if !errors.Is(err, ErrProfileLoad) { if !errors.Is(err, ErrProfileLoad) {
t.Fatalf("expected ErrProfileLoad, got %v", err) t.Fatalf("expected ErrProfileLoad, got %v", err)
@@ -2135,7 +2403,8 @@ func TestRunnerPrepareBackendResolutionAndCredentialPrecedence(t *testing.T) {
t.Setenv("REQUEST_KEY", "request-secret") t.Setenv("REQUEST_KEY", "request-secret")
runner := NewRunner(promptRepo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ runner := NewRunner(promptRepo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", BackendID: "custom", Model: "model", APIKeyEnv: "PROFILE_KEY"}, "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{ prepared, err := runner.Prepare(context.Background(), domain.RunRequest{
PromptID: "p", ProfileID: "exec", Inputs: singleInputRef(), PromptID: "p", ProfileID: "exec", Inputs: singleInputRef(),
Execution: &domain.ExecutionTargetOverride{APIKeyEnv: "REQUEST_KEY"}, Execution: &domain.ExecutionTargetOverride{APIKeyEnv: "REQUEST_KEY"},
@@ -2152,7 +2421,8 @@ func TestRunnerPrepareBackendResolutionAndCredentialPrecedence(t *testing.T) {
t.Setenv("BACKEND_KEY", "backend-secret") t.Setenv("BACKEND_KEY", "backend-secret")
runner := NewRunner(promptRepo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ runner := NewRunner(promptRepo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", BackendID: "custom", Model: "model", APIKeyRequired: true}, "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()}) _, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
if !errors.Is(err, ErrAPIKeyRequired) { if !errors.Is(err, ErrAPIKeyRequired) {
t.Fatalf("expected ErrAPIKeyRequired, got %v", err) t.Fatalf("expected ErrAPIKeyRequired, got %v", err)
@@ -2203,6 +2473,6 @@ func newMinimalRunner(promptRepo *fakePromptRepo, execRepo *fakeExecutionProfile
defaultArtifactReader(), defaultArtifactReader(),
defaultRenderer(), defaultRenderer(),
&fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}},
nil) nil, nil)
} }