548 lines
17 KiB
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()
|
|
}
|