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
**Status:** Pending.
**Status:** Complete.
### Goal

View File

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

View File

@@ -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
}

View File

@@ -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)
}