From e4e824ed412b733a9b6211a7844796a449e9f89f Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Fri, 31 Jul 2026 03:58:18 +0000 Subject: [PATCH] Define prompt execution contract --- internal/promptexec/copy.go | 80 ++++++++ internal/promptexec/promptexec.go | 267 +++++++++++++++++++++++++ internal/promptexec/promptexec_test.go | 262 ++++++++++++++++++++++++ 3 files changed, 609 insertions(+) create mode 100644 internal/promptexec/copy.go create mode 100644 internal/promptexec/promptexec.go create mode 100644 internal/promptexec/promptexec_test.go diff --git a/internal/promptexec/copy.go b/internal/promptexec/copy.go new file mode 100644 index 0000000..8a05611 --- /dev/null +++ b/internal/promptexec/copy.go @@ -0,0 +1,80 @@ +package promptexec + +import ( + "strings" + "unicode/utf8" +) + +func copyPreparation(value Preparation) Preparation { + value.InputHashes = copyStringMap(value.InputHashes) + return value +} + +func copyPreparationDebug(value *PreparationDebug) *PreparationDebug { + if value == nil { + return nil + } + copy := *value + copy.RenderedMessages = append([]RenderedMessage(nil), value.RenderedMessages...) + copy.StructuredSchema = append([]byte(nil), value.StructuredSchema...) + copy.ParametersJSON = append([]byte(nil), value.ParametersJSON...) + return © +} + +func copyExecution(value Execution) Execution { + value.InputHashes = copyStringMap(value.InputHashes) + value.Validation.Diagnostics = boundDiagnostics(value.Validation.Diagnostics) + value.RawOutput = append([]byte(nil), value.RawOutput...) + value.Debug = copyExecutionDebug(value.Debug) + return value +} + +func copyExecutionDebug(value *ExecutionDebug) *ExecutionDebug { + if value == nil { + return nil + } + copy := *value + copy.RawOutput = append([]byte(nil), value.RawOutput...) + copy.ValidationDiagnostics = boundDiagnostics(value.ValidationDiagnostics) + return © +} + +func copyStringMap(value map[string]string) map[string]string { + if value == nil { + return nil + } + copy := make(map[string]string, len(value)) + for key, item := range value { + copy[key] = item + } + return copy +} + +func boundDiagnostics(values []string) []string { + if len(values) == 0 { + return nil + } + if len(values) > maxValidationDiagnostics { + values = values[:maxValidationDiagnostics] + } + bounded := make([]string, len(values)) + for index, value := range values { + bounded[index] = boundText(value, maxDiagnosticBytes) + } + return bounded +} + +func boundText(value string, limit int) string { + if limit <= 0 || value == "" { + return "" + } + value = strings.ToValidUTF8(value, "�") + if len(value) <= limit { + return value + } + value = value[:limit] + for len(value) > 0 && !utf8.ValidString(value) { + value = value[:len(value)-1] + } + return value +} diff --git a/internal/promptexec/promptexec.go b/internal/promptexec/promptexec.go new file mode 100644 index 0000000..71b78ac --- /dev/null +++ b/internal/promptexec/promptexec.go @@ -0,0 +1,267 @@ +// Package promptexec defines Weatherreporter's provider-neutral prompt execution contract. +package promptexec + +import ( + "context" + "errors" + "fmt" + "time" +) + +const ( + maxValidationDiagnostics = 10 + maxDiagnosticBytes = 1024 + maxErrorMessageBytes = 2048 +) + +// Executor inspects and executes configured prompts without exposing provider types. +// Inspection is side-effect-free. Execute invokes prepared exactly once after a +// successful preparation and before provider execution. If prepared returns an +// error, Execute must not call the provider. Completed validation rejection is +// returned as an Execution with a failed Validation status; operational failures +// return no Execution. Sensitive debug values are populated only when requested. +type Executor interface { + InspectPrompt(context.Context, string, string) (PromptInspection, error) + InspectProfile(context.Context, string) (ProfileInspection, error) + Execute(context.Context, ExecuteRequest, PreparationCallback) (*Execution, error) +} + +// PromptInspection describes one exact prompt definition without selecting a profile. +type PromptInspection struct { + PromptID string + PromptVersion string + PromptHash string + DefaultProfileID string + Inputs []InputDefinition + Output OutputContract +} + +// InputDefinition describes one declared prompt input. +type InputDefinition struct { + Name string + Required bool + ContentType string + Description string +} + +// OutputContract summarizes the output requirements declared by a prompt. +type OutputContract struct { + Format string + ValidationMode string + SchemaPath string +} + +// ProfileInspection describes the safe, selected execution identity for one profile. +type ProfileInspection struct { + ProfileID string + BackendID string + ModelName string + CredentialRequired bool +} + +// ExecuteRequest selects one exact prompt execution. DataPackage is the exact +// YAML input; implementations must copy it before retaining it. DataPackagePath +// is provenance for the inline input, not a provider-readable file reference. +type ExecuteRequest struct { + PromptID string + PromptVersion string + ProfileID string + DataPackage []byte + DataPackagePath string + CaptureDebug bool +} + +// PreparationCallback receives safe preparation provenance before provider work. +// The callback receives independent copies which it may retain or mutate. +type PreparationCallback func(Preparation, *PreparationDebug) error + +// Preparation contains non-sensitive provenance from a completed preparation. +type Preparation struct { + PromptID string + PromptVersion string + PromptHash string + RenderedPromptHash string + InputHashes map[string]string + ProfileID string + BackendID string + ModelName string + Output OutputContract + StartedAt time.Time + EndedAt time.Time + Duration time.Duration + DataPackagePath string +} + +// PreparationDebug contains content-rich preparation details for an explicitly +// enabled sensitive-debug destination. It must never be persisted routinely. +type PreparationDebug struct { + RenderedMessages []RenderedMessage + StructuredSchema []byte + Endpoint string + ParametersJSON []byte +} + +// RenderedMessage is one fully rendered model message for sensitive debugging. +type RenderedMessage struct { + Role string + Content string +} + +// Execution contains the completed result of one provider run. RawOutput is +// the generated content, not a provider transport response body. It is copied +// before return and must be persisted separately from routine metadata. +type Execution struct { + RunID string + PromptID string + PromptVersion string + PromptHash string + RenderedPromptHash string + InputHashes map[string]string + ProfileID string + BackendID string + ModelName string + GeneratedHash string + Usage TokenUsage + StartedAt time.Time + EndedAt time.Time + Duration time.Duration + Validation Validation + DataPackagePath string + RawOutput []byte + Debug *ExecutionDebug +} + +// TokenUsage is provider-reported token accounting. +type TokenUsage struct { + PromptTokens int + CompletionTokens int + TotalTokens int + CachedTokens int + CacheWriteTokens int +} + +// Validation records a completed output validation check. +type Validation struct { + Status ValidationStatus + Mode string + SchemaPath string + Diagnostics []string +} + +// ValidationStatus identifies the completed validation state. +type ValidationStatus string + +const ( + ValidationPassed ValidationStatus = "passed" + ValidationFailed ValidationStatus = "failed" + ValidationSkipped ValidationStatus = "skipped" +) + +// ExecutionDebug contains content-rich execution details for explicitly enabled +// sensitive debugging. It must never be persisted routinely. +type ExecutionDebug struct { + RawOutput []byte + ValidationDiagnostics []string +} + +// ErrorCategory classifies a project-owned operational failure. +type ErrorCategory string + +const ( + InvalidConfiguration ErrorCategory = "invalid_configuration" + InvalidRequest ErrorCategory = "invalid_request" + PromptNotFound ErrorCategory = "prompt_not_found" + PromptLoad ErrorCategory = "prompt_load" + ProfileNotFound ErrorCategory = "profile_not_found" + ProfileLoad ErrorCategory = "profile_load" + MissingCredential ErrorCategory = "missing_credential" + ArtifactLoad ErrorCategory = "artifact_load" + PromptRender ErrorCategory = "prompt_render" + Capacity ErrorCategory = "capacity" + Generation ErrorCategory = "generation" + OperationalValidation ErrorCategory = "operational_validation" + ValidationRejected ErrorCategory = "validation_rejected" + Canceled ErrorCategory = "canceled" + DeadlineExceeded ErrorCategory = "deadline_exceeded" +) + +// Error is a bounded safe error suitable for workflow and persistence records. +// Its optional cause remains available to errors.Is and errors.As but is never +// included in Error's text. +type Error struct { + category ErrorCategory + message string + cause error +} + +// NewError returns a classified error with a bounded, Weatherreporter-owned message. +func NewError(category ErrorCategory, message string, cause error) *Error { + messageLimit := maxErrorMessageBytes - len(category) - len(": ") + return &Error{category: category, message: boundText(message, messageLimit), cause: cause} +} + +func (e *Error) Error() string { + if e == nil { + return "" + } + if e.message == "" { + return string(e.category) + } + return fmt.Sprintf("%s: %s", e.category, e.message) +} + +// Unwrap preserves an underlying error identity without exposing its text. +func (e *Error) Unwrap() error { + if e == nil { + return nil + } + return e.cause +} + +// Category returns the stable classification. +func (e *Error) Category() ErrorCategory { + if e == nil { + return "" + } + return e.category +} + +// CapacityError adds the safe backend identity to a capacity failure. +type CapacityError struct { + BackendID string + Err *Error +} + +// NewCapacityError returns a classified capacity failure for backendID. +func NewCapacityError(backendID string, message string, cause error) *CapacityError { + return &CapacityError{BackendID: boundText(backendID, 256), Err: NewError(Capacity, message, cause)} +} + +func (e *CapacityError) Error() string { + if e == nil { + return "" + } + if e.BackendID == "" { + return e.Err.Error() + } + return boundText(fmt.Sprintf("%s (backend %q)", e.Err.Error(), e.BackendID), maxErrorMessageBytes) +} + +func (e *CapacityError) Unwrap() error { + if e == nil { + return nil + } + return e.Err +} + +// Category returns Capacity for every capacity error. +func (e *CapacityError) Category() ErrorCategory { return Capacity } + +// CategoryOf returns the classification carried by err, including wrapped errors. +func CategoryOf(err error) ErrorCategory { + var categorized interface{ Category() ErrorCategory } + if errors.As(err, &categorized) { + return categorized.Category() + } + return "" +} diff --git a/internal/promptexec/promptexec_test.go b/internal/promptexec/promptexec_test.go new file mode 100644 index 0000000..cc05139 --- /dev/null +++ b/internal/promptexec/promptexec_test.go @@ -0,0 +1,262 @@ +package promptexec + +import ( + "context" + "errors" + "strings" + "testing" + "unicode/utf8" +) + +type fakeExecutor struct{} + +func (fakeExecutor) InspectPrompt(context.Context, string, string) (PromptInspection, error) { + return PromptInspection{}, nil +} + +func (fakeExecutor) InspectProfile(context.Context, string) (ProfileInspection, error) { + return ProfileInspection{}, nil +} + +func (fakeExecutor) Execute(context.Context, ExecuteRequest, PreparationCallback) (*Execution, error) { + return nil, nil +} + +var _ Executor = fakeExecutor{} + +type lifecycleExecutor struct { + providerCalled bool + operationalFailure error + validationRejected bool +} + +func (executor *lifecycleExecutor) InspectPrompt(context.Context, string, string) (PromptInspection, error) { + return PromptInspection{}, nil +} + +func (executor *lifecycleExecutor) InspectProfile(context.Context, string) (ProfileInspection, error) { + return ProfileInspection{}, nil +} + +func (executor *lifecycleExecutor) Execute(_ context.Context, request ExecuteRequest, callback PreparationCallback) (*Execution, error) { + if executor.operationalFailure != nil { + return nil, executor.operationalFailure + } + preparation := Preparation{PromptID: request.PromptID, PromptVersion: request.PromptVersion, DataPackagePath: request.DataPackagePath} + var debug *PreparationDebug + if request.CaptureDebug { + debug = &PreparationDebug{RenderedMessages: []RenderedMessage{{Role: "user", Content: "sensitive rendered message"}}} + } + if callback != nil { + if err := callback(copyPreparation(preparation), copyPreparationDebug(debug)); err != nil { + return nil, err + } + } + executor.providerCalled = true + status := ValidationPassed + if executor.validationRejected { + status = ValidationFailed + } + result := Execution{PromptID: request.PromptID, PromptVersion: request.PromptVersion, DataPackagePath: request.DataPackagePath, Validation: Validation{Status: status}, RawOutput: []byte("generated content")} + if request.CaptureDebug { + result.Debug = &ExecutionDebug{RawOutput: []byte("generated content")} + } + return &result, nil +} + +func TestExecutorLifecycleFixtures(t *testing.T) { + request := ExecuteRequest{PromptID: "weather.daily_generated_text", PromptVersion: "1.0.0", DataPackagePath: "data_package.yaml"} + t.Run("callback failure prevents provider execution", func(t *testing.T) { + executor := &lifecycleExecutor{} + callbackError := errors.New("persistence failed") + result, err := executor.Execute(context.Background(), request, func(Preparation, *PreparationDebug) error { return callbackError }) + if result != nil || !errors.Is(err, callbackError) || executor.providerCalled { + t.Fatalf("result/error/provider = %#v/%v/%t", result, err, executor.providerCalled) + } + }) + t.Run("validation rejection is completed result", func(t *testing.T) { + executor := &lifecycleExecutor{validationRejected: true} + result, err := executor.Execute(context.Background(), request, nil) + if err != nil || result == nil || result.Validation.Status != ValidationFailed || !executor.providerCalled { + t.Fatalf("result/error/provider = %#v/%v/%t", result, err, executor.providerCalled) + } + }) + t.Run("operational failure has no completed result", func(t *testing.T) { + failure := NewError(Generation, "generation failed", nil) + executor := &lifecycleExecutor{operationalFailure: failure} + result, err := executor.Execute(context.Background(), request, nil) + if result != nil || !errors.Is(err, failure) || executor.providerCalled { + t.Fatalf("result/error/provider = %#v/%v/%t", result, err, executor.providerCalled) + } + }) + t.Run("debug requires explicit request", func(t *testing.T) { + executor := &lifecycleExecutor{} + var callbackDebug *PreparationDebug + result, err := executor.Execute(context.Background(), request, func(_ Preparation, debug *PreparationDebug) error { + callbackDebug = debug + return nil + }) + if err != nil || callbackDebug != nil || result.Debug != nil { + t.Fatalf("debug = %#v/%#v, error = %v", callbackDebug, result.Debug, err) + } + request.CaptureDebug = true + result, err = executor.Execute(context.Background(), request, func(_ Preparation, debug *PreparationDebug) error { + callbackDebug = debug + return nil + }) + if err != nil || callbackDebug == nil || result.Debug == nil { + t.Fatalf("debug = %#v/%#v, error = %v", callbackDebug, result.Debug, err) + } + }) +} + +func TestErrorCategoriesAndCapacityError(t *testing.T) { + categories := []ErrorCategory{ + InvalidConfiguration, + InvalidRequest, + PromptNotFound, + PromptLoad, + ProfileNotFound, + ProfileLoad, + MissingCredential, + ArtifactLoad, + PromptRender, + Capacity, + Generation, + OperationalValidation, + ValidationRejected, + Canceled, + DeadlineExceeded, + } + cause := errors.New("dependency details must not become safe error text") + for _, category := range categories { + t.Run(string(category), func(t *testing.T) { + err := NewError(category, "safe workflow failure", cause) + if err.Category() != category || CategoryOf(err) != category { + t.Fatalf("category = %q/%q, want %q", err.Category(), CategoryOf(err), category) + } + if !errors.Is(err, cause) { + t.Fatal("errors.Is() = false, want preserved cause") + } + if strings.Contains(err.Error(), cause.Error()) { + t.Fatalf("error leaks cause: %q", err) + } + }) + } + + capacity := NewCapacityError("local", "safe capacity failure", cause) + if capacity.Category() != Capacity || CategoryOf(capacity) != Capacity || capacity.BackendID != "local" { + t.Fatalf("capacity error = %#v", capacity) + } + if !errors.Is(capacity, cause) { + t.Fatal("capacity error does not preserve cause") + } +} + +func TestBoundDiagnosticAndErrorText(t *testing.T) { + longUTF8 := strings.Repeat("é", maxDiagnosticBytes) + values := make([]string, maxValidationDiagnostics+2) + for index := range values { + values[index] = longUTF8 + } + values[0] = string([]byte{'a', 0xff, 'b'}) + bounded := boundDiagnostics(values) + if len(bounded) != maxValidationDiagnostics { + t.Fatalf("diagnostics length = %d, want %d", len(bounded), maxValidationDiagnostics) + } + for index, value := range bounded { + if len(value) > maxDiagnosticBytes || !utf8.ValidString(value) { + t.Fatalf("diagnostic %d = %q, want valid UTF-8 within %d bytes", index, value, maxDiagnosticBytes) + } + } + if bounded[0] != "a�b" { + t.Fatalf("invalid UTF-8 diagnostic = %q, want replacement", bounded[0]) + } + + err := NewError(Generation, strings.Repeat("é", maxErrorMessageBytes), nil) + if len(err.Error()) > maxErrorMessageBytes || !utf8.ValidString(err.Error()) { + t.Fatalf("error = %q, want valid UTF-8 within %d bytes", err, maxErrorMessageBytes) + } + capacity := NewCapacityError(strings.Repeat("x", 300), strings.Repeat("é", maxErrorMessageBytes), nil) + if len(capacity.Error()) > maxErrorMessageBytes || !utf8.ValidString(capacity.Error()) { + t.Fatalf("capacity error = %q, want valid UTF-8 within %d bytes", capacity, maxErrorMessageBytes) + } +} + +func TestContractCopiesMutableValues(t *testing.T) { + preparation := Preparation{InputHashes: map[string]string{"data_package": "input-hash"}} + preparationDebug := &PreparationDebug{ + RenderedMessages: []RenderedMessage{{Role: "user", Content: "rendered input"}}, + StructuredSchema: []byte("schema body"), + ParametersJSON: []byte(`{"temperature":0.2}`), + } + execution := Execution{ + InputHashes: map[string]string{"data_package": "input-hash"}, + Validation: Validation{Diagnostics: []string{"validation detail"}}, + RawOutput: []byte("generated output"), + Debug: &ExecutionDebug{ + RawOutput: []byte("provider output"), + ValidationDiagnostics: []string{"detailed validation"}, + }, + } + + preparationCopy := copyPreparation(preparation) + debugCopy := copyPreparationDebug(preparationDebug) + executionCopy := copyExecution(execution) + preparation.InputHashes["data_package"] = "changed" + preparationDebug.RenderedMessages[0].Content = "changed" + preparationDebug.StructuredSchema[0] = 'x' + preparationDebug.ParametersJSON[0] = 'x' + execution.InputHashes["data_package"] = "changed" + execution.Validation.Diagnostics[0] = "changed" + execution.RawOutput[0] = 'x' + execution.Debug.RawOutput[0] = 'x' + execution.Debug.ValidationDiagnostics[0] = "changed" + + if preparationCopy.InputHashes["data_package"] != "input-hash" { + t.Fatalf("preparation copy = %#v", preparationCopy) + } + if debugCopy.RenderedMessages[0].Content != "rendered input" || string(debugCopy.StructuredSchema) != "schema body" || string(debugCopy.ParametersJSON) != `{"temperature":0.2}` { + t.Fatalf("preparation debug copy = %#v", debugCopy) + } + if executionCopy.InputHashes["data_package"] != "input-hash" || executionCopy.Validation.Diagnostics[0] != "validation detail" || string(executionCopy.RawOutput) != "generated output" || string(executionCopy.Debug.RawOutput) != "provider output" || executionCopy.Debug.ValidationDiagnostics[0] != "detailed validation" { + t.Fatalf("execution copy = %#v", executionCopy) + } +} + +func TestSafeContractValuesExcludeSensitiveFields(t *testing.T) { + preparation := Preparation{ + PromptID: "weather.daily_generated_text", + PromptVersion: "1.0.0", + PromptHash: "prompt-hash", + RenderedPromptHash: "rendered-hash", + InputHashes: map[string]string{"data_package": "input-hash"}, + ProfileID: "configured-profile", + BackendID: "local", + ModelName: "model-name", + Output: OutputContract{Format: "json", ValidationMode: "json_schema", SchemaPath: "daily.generated_text.schema.json"}, + DataPackagePath: "data-packages/daily/data_package.yaml", + } + execution := Execution{ + RunID: "run-id", + PromptID: preparation.PromptID, + PromptVersion: preparation.PromptVersion, + PromptHash: preparation.PromptHash, + RenderedPromptHash: preparation.RenderedPromptHash, + InputHashes: copyStringMap(preparation.InputHashes), + ProfileID: preparation.ProfileID, + BackendID: preparation.BackendID, + ModelName: preparation.ModelName, + GeneratedHash: "generated-hash", + RawOutput: []byte("generated content"), + } + text := preparation.PromptID + preparation.PromptVersion + preparation.PromptHash + preparation.RenderedPromptHash + preparation.ProfileID + preparation.BackendID + preparation.ModelName + preparation.Output.SchemaPath + preparation.DataPackagePath + execution.RunID + execution.GeneratedHash + for _, unwanted := range []string{"https://provider.example", "API_KEY_ENV", "rendered message", "schema body", "input body", "provider response body", "full parameters"} { + if strings.Contains(text, unwanted) { + t.Fatalf("safe values contain %q: %s", unwanted, text) + } + } + if string(execution.RawOutput) != "generated content" { + t.Fatalf("raw output = %q, want generated content", execution.RawOutput) + } +}