Honor cancellation while rendering prompts
This commit is contained in:
@@ -101,8 +101,17 @@ implemented reader behavior and failures.
|
|||||||
## Rendering
|
## Rendering
|
||||||
|
|
||||||
`internal/prompt` renders definition messages as Go templates using named
|
`internal/prompt` renders definition messages as Go templates using named
|
||||||
artifacts and variables. It carries message roles, session IDs, and cache
|
artifacts and variables. Within one render, each referenced artifact body is
|
||||||
control into the rendered prompt. The
|
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.
|
[renderer tests](../../internal/prompt/renderer_test.go) own rendering behavior.
|
||||||
|
|
||||||
## Schemas And Output Validation
|
## Schemas And Output Validation
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"strings"
|
||||||
"text/template"
|
"text/template"
|
||||||
|
|
||||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||||
@@ -18,6 +19,8 @@ var (
|
|||||||
ErrInvalidMessageRole = errors.New("invalid or empty message role")
|
ErrInvalidMessageRole = errors.New("invalid or empty message role")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const artifactTextChunkSize = 64 * 1024
|
||||||
|
|
||||||
type goRenderer struct{}
|
type goRenderer struct{}
|
||||||
|
|
||||||
func NewGoRenderer() Renderer {
|
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) {
|
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 {
|
if definition == nil {
|
||||||
return nil, fmt.Errorf("%w: nil prompt definition", ErrRenderFailure)
|
return nil, fmt.Errorf("%w: nil prompt definition", ErrRenderFailure)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 1. Verify required inputs
|
|
||||||
for _, in := range definition.Inputs {
|
for _, in := range definition.Inputs {
|
||||||
if !in.Required {
|
if !in.Required {
|
||||||
continue
|
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{
|
funcs := template.FuncMap{
|
||||||
"input": func(name string) (string, error) {
|
"input": resolver.resolve,
|
||||||
art, ok := inputs[name]
|
|
||||||
if !ok || art == nil {
|
|
||||||
return "", fmt.Errorf("%w: %s", ErrUnknownInput, name)
|
|
||||||
}
|
|
||||||
return string(art.Body), nil
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
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 {
|
if err != nil {
|
||||||
return nil, err
|
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 {
|
for i, tmplMsg := range definition.Templates {
|
||||||
select {
|
if err := ctx.Err(); err != nil {
|
||||||
case <-ctx.Done():
|
return nil, err
|
||||||
return nil, ctx.Err()
|
|
||||||
default:
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if tmplMsg.Role == "" {
|
if tmplMsg.Role == "" {
|
||||||
return nil, fmt.Errorf("%w: message %d", ErrInvalidMessageRole, i)
|
return nil, fmt.Errorf("%w: message %d", ErrInvalidMessageRole, i)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse and execute template
|
if err := ctx.Err(); err != nil {
|
||||||
tmpl, err := template.New(fmt.Sprintf("msg_%d", i)).Funcs(funcs).Option("missingkey=error").Parse(tmplMsg.Content)
|
return nil, err
|
||||||
if err != nil {
|
}
|
||||||
return nil, fmt.Errorf("%w: message %d: %v", ErrInvalidTemplate, i, 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
|
var buf bytes.Buffer
|
||||||
if err := tmpl.Execute(&buf, vars); err != nil {
|
executeErr := tmpl.Execute(&buf, vars)
|
||||||
return nil, fmt.Errorf("%w: message %d: %w", ErrRenderFailure, i, err)
|
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{
|
renderedMessages = append(renderedMessages, domain.RenderedMessage{
|
||||||
@@ -85,6 +100,13 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini
|
|||||||
Content: buf.String(),
|
Content: buf.String(),
|
||||||
CacheControl: cloneCacheControl(tmplMsg.CacheControl),
|
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{
|
return &domain.RenderedPrompt{
|
||||||
@@ -93,15 +115,75 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func renderSessionID(raw string, funcs template.FuncMap, vars map[string]string) (string, error) {
|
type artifactTextResolver struct {
|
||||||
tmpl, err := template.New("session_id").Funcs(funcs).Option("missingkey=error").Parse(raw)
|
ctx context.Context
|
||||||
if err != nil {
|
inputs map[string]*domain.Artifact
|
||||||
return "", fmt.Errorf("%w: session_id: %v", ErrInvalidTemplate, err)
|
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
|
var buf bytes.Buffer
|
||||||
if err := tmpl.Execute(&buf, vars); err != nil {
|
executeErr := tmpl.Execute(&buf, vars)
|
||||||
return "", fmt.Errorf("%w: session_id: %w", ErrRenderFailure, err)
|
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())
|
sessionID, err := domain.NormalizeSessionID(buf.String())
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package prompt
|
package prompt
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"strings"
|
"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