Compose appended request messages
This commit is contained in:
@@ -3,6 +3,7 @@ package promptkit_test
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
|
"reflect"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -21,6 +22,85 @@ func TestAppendedMessageRoleConstantsAreStrings(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestAppendedMessagesComposeFrozenEffectivePrompt(t *testing.T) {
|
||||||
|
cacheControl := &promptkit.CacheControl{Type: promptkit.CacheControlEphemeral, TTL: "1h"}
|
||||||
|
appended := []promptkit.RenderedMessage{
|
||||||
|
{Role: " \tAsSiStAnT\n", Content: "previous response"},
|
||||||
|
{Role: promptkit.RoleUser, Content: "corrective request", CacheControl: cacheControl},
|
||||||
|
}
|
||||||
|
client := &fakeLLMClient{response: &promptkit.GenerateResponse{Content: "output"}}
|
||||||
|
engine, err := promptkit.NewEngine(
|
||||||
|
promptkit.Config{},
|
||||||
|
promptkit.WithPromptFS(contractPromptFS("prompt", "profile", "configured prompt"), "."),
|
||||||
|
promptkit.WithProfiles(promptkit.Profile{ID: "profile", Endpoint: "http://example.test/v1", Model: "model"}),
|
||||||
|
promptkit.WithLLMClient(client),
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("construct engine: %v", err)
|
||||||
|
}
|
||||||
|
request := promptkit.RunRequest{PromptID: "prompt", AppendedMessages: appended}
|
||||||
|
|
||||||
|
plain, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("prepare plain request: %v", err)
|
||||||
|
}
|
||||||
|
prepared, err := engine.Prepare(context.Background(), request)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("prepare appended request: %v", err)
|
||||||
|
}
|
||||||
|
expected := []promptkit.RenderedMessage{
|
||||||
|
{Role: promptkit.RoleUser, Content: "configured prompt"},
|
||||||
|
{Role: promptkit.RoleAssistant, Content: "previous response"},
|
||||||
|
{Role: promptkit.RoleUser, Content: "corrective request", CacheControl: &promptkit.CacheControl{Type: promptkit.CacheControlEphemeral, TTL: "1h"}},
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(prepared.Messages, expected) || prepared.PromptHash != plain.PromptHash || prepared.RenderedPromptHash == plain.RenderedPromptHash {
|
||||||
|
t.Fatalf("prepared effective prompt = %#v, plain = %#v", prepared, plain)
|
||||||
|
}
|
||||||
|
|
||||||
|
equivalent, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt", AppendedMessages: []promptkit.RenderedMessage{
|
||||||
|
{Role: promptkit.RoleAssistant, Content: "previous response"},
|
||||||
|
{Role: promptkit.RoleUser, Content: "corrective request", CacheControl: &promptkit.CacheControl{Type: promptkit.CacheControlEphemeral, TTL: "1h"}},
|
||||||
|
}})
|
||||||
|
if err != nil || equivalent.RenderedPromptHash != prepared.RenderedPromptHash {
|
||||||
|
t.Fatalf("normalized-equivalent request = (%#v, %v), want matching rendered hash %q", equivalent, err, prepared.RenderedPromptHash)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, variant := range []promptkit.RunRequest{
|
||||||
|
{PromptID: "prompt", AppendedMessages: []promptkit.RenderedMessage{{Role: promptkit.RoleAssistant, Content: "changed"}, expected[2]}},
|
||||||
|
{PromptID: "prompt", AppendedMessages: []promptkit.RenderedMessage{expected[2], expected[1]}},
|
||||||
|
{PromptID: "prompt", AppendedMessages: []promptkit.RenderedMessage{{Role: promptkit.RoleAssistant, Content: "previous response"}, {Role: promptkit.RoleUser, Content: "corrective request"}}},
|
||||||
|
} {
|
||||||
|
variantPrepared, prepareErr := engine.Prepare(context.Background(), variant)
|
||||||
|
if prepareErr != nil || variantPrepared.RenderedPromptHash == prepared.RenderedPromptHash || variantPrepared.PromptHash != prepared.PromptHash {
|
||||||
|
t.Fatalf("variant preparation = (%#v, %v)", variantPrepared, prepareErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := engine.Run(context.Background(), request)
|
||||||
|
if err != nil || result.RenderedPromptHash != prepared.RenderedPromptHash || !reflect.DeepEqual(client.requests[0].Prompt.Messages, expected) {
|
||||||
|
t.Fatalf("run result = (%#v, %v), request = %#v", result, err, client.requests)
|
||||||
|
}
|
||||||
|
|
||||||
|
execution, err := engine.PrepareExecution(context.Background(), request)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("prepare execution: %v", err)
|
||||||
|
}
|
||||||
|
appended[0].Content = "changed caller content"
|
||||||
|
cacheControl.TTL = ""
|
||||||
|
firstDetails := execution.Details()
|
||||||
|
firstDetails.Messages[1].Content = "changed details content"
|
||||||
|
firstDetails.Messages[2].CacheControl.Type = "changed details cache"
|
||||||
|
secondDetails := execution.Details()
|
||||||
|
if !reflect.DeepEqual(secondDetails.Messages, expected) || secondDetails.RenderedPromptHash != prepared.RenderedPromptHash {
|
||||||
|
t.Fatalf("prepared execution details = %#v, want frozen %#v", secondDetails, expected)
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err = engine.RunPrepared(context.Background(), execution)
|
||||||
|
if err != nil || result.RenderedPromptHash != prepared.RenderedPromptHash || !reflect.DeepEqual(client.requests[1].Prompt.Messages, expected) {
|
||||||
|
t.Fatalf("prepared result = (%#v, %v), requests = %#v", result, err, client.requests)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestInvalidAppendedMessagesFailBeforeSourceOrModelWork(t *testing.T) {
|
func TestInvalidAppendedMessagesFailBeforeSourceOrModelWork(t *testing.T) {
|
||||||
promptSource := &inspectionCountingFS{}
|
promptSource := &inspectionCountingFS{}
|
||||||
var modelCalls atomic.Int64
|
var modelCalls atomic.Int64
|
||||||
|
|||||||
@@ -238,6 +238,8 @@ The runner may still ignore the new domain field until Stage 3.
|
|||||||
|
|
||||||
## Stage 3: Compose, Freeze, Hash, Execute, And Repair The Effective Prompt
|
## Stage 3: Compose, Freeze, Hash, Execute, And Repair The Effective Prompt
|
||||||
|
|
||||||
|
**Status:** Complete
|
||||||
|
|
||||||
### Objective
|
### Objective
|
||||||
|
|
||||||
Make appended messages part of the effective rendered prompt everywhere after
|
Make appended messages part of the effective rendered prompt everywhere after
|
||||||
|
|||||||
@@ -39,6 +39,23 @@ func TestBuiltInGenerationError(t *testing.T) {
|
|||||||
assertGenerationError(t, err, http.StatusServiceUnavailable, "", "", "")
|
assertGenerationError(t, err, http.StatusServiceUnavailable, "", "", "")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestBuiltInGenerationErrorWithAppendedMessages(t *testing.T) {
|
||||||
|
const messageMarker = "combined-request-provider-marker"
|
||||||
|
engine := newBuiltInGenerationErrorEngine(t, http.StatusUnprocessableEntity,
|
||||||
|
`{"error":{"message":"`+messageMarker+`"}}`)
|
||||||
|
request := generationErrorRunRequest()
|
||||||
|
request.AppendedMessages = []promptkit.RenderedMessage{
|
||||||
|
{Role: promptkit.RoleAssistant, Content: "previous response"},
|
||||||
|
{Role: promptkit.RoleUser, Content: "consumer correction"},
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := engine.Run(context.Background(), request)
|
||||||
|
if result != nil {
|
||||||
|
t.Fatalf("Run result = %#v, want nil", result)
|
||||||
|
}
|
||||||
|
assertGenerationError(t, err, http.StatusUnprocessableEntity, "", "", messageMarker)
|
||||||
|
}
|
||||||
|
|
||||||
func TestBuiltInRepairGenerationError(t *testing.T) {
|
func TestBuiltInRepairGenerationError(t *testing.T) {
|
||||||
const (
|
const (
|
||||||
codeMarker = "repair-code-marker"
|
codeMarker = "repair-code-marker"
|
||||||
|
|||||||
@@ -59,20 +59,24 @@ func NormalizeCacheControl(control *CacheControl) (*CacheControl, error) {
|
|||||||
// CloneRenderedMessages returns a deep copy of rendered messages.
|
// CloneRenderedMessages returns a deep copy of rendered messages.
|
||||||
func CloneRenderedMessages(messages []RenderedMessage) []RenderedMessage {
|
func CloneRenderedMessages(messages []RenderedMessage) []RenderedMessage {
|
||||||
cloned := make([]RenderedMessage, len(messages))
|
cloned := make([]RenderedMessage, len(messages))
|
||||||
for index, message := range messages {
|
copyRenderedMessages(cloned, messages)
|
||||||
cloned[index] = message
|
|
||||||
if message.CacheControl != nil {
|
|
||||||
cacheControl := *message.CacheControl
|
|
||||||
cloned[index].CacheControl = &cacheControl
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return cloned
|
return cloned
|
||||||
}
|
}
|
||||||
|
|
||||||
// ConcatRenderedMessages returns an independently owned concatenation of messages.
|
// ConcatRenderedMessages returns an independently owned concatenation of messages.
|
||||||
func ConcatRenderedMessages(prefix, suffix []RenderedMessage) []RenderedMessage {
|
func ConcatRenderedMessages(prefix, suffix []RenderedMessage) []RenderedMessage {
|
||||||
messages := make([]RenderedMessage, 0, len(prefix)+len(suffix))
|
messages := make([]RenderedMessage, len(prefix)+len(suffix))
|
||||||
messages = append(messages, CloneRenderedMessages(prefix)...)
|
copyRenderedMessages(messages, prefix)
|
||||||
messages = append(messages, CloneRenderedMessages(suffix)...)
|
copyRenderedMessages(messages[len(prefix):], suffix)
|
||||||
return messages
|
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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -363,6 +363,7 @@ func TestOpenAICompatibleClientRequestMapping(t *testing.T) {
|
|||||||
run func(*testing.T)
|
run func(*testing.T)
|
||||||
}{
|
}{
|
||||||
{name: "complete request and response mapping", run: checkCompleteRequestAndResponseMapping},
|
{name: "complete request and response mapping", run: checkCompleteRequestAndResponseMapping},
|
||||||
|
{name: "canonical role and exact content", run: checkCanonicalRoleAndExactContentMapping},
|
||||||
{name: "cache-controlled message", run: checkCacheControlledMessageMapping},
|
{name: "cache-controlled message", run: checkCacheControlledMessageMapping},
|
||||||
{name: "empty cache-control TTL", run: checkEmptyCacheControlTTLOmission},
|
{name: "empty cache-control TTL", run: checkEmptyCacheControlTTLOmission},
|
||||||
{name: "session ID", run: checkSessionIDMapping},
|
{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) {
|
func checkCompleteRequestAndResponseMapping(t *testing.T) {
|
||||||
provider := newRecordingProvider(t)
|
provider := newRecordingProvider(t)
|
||||||
provider.respond(http.StatusOK, `{
|
provider.respond(http.StatusOK, `{
|
||||||
|
|||||||
@@ -213,16 +213,7 @@ func clonePreparedRun(source *domain.PreparedRun) (*domain.PreparedRun, error) {
|
|||||||
copied.InputHashes[name] = hash
|
copied.InputHashes[name] = hash
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if source.Messages != nil {
|
copied.Messages = domain.CloneRenderedMessages(source.Messages)
|
||||||
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
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if source.StructuredOutput != nil {
|
if source.StructuredOutput != nil {
|
||||||
structuredOutput := *source.StructuredOutput
|
structuredOutput := *source.StructuredOutput
|
||||||
|
|||||||
@@ -55,8 +55,7 @@ func (r *defaultOutputRepairer) Repair(ctx context.Context, req RepairRequest) (
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
messages := make([]domain.RenderedMessage, len(req.OriginalMessages), len(req.OriginalMessages)+2)
|
messages := domain.ConcatRenderedMessages(req.OriginalMessages, nil)
|
||||||
copy(messages, req.OriginalMessages)
|
|
||||||
if strings.TrimSpace(req.PreviousOutput) != "" {
|
if strings.TrimSpace(req.PreviousOutput) != "" {
|
||||||
messages = append(messages, domain.RenderedMessage{
|
messages = append(messages, domain.RenderedMessage{
|
||||||
Role: domain.RoleAssistant,
|
Role: domain.RoleAssistant,
|
||||||
|
|||||||
@@ -30,6 +30,8 @@ func TestDefaultOutputRepairerBuildsFullContextRequest(t *testing.T) {
|
|||||||
original := []domain.RenderedMessage{
|
original := []domain.RenderedMessage{
|
||||||
{Role: "system", Content: "Follow the task.", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral, TTL: "1h"}},
|
{Role: "system", Content: "Follow the task.", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral, TTL: "1h"}},
|
||||||
{Role: "user", Content: "Summarize the report."},
|
{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...)
|
before := append([]domain.RenderedMessage(nil), original...)
|
||||||
previous := strings.Repeat("candidate ", 12_000)
|
previous := strings.Repeat("candidate ", 12_000)
|
||||||
|
|||||||
@@ -4,10 +4,13 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
|
"encoding/binary"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"hash"
|
||||||
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -444,9 +447,11 @@ func (r *Runner) completePreparationWithStructuredOutput(
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("%w: %w", ErrPromptRender, err)
|
return nil, fmt.Errorf("%w: %w", ErrPromptRender, err)
|
||||||
}
|
}
|
||||||
|
effectivePrompt := *renderedPrompt
|
||||||
if state.directSessionID != "" {
|
if state.directSessionID != "" {
|
||||||
renderedPrompt.SessionID = state.directSessionID
|
effectivePrompt.SessionID = state.directSessionID
|
||||||
}
|
}
|
||||||
|
effectivePrompt.Messages = domain.ConcatRenderedMessages(renderedPrompt.Messages, req.AppendedMessages)
|
||||||
|
|
||||||
end := time.Now().UTC()
|
end := time.Now().UTC()
|
||||||
effectiveModel := state.effectiveModel
|
effectiveModel := state.effectiveModel
|
||||||
@@ -462,9 +467,9 @@ func (r *Runner) completePreparationWithStructuredOutput(
|
|||||||
OutputContract: state.effectiveContract,
|
OutputContract: state.effectiveContract,
|
||||||
StructuredOutput: structuredOutput,
|
StructuredOutput: structuredOutput,
|
||||||
InputHashes: inputHashes,
|
InputHashes: inputHashes,
|
||||||
SessionID: renderedPrompt.SessionID,
|
SessionID: effectivePrompt.SessionID,
|
||||||
RenderedPromptHash: hashRenderedPrompt(*renderedPrompt),
|
RenderedPromptHash: hashRenderedPrompt(effectivePrompt),
|
||||||
Messages: renderedPrompt.Messages,
|
Messages: effectivePrompt.Messages,
|
||||||
StartTime: state.start,
|
StartTime: state.start,
|
||||||
EndTime: end,
|
EndTime: end,
|
||||||
DurationMS: end.Sub(state.start).Milliseconds(),
|
DurationMS: end.Sub(state.start).Milliseconds(),
|
||||||
@@ -714,29 +719,38 @@ func resolveOutputContract(def *domain.PromptDefinition, override *domain.Output
|
|||||||
return contract, nil
|
return contract, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const renderedPromptHashVersion = "promptkit/rendered-prompt/v2\x00"
|
||||||
|
|
||||||
func hashRenderedPrompt(p domain.RenderedPrompt) string {
|
func hashRenderedPrompt(p domain.RenderedPrompt) string {
|
||||||
var b strings.Builder
|
hasher := sha256.New()
|
||||||
if p.SessionID != "" {
|
var lengthBuffer [8]byte
|
||||||
b.WriteString("session_id=")
|
_, _ = io.WriteString(hasher, renderedPromptHashVersion)
|
||||||
b.WriteString(p.SessionID)
|
writeRenderedPromptHashString(hasher, &lengthBuffer, p.SessionID)
|
||||||
b.WriteString("\n---\n")
|
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
|
||||||
}
|
}
|
||||||
for _, msg := range p.Messages {
|
lengthBuffer[0] = 1
|
||||||
b.WriteString(msg.Role)
|
_, _ = hasher.Write(lengthBuffer[:1])
|
||||||
b.WriteByte('\n')
|
writeRenderedPromptHashString(hasher, &lengthBuffer, string(message.CacheControl.Type))
|
||||||
b.WriteString(msg.Content)
|
writeRenderedPromptHashString(hasher, &lengthBuffer, message.CacheControl.TTL)
|
||||||
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)
|
|
||||||
}
|
}
|
||||||
}
|
return hex.EncodeToString(hasher.Sum(nil))
|
||||||
b.WriteString("\n---\n")
|
}
|
||||||
}
|
|
||||||
h := sha256.Sum256([]byte(b.String()))
|
func writeRenderedPromptHashLength(hasher hash.Hash, buffer *[8]byte, length int) {
|
||||||
return hex.EncodeToString(h[:])
|
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 {
|
func buildOutputArtifact(content string, format domain.OutputFormat) domain.Artifact {
|
||||||
|
|||||||
@@ -1110,14 +1110,18 @@ func TestDeriveStructuredSchemaName(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
|
func TestHashRenderedPromptIncludesEveryContractedField(t *testing.T) {
|
||||||
uncached := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
base := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
||||||
{Role: "system", Content: "sys"},
|
{Role: domain.RoleSystem, Content: "sys"},
|
||||||
{Role: "user", Content: "usr"},
|
{Role: domain.RoleUser, Content: "usr"},
|
||||||
}}
|
}}
|
||||||
wantLegacyHash := hashString("system\nsys\n---\nuser\nusr\n---\n")
|
identical := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
||||||
if got := hashRenderedPrompt(uncached); got != wantLegacyHash {
|
{Role: domain.RoleSystem, Content: "sys"},
|
||||||
t.Fatalf("expected no-cache hash to preserve legacy input, got %q want %q", got, wantLegacyHash)
|
{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{
|
withCache := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
||||||
@@ -1129,7 +1133,7 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
|
|||||||
TTL: "1h",
|
TTL: "1h",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{Role: "user", Content: "usr"},
|
{Role: domain.RoleUser, Content: "usr"},
|
||||||
}}
|
}}
|
||||||
alsoWithCache := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
alsoWithCache := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
||||||
{
|
{
|
||||||
@@ -1140,7 +1144,7 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
|
|||||||
TTL: "1h",
|
TTL: "1h",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{Role: "user", Content: "usr"},
|
{Role: domain.RoleUser, Content: "usr"},
|
||||||
}}
|
}}
|
||||||
withoutTTL := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
withoutTTL := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
||||||
{
|
{
|
||||||
@@ -1150,11 +1154,11 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
|
|||||||
Type: domain.CacheControlEphemeral,
|
Type: domain.CacheControlEphemeral,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{Role: "user", Content: "usr"},
|
{Role: domain.RoleUser, Content: "usr"},
|
||||||
}}
|
}}
|
||||||
|
|
||||||
cachedHash := hashRenderedPrompt(withCache)
|
cachedHash := hashRenderedPrompt(withCache)
|
||||||
if cachedHash == hashRenderedPrompt(uncached) {
|
if cachedHash == baseHash {
|
||||||
t.Fatal("expected cache control to change rendered prompt hash")
|
t.Fatal("expected cache control to change rendered prompt hash")
|
||||||
}
|
}
|
||||||
if cachedHash != hashRenderedPrompt(alsoWithCache) {
|
if cachedHash != hashRenderedPrompt(alsoWithCache) {
|
||||||
@@ -1163,6 +1167,29 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
|
|||||||
if cachedHash == hashRenderedPrompt(withoutTTL) {
|
if cachedHash == hashRenderedPrompt(withoutTTL) {
|
||||||
t.Fatal("expected ttl changes to affect rendered prompt hash")
|
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) {
|
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) {
|
func TestRunnerRunSuccessful(t *testing.T) {
|
||||||
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatMarkdown, domain.ValidationBasic, 0)}
|
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatMarkdown, domain.ValidationBasic, 0)}
|
||||||
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}
|
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}
|
||||||
|
|||||||
Reference in New Issue
Block a user