Honor cancellation while rendering prompts
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
package prompt
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
@@ -361,3 +362,186 @@ func TestGoRenderer_Render(t *testing.T) {
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
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()
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user