diff --git a/docs/roadmap/implementation.md b/docs/roadmap/implementation.md index 32045292..d666a4f6 100644 --- a/docs/roadmap/implementation.md +++ b/docs/roadmap/implementation.md @@ -344,7 +344,7 @@ git diff --check - Capacity and context classifications remain unchanged and more specific. - Security tests prove selected credentials and bearer tokens are not leaked. -## Stage 5: Add Adapter-Level Structured Repair Support +## Stage 5: Add Adapter-Level Structured Repair Support ✅ ### Goal diff --git a/internal/framework/contracts/contracts.go b/internal/framework/contracts/contracts.go index 00fac69a..0465d954 100644 --- a/internal/framework/contracts/contracts.go +++ b/internal/framework/contracts/contracts.go @@ -9,13 +9,14 @@ import ( ) type StructuredCompletionRequest struct { - StageName string `json:"stage_name"` - PromptID string `json:"prompt_id,omitempty"` - PromptVersion string `json:"prompt_version,omitempty"` - ProfileID string `json:"profile_id,omitempty"` - SessionID string `json:"session_id,omitempty"` - Inputs LLMInputSet `json:"inputs,omitempty"` - Vars map[string]any `json:"vars,omitempty"` + StageName string `json:"stage_name"` + PromptID string `json:"prompt_id,omitempty"` + PromptVersion string `json:"prompt_version,omitempty"` + ProfileID string `json:"profile_id,omitempty"` + SessionID string `json:"session_id,omitempty"` + Inputs LLMInputSet `json:"inputs,omitempty"` + Vars map[string]any `json:"vars,omitempty"` + StructuredOutputRepairAttempts *int `json:"structured_output_repair_attempts,omitempty"` } type StructuredCompletionResponse struct { @@ -26,6 +27,7 @@ type StructuredCompletionResponse struct { PromptTokens int `json:"prompt_tokens,omitempty"` CompletionTokens int `json:"completion_tokens,omitempty"` TotalTokens int `json:"total_tokens,omitempty"` + RepairAttempts int `json:"repair_attempts,omitempty"` Debug *LLMDebugMaterial `json:"debug,omitempty"` } diff --git a/internal/framework/llm/promptkit_client.go b/internal/framework/llm/promptkit_client.go index 80e5487b..e035d60f 100644 --- a/internal/framework/llm/promptkit_client.go +++ b/internal/framework/llm/promptkit_client.go @@ -116,6 +116,9 @@ func (c *PromptKitClient) CompleteStructured(ctx context.Context, req contracts. if err := validateOutputTarget(out); err != nil { return contracts.StructuredCompletionResponse{}, err } + if req.StructuredOutputRepairAttempts != nil && (*req.StructuredOutputRepairAttempts < 0 || *req.StructuredOutputRepairAttempts > 3) { + return contracts.StructuredCompletionResponse{}, fmt.Errorf("structured output repair attempts must be between zero and three") + } promptID := strings.TrimSpace(req.PromptID) if promptID == "" { return contracts.StructuredCompletionResponse{}, fmt.Errorf("structured completion prompt_id must not be empty") @@ -138,6 +141,15 @@ func (c *PromptKitClient) CompleteStructured(ctx context.Context, req contracts. Vars: promptKitVars(req, sessionID), Execution: execution, } + if req.StructuredOutputRepairAttempts != nil { + inspection, err := c.engine.InspectPrompt(ctx, promptID, strings.TrimSpace(req.PromptVersion)) + if err != nil { + return contracts.StructuredCompletionResponse{}, fmt.Errorf("inspect PromptKit prompt %q: %v", promptID, redactPromptKitError(err)) + } + contract := inspection.OutputContract + contract.RepairAttempts = *req.StructuredOutputRepairAttempts + runReq.Validation = &contract + } prepared, err := c.engine.PrepareExecution(ctx, runReq) if err != nil { if ctxErr := ctx.Err(); ctxErr != nil { @@ -242,6 +254,7 @@ func (c *PromptKitClient) responseFromResult(result *promptkit.RunResult, prepar PromptTokens: result.Usage.PromptTokens, CompletionTokens: result.Usage.CompletionTokens, TotalTokens: result.Usage.TotalTokens, + RepairAttempts: result.Validation.RepairAttempts, Debug: promptKitDebugMaterial(prepared, result), } } diff --git a/internal/framework/llm/promptkit_client_test.go b/internal/framework/llm/promptkit_client_test.go index 0b4241d8..1a6d9a3a 100644 --- a/internal/framework/llm/promptkit_client_test.go +++ b/internal/framework/llm/promptkit_client_test.go @@ -667,6 +667,20 @@ func TestPromptKitClientMapsHTTPGenerationErrorToApplicationBoundary(t *testing. } } +func TestPromptKitClientRejectsInvalidRepairAttemptOverride(t *testing.T) { + for _, attempts := range []int{-1, 4} { + t.Run("invalid", func(t *testing.T) { + fake := &fakePromptKitLLM{content: `{"ok":true}`} + client := newTestPromptKitClient(t, fake) + var out map[string]any + _, err := client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{PromptID: "adapter.test", StructuredOutputRepairAttempts: &attempts}, &out) + if err == nil || fake.calls != 0 { + t.Fatalf("error=%v calls=%d", err, fake.calls) + } + }) + } +} + func TestPromptKitClientCheckpointFingerprintTracksFallbackProfileAssets(t *testing.T) { fingerprintFor := func(content string) CheckpointFingerprint { t.Helper()