From ed0f7527d55182cca6698dbe6f7ee3b7478818ae Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Wed, 26 Aug 2026 02:10:01 +0000 Subject: [PATCH] Add appended message request validation --- appended_messages_contract_test.go | 65 +++++++++++++++++ convert.go | 63 ++++++++++++++--- convert_appended_messages_test.go | 108 +++++++++++++++++++++++++++++ docs/roadmap/implementation.md | 2 + engine_test.go | 11 ++- formatting.go | 3 +- internal/domain/domain.go | 19 ++--- message_roles.go | 14 ++++ types.go | 16 +++-- 9 files changed, 277 insertions(+), 24 deletions(-) create mode 100644 appended_messages_contract_test.go create mode 100644 convert_appended_messages_test.go create mode 100644 message_roles.go diff --git a/appended_messages_contract_test.go b/appended_messages_contract_test.go new file mode 100644 index 0000000..b6b37d7 --- /dev/null +++ b/appended_messages_contract_test.go @@ -0,0 +1,65 @@ +package promptkit_test + +import ( + "context" + "errors" + "sync/atomic" + "testing" + + "gitea.maximumdirect.net/eric/promptkit" +) + +func TestAppendedMessageRoleConstantsAreStrings(t *testing.T) { + var ( + developer string = promptkit.RoleDeveloper + system string = promptkit.RoleSystem + user string = promptkit.RoleUser + assistant string = promptkit.RoleAssistant + ) + if developer != "developer" || system != "system" || user != "user" || assistant != "assistant" { + t.Fatalf("unexpected role constants: %q %q %q %q", developer, system, user, assistant) + } +} + +func TestInvalidAppendedMessagesFailBeforeSourceOrModelWork(t *testing.T) { + promptSource := &inspectionCountingFS{} + var modelCalls atomic.Int64 + engine, err := promptkit.NewEngine( + promptkit.Config{}, + promptkit.WithPromptFS(promptSource, "."), + promptkit.WithLLMClient(countingLLMClient{calls: &modelCalls}), + ) + if err != nil { + t.Fatalf("construct engine: %v", err) + } + request := promptkit.RunRequest{ + PromptID: "unreached", + AppendedMessages: []promptkit.RenderedMessage{{ + Role: "unsupported-role", + Content: "sensitive appended content", + }}, + } + + operations := []struct { + name string + run func() error + }{ + {name: "Prepare", run: func() error { _, err := engine.Prepare(context.Background(), request); return err }}, + {name: "PrepareExecution", run: func() error { _, err := engine.PrepareExecution(context.Background(), request); return err }}, + {name: "Run", run: func() error { _, err := engine.Run(context.Background(), request); return err }}, + } + for _, operation := range operations { + t.Run(operation.name, func(t *testing.T) { + err := operation.run() + if !errors.Is(err, promptkit.ErrInvalidRequest) { + t.Fatalf("error = %v, want ErrInvalidRequest", err) + } + }) + } + if promptSource.opens.Load() != 0 { + t.Fatalf("invalid appended message opened prompt sources %d times", promptSource.opens.Load()) + } + if modelCalls.Load() != 0 { + t.Fatalf("invalid appended message invoked the model %d times", modelCalls.Load()) + } +} diff --git a/convert.go b/convert.go index 0b75d9d..6122f9b 100644 --- a/convert.go +++ b/convert.go @@ -1,7 +1,9 @@ package promptkit import ( + "fmt" "reflect" + "unicode/utf8" "gitea.maximumdirect.net/eric/promptkit/internal/domain" "gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue" @@ -12,19 +14,52 @@ func toDomainRunRequest(req RunRequest) (domain.RunRequest, error) { if err != nil { return domain.RunRequest{}, err } + appendedMessages, err := toDomainAppendedMessages(req.AppendedMessages) + if err != nil { + return domain.RunRequest{}, err + } return domain.RunRequest{ - PromptID: req.PromptID, - PromptVersion: req.PromptVersion, - ProfileID: req.ProfileID, - SessionID: req.SessionID, - APIKey: req.APIKey, - Inputs: toDomainArtifactRefMap(req.Inputs), - Vars: copyStringMap(req.Vars), - Execution: execution, - Validation: toDomainOutputContractPtr(req.Validation), + PromptID: req.PromptID, + PromptVersion: req.PromptVersion, + ProfileID: req.ProfileID, + SessionID: req.SessionID, + APIKey: req.APIKey, + Inputs: toDomainArtifactRefMap(req.Inputs), + Vars: copyStringMap(req.Vars), + Execution: execution, + Validation: toDomainOutputContractPtr(req.Validation), + AppendedMessages: appendedMessages, }, nil } +func toDomainAppendedMessages(messages []RenderedMessage) ([]domain.RenderedMessage, error) { + if messages == nil { + return nil, nil + } + + converted := make([]domain.RenderedMessage, len(messages)) + for index, message := range messages { + if !utf8.ValidString(message.Content) { + return nil, fmt.Errorf("appended message %d content must be valid UTF-8", index) + } + role, err := domain.NormalizeMessageRole(message.Role) + if err != nil { + return nil, fmt.Errorf("appended message %d role: %w", index, err) + } + + cacheControl, err := domain.NormalizeCacheControl(toDomainCacheControl(message.CacheControl)) + if err != nil { + return nil, fmt.Errorf("appended message %d cache_control: %w", index, err) + } + converted[index] = domain.RenderedMessage{ + Role: role, + Content: message.Content, + CacheControl: cacheControl, + } + } + return converted, nil +} + func fromDomainPreparedRun(prepared *domain.PreparedRun) *PreparedRun { if prepared == nil { return nil @@ -285,6 +320,16 @@ func fromDomainRenderedMessages(messages []domain.RenderedMessage) []RenderedMes return out } +func toDomainCacheControl(cacheControl *CacheControl) *domain.CacheControl { + if cacheControl == nil { + return nil + } + return &domain.CacheControl{ + Type: domain.CacheControlType(cacheControl.Type), + TTL: cacheControl.TTL, + } +} + func fromDomainCacheControl(cacheControl *domain.CacheControl) *CacheControl { if cacheControl == nil { return nil diff --git a/convert_appended_messages_test.go b/convert_appended_messages_test.go new file mode 100644 index 0000000..35fee03 --- /dev/null +++ b/convert_appended_messages_test.go @@ -0,0 +1,108 @@ +package promptkit + +import ( + "reflect" + "strings" + "testing" + + "gitea.maximumdirect.net/eric/promptkit/internal/domain" +) + +func TestToDomainAppendedMessages(t *testing.T) { + tests := []struct { + name string + messages []RenderedMessage + want []domain.RenderedMessage + wantErr string + privateRole string + privateText string + }{ + { + name: "normalizes supported roles", + messages: []RenderedMessage{ + {Role: " DeVeLoPeR ", Content: "developer"}, + {Role: "SyStEm", Content: "system"}, + {Role: "\tUsEr\n", Content: "user"}, + {Role: "assistant", Content: "assistant"}, + }, + want: []domain.RenderedMessage{ + {Role: domain.RoleDeveloper, Content: "developer"}, + {Role: domain.RoleSystem, Content: "system"}, + {Role: domain.RoleUser, Content: "user"}, + {Role: domain.RoleAssistant, Content: "assistant"}, + }, + }, + { + name: "invalid role UTF-8", + messages: []RenderedMessage{{Role: string([]byte{0xff}), Content: "private-content"}}, + wantErr: "appended message 0 role", + }, + { + name: "invalid content UTF-8", + messages: []RenderedMessage{{Role: "private-role", Content: string([]byte{0xff})}}, + wantErr: "appended message 0 content", + }, + { + name: "unsupported role", + messages: []RenderedMessage{{Role: "private-role", Content: "private-content"}}, + wantErr: "appended message 0 role", + privateRole: "private-role", + privateText: "private-content", + }, + { + name: "invalid cache control", + messages: []RenderedMessage{{Role: RoleUser, Content: "private-content", CacheControl: &CacheControl{Type: "private-cache"}}}, + wantErr: "appended message 0 cache_control", + privateText: "private-content", + }, + { + name: "empty and whitespace content", + messages: []RenderedMessage{ + {Role: RoleUser, Content: ""}, + {Role: RoleAssistant, Content: " \t\n "}, + }, + want: []domain.RenderedMessage{ + {Role: domain.RoleUser, Content: ""}, + {Role: domain.RoleAssistant, Content: " \t\n "}, + }, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + got, err := toDomainAppendedMessages(test.messages) + if test.wantErr != "" { + if err == nil || !strings.Contains(err.Error(), test.wantErr) { + t.Fatalf("error = %v, want %q", err, test.wantErr) + } + for _, privateValue := range []string{test.privateRole, test.privateText} { + if privateValue != "" && strings.Contains(err.Error(), privateValue) { + t.Fatalf("error exposed appended message data %q: %v", privateValue, err) + } + } + return + } + if err != nil { + t.Fatalf("toDomainAppendedMessages() error = %v", err) + } + if !reflect.DeepEqual(got, test.want) { + t.Fatalf("toDomainAppendedMessages() = %#v, want %#v", got, test.want) + } + }) + } +} + +func TestToDomainAppendedMessagesCopiesCacheControl(t *testing.T) { + cacheControl := &CacheControl{Type: CacheControlEphemeral, TTL: "1h"} + messages := []RenderedMessage{{Role: RoleUser, Content: "content", CacheControl: cacheControl}} + request, err := toDomainRunRequest(RunRequest{AppendedMessages: messages}) + if err != nil { + t.Fatalf("toDomainRunRequest() error = %v", err) + } + + messages[0].Content = "changed" + cacheControl.TTL = "" + if request.AppendedMessages[0].Content != "content" || request.AppendedMessages[0].CacheControl.TTL != "1h" { + t.Fatalf("domain request did not retain an independent appended-message copy: %#v", request.AppendedMessages) + } +} diff --git a/docs/roadmap/implementation.md b/docs/roadmap/implementation.md index b0288fa..948a514 100644 --- a/docs/roadmap/implementation.md +++ b/docs/roadmap/implementation.md @@ -153,6 +153,8 @@ field. ## Stage 2: Add The Public Request Contract And Boundary Validation +**Status:** Complete + ### Objective Expose the smallest public API for appended messages and convert it into a diff --git a/engine_test.go b/engine_test.go index a60ef34..edd109c 100644 --- a/engine_test.go +++ b/engine_test.go @@ -220,6 +220,14 @@ func TestRunRequestFormattingRedactsDirectAPIKey(t *testing.T) { "transcript": promptkit.InlineWithURI(inputURI, inputBody), }, Vars: map[string]string{"audience": variableValue}, + AppendedMessages: []promptkit.RenderedMessage{{ + Role: "private-role-sentinel", + Content: "private-appended-content-sentinel", + CacheControl: &promptkit.CacheControl{ + Type: "private-cache-sentinel", + TTL: "private-ttl-sentinel", + }, + }}, } for _, formatted := range []string{ @@ -229,7 +237,7 @@ func TestRunRequestFormattingRedactsDirectAPIKey(t *testing.T) { fmt.Sprintf("%+v", req), fmt.Sprintf("%#v", req), } { - for _, privateValue := range []string{secret, inputURI, inputBody, variableValue} { + for _, privateValue := range []string{secret, inputURI, inputBody, variableValue, "private-role-sentinel", "private-appended-content-sentinel", "private-cache-sentinel", "private-ttl-sentinel"} { if strings.Contains(formatted, privateValue) { t.Fatalf("formatted RunRequest leaked private value %q: %s", privateValue, formatted) } @@ -240,6 +248,7 @@ func TestRunRequestFormattingRedactsDirectAPIKey(t *testing.T) { "APIKeySet:true", "Inputs:1", "Vars:1", + "AppendedMessages:1", } { if !strings.Contains(formatted, summary) { t.Fatalf("formatted RunRequest omitted structural summary %q: %s", summary, formatted) diff --git a/formatting.go b/formatting.go index 1985836..e9958ab 100644 --- a/formatting.go +++ b/formatting.go @@ -18,7 +18,7 @@ func (r RunRequest) GoString() string { func (r RunRequest) redactedString() string { return fmt.Sprintf( - "promptkit.RunRequest{PromptID:%q PromptVersion:%q ProfileID:%q APIKeySet:%t Inputs:%d Vars:%d ExecutionSet:%t ValidationSet:%t}", + "promptkit.RunRequest{PromptID:%q PromptVersion:%q ProfileID:%q APIKeySet:%t Inputs:%d Vars:%d ExecutionSet:%t ValidationSet:%t AppendedMessages:%d}", r.PromptID, r.PromptVersion, r.ProfileID, @@ -27,6 +27,7 @@ func (r RunRequest) redactedString() string { len(r.Vars), r.Execution != nil, r.Validation != nil, + len(r.AppendedMessages), ) } diff --git a/internal/domain/domain.go b/internal/domain/domain.go index 7c28eac..a6a0410 100644 --- a/internal/domain/domain.go +++ b/internal/domain/domain.go @@ -60,15 +60,16 @@ type CacheControl struct { // RunRequest represents a request to generate a single artifact. type RunRequest struct { - PromptID string - PromptVersion string - ProfileID string - SessionID string - APIKey string `json:"-" yaml:"-"` - Inputs map[string]ArtifactRef - Vars map[string]string - Execution *ExecutionTargetOverride - Validation *OutputContract + PromptID string + PromptVersion string + ProfileID string + SessionID string + APIKey string `json:"-" yaml:"-"` + Inputs map[string]ArtifactRef + Vars map[string]string + Execution *ExecutionTargetOverride + Validation *OutputContract + AppendedMessages []RenderedMessage } // RunResult represents the complete result of a prompt execution run. diff --git a/message_roles.go b/message_roles.go new file mode 100644 index 0000000..cb691d3 --- /dev/null +++ b/message_roles.go @@ -0,0 +1,14 @@ +package promptkit + +import "gitea.maximumdirect.net/eric/promptkit/internal/domain" + +const ( + // RoleDeveloper identifies a developer instruction message. + RoleDeveloper = domain.RoleDeveloper + // RoleSystem identifies a system instruction message. + RoleSystem = domain.RoleSystem + // RoleUser identifies a user message. + RoleUser = domain.RoleUser + // RoleAssistant identifies an assistant message. + RoleAssistant = domain.RoleAssistant +) diff --git a/types.go b/types.go index 732bc06..15a31bf 100644 --- a/types.go +++ b/types.go @@ -123,6 +123,13 @@ type RunRequest struct { // Validation optionally replaces the prompt's complete output contract. It // does not merge individual fields. Nil uses the prompt contract. Validation *OutputContract + // AppendedMessages are already-rendered messages appended after every prompt + // definition message. Promptkit neither templates nor resolves files in + // them, and preserves valid content exactly. Nil and empty slices are + // equivalent. Prepare, PrepareExecution, and Run validate and copy the + // messages before source or model work; malformed values return an error + // matching ErrInvalidRequest. + AppendedMessages []RenderedMessage } // PreparedRun contains prepared prompt execution state returned by @@ -629,12 +636,13 @@ type RenderedPrompt struct { Messages []RenderedMessage `json:"messages"` } -// RenderedMessage is a rendered chat message and has a stable JSON -// representation. +// RenderedMessage is a prepared, provider-bound text chat message and has a +// stable JSON representation. Its role must be one of [RoleDeveloper], +// [RoleSystem], [RoleUser], or [RoleAssistant]. type RenderedMessage struct { - // Role is the definition-supplied chat role. + // Role is the provider-bound chat role. Role string `json:"role"` - // Content is the rendered message text. + // Content is the provider-bound message text and may be empty or whitespace. Content string `json:"content"` // CacheControl is optional provider cache metadata. CacheControl *CacheControl `json:"cache_control,omitempty"`