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

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