146 lines
6.2 KiB
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())
|
|
}
|
|
}
|