Run bounded output repair attempts
This commit is contained in:
@@ -184,10 +184,10 @@ func (r *Runner) executePreparedRun(
|
||||
prepared.StructuredOutput,
|
||||
))
|
||||
if err != nil {
|
||||
if errors.Is(err, llm.ErrInvalidRequest) {
|
||||
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
|
||||
}
|
||||
return nil, fmt.Errorf("%w: %w", ErrLLMGenerate, err)
|
||||
return nil, wrapGenerationError(err)
|
||||
}
|
||||
if genResp == nil {
|
||||
return nil, fmt.Errorf("%w: model returned nil response", ErrLLMGenerate)
|
||||
}
|
||||
usage := genResp.Usage
|
||||
|
||||
@@ -203,6 +203,7 @@ func (r *Runner) executePreparedRun(
|
||||
attemptsUsed++
|
||||
|
||||
repairResp, repairErr := r.repairer.Repair(ctx, RepairRequest{
|
||||
OriginalMessages: prepared.Messages,
|
||||
PreviousOutput: genResp.Content,
|
||||
ValidationErrors: validationResult.Errors,
|
||||
SessionID: prepared.SessionID,
|
||||
@@ -214,10 +215,10 @@ func (r *Runner) executePreparedRun(
|
||||
Mode: prepared.OutputContract.ValidationMode,
|
||||
})
|
||||
if repairErr != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrValidation, repairErr)
|
||||
return nil, wrapGenerationError(repairErr)
|
||||
}
|
||||
if repairResp == nil {
|
||||
return nil, fmt.Errorf("%w: repairer returned nil response", ErrValidation)
|
||||
return nil, fmt.Errorf("%w: repairer returned nil response", ErrLLMGenerate)
|
||||
}
|
||||
|
||||
genResp = repairResp
|
||||
@@ -257,6 +258,13 @@ func (r *Runner) executePreparedRun(
|
||||
}, nil
|
||||
}
|
||||
|
||||
func wrapGenerationError(err error) error {
|
||||
if errors.Is(err, llm.ErrInvalidRequest) {
|
||||
return fmt.Errorf("%w: %w", ErrInvalidRequest, err)
|
||||
}
|
||||
return fmt.Errorf("%w: %w", ErrLLMGenerate, err)
|
||||
}
|
||||
|
||||
func addTokenUsage(total, next domain.TokenUsage) domain.TokenUsage {
|
||||
return domain.TokenUsage{
|
||||
PromptTokens: total.PromptTokens + next.PromptTokens,
|
||||
@@ -524,7 +532,12 @@ func (r *Runner) shouldAttemptRepair(contract domain.OutputContract, validationR
|
||||
if validationResult.Status != domain.ValidationFailed {
|
||||
return false
|
||||
}
|
||||
return contract.ValidationMode == domain.ValidationJSON || contract.ValidationMode == domain.ValidationJSONSchema
|
||||
switch contract.ValidationMode {
|
||||
case domain.ValidationBasic, domain.ValidationJSON, domain.ValidationJSONSchema:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func mergeExecutionTarget(base domain.ExecutionTarget, override domain.ExecutionTarget) domain.ExecutionTarget {
|
||||
|
||||
Reference in New Issue
Block a user