Add Runner.Prepare and refactor Runner.Run to reuse pre-LLM prepare flow

This commit is contained in:
2026-05-06 15:12:58 +00:00
parent 48c06218dc
commit 44056d7b8f
3 changed files with 420 additions and 61 deletions

View File

@@ -74,10 +74,6 @@ func NewRunnerWithRepairer(
}
func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunResult, error) {
if strings.TrimSpace(req.PromptID) == "" {
return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidRequest)
}
runID, err := newRunID()
if err != nil {
return nil, fmt.Errorf("failed to create run id: %w", err)
@@ -85,6 +81,85 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
start := time.Now().UTC()
prepared, err := r.Prepare(ctx, req)
if err != nil {
return nil, err
}
genResp, err := r.llm.Generate(ctx, domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: prepared.Messages},
Target: prepared.EffectiveModelParams,
})
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrLLMGenerate, err)
}
outputArtifact := buildOutputArtifact(genResp.Content, prepared.OutputContract.Format)
validationResult, err := r.validateOutput(ctx, &outputArtifact, prepared.OutputContract, 0)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, err)
}
if r.shouldAttemptRepair(prepared.OutputContract, validationResult) {
attemptsUsed := 0
for attemptsUsed < prepared.OutputContract.RepairAttempts && validationResult.Status == domain.ValidationFailed {
attemptsUsed++
repairResp, repairErr := r.repairer.Repair(ctx, RepairRequest{
PreviousOutput: genResp.Content,
ValidationErrors: validationResult.Errors,
Target: prepared.EffectiveModelParams,
Attempt: attemptsUsed,
MaxAttempts: prepared.OutputContract.RepairAttempts,
Mode: prepared.OutputContract.ValidationMode,
})
if repairErr != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, repairErr)
}
if repairResp == nil {
return nil, fmt.Errorf("%w: repairer returned nil response", ErrValidation)
}
genResp = repairResp
outputArtifact = buildOutputArtifact(genResp.Content, prepared.OutputContract.Format)
validationResult, err = r.validateOutput(ctx, &outputArtifact, prepared.OutputContract, attemptsUsed)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, err)
}
}
}
end := time.Now().UTC()
return &domain.RunResult{
RunID: runID,
Artifact: outputArtifact,
RawOutput: genResp.Content,
Validation: validationResult,
PromptID: prepared.PromptID,
PromptVersion: prepared.PromptVersion,
PromptHash: prepared.PromptHash,
RenderedPromptHash: prepared.RenderedPromptHash,
SelectedProfileID: prepared.SelectedProfileID,
ModelName: prepared.EffectiveModelParams.Model,
Endpoint: prepared.EffectiveModelParams.Endpoint,
EffectiveModelParams: prepared.EffectiveModelParams,
InputHashes: prepared.InputHashes,
Usage: genResp.Usage,
StartTime: start,
EndTime: end,
Duration: end.Sub(start),
}, nil
}
func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.PreparedRun, error) {
if strings.TrimSpace(req.PromptID) == "" {
return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidRequest)
}
start := time.Now().UTC()
def, err := r.promptDefs.GetPromptDefinition(ctx, req.PromptID, req.PromptVersion)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err)
@@ -93,6 +168,7 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
if err != nil {
return nil, fmt.Errorf("%w: failed to hash prompt definition: %v", ErrProfileLoad, err)
}
selectedProfileID := strings.TrimSpace(req.ProfileID)
if selectedProfileID == "" {
selectedProfileID = strings.TrimSpace(def.DefaultProfile)
@@ -100,10 +176,12 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
if selectedProfileID == "" {
return nil, fmt.Errorf("%w: profile id is required either in request or prompt default_profile", ErrInvalidRequest)
}
execProfile, err := r.profiles.GetProfile(ctx, selectedProfileID)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err)
}
effectiveModel := resolveExecutionTarget(execProfile, req.Execution)
if strings.TrimSpace(effectiveModel.Endpoint) == "" {
return nil, fmt.Errorf("%w: execution endpoint is required", ErrInvalidRequest)
@@ -114,6 +192,7 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
if err := validateAPIKeyEnv(effectiveModel.APIKeyEnv); err != nil {
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
}
effectiveContract := resolveOutputContract(def, req.Validation)
resolvedInputs := make(map[string]*domain.Artifact, len(req.Inputs))
@@ -135,72 +214,20 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
return nil, fmt.Errorf("%w: %w", ErrPromptRender, err)
}
renderedPromptHash := hashRenderedPrompt(*renderedPrompt)
genResp, err := r.llm.Generate(ctx, domain.GenerateRequest{
Prompt: *renderedPrompt,
Target: effectiveModel,
})
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrLLMGenerate, err)
}
outputArtifact := buildOutputArtifact(genResp.Content, effectiveContract.Format)
validationResult, err := r.validateOutput(ctx, &outputArtifact, effectiveContract, 0)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, err)
}
if r.shouldAttemptRepair(effectiveContract, validationResult) {
attemptsUsed := 0
for attemptsUsed < effectiveContract.RepairAttempts && validationResult.Status == domain.ValidationFailed {
attemptsUsed++
repairResp, repairErr := r.repairer.Repair(ctx, RepairRequest{
PreviousOutput: genResp.Content,
ValidationErrors: validationResult.Errors,
Target: effectiveModel,
Attempt: attemptsUsed,
MaxAttempts: effectiveContract.RepairAttempts,
Mode: effectiveContract.ValidationMode,
})
if repairErr != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, repairErr)
}
if repairResp == nil {
return nil, fmt.Errorf("%w: repairer returned nil response", ErrValidation)
}
genResp = repairResp
outputArtifact = buildOutputArtifact(genResp.Content, effectiveContract.Format)
validationResult, err = r.validateOutput(ctx, &outputArtifact, effectiveContract, attemptsUsed)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, err)
}
}
}
end := time.Now().UTC()
return &domain.RunResult{
RunID: runID,
Artifact: outputArtifact,
RawOutput: genResp.Content,
Validation: validationResult,
return &domain.PreparedRun{
PromptID: def.ID,
PromptVersion: def.Version,
PromptHash: promptDefinitionHash,
RenderedPromptHash: renderedPromptHash,
SelectedProfileID: selectedProfileID,
ModelName: effectiveModel.Model,
Endpoint: effectiveModel.Endpoint,
EffectiveModelParams: effectiveModel,
OutputContract: effectiveContract,
InputHashes: inputHashes,
Usage: genResp.Usage,
RenderedPromptHash: hashRenderedPrompt(*renderedPrompt),
Messages: renderedPrompt.Messages,
StartTime: start,
EndTime: end,
Duration: end.Sub(start),
DurationMS: end.Sub(start).Milliseconds(),
}, nil
}