Compose appended request messages

This commit is contained in:
2026-08-26 02:16:16 +00:00
parent ed0f7527d5
commit b0999112cc
10 changed files with 291 additions and 57 deletions

View File

@@ -59,20 +59,24 @@ func NormalizeCacheControl(control *CacheControl) (*CacheControl, error) {
// CloneRenderedMessages returns a deep copy of rendered messages.
func CloneRenderedMessages(messages []RenderedMessage) []RenderedMessage {
cloned := make([]RenderedMessage, len(messages))
for index, message := range messages {
cloned[index] = message
if message.CacheControl != nil {
cacheControl := *message.CacheControl
cloned[index].CacheControl = &cacheControl
}
}
copyRenderedMessages(cloned, messages)
return cloned
}
// ConcatRenderedMessages returns an independently owned concatenation of messages.
func ConcatRenderedMessages(prefix, suffix []RenderedMessage) []RenderedMessage {
messages := make([]RenderedMessage, 0, len(prefix)+len(suffix))
messages = append(messages, CloneRenderedMessages(prefix)...)
messages = append(messages, CloneRenderedMessages(suffix)...)
messages := make([]RenderedMessage, len(prefix)+len(suffix))
copyRenderedMessages(messages, prefix)
copyRenderedMessages(messages[len(prefix):], suffix)
return messages
}
func copyRenderedMessages(destination, source []RenderedMessage) {
for index, message := range source {
destination[index] = message
if message.CacheControl != nil {
cacheControl := *message.CacheControl
destination[index].CacheControl = &cacheControl
}
}
}

View File

@@ -363,6 +363,7 @@ func TestOpenAICompatibleClientRequestMapping(t *testing.T) {
run func(*testing.T)
}{
{name: "complete request and response mapping", run: checkCompleteRequestAndResponseMapping},
{name: "canonical role and exact content", run: checkCanonicalRoleAndExactContentMapping},
{name: "cache-controlled message", run: checkCacheControlledMessageMapping},
{name: "empty cache-control TTL", run: checkEmptyCacheControlTTLOmission},
{name: "session ID", run: checkSessionIDMapping},
@@ -377,6 +378,34 @@ func TestOpenAICompatibleClientRequestMapping(t *testing.T) {
}
}
func checkCanonicalRoleAndExactContentMapping(t *testing.T) {
provider := newRecordingProvider(t)
client := newProviderClient(t, provider, OpenAICompatibleConfig{})
const content = " exact content\nwith whitespace "
_, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{
Role: domain.RoleDeveloper,
Content: content,
}}},
Target: domain.ExecutionTarget{Model: "model"},
})
if err != nil {
t.Fatalf("Generate() error = %v", err)
}
var messages []struct {
Role string `json:"role"`
Content string `json:"content"`
}
request := provider.lastRequest(t)
if !request.decodeField(t, "messages", &messages) || len(messages) != 1 {
t.Fatalf("messages = %#v, want one message", messages)
}
if messages[0].Role != domain.RoleDeveloper || messages[0].Content != content {
t.Fatalf("message = %#v, want role %q and exact content %q", messages[0], domain.RoleDeveloper, content)
}
}
func checkCompleteRequestAndResponseMapping(t *testing.T) {
provider := newRecordingProvider(t)
provider.respond(http.StatusOK, `{

View File

@@ -213,16 +213,7 @@ func clonePreparedRun(source *domain.PreparedRun) (*domain.PreparedRun, error) {
copied.InputHashes[name] = hash
}
}
if source.Messages != nil {
copied.Messages = make([]domain.RenderedMessage, len(source.Messages))
for i, message := range source.Messages {
copied.Messages[i] = message
if message.CacheControl != nil {
cacheControl := *message.CacheControl
copied.Messages[i].CacheControl = &cacheControl
}
}
}
copied.Messages = domain.CloneRenderedMessages(source.Messages)
if source.StructuredOutput != nil {
structuredOutput := *source.StructuredOutput

View File

@@ -55,8 +55,7 @@ func (r *defaultOutputRepairer) Repair(ctx context.Context, req RepairRequest) (
return nil, err
}
messages := make([]domain.RenderedMessage, len(req.OriginalMessages), len(req.OriginalMessages)+2)
copy(messages, req.OriginalMessages)
messages := domain.ConcatRenderedMessages(req.OriginalMessages, nil)
if strings.TrimSpace(req.PreviousOutput) != "" {
messages = append(messages, domain.RenderedMessage{
Role: domain.RoleAssistant,

View File

@@ -30,6 +30,8 @@ func TestDefaultOutputRepairerBuildsFullContextRequest(t *testing.T) {
original := []domain.RenderedMessage{
{Role: "system", Content: "Follow the task.", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral, TTL: "1h"}},
{Role: "user", Content: "Summarize the report."},
{Role: "assistant", Content: "Consumer-supplied previous response."},
{Role: "user", Content: "Consumer-supplied correction."},
}
before := append([]domain.RenderedMessage(nil), original...)
previous := strings.Repeat("candidate ", 12_000)

View File

@@ -4,10 +4,13 @@ import (
"context"
"crypto/rand"
"crypto/sha256"
"encoding/binary"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"hash"
"io"
"os"
"strings"
"time"
@@ -444,9 +447,11 @@ func (r *Runner) completePreparationWithStructuredOutput(
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrPromptRender, err)
}
effectivePrompt := *renderedPrompt
if state.directSessionID != "" {
renderedPrompt.SessionID = state.directSessionID
effectivePrompt.SessionID = state.directSessionID
}
effectivePrompt.Messages = domain.ConcatRenderedMessages(renderedPrompt.Messages, req.AppendedMessages)
end := time.Now().UTC()
effectiveModel := state.effectiveModel
@@ -462,9 +467,9 @@ func (r *Runner) completePreparationWithStructuredOutput(
OutputContract: state.effectiveContract,
StructuredOutput: structuredOutput,
InputHashes: inputHashes,
SessionID: renderedPrompt.SessionID,
RenderedPromptHash: hashRenderedPrompt(*renderedPrompt),
Messages: renderedPrompt.Messages,
SessionID: effectivePrompt.SessionID,
RenderedPromptHash: hashRenderedPrompt(effectivePrompt),
Messages: effectivePrompt.Messages,
StartTime: state.start,
EndTime: end,
DurationMS: end.Sub(state.start).Milliseconds(),
@@ -714,29 +719,38 @@ func resolveOutputContract(def *domain.PromptDefinition, override *domain.Output
return contract, nil
}
const renderedPromptHashVersion = "promptkit/rendered-prompt/v2\x00"
func hashRenderedPrompt(p domain.RenderedPrompt) string {
var b strings.Builder
if p.SessionID != "" {
b.WriteString("session_id=")
b.WriteString(p.SessionID)
b.WriteString("\n---\n")
}
for _, msg := range p.Messages {
b.WriteString(msg.Role)
b.WriteByte('\n')
b.WriteString(msg.Content)
if msg.CacheControl != nil {
b.WriteString("\ncache_control.type=")
b.WriteString(string(msg.CacheControl.Type))
if msg.CacheControl.TTL != "" {
b.WriteString("\ncache_control.ttl=")
b.WriteString(msg.CacheControl.TTL)
}
hasher := sha256.New()
var lengthBuffer [8]byte
_, _ = io.WriteString(hasher, renderedPromptHashVersion)
writeRenderedPromptHashString(hasher, &lengthBuffer, p.SessionID)
writeRenderedPromptHashLength(hasher, &lengthBuffer, len(p.Messages))
for _, message := range p.Messages {
writeRenderedPromptHashString(hasher, &lengthBuffer, message.Role)
writeRenderedPromptHashString(hasher, &lengthBuffer, message.Content)
if message.CacheControl == nil {
lengthBuffer[0] = 0
_, _ = hasher.Write(lengthBuffer[:1])
continue
}
b.WriteString("\n---\n")
lengthBuffer[0] = 1
_, _ = hasher.Write(lengthBuffer[:1])
writeRenderedPromptHashString(hasher, &lengthBuffer, string(message.CacheControl.Type))
writeRenderedPromptHashString(hasher, &lengthBuffer, message.CacheControl.TTL)
}
h := sha256.Sum256([]byte(b.String()))
return hex.EncodeToString(h[:])
return hex.EncodeToString(hasher.Sum(nil))
}
func writeRenderedPromptHashLength(hasher hash.Hash, buffer *[8]byte, length int) {
binary.BigEndian.PutUint64(buffer[:], uint64(length))
_, _ = hasher.Write(buffer[:])
}
func writeRenderedPromptHashString(hasher hash.Hash, buffer *[8]byte, value string) {
writeRenderedPromptHashLength(hasher, buffer, len(value))
_, _ = io.WriteString(hasher, value)
}
func buildOutputArtifact(content string, format domain.OutputFormat) domain.Artifact {

View File

@@ -1110,14 +1110,18 @@ func TestDeriveStructuredSchemaName(t *testing.T) {
}
}
func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
uncached := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
{Role: "system", Content: "sys"},
{Role: "user", Content: "usr"},
func TestHashRenderedPromptIncludesEveryContractedField(t *testing.T) {
base := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
{Role: domain.RoleSystem, Content: "sys"},
{Role: domain.RoleUser, Content: "usr"},
}}
wantLegacyHash := hashString("system\nsys\n---\nuser\nusr\n---\n")
if got := hashRenderedPrompt(uncached); got != wantLegacyHash {
t.Fatalf("expected no-cache hash to preserve legacy input, got %q want %q", got, wantLegacyHash)
identical := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
{Role: domain.RoleSystem, Content: "sys"},
{Role: domain.RoleUser, Content: "usr"},
}}
baseHash := hashRenderedPrompt(base)
if baseHash != hashRenderedPrompt(identical) {
t.Fatal("identical rendered prompts must have the same hash")
}
withCache := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
@@ -1129,7 +1133,7 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
TTL: "1h",
},
},
{Role: "user", Content: "usr"},
{Role: domain.RoleUser, Content: "usr"},
}}
alsoWithCache := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
{
@@ -1140,7 +1144,7 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
TTL: "1h",
},
},
{Role: "user", Content: "usr"},
{Role: domain.RoleUser, Content: "usr"},
}}
withoutTTL := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
{
@@ -1150,11 +1154,11 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
Type: domain.CacheControlEphemeral,
},
},
{Role: "user", Content: "usr"},
{Role: domain.RoleUser, Content: "usr"},
}}
cachedHash := hashRenderedPrompt(withCache)
if cachedHash == hashRenderedPrompt(uncached) {
if cachedHash == baseHash {
t.Fatal("expected cache control to change rendered prompt hash")
}
if cachedHash != hashRenderedPrompt(alsoWithCache) {
@@ -1163,6 +1167,29 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
if cachedHash == hashRenderedPrompt(withoutTTL) {
t.Fatal("expected ttl changes to affect rendered prompt hash")
}
variants := []struct {
name string
prompt domain.RenderedPrompt
}{
{name: "role", prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: domain.RoleDeveloper, Content: "sys"}, {Role: domain.RoleUser, Content: "usr"}}}},
{name: "content", prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: domain.RoleSystem, Content: "changed"}, {Role: domain.RoleUser, Content: "usr"}}}},
{name: "order", prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: domain.RoleUser, Content: "usr"}, {Role: domain.RoleSystem, Content: "sys"}}}},
{name: "session", prompt: domain.RenderedPrompt{SessionID: "session", Messages: base.Messages}},
}
for _, variant := range variants {
t.Run(variant.name, func(t *testing.T) {
if hashRenderedPrompt(variant.prompt) == baseHash {
t.Fatalf("%s did not change the rendered-prompt hash", variant.name)
}
})
}
oneMessage := domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: domain.RoleUser, Content: "x\n---\nassistant\ny"}}}
twoMessages := domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: domain.RoleUser, Content: "x"}, {Role: domain.RoleAssistant, Content: "y"}}}
if hashRenderedPrompt(oneMessage) == hashRenderedPrompt(twoMessages) {
t.Fatal("length-framed hashing did not distinguish legacy separator collision")
}
}
func TestHashRenderedPromptIncludesSessionIDWhenPresent(t *testing.T) {
@@ -1204,6 +1231,75 @@ func TestHashRenderedPromptIncludesSessionIDWhenPresent(t *testing.T) {
}
}
func TestRunnerComposesAppendedMessagesIntoEveryRenderedPrompt(t *testing.T) {
renderer := &fakeRenderer{rendered: &domain.RenderedPrompt{
SessionID: "rendered-session",
Messages: []domain.RenderedMessage{
{Role: domain.RoleSystem, Content: "ordinary prefix", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral, TTL: "1h"}},
{Role: domain.RoleUser, Content: "ordinary request"},
},
}}
client := &fakeLLM{resp: &domain.GenerateResponse{Content: "output"}}
runner := NewRunner(
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)},
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
nil,
defaultArtifactReader(),
renderer,
client,
nil,
nil,
)
appended := []domain.RenderedMessage{
{Role: domain.RoleAssistant, Content: "consumer response"},
{Role: domain.RoleUser, Content: "consumer correction", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral}},
}
request := domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef(), AppendedMessages: appended}
plain, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
if err != nil {
t.Fatalf("prepare without appended messages: %v", err)
}
empty, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef(), AppendedMessages: []domain.RenderedMessage{}})
if err != nil {
t.Fatalf("prepare with empty appended messages: %v", err)
}
if !reflect.DeepEqual(empty.Messages, plain.Messages) || empty.RenderedPromptHash != plain.RenderedPromptHash {
t.Fatalf("empty appended messages changed preparation: empty=%#v plain=%#v", empty, plain)
}
prepared, err := runner.Prepare(context.Background(), request)
if err != nil {
t.Fatalf("prepare with appended messages: %v", err)
}
if prepared.PromptHash != plain.PromptHash || prepared.RenderedPromptHash == plain.RenderedPromptHash {
t.Fatalf("hashes with appended messages = (%q, %q), plain = (%q, %q)", prepared.PromptHash, prepared.RenderedPromptHash, plain.PromptHash, plain.RenderedPromptHash)
}
expected := []domain.RenderedMessage{
{Role: domain.RoleSystem, Content: "ordinary prefix", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral, TTL: "1h"}},
{Role: domain.RoleUser, Content: "ordinary request"},
{Role: domain.RoleAssistant, Content: "consumer response"},
{Role: domain.RoleUser, Content: "consumer correction", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral}},
}
if !reflect.DeepEqual(prepared.Messages, expected) {
t.Fatalf("prepared messages = %#v, want %#v", prepared.Messages, expected)
}
result, err := runner.Run(context.Background(), request)
if err != nil {
t.Fatalf("run with appended messages: %v", err)
}
if result.RenderedPromptHash != prepared.RenderedPromptHash || !reflect.DeepEqual(client.lastReq.Prompt.Messages, expected) {
t.Fatalf("run did not use the prepared effective prompt: result=%#v request=%#v prepared=%#v", result, client.lastReq.Prompt.Messages, prepared)
}
renderer.rendered.Messages[0].Content = "changed source"
appended[1].CacheControl.Type = "changed caller value"
if !reflect.DeepEqual(prepared.Messages, expected) {
t.Fatalf("prepared messages changed after source or caller mutation: %#v", prepared.Messages)
}
}
func TestRunnerRunSuccessful(t *testing.T) {
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatMarkdown, domain.ValidationBasic, 0)}
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}