Enable bounded output repair in the engine

This commit is contained in:
2026-08-25 09:55:19 +00:00
parent ee99dc9478
commit ae6f1a9865
7 changed files with 213 additions and 31 deletions

View File

@@ -256,6 +256,50 @@ func TestPreparedExecutionLifecycleAndEngineBinding(t *testing.T) {
}
}
func TestPreparedExecutionRepairsEmptyBasicOutput(t *testing.T) {
client := &preparedRecordingClient{responses: []*promptkit.GenerateResponse{
{
Content: "",
Usage: promptkit.TokenUsage{PromptTokens: 3, CompletionTokens: 5, TotalTokens: 8},
},
{
Content: "Corrected summary.",
Usage: promptkit.TokenUsage{PromptTokens: 7, CompletionTokens: 11, TotalTokens: 18},
},
}}
engine := newPreparedContractEngine(t, client, "Summarize the source.")
prepared, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{
PromptID: "prepared",
Validation: &promptkit.OutputContract{
Format: promptkit.FormatMarkdown,
ValidationMode: promptkit.ValidationBasic,
RepairAttempts: 1,
},
})
if err != nil {
t.Fatalf("prepare execution: %v", err)
}
details := prepared.Details()
result, err := engine.RunPrepared(context.Background(), prepared)
if err != nil {
t.Fatalf("run prepared: %v", err)
}
if result.RawOutput != "Corrected summary." || result.Validation.Status != promptkit.ValidationPassed ||
result.Validation.RepairAttempts != 1 || result.Usage != (promptkit.TokenUsage{PromptTokens: 10, CompletionTokens: 16, TotalTokens: 26}) {
t.Fatalf("repaired result = %+v", result)
}
requests := client.snapshot()
if len(requests) != 2 || len(requests[1].Prompt.Messages) != len(details.Messages)+1 ||
!reflect.DeepEqual(requests[1].Prompt.Messages[:len(details.Messages)], details.Messages) ||
requests[1].Prompt.Messages[len(requests[1].Prompt.Messages)-1].Role != "user" {
t.Fatalf("prepared repair requests = %#v", requests)
}
if _, err := engine.RunPrepared(context.Background(), prepared); !errors.Is(err, promptkit.ErrInvalidRequest) {
t.Fatalf("second RunPrepared error = %v, want ErrInvalidRequest", err)
}
}
func TestPreparedExecutionConcurrentClaimAllowsOneGeneration(t *testing.T) {
release := make(chan struct{})
client := &preparedRecordingClient{
@@ -671,12 +715,13 @@ func (r *mutablePreparedArtifactReader) callCount() int {
}
type preparedRecordingClient struct {
mu sync.Mutex
response *promptkit.GenerateResponse
err error
requests []promptkit.GenerateRequest
started chan struct{}
release <-chan struct{}
mu sync.Mutex
response *promptkit.GenerateResponse
responses []*promptkit.GenerateResponse
err error
requests []promptkit.GenerateRequest
started chan struct{}
release <-chan struct{}
}
func (c *preparedRecordingClient) Generate(
@@ -700,6 +745,12 @@ func (c *preparedRecordingClient) Generate(
if c.err != nil {
return nil, c.err
}
if len(c.responses) > 0 {
if index := len(c.requests) - 1; index < len(c.responses) {
return c.responses[index], nil
}
return nil, fmt.Errorf("no response configured for generation %d", len(c.requests))
}
return c.response, nil
}