Honor cancellation while rendering prompts

This commit is contained in:
2026-08-11 22:48:20 +00:00
parent a04a3bbc5f
commit e83a3ce179
3 changed files with 304 additions and 29 deletions

View File

@@ -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()
}