Run bounded output repair attempts
This commit is contained in:
@@ -3,10 +3,12 @@ package usecase
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/llm"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/validate"
|
||||
)
|
||||
|
||||
@@ -385,19 +387,20 @@ func TestRunnerRunPreparedUsesFrozenValidationForInitialAndRepairOutputs(t *test
|
||||
admitter := &fakeRunAdmitter{}
|
||||
reader := defaultArtifactReader()
|
||||
renderer := defaultRenderer()
|
||||
client := &sequenceLLM{responses: []*domain.GenerateResponse{{
|
||||
Content: `{"broken":true}`,
|
||||
Usage: domain.TokenUsage{
|
||||
PromptTokens: 13, CompletionTokens: 17, TotalTokens: 19,
|
||||
CachedTokens: 23, CacheWriteTokens: 29,
|
||||
},
|
||||
}}}
|
||||
runner := NewRunnerWithRepairer(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatJSON, domain.ValidationJSON, 1)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
reader,
|
||||
renderer,
|
||||
&fakeLLM{resp: &domain.GenerateResponse{
|
||||
Content: `{"broken":true}`,
|
||||
Usage: domain.TokenUsage{
|
||||
PromptTokens: 13, CompletionTokens: 17, TotalTokens: 19,
|
||||
CachedTokens: 23, CacheWriteTokens: 29,
|
||||
},
|
||||
}},
|
||||
client,
|
||||
validator,
|
||||
repairer,
|
||||
admitter,
|
||||
@@ -421,6 +424,10 @@ func TestRunnerRunPreparedUsesFrozenValidationForInitialAndRepairOutputs(t *test
|
||||
if validator.directValidateCalls != 0 || repairer.calls != 1 {
|
||||
t.Fatalf("validation/repair calls=(direct=%d repair=%d), want (0, 1)", validator.directValidateCalls, repairer.calls)
|
||||
}
|
||||
if len(client.requests) != 1 || len(repairer.reqs) != 1 ||
|
||||
!reflect.DeepEqual(repairer.reqs[0].OriginalMessages, client.requests[0].Prompt.Messages) {
|
||||
t.Fatalf("initial and repair messages = (%#v, %#v)", client.requests, repairer.reqs)
|
||||
}
|
||||
if result.Validation.Status != domain.ValidationPassed || result.Validation.RepairAttempts != 1 {
|
||||
t.Fatalf("unexpected repaired validation result: %+v", result.Validation)
|
||||
}
|
||||
@@ -443,6 +450,7 @@ func TestRunnerRunPreparedReleasesAdmissionAcrossExecutionErrors(t *testing.T) {
|
||||
generationFailure := errors.New("generation failed")
|
||||
validationFailure := errors.New("validation failed")
|
||||
repairFailure := errors.New("repair failed")
|
||||
repairInvalidRequest := fmt.Errorf("repair request: %w", llm.ErrInvalidRequest)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -450,6 +458,7 @@ func TestRunnerRunPreparedReleasesAdmissionAcrossExecutionErrors(t *testing.T) {
|
||||
validation *recordingPreparedValidation
|
||||
repairer *fakeRepairer
|
||||
wantError error
|
||||
wantSource error
|
||||
}{
|
||||
{
|
||||
name: "generation failure",
|
||||
@@ -474,8 +483,51 @@ func TestRunnerRunPreparedReleasesAdmissionAcrossExecutionErrors(t *testing.T) {
|
||||
IsValid: false,
|
||||
}},
|
||||
},
|
||||
repairer: &fakeRepairer{err: repairFailure},
|
||||
wantError: ErrValidation,
|
||||
repairer: &fakeRepairer{err: repairFailure},
|
||||
wantError: ErrLLMGenerate,
|
||||
wantSource: repairFailure,
|
||||
},
|
||||
{
|
||||
name: "repair invalid request",
|
||||
validation: &recordingPreparedValidation{
|
||||
results: []domain.ValidationResult{{
|
||||
Status: domain.ValidationFailed,
|
||||
Mode: domain.ValidationJSON,
|
||||
Errors: []string{"invalid"},
|
||||
IsValid: false,
|
||||
}},
|
||||
},
|
||||
repairer: &fakeRepairer{err: repairInvalidRequest},
|
||||
wantError: ErrInvalidRequest,
|
||||
wantSource: repairInvalidRequest,
|
||||
},
|
||||
{
|
||||
name: "repair cancellation",
|
||||
validation: &recordingPreparedValidation{
|
||||
results: []domain.ValidationResult{{
|
||||
Status: domain.ValidationFailed,
|
||||
Mode: domain.ValidationJSON,
|
||||
Errors: []string{"invalid"},
|
||||
IsValid: false,
|
||||
}},
|
||||
},
|
||||
repairer: &fakeRepairer{err: context.Canceled},
|
||||
wantError: ErrLLMGenerate,
|
||||
wantSource: context.Canceled,
|
||||
},
|
||||
{
|
||||
name: "repair deadline",
|
||||
validation: &recordingPreparedValidation{
|
||||
results: []domain.ValidationResult{{
|
||||
Status: domain.ValidationFailed,
|
||||
Mode: domain.ValidationJSON,
|
||||
Errors: []string{"invalid"},
|
||||
IsValid: false,
|
||||
}},
|
||||
},
|
||||
repairer: &fakeRepairer{err: context.DeadlineExceeded},
|
||||
wantError: ErrLLMGenerate,
|
||||
wantSource: context.DeadlineExceeded,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -514,6 +566,9 @@ func TestRunnerRunPreparedReleasesAdmissionAcrossExecutionErrors(t *testing.T) {
|
||||
if result != nil || !errors.Is(err, test.wantError) {
|
||||
t.Fatalf("run prepared=(%+v, %v), want %v", result, err, test.wantError)
|
||||
}
|
||||
if test.wantSource != nil && !errors.Is(err, test.wantSource) {
|
||||
t.Fatalf("run prepared error = %v, want source %v", err, test.wantSource)
|
||||
}
|
||||
if len(admitter.backendIDs) != 1 || admitter.releaseCalls != 1 {
|
||||
t.Fatalf(
|
||||
"admission calls=%#v releases=%d, want one each",
|
||||
|
||||
Reference in New Issue
Block a user