Files
promptkit/appended_messages_contract_test.go

146 lines
6.2 KiB
Go

package promptkit_test
import (
"context"
"errors"
"reflect"
"sync/atomic"
"testing"
"gitea.maximumdirect.net/eric/promptkit"
)
func TestAppendedMessageRoleConstantsAreStrings(t *testing.T) {
var (
developer string = promptkit.RoleDeveloper
system string = promptkit.RoleSystem
user string = promptkit.RoleUser
assistant string = promptkit.RoleAssistant
)
if developer != "developer" || system != "system" || user != "user" || assistant != "assistant" {
t.Fatalf("unexpected role constants: %q %q %q %q", developer, system, user, assistant)
}
}
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) {
promptSource := &inspectionCountingFS{}
var modelCalls atomic.Int64
engine, err := promptkit.NewEngine(
promptkit.Config{},
promptkit.WithPromptFS(promptSource, "."),
promptkit.WithLLMClient(countingLLMClient{calls: &modelCalls}),
)
if err != nil {
t.Fatalf("construct engine: %v", err)
}
request := promptkit.RunRequest{
PromptID: "unreached",
AppendedMessages: []promptkit.RenderedMessage{{
Role: "unsupported-role",
Content: "sensitive appended content",
}},
}
operations := []struct {
name string
run func() error
}{
{name: "Prepare", run: func() error { _, err := engine.Prepare(context.Background(), request); return err }},
{name: "PrepareExecution", run: func() error { _, err := engine.PrepareExecution(context.Background(), request); return err }},
{name: "Run", run: func() error { _, err := engine.Run(context.Background(), request); return err }},
}
for _, operation := range operations {
t.Run(operation.name, func(t *testing.T) {
err := operation.run()
if !errors.Is(err, promptkit.ErrInvalidRequest) {
t.Fatalf("error = %v, want ErrInvalidRequest", err)
}
})
}
if promptSource.opens.Load() != 0 {
t.Fatalf("invalid appended message opened prompt sources %d times", promptSource.opens.Load())
}
if modelCalls.Load() != 0 {
t.Fatalf("invalid appended message invoked the model %d times", modelCalls.Load())
}
}