Files
promptkit/internal/prompt/renderer_test.go

548 lines
17 KiB
Go

package prompt
import (
"bytes"
"context"
"errors"
"strings"
"testing"
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
)
func TestGoRenderer_Render(t *testing.T) {
renderer := NewGoRenderer()
ctx := context.Background()
inputs := map[string]*domain.Artifact{
"transcript": {Body: []byte("The quick brown fox.")},
}
vars := map[string]string{
"role": "helpful assistant",
"tone": "concise",
}
t.Run("rendering inline message content", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "Analyze this: {{input \"transcript\"}}"},
},
}
res, err := renderer.Render(ctx, def, inputs, vars)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(res.Messages) != 1 {
t.Fatalf("expected 1 message, got %d", len(res.Messages))
}
if res.Messages[0].Content != "Analyze this: The quick brown fox." {
t.Fatalf("unexpected rendered content: %q", res.Messages[0].Content)
}
})
t.Run("rendering file-backed message content loaded into prompt definition", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "From file: {{input \"transcript\"}}", ContentFile: "/tmp/user.tmpl"},
},
}
res, err := renderer.Render(ctx, def, inputs, vars)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if got := res.Messages[0].Content; got != "From file: The quick brown fox." {
t.Fatalf("unexpected file-backed render result: %q", got)
}
})
t.Run("rendering system and user messages", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "system", Content: "You are a {{.role}}."},
{Role: "user", Content: "Analyze this: {{input \"transcript\"}}"},
},
}
res, err := renderer.Render(ctx, def, inputs, vars)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(res.Messages) != 2 {
t.Fatalf("expected 2 messages, got %d", len(res.Messages))
}
if res.Messages[0].Role != "system" || res.Messages[1].Role != "user" {
t.Fatalf("unexpected roles: %#v", res.Messages)
}
})
t.Run("copying cache control to rendered messages", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{
Role: "system",
Content: "You are concise.",
CacheControl: &domain.CacheControl{
Type: domain.CacheControlEphemeral,
TTL: "1h",
},
},
{Role: "user", Content: "Analyze this: {{input \"transcript\"}}"},
},
}
res, err := renderer.Render(ctx, def, inputs, vars)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(res.Messages) != 2 {
t.Fatalf("expected 2 messages, got %d", len(res.Messages))
}
if res.Messages[0].CacheControl == nil {
t.Fatal("expected rendered cache control")
}
if res.Messages[0].CacheControl.Type != domain.CacheControlEphemeral {
t.Fatalf("unexpected cache control type: %q", res.Messages[0].CacheControl.Type)
}
if res.Messages[0].CacheControl.TTL != "1h" {
t.Fatalf("unexpected cache control ttl: %q", res.Messages[0].CacheControl.TTL)
}
if res.Messages[1].CacheControl != nil {
t.Fatalf("expected no cache control on second message, got %#v", res.Messages[1].CacheControl)
}
})
t.Run("rendered cache control does not alias source template", func(t *testing.T) {
source := &domain.CacheControl{Type: domain.CacheControlEphemeral, TTL: "1h"}
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "system", Content: "You are concise.", CacheControl: source},
},
}
res, err := renderer.Render(ctx, def, inputs, vars)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if res.Messages[0].CacheControl == source {
t.Fatal("expected rendered cache control to be cloned")
}
res.Messages[0].CacheControl.TTL = ""
if source.TTL != "1h" {
t.Fatalf("source cache control was mutated, ttl=%q", source.TTL)
}
})
t.Run("accessing vars", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "system", Content: "Speak in a {{.tone}} tone."},
},
}
res, err := renderer.Render(ctx, def, inputs, vars)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if res.Messages[0].Content != "Speak in a concise tone." {
t.Fatalf("unexpected vars rendering: %q", res.Messages[0].Content)
}
})
t.Run("rendering session id from vars", func(t *testing.T) {
def := &domain.PromptDefinition{
SessionID: " {{ .session_id }} ",
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "system", Content: "Speak in a {{.tone}} tone."},
},
}
res, err := renderer.Render(ctx, def, inputs, map[string]string{
"tone": "concise",
"session_id": "agent-session-123",
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if res.SessionID != "agent-session-123" {
t.Fatalf("unexpected session id: %q", res.SessionID)
}
})
t.Run("empty rendered session id is omitted", func(t *testing.T) {
def := &domain.PromptDefinition{
SessionID: " ",
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "system", Content: "Speak in a {{.tone}} tone."},
},
}
res, err := renderer.Render(ctx, def, inputs, vars)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if res.SessionID != "" {
t.Fatalf("expected empty session id, got %q", res.SessionID)
}
})
t.Run("missing session id var fails rendering", func(t *testing.T) {
def := &domain.PromptDefinition{
SessionID: "{{ .session_id }}",
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "system", Content: "Speak in a {{.tone}} tone."},
},
}
_, err := renderer.Render(ctx, def, inputs, vars)
if !errors.Is(err, ErrRenderFailure) {
t.Fatalf("expected ErrRenderFailure, got %v", err)
}
})
t.Run("too long rendered session id fails rendering", func(t *testing.T) {
def := &domain.PromptDefinition{
SessionID: "{{ .session_id }}",
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "system", Content: "Speak in a {{.tone}} tone."},
},
}
_, err := renderer.Render(ctx, def, inputs, map[string]string{
"tone": "concise",
"session_id": strings.Repeat("x", domain.SessionIDMaxLength+1),
})
if !errors.Is(err, ErrRenderFailure) {
t.Fatalf("expected ErrRenderFailure, got %v", err)
}
})
t.Run("malformed rendered session id fails rendering", func(t *testing.T) {
def := &domain.PromptDefinition{
SessionID: "{{ .session_id }}",
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "system", Content: "Speak in a {{.tone}} tone."},
},
}
_, err := renderer.Render(ctx, def, inputs, map[string]string{
"tone": "concise",
"session_id": "session" + string([]byte{0xff}),
})
if !errors.Is(err, ErrRenderFailure) {
t.Fatalf("expected ErrRenderFailure, got %v", err)
}
})
t.Run("inserting required input artifact", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "{{input \"transcript\"}}"},
},
}
res, err := renderer.Render(ctx, def, inputs, vars)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if res.Messages[0].Content != "The quick brown fox." {
t.Fatalf("unexpected required input rendering: %q", res.Messages[0].Content)
}
})
t.Run("optional input absent and not referenced", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{
{Name: "transcript", Required: true},
{Name: "glossary", Required: false},
},
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "Transcript: {{input \"transcript\"}}"},
},
}
res, err := renderer.Render(ctx, def, inputs, vars)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(res.Messages) != 1 {
t.Fatalf("expected one rendered message, got %d", len(res.Messages))
}
})
t.Run("optional input absent but referenced, expecting failure", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{
{Name: "transcript", Required: true},
{Name: "glossary", Required: false},
},
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "Glossary: {{input \"glossary\"}}"},
},
}
_, err := renderer.Render(ctx, def, inputs, vars)
if !errors.Is(err, ErrRenderFailure) {
t.Fatalf("expected ErrRenderFailure, got %v", err)
}
if !errors.Is(err, ErrUnknownInput) {
t.Fatalf("expected ErrUnknownInput, got %v", err)
}
})
t.Run("required input missing, expecting failure", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "Analyze this: {{input \"transcript\"}}"},
},
}
_, err := renderer.Render(ctx, def, map[string]*domain.Artifact{}, vars)
if !errors.Is(err, ErrMissingRequiredInput) {
t.Fatalf("expected ErrMissingRequiredInput, got %v", err)
}
})
t.Run("invalid template syntax", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "Hello {{.unclosed"},
},
}
_, err := renderer.Render(ctx, def, inputs, vars)
if !errors.Is(err, ErrInvalidTemplate) {
t.Fatalf("expected ErrInvalidTemplate, got %v", err)
}
})
t.Run("unknown input reference", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "Hello {{input \"ghost\"}}"},
},
}
_, err := renderer.Render(ctx, def, inputs, vars)
if !errors.Is(err, ErrRenderFailure) {
t.Fatalf("expected ErrRenderFailure, got %v", err)
}
if !errors.Is(err, ErrUnknownInput) {
t.Fatalf("expected ErrUnknownInput, got %v", err)
}
})
t.Run("empty message role", func(t *testing.T) {
def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
Templates: []domain.PromptMessageTemplate{
{Role: "", Content: "Hello"},
},
}
_, err := renderer.Render(ctx, def, inputs, vars)
if !errors.Is(err, ErrInvalidMessageRole) {
t.Fatalf("expected ErrInvalidMessageRole, got %v", err)
}
})
}
func TestGoRendererCancellation(t *testing.T) {
t.Run("before session parsing", func(t *testing.T) {
definition := &domain.PromptDefinition{
SessionID: "{{ malformed",
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "not rendered"},
},
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
result, err := NewGoRenderer().Render(ctx, definition, nil, nil)
if result != nil || !errors.Is(err, context.Canceled) {
t.Fatalf("result=%#v err=%v, want nil/context.Canceled", result, err)
}
if errors.Is(err, ErrInvalidTemplate) {
t.Fatalf("pre-canceled render parsed the malformed session: %v", err)
}
result, err = NewGoRenderer().Render(context.Background(), definition, nil, nil)
if result != nil || !errors.Is(err, ErrInvalidTemplate) {
t.Fatalf("active render result=%#v err=%v, want nil/ErrInvalidTemplate", result, err)
}
})
t.Run("during artifact text conversion", func(t *testing.T) {
ctx := newCancelOnCheckContext(3)
body := bytes.Repeat([]byte("x"), artifactTextChunkSize*2)
original := append([]byte(nil), body...)
resolver := newArtifactTextResolver(ctx, map[string]*domain.Artifact{
"document": {Body: body},
})
text, err := resolver.resolve("document")
if text != "" || !errors.Is(err, context.Canceled) {
t.Fatalf("text length=%d err=%v, want empty/context.Canceled", len(text), err)
}
if _, published := resolver.textByName["document"]; published {
t.Fatal("canceled conversion published partial artifact text")
}
if !bytes.Equal(body, original) {
t.Fatal("resolver mutated the artifact body")
}
})
t.Run("after final message execution", func(t *testing.T) {
definition := &domain.PromptDefinition{
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "fully rendered"},
},
}
counter := &checkCountingContext{Context: context.Background()}
if _, err := NewGoRenderer().Render(counter, definition, nil, nil); err != nil {
t.Fatalf("count render checkpoints: %v", err)
}
// The final three checks occur after template execution, after the
// message is assembled, and immediately before publication.
ctx := newCancelOnCheckContext(counter.checks - 2)
result, err := NewGoRenderer().Render(ctx, definition, nil, nil)
if result != nil || !errors.Is(err, context.Canceled) {
t.Fatalf("result=%#v err=%v, want nil/context.Canceled", result, err)
}
})
}
func TestGoRendererArtifactTextLifecycle(t *testing.T) {
body := []byte{'a', 0xff, 'b', 0xfe}
original := append([]byte(nil), body...)
artifact := &domain.Artifact{Body: body}
inputs := map[string]*domain.Artifact{"document": artifact}
definition := &domain.PromptDefinition{
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "{{input \"document\"}}|{{input \"document\"}}"},
},
}
first, err := NewGoRenderer().Render(context.Background(), definition, inputs, nil)
if err != nil {
t.Fatalf("first render: %v", err)
}
wantFirst := append(append(append([]byte(nil), body...), '|'), body...)
if !bytes.Equal([]byte(first.Messages[0].Content), wantFirst) {
t.Fatalf("rendered bytes=%v, want %v", []byte(first.Messages[0].Content), wantFirst)
}
if !bytes.Equal(body, original) {
t.Fatalf("renderer mutated artifact body: got %v want %v", body, original)
}
body[0] = 'z'
if bytes.Equal([]byte(first.Messages[0].Content), append(append(append([]byte(nil), body...), '|'), body...)) {
t.Fatal("completed render aliases the artifact body")
}
second, err := NewGoRenderer().Render(context.Background(), definition, inputs, nil)
if err != nil {
t.Fatalf("second render: %v", err)
}
wantSecond := append(append(append([]byte(nil), body...), '|'), body...)
if !bytes.Equal([]byte(second.Messages[0].Content), wantSecond) {
t.Fatalf("second render reused text from another call: got %v want %v", []byte(second.Messages[0].Content), wantSecond)
}
nilInputs := map[string]*domain.Artifact{"document": nil}
result, err := NewGoRenderer().Render(context.Background(), definition, nilInputs, nil)
if result != nil || !errors.Is(err, ErrUnknownInput) || !errors.Is(err, ErrRenderFailure) {
t.Fatalf("nil input result=%#v err=%v, want ErrUnknownInput and ErrRenderFailure", result, err)
}
}
func BenchmarkGoRendererArtifactReferences(b *testing.B) {
body := bytes.Repeat([]byte("document content "), (artifactTextChunkSize*4)/len("document content "))
inputs := map[string]*domain.Artifact{
"document": {Body: body},
}
tests := []struct {
name string
definition *domain.PromptDefinition
}{
{
name: "one reference",
definition: &domain.PromptDefinition{
Templates: []domain.PromptMessageTemplate{
{Role: "user", Content: "{{input \"document\"}}"},
},
},
},
{
name: "repeated across session and messages",
definition: &domain.PromptDefinition{
SessionID: "document-{{len (input \"document\")}}",
Templates: []domain.PromptMessageTemplate{
{Role: "system", Content: "{{input \"document\"}}"},
{Role: "user", Content: "{{input \"document\"}} {{input \"document\"}}"},
},
},
},
}
for _, tc := range tests {
b.Run(tc.name, func(b *testing.B) {
renderer := NewGoRenderer()
b.ReportAllocs()
b.SetBytes(int64(len(body)))
for range b.N {
if _, err := renderer.Render(context.Background(), tc.definition, inputs, nil); err != nil {
b.Fatal(err)
}
}
})
}
}
type checkCountingContext struct {
context.Context
checks int
}
func (c *checkCountingContext) Err() error {
c.checks++
return c.Context.Err()
}
type cancelOnCheckContext struct {
context.Context
cancel context.CancelFunc
remaining int
}
func newCancelOnCheckContext(checks int) *cancelOnCheckContext {
ctx, cancel := context.WithCancel(context.Background())
return &cancelOnCheckContext{Context: ctx, cancel: cancel, remaining: checks}
}
func (c *cancelOnCheckContext) Err() error {
if c.Context.Err() == nil {
c.remaining--
if c.remaining == 0 {
c.cancel()
}
}
return c.Context.Err()
}