From 1a7b20c76633610e4431ecb879007a4803efa0d3 Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Wed, 26 Aug 2026 23:50:24 +0000 Subject: [PATCH] Adapt corrections through PromptKit --- docs/internal/llm.md | 8 ++ docs/roadmap/implementation.md | 2 +- .../contracts/completion_request_debug.go | 84 ++++++++++++++++++ .../framework/contracts/correction_test.go | 35 ++++++++ internal/framework/llm/promptkit_client.go | 33 +++++-- .../framework/llm/promptkit_client_test.go | 85 +++++++++++++++++++ internal/framework/pipeline/debug.go | 43 +--------- internal/framework/pipeline/debug_test.go | 26 ++++++ 8 files changed, 267 insertions(+), 49 deletions(-) create mode 100644 internal/framework/contracts/completion_request_debug.go diff --git a/docs/internal/llm.md b/docs/internal/llm.md index 1a37a392..8021f938 100644 --- a/docs/internal/llm.md +++ b/docs/internal/llm.md @@ -42,6 +42,14 @@ observability. The adapter returns PromptKit’s validated raw bytes rather than re-encoding the decoded target. An empty optional material is represented as one space so its named input is retained by PromptKit. +When a request includes semantic correction material, the adapter validates and +defensively copies it before preparation, then appends exactly two messages +after the ordinarily rendered prompt: the prior response as an assistant +message and the correction guidance as a user message. Requests without a +correction do not add messages or introduce caller roles. Ordinary request +summaries record correction byte counts and digests only; complete messages are +available solely in an explicitly requested debug trace. + Client construction may also receive a run-wide reasoning-effort override from the CLI factory boundary. The adapter copies the caller-owned pointer and creates a fresh PromptKit execution override for each request: a nil pointer diff --git a/docs/roadmap/implementation.md b/docs/roadmap/implementation.md index f9d38a20..81bc8cd7 100644 --- a/docs/roadmap/implementation.md +++ b/docs/roadmap/implementation.md @@ -269,7 +269,7 @@ Capability is explicit, fingerprinted, and testable, while existing production pipelines remain executable pending their migrations. No dormant feature gate or unused preparation check is introduced. This stage is one Terra prompt. -## Stage 5 — Adapt Corrections Through PromptKit +## Stage 5 — Adapt Corrections Through PromptKit ✅ ### Goal diff --git a/internal/framework/contracts/completion_request_debug.go b/internal/framework/contracts/completion_request_debug.go new file mode 100644 index 00000000..5598524c --- /dev/null +++ b/internal/framework/contracts/completion_request_debug.go @@ -0,0 +1,84 @@ +package contracts + +import ( + "crypto/sha256" + "encoding/hex" + "fmt" +) + +// DebugStructuredCompletionRequest is the content-safe representation of a +// structured completion request for ordinary diagnostics and summaries. +// Detailed prompt material remains available only through the explicitly +// requested LLM debug trace. +type DebugStructuredCompletionRequest struct { + StageName string `json:"stage_name,omitempty"` + PromptID string `json:"prompt_id,omitempty"` + PromptVersion string `json:"prompt_version,omitempty"` + ProfileID string `json:"profile_id,omitempty"` + SessionID string `json:"session_id,omitempty"` + InputCount int `json:"input_count"` + VariableCount int `json:"variable_count"` + StructuredOutputRepairAttempts *int `json:"structured_output_repair_attempts,omitempty"` + Correction *DebugSemanticCorrection `json:"correction,omitempty"` +} + +// DebugSemanticCorrection records only safe correction metadata. It never +// exposes the assistant response or user guidance text. +type DebugSemanticCorrection struct { + AssistantResponseBytes int `json:"assistant_response_bytes"` + AssistantResponseDigest string `json:"assistant_response_digest"` + UserGuidanceBytes int `json:"user_guidance_bytes"` + UserGuidanceDigest string `json:"user_guidance_digest"` +} + +// DebugSummary returns a content-safe representation suitable for ordinary +// diagnostics. It does not validate or retain correction content. +func (request StructuredCompletionRequest) DebugSummary() DebugStructuredCompletionRequest { + summary := DebugStructuredCompletionRequest{ + StageName: request.StageName, + PromptID: request.PromptID, + PromptVersion: request.PromptVersion, + ProfileID: request.ProfileID, + SessionID: request.SessionID, + InputCount: len(request.Inputs), + VariableCount: len(request.Vars), + } + if request.StructuredOutputRepairAttempts != nil { + attempts := *request.StructuredOutputRepairAttempts + summary.StructuredOutputRepairAttempts = &attempts + } + if request.Correction != nil { + summary.Correction = request.Correction.DebugSummary() + } + return summary +} + +// DebugSummary returns content-safe correction metadata suitable for ordinary +// diagnostics. +func (correction *SemanticCorrection) DebugSummary() *DebugSemanticCorrection { + if correction == nil { + return nil + } + return &DebugSemanticCorrection{ + AssistantResponseBytes: len(correction.AssistantResponse), + AssistantResponseDigest: debugContentDigest(correction.AssistantResponse), + UserGuidanceBytes: len(correction.UserGuidance), + UserGuidanceDigest: debugContentDigest([]byte(correction.UserGuidance)), + } +} + +// String prevents ordinary request formatting from exposing correction +// content. Use the explicitly requested debug trace for complete messages. +func (request StructuredCompletionRequest) String() string { + return fmt.Sprintf("%+v", request.DebugSummary()) +} + +// GoString gives %#v formatting the same content-safe behavior as String. +func (request StructuredCompletionRequest) GoString() string { + return request.String() +} + +func debugContentDigest(content []byte) string { + sum := sha256.Sum256(content) + return "sha256:" + hex.EncodeToString(sum[:]) +} diff --git a/internal/framework/contracts/correction_test.go b/internal/framework/contracts/correction_test.go index 03708879..afb178ca 100644 --- a/internal/framework/contracts/correction_test.go +++ b/internal/framework/contracts/correction_test.go @@ -3,6 +3,7 @@ package contracts import ( "bytes" "encoding/json" + "fmt" "strings" "testing" ) @@ -139,3 +140,37 @@ func TestCloneStructuredCompletionRequestOwnsCorrection(t *testing.T) { t.Fatalf("cloned repair attempts = %v, want 2", clone.StructuredOutputRepairAttempts) } } + +func TestStructuredCompletionRequestDebugSummaryOmitsCorrectionContent(t *testing.T) { + const assistantResponse = `{"secret":"assistant response"}` + const userGuidance = "secret user guidance" + correction, err := NewSemanticCorrection([]byte(assistantResponse), userGuidance) + if err != nil { + t.Fatalf("NewSemanticCorrection() error = %v", err) + } + request := StructuredCompletionRequest{ + Inputs: LLMInputSet{"source": NewLLMInputMaterial("source", "application/json", []byte(`{"source":true}`), "", "")}, + Vars: map[string]any{"custom": "value"}, + Correction: correction, + } + + summary := request.DebugSummary() + if summary.InputCount != 1 || summary.VariableCount != 1 || summary.Correction == nil { + t.Fatalf("debug summary = %#v, want input, variable, and correction metadata", summary) + } + if summary.Correction.AssistantResponseBytes != len(assistantResponse) || summary.Correction.UserGuidanceBytes != len(userGuidance) || + summary.Correction.AssistantResponseDigest == "" || summary.Correction.UserGuidanceDigest == "" { + t.Fatalf("correction summary = %#v, want byte counts and digests", summary.Correction) + } + encoded, err := json.Marshal(summary) + if err != nil { + t.Fatalf("marshal debug summary: %v", err) + } + for _, rendered := range []string{string(encoded), fmt.Sprintf("%+v", request), fmt.Sprintf("%#v", request)} { + for _, secret := range []string{assistantResponse, userGuidance} { + if strings.Contains(rendered, secret) { + t.Fatalf("content-safe request rendering leaked %q: %s", secret, rendered) + } + } + } +} diff --git a/internal/framework/llm/promptkit_client.go b/internal/framework/llm/promptkit_client.go index d8a337f5..293055b6 100644 --- a/internal/framework/llm/promptkit_client.go +++ b/internal/framework/llm/promptkit_client.go @@ -119,6 +119,10 @@ func (c *PromptKitClient) CompleteStructured(ctx context.Context, req contracts. 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") } + appendedMessages, err := promptKitCorrectionMessages(req.Correction) + if err != nil { + return contracts.StructuredCompletionResponse{}, fmt.Errorf("structured completion correction: %w", err) + } promptID := strings.TrimSpace(req.PromptID) if promptID == "" { return contracts.StructuredCompletionResponse{}, fmt.Errorf("structured completion prompt_id must not be empty") @@ -133,13 +137,14 @@ func (c *PromptKitClient) CompleteStructured(ctx context.Context, req contracts. } runReq := promptkit.RunRequest{ - PromptID: promptID, - PromptVersion: strings.TrimSpace(req.PromptVersion), - ProfileID: strings.TrimSpace(req.ProfileID), - SessionID: sessionID, - Inputs: promptKitInputs(req.Inputs), - Vars: promptKitVars(req, sessionID), - Execution: execution, + PromptID: promptID, + PromptVersion: strings.TrimSpace(req.PromptVersion), + ProfileID: strings.TrimSpace(req.ProfileID), + SessionID: sessionID, + Inputs: promptKitInputs(req.Inputs), + Vars: promptKitVars(req, sessionID), + Execution: execution, + AppendedMessages: appendedMessages, } if req.StructuredOutputRepairAttempts != nil { inspection, err := c.engine.InspectPrompt(ctx, promptID, strings.TrimSpace(req.PromptVersion)) @@ -213,6 +218,20 @@ func (c *PromptKitClient) CompleteStructured(ctx context.Context, req contracts. return response, nil } +func promptKitCorrectionMessages(correction *contracts.SemanticCorrection) ([]promptkit.RenderedMessage, error) { + owned, err := contracts.CloneSemanticCorrection(correction) + if err != nil { + return nil, err + } + if owned == nil { + return nil, nil + } + return []promptkit.RenderedMessage{ + {Role: promptkit.RoleAssistant, Content: string(owned.AssistantResponse)}, + {Role: promptkit.RoleUser, Content: owned.UserGuidance}, + }, nil +} + func promptKitDebugGenerationError(prepared *promptkit.PreparedRun, generationErr *promptkit.GenerationError) *contracts.LLMDebugResponse { if generationErr == nil { return nil diff --git a/internal/framework/llm/promptkit_client_test.go b/internal/framework/llm/promptkit_client_test.go index ce44f557..1b7ff767 100644 --- a/internal/framework/llm/promptkit_client_test.go +++ b/internal/framework/llm/promptkit_client_test.go @@ -1,6 +1,7 @@ package llm import ( + "bytes" "context" "encoding/json" "errors" @@ -116,6 +117,74 @@ func TestPromptKitClientMapsPromptRequestAndUnmarshalsOutput(t *testing.T) { } } +func TestPromptKitClientAppendsSemanticCorrectionAfterRenderedPrompt(t *testing.T) { + fake := &fakePromptKitLLM{content: `{"ok":true}`} + client := newTestPromptKitClient(t, fake) + request := contracts.StructuredCompletionRequest{ + PromptID: "adapter.direct-session", + ProfileID: "explicit-profile", + SessionID: "correction-session", + Inputs: contracts.LLMInputSet{ + "transcript": contracts.NewLLMInputMaterial("transcript", "application/json", []byte(`{"source":true}`), "", ""), + }, + Vars: map[string]any{"custom": "value"}, + } + + var ordinary map[string]any + if _, err := client.CompleteStructured(context.Background(), request, &ordinary); err != nil { + t.Fatalf("ordinary CompleteStructured() error = %v", err) + } + ordinaryMessages := append([]promptkit.RenderedMessage(nil), fake.lastRequest().Prompt.Messages...) + if len(ordinaryMessages) != 1 { + t.Fatalf("ordinary rendered messages = %#v, want only the declared prompt message", ordinaryMessages) + } + + correction, err := contracts.NewSemanticCorrection([]byte(`{"previous":"response"}`), "Return the corrected JSON object.") + if err != nil { + t.Fatalf("NewSemanticCorrection() error = %v", err) + } + request.Correction = correction + var corrected map[string]any + if _, err := client.CompleteStructured(context.Background(), request, &corrected); err != nil { + t.Fatalf("corrected CompleteStructured() error = %v", err) + } + correctedMessages := fake.lastRequest().Prompt.Messages + if len(correctedMessages) != len(ordinaryMessages)+2 { + t.Fatalf("corrected message count = %d, want %d", len(correctedMessages), len(ordinaryMessages)+2) + } + if !reflect.DeepEqual(correctedMessages[:len(ordinaryMessages)], ordinaryMessages) { + t.Fatalf("ordinary rendered prefix changed: got %#v, want %#v", correctedMessages[:len(ordinaryMessages)], ordinaryMessages) + } + if got, want := correctedMessages[len(ordinaryMessages)], (promptkit.RenderedMessage{Role: promptkit.RoleAssistant, Content: `{"previous":"response"}`}); got != want { + t.Fatalf("assistant correction message = %#v, want %#v", got, want) + } + if got, want := correctedMessages[len(ordinaryMessages)+1], (promptkit.RenderedMessage{Role: promptkit.RoleUser, Content: "Return the corrected JSON object."}); got != want { + t.Fatalf("user correction message = %#v, want %#v", got, want) + } +} + +func TestPromptKitClientRejectsInvalidCorrectionsBeforePromptPreparation(t *testing.T) { + const sensitiveResponse = "assistant-response-must-not-appear-in-errors" + for _, correction := range []*contracts.SemanticCorrection{ + {AssistantResponse: []byte(sensitiveResponse), UserGuidance: " \t"}, + {AssistantResponse: bytes.Repeat([]byte(sensitiveResponse), contracts.MaxAssistantResponseBytes/len(sensitiveResponse)+1), UserGuidance: "Use a smaller response."}, + } { + fake := &fakePromptKitLLM{content: `{"ok":true}`} + client := newTestPromptKitClient(t, fake) + var out map[string]any + _, err := client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{Correction: correction}, &out) + if err == nil || !strings.Contains(err.Error(), "correction") { + t.Fatalf("CompleteStructured() error = %v, want correction validation failure", err) + } + if strings.Contains(err.Error(), sensitiveResponse) { + t.Fatalf("correction validation error leaked response content: %v", err) + } + if got := atomic.LoadInt32(&fake.calls); got != 0 { + t.Fatalf("provider calls = %d, want no provider call after invalid correction", got) + } + } +} + func TestPromptKitClientUsesOnePreparedSnapshotForDebugAndGeneration(t *testing.T) { const initialPrompt = `id: snapshot.test version: "v1" @@ -960,6 +1029,10 @@ func TestPromptKitClientRepairsStructuredOutputAndReportsCumulativeUsage(t *test {Content: `{"ok":true}`, Usage: promptkit.TokenUsage{PromptTokens: 7, CompletionTokens: 11, TotalTokens: 18}}, }} client := newTestPromptKitClient(t, fake) + correction, err := contracts.NewSemanticCorrection([]byte(`{"bad":true}`), "Return the required ok field.") + if err != nil { + t.Fatalf("NewSemanticCorrection() error = %v", err) + } var out struct { OK bool `json:"ok"` @@ -968,6 +1041,7 @@ func TestPromptKitClientRepairsStructuredOutputAndReportsCumulativeUsage(t *test PromptID: "adapter.test", StructuredOutputRepairAttempts: &attempts, SessionID: "repair-test", + Correction: correction, Inputs: contracts.LLMInputSet{ "transcript": contracts.NewLLMInputMaterial("transcript", "application/json", []byte(`{"source":true}`), "", ""), }, @@ -978,6 +1052,17 @@ func TestPromptKitClientRepairsStructuredOutputAndReportsCumulativeUsage(t *test if got := atomic.LoadInt32(&fake.calls); got != 2 { t.Fatalf("provider calls = %d, want initial generation and one repair", got) } + requests := fake.requestsSnapshot() + if len(requests) != 2 || len(requests[0].Prompt.Messages) < 3 { + t.Fatalf("repair requests = %#v, want correction messages on the initial prepared request", requests) + } + messages := requests[0].Prompt.Messages + if got, want := messages[len(messages)-2], (promptkit.RenderedMessage{Role: promptkit.RoleAssistant, Content: `{"bad":true}`}); got != want { + t.Fatalf("repair assistant correction = %#v, want %#v", got, want) + } + if got, want := messages[len(messages)-1], (promptkit.RenderedMessage{Role: promptkit.RoleUser, Content: "Return the required ok field."}); got != want { + t.Fatalf("repair user correction = %#v, want %#v", got, want) + } if response.RepairAttempts != 1 || response.PromptTokens != 10 || response.CompletionTokens != 16 || response.TotalTokens != 26 { t.Fatalf("response repair and usage = %#v, want one repair and PromptKit cumulative usage", response) } diff --git a/internal/framework/pipeline/debug.go b/internal/framework/pipeline/debug.go index 15c607dd..7a336266 100644 --- a/internal/framework/pipeline/debug.go +++ b/internal/framework/pipeline/debug.go @@ -110,24 +110,7 @@ type debugSerializedOutput struct { Content debugBinaryEnvelope `json:"content"` } -type debugLLMInputMaterial struct { - Name string `json:"name"` - MediaType string `json:"media_type,omitempty"` - Content string `json:"content_base64,omitempty"` - Digest string `json:"digest,omitempty"` - OriginURI string `json:"origin_uri,omitempty"` - SizeBytes int64 `json:"size_bytes,omitempty"` -} - -type debugStructuredCompletionRequest 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 map[string]debugLLMInputMaterial `json:"inputs,omitempty"` - Vars map[string]any `json:"vars,omitempty"` -} +type debugStructuredCompletionRequest = contracts.DebugStructuredCompletionRequest type debugStructuredCompletionResponse struct { Content string `json:"content,omitempty"` @@ -567,29 +550,7 @@ func debugOutputFiles(files []contracts.OutputFile) []debugOutputFile { } func debugCompletionRequest(req contracts.StructuredCompletionRequest) debugStructuredCompletionRequest { - inputs := make(map[string]debugLLMInputMaterial, len(req.Inputs)) - for key, material := range req.Inputs { - inputs[key] = debugLLMInputMaterial{ - Name: material.Name, - MediaType: material.MediaType, - Content: base64.StdEncoding.EncodeToString(redactSecretBytes(material.Content)), - Digest: material.Digest, - OriginURI: material.OriginURI, - SizeBytes: material.SizeBytes, - } - } - if len(inputs) == 0 { - inputs = nil - } - return debugStructuredCompletionRequest{ - StageName: req.StageName, - PromptID: req.PromptID, - PromptVersion: req.PromptVersion, - ProfileID: req.ProfileID, - SessionID: req.SessionID, - Inputs: inputs, - Vars: redactSensitiveMap(req.Vars), - } + return req.DebugSummary() } func debugCompletionResponse(response contracts.StructuredCompletionResponse) debugStructuredCompletionResponse { diff --git a/internal/framework/pipeline/debug_test.go b/internal/framework/pipeline/debug_test.go index f570118f..a46ee5a4 100644 --- a/internal/framework/pipeline/debug_test.go +++ b/internal/framework/pipeline/debug_test.go @@ -47,6 +47,32 @@ func TestDebugLLMPathsKeepDotIdentitiesDistinct(t *testing.T) { } } +func TestDebugCompletionRequestOmitsCorrectionContent(t *testing.T) { + const assistantResponse = `{"secret":"assistant response"}` + const userGuidance = "secret user guidance" + correction, err := contracts.NewSemanticCorrection([]byte(assistantResponse), userGuidance) + if err != nil { + t.Fatalf("NewSemanticCorrection() error = %v", err) + } + summary := debugCompletionRequest(contracts.StructuredCompletionRequest{ + Inputs: contracts.LLMInputSet{"source": contracts.NewLLMInputMaterial("source", "application/json", []byte(`{"source":true}`), "", "")}, + Vars: map[string]any{"custom": "value"}, + Correction: correction, + }) + encoded, err := json.Marshal(summary) + if err != nil { + t.Fatalf("marshal completion summary: %v", err) + } + for _, secret := range []string{assistantResponse, userGuidance} { + if strings.Contains(string(encoded), secret) { + t.Fatalf("debug completion summary leaked %q: %s", secret, encoded) + } + } + if summary.InputCount != 1 || summary.VariableCount != 1 || summary.Correction == nil { + t.Fatalf("debug completion summary = %#v, want counts and correction metadata", summary) + } +} + func TestDebugSourceDocumentPreservesUnitReferences(t *testing.T) { doc := validSourceDocument() envelope := debugSourceDocumentEnvelope(doc)