From e83a3ce179a4ec20c5d6ed08ac53e2293040a4e8 Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Tue, 11 Aug 2026 22:48:20 +0000 Subject: [PATCH] Honor cancellation while rendering prompts --- docs/internal/sources.md | 13 ++- internal/prompt/go_renderer.go | 136 ++++++++++++++++++----- internal/prompt/renderer_test.go | 184 +++++++++++++++++++++++++++++++ 3 files changed, 304 insertions(+), 29 deletions(-) diff --git a/docs/internal/sources.md b/docs/internal/sources.md index 7f7a47b..3989b68 100644 --- a/docs/internal/sources.md +++ b/docs/internal/sources.md @@ -101,8 +101,17 @@ implemented reader behavior and failures. ## Rendering `internal/prompt` renders definition messages as Go templates using named -artifacts and variables. It carries message roles, session IDs, and cache -control into the rendered prompt. The +artifacts and variables. Within one render, each referenced artifact body is +converted to text lazily and cached by input name for reuse across the session +and every message; the cache is not shared across renders. Conversion uses +bounded chunks and preserves the artifact bytes exactly. + +Session and message parsing and execution remain synchronous. The renderer +checks cancellation before and after each parse and execution boundary, +between artifact conversion chunks, around each message, and before publishing +the complete prompt. It cannot interrupt template work already in progress and +never publishes a partial prompt after observing cancellation. It carries +message roles, session IDs, and cache control into the rendered prompt. The [renderer tests](../../internal/prompt/renderer_test.go) own rendering behavior. ## Schemas And Output Validation diff --git a/internal/prompt/go_renderer.go b/internal/prompt/go_renderer.go index 99345de..be67c76 100644 --- a/internal/prompt/go_renderer.go +++ b/internal/prompt/go_renderer.go @@ -5,6 +5,7 @@ import ( "context" "errors" "fmt" + "strings" "text/template" "gitea.maximumdirect.net/eric/promptkit/internal/domain" @@ -18,6 +19,8 @@ var ( ErrInvalidMessageRole = errors.New("invalid or empty message role") ) +const artifactTextChunkSize = 64 * 1024 + type goRenderer struct{} func NewGoRenderer() Renderer { @@ -25,11 +28,13 @@ func NewGoRenderer() Renderer { } func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefinition, inputs map[string]*domain.Artifact, vars map[string]string) (*domain.RenderedPrompt, error) { + if err := ctx.Err(); err != nil { + return nil, err + } if definition == nil { return nil, fmt.Errorf("%w: nil prompt definition", ErrRenderFailure) } - // 1. Verify required inputs for _, in := range definition.Inputs { if !in.Required { continue @@ -40,44 +45,54 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini } } - // 2. Setup template functions + resolver := newArtifactTextResolver(ctx, inputs) funcs := template.FuncMap{ - "input": func(name string) (string, error) { - art, ok := inputs[name] - if !ok || art == nil { - return "", fmt.Errorf("%w: %s", ErrUnknownInput, name) - } - return string(art.Body), nil - }, + "input": resolver.resolve, } - sessionID, err := renderSessionID(definition.SessionID, funcs, vars) + if err := ctx.Err(); err != nil { + return nil, err + } + sessionID, err := renderSessionID(ctx, definition.SessionID, funcs, vars) if err != nil { return nil, err } + if err := ctx.Err(); err != nil { + return nil, err + } - var renderedMessages []domain.RenderedMessage + renderedMessages := make([]domain.RenderedMessage, 0, len(definition.Templates)) for i, tmplMsg := range definition.Templates { - select { - case <-ctx.Done(): - return nil, ctx.Err() - default: + if err := ctx.Err(); err != nil { + return nil, err } if tmplMsg.Role == "" { return nil, fmt.Errorf("%w: message %d", ErrInvalidMessageRole, i) } - // Parse and execute template - tmpl, err := template.New(fmt.Sprintf("msg_%d", i)).Funcs(funcs).Option("missingkey=error").Parse(tmplMsg.Content) - if err != nil { - return nil, fmt.Errorf("%w: message %d: %v", ErrInvalidTemplate, i, err) + if err := ctx.Err(); err != nil { + return nil, err + } + tmpl, parseErr := template.New(fmt.Sprintf("msg_%d", i)).Funcs(funcs).Option("missingkey=error").Parse(tmplMsg.Content) + if err := ctx.Err(); err != nil { + return nil, err + } + if parseErr != nil { + return nil, fmt.Errorf("%w: message %d: %v", ErrInvalidTemplate, i, parseErr) } + if err := ctx.Err(); err != nil { + return nil, err + } var buf bytes.Buffer - if err := tmpl.Execute(&buf, vars); err != nil { - return nil, fmt.Errorf("%w: message %d: %w", ErrRenderFailure, i, err) + executeErr := tmpl.Execute(&buf, vars) + if err := ctx.Err(); err != nil { + return nil, err + } + if executeErr != nil { + return nil, fmt.Errorf("%w: message %d: %w", ErrRenderFailure, i, executeErr) } renderedMessages = append(renderedMessages, domain.RenderedMessage{ @@ -85,6 +100,13 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini Content: buf.String(), CacheControl: cloneCacheControl(tmplMsg.CacheControl), }) + if err := ctx.Err(); err != nil { + return nil, err + } + } + + if err := ctx.Err(); err != nil { + return nil, err } return &domain.RenderedPrompt{ @@ -93,15 +115,75 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini }, nil } -func renderSessionID(raw string, funcs template.FuncMap, vars map[string]string) (string, error) { - tmpl, err := template.New("session_id").Funcs(funcs).Option("missingkey=error").Parse(raw) - if err != nil { - return "", fmt.Errorf("%w: session_id: %v", ErrInvalidTemplate, err) +type artifactTextResolver struct { + ctx context.Context + inputs map[string]*domain.Artifact + textByName map[string]string +} + +func newArtifactTextResolver(ctx context.Context, inputs map[string]*domain.Artifact) *artifactTextResolver { + return &artifactTextResolver{ + ctx: ctx, + inputs: inputs, + textByName: make(map[string]string), + } +} + +func (r *artifactTextResolver) resolve(name string) (string, error) { + if err := r.ctx.Err(); err != nil { + return "", err + } + artifact, ok := r.inputs[name] + if !ok || artifact == nil { + return "", fmt.Errorf("%w: %s", ErrUnknownInput, name) + } + if text, ok := r.textByName[name]; ok { + return text, nil } + var builder strings.Builder + builder.Grow(len(artifact.Body)) + for start := 0; start < len(artifact.Body); start += artifactTextChunkSize { + if err := r.ctx.Err(); err != nil { + return "", err + } + end := min(start+artifactTextChunkSize, len(artifact.Body)) + _, _ = builder.Write(artifact.Body[start:end]) + if err := r.ctx.Err(); err != nil { + return "", err + } + } + if err := r.ctx.Err(); err != nil { + return "", err + } + + text := builder.String() + r.textByName[name] = text + return text, nil +} + +func renderSessionID(ctx context.Context, raw string, funcs template.FuncMap, vars map[string]string) (string, error) { + if err := ctx.Err(); err != nil { + return "", err + } + tmpl, parseErr := template.New("session_id").Funcs(funcs).Option("missingkey=error").Parse(raw) + if err := ctx.Err(); err != nil { + return "", err + } + if parseErr != nil { + return "", fmt.Errorf("%w: session_id: %v", ErrInvalidTemplate, parseErr) + } + + if err := ctx.Err(); err != nil { + return "", err + } var buf bytes.Buffer - if err := tmpl.Execute(&buf, vars); err != nil { - return "", fmt.Errorf("%w: session_id: %w", ErrRenderFailure, err) + executeErr := tmpl.Execute(&buf, vars) + if err := ctx.Err(); err != nil { + return "", err + } + if executeErr != nil { + return "", fmt.Errorf("%w: session_id: %w", ErrRenderFailure, executeErr) } sessionID, err := domain.NormalizeSessionID(buf.String()) diff --git a/internal/prompt/renderer_test.go b/internal/prompt/renderer_test.go index d8e53b0..43a0eff 100644 --- a/internal/prompt/renderer_test.go +++ b/internal/prompt/renderer_test.go @@ -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() +}