Enable bounded output repair in the engine
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user