Add appended message request validation
This commit is contained in:
65
appended_messages_contract_test.go
Normal file
65
appended_messages_contract_test.go
Normal file
@@ -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())
|
||||||
|
}
|
||||||
|
}
|
||||||
63
convert.go
63
convert.go
@@ -1,7 +1,9 @@
|
|||||||
package promptkit
|
package promptkit
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"unicode/utf8"
|
||||||
|
|
||||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||||
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
|
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
|
||||||
@@ -12,19 +14,52 @@ func toDomainRunRequest(req RunRequest) (domain.RunRequest, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return domain.RunRequest{}, err
|
return domain.RunRequest{}, err
|
||||||
}
|
}
|
||||||
|
appendedMessages, err := toDomainAppendedMessages(req.AppendedMessages)
|
||||||
|
if err != nil {
|
||||||
|
return domain.RunRequest{}, err
|
||||||
|
}
|
||||||
return domain.RunRequest{
|
return domain.RunRequest{
|
||||||
PromptID: req.PromptID,
|
PromptID: req.PromptID,
|
||||||
PromptVersion: req.PromptVersion,
|
PromptVersion: req.PromptVersion,
|
||||||
ProfileID: req.ProfileID,
|
ProfileID: req.ProfileID,
|
||||||
SessionID: req.SessionID,
|
SessionID: req.SessionID,
|
||||||
APIKey: req.APIKey,
|
APIKey: req.APIKey,
|
||||||
Inputs: toDomainArtifactRefMap(req.Inputs),
|
Inputs: toDomainArtifactRefMap(req.Inputs),
|
||||||
Vars: copyStringMap(req.Vars),
|
Vars: copyStringMap(req.Vars),
|
||||||
Execution: execution,
|
Execution: execution,
|
||||||
Validation: toDomainOutputContractPtr(req.Validation),
|
Validation: toDomainOutputContractPtr(req.Validation),
|
||||||
|
AppendedMessages: appendedMessages,
|
||||||
}, nil
|
}, 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 {
|
func fromDomainPreparedRun(prepared *domain.PreparedRun) *PreparedRun {
|
||||||
if prepared == nil {
|
if prepared == nil {
|
||||||
return nil
|
return nil
|
||||||
@@ -285,6 +320,16 @@ func fromDomainRenderedMessages(messages []domain.RenderedMessage) []RenderedMes
|
|||||||
return out
|
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 {
|
func fromDomainCacheControl(cacheControl *domain.CacheControl) *CacheControl {
|
||||||
if cacheControl == nil {
|
if cacheControl == nil {
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
108
convert_appended_messages_test.go
Normal file
108
convert_appended_messages_test.go
Normal file
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -153,6 +153,8 @@ field.
|
|||||||
|
|
||||||
## Stage 2: Add The Public Request Contract And Boundary Validation
|
## Stage 2: Add The Public Request Contract And Boundary Validation
|
||||||
|
|
||||||
|
**Status:** Complete
|
||||||
|
|
||||||
### Objective
|
### Objective
|
||||||
|
|
||||||
Expose the smallest public API for appended messages and convert it into a
|
Expose the smallest public API for appended messages and convert it into a
|
||||||
|
|||||||
@@ -220,6 +220,14 @@ func TestRunRequestFormattingRedactsDirectAPIKey(t *testing.T) {
|
|||||||
"transcript": promptkit.InlineWithURI(inputURI, inputBody),
|
"transcript": promptkit.InlineWithURI(inputURI, inputBody),
|
||||||
},
|
},
|
||||||
Vars: map[string]string{"audience": variableValue},
|
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{
|
for _, formatted := range []string{
|
||||||
@@ -229,7 +237,7 @@ func TestRunRequestFormattingRedactsDirectAPIKey(t *testing.T) {
|
|||||||
fmt.Sprintf("%+v", req),
|
fmt.Sprintf("%+v", req),
|
||||||
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) {
|
if strings.Contains(formatted, privateValue) {
|
||||||
t.Fatalf("formatted RunRequest leaked private value %q: %s", privateValue, formatted)
|
t.Fatalf("formatted RunRequest leaked private value %q: %s", privateValue, formatted)
|
||||||
}
|
}
|
||||||
@@ -240,6 +248,7 @@ func TestRunRequestFormattingRedactsDirectAPIKey(t *testing.T) {
|
|||||||
"APIKeySet:true",
|
"APIKeySet:true",
|
||||||
"Inputs:1",
|
"Inputs:1",
|
||||||
"Vars:1",
|
"Vars:1",
|
||||||
|
"AppendedMessages:1",
|
||||||
} {
|
} {
|
||||||
if !strings.Contains(formatted, summary) {
|
if !strings.Contains(formatted, summary) {
|
||||||
t.Fatalf("formatted RunRequest omitted structural summary %q: %s", summary, formatted)
|
t.Fatalf("formatted RunRequest omitted structural summary %q: %s", summary, formatted)
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ func (r RunRequest) GoString() string {
|
|||||||
|
|
||||||
func (r RunRequest) redactedString() string {
|
func (r RunRequest) redactedString() string {
|
||||||
return fmt.Sprintf(
|
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.PromptID,
|
||||||
r.PromptVersion,
|
r.PromptVersion,
|
||||||
r.ProfileID,
|
r.ProfileID,
|
||||||
@@ -27,6 +27,7 @@ func (r RunRequest) redactedString() string {
|
|||||||
len(r.Vars),
|
len(r.Vars),
|
||||||
r.Execution != nil,
|
r.Execution != nil,
|
||||||
r.Validation != nil,
|
r.Validation != nil,
|
||||||
|
len(r.AppendedMessages),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -60,15 +60,16 @@ type CacheControl struct {
|
|||||||
|
|
||||||
// RunRequest represents a request to generate a single artifact.
|
// RunRequest represents a request to generate a single artifact.
|
||||||
type RunRequest struct {
|
type RunRequest struct {
|
||||||
PromptID string
|
PromptID string
|
||||||
PromptVersion string
|
PromptVersion string
|
||||||
ProfileID string
|
ProfileID string
|
||||||
SessionID string
|
SessionID string
|
||||||
APIKey string `json:"-" yaml:"-"`
|
APIKey string `json:"-" yaml:"-"`
|
||||||
Inputs map[string]ArtifactRef
|
Inputs map[string]ArtifactRef
|
||||||
Vars map[string]string
|
Vars map[string]string
|
||||||
Execution *ExecutionTargetOverride
|
Execution *ExecutionTargetOverride
|
||||||
Validation *OutputContract
|
Validation *OutputContract
|
||||||
|
AppendedMessages []RenderedMessage
|
||||||
}
|
}
|
||||||
|
|
||||||
// RunResult represents the complete result of a prompt execution run.
|
// RunResult represents the complete result of a prompt execution run.
|
||||||
|
|||||||
14
message_roles.go
Normal file
14
message_roles.go
Normal file
@@ -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
|
||||||
|
)
|
||||||
16
types.go
16
types.go
@@ -123,6 +123,13 @@ type RunRequest struct {
|
|||||||
// Validation optionally replaces the prompt's complete output contract. It
|
// Validation optionally replaces the prompt's complete output contract. It
|
||||||
// does not merge individual fields. Nil uses the prompt contract.
|
// does not merge individual fields. Nil uses the prompt contract.
|
||||||
Validation *OutputContract
|
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
|
// PreparedRun contains prepared prompt execution state returned by
|
||||||
@@ -629,12 +636,13 @@ type RenderedPrompt struct {
|
|||||||
Messages []RenderedMessage `json:"messages"`
|
Messages []RenderedMessage `json:"messages"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// RenderedMessage is a rendered chat message and has a stable JSON
|
// RenderedMessage is a prepared, provider-bound text chat message and has a
|
||||||
// representation.
|
// stable JSON representation. Its role must be one of [RoleDeveloper],
|
||||||
|
// [RoleSystem], [RoleUser], or [RoleAssistant].
|
||||||
type RenderedMessage struct {
|
type RenderedMessage struct {
|
||||||
// Role is the definition-supplied chat role.
|
// Role is the provider-bound chat role.
|
||||||
Role string `json:"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"`
|
Content string `json:"content"`
|
||||||
// CacheControl is optional provider cache metadata.
|
// CacheControl is optional provider cache metadata.
|
||||||
CacheControl *CacheControl `json:"cache_control,omitempty"`
|
CacheControl *CacheControl `json:"cache_control,omitempty"`
|
||||||
|
|||||||
Reference in New Issue
Block a user