627 lines
20 KiB
Go
627 lines
20 KiB
Go
package promptdef
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"io/fs"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"sync"
|
|
"testing"
|
|
"testing/fstest"
|
|
|
|
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
|
)
|
|
|
|
func TestFilesystemRepository_GetPromptDefinition(t *testing.T) {
|
|
tmpDir := t.TempDir()
|
|
if err := copyTree("testdata", tmpDir); err != nil {
|
|
t.Fatalf("failed to copy testdata: %v", err)
|
|
}
|
|
|
|
repo := NewFilesystemRepository(tmpDir)
|
|
ctx := context.Background()
|
|
|
|
t.Run("valid inline prompt", func(t *testing.T) {
|
|
p, err := repo.GetPromptDefinition(ctx, "valid-inline", "")
|
|
if err != nil {
|
|
t.Fatalf("expected no error, got %v", err)
|
|
}
|
|
if p.ID != "valid-inline" {
|
|
t.Fatalf("unexpected id: %q", p.ID)
|
|
}
|
|
if p.Version != "1.0.0" {
|
|
t.Fatalf("unexpected version: %q", p.Version)
|
|
}
|
|
if p.OutputFormat != domain.FormatMarkdown {
|
|
t.Fatalf("unexpected output format: %q", p.OutputFormat)
|
|
}
|
|
if p.Validation.ValidationMode != domain.ValidationBasic {
|
|
t.Fatalf("unexpected validation mode: %q", p.Validation.ValidationMode)
|
|
}
|
|
if len(p.Templates) != 2 {
|
|
t.Fatalf("expected 2 messages, got %d", len(p.Templates))
|
|
}
|
|
if len(p.Inputs) != 1 {
|
|
t.Fatalf("expected 1 input, got %d", len(p.Inputs))
|
|
}
|
|
if p.Inputs[0].ContentType != "text/markdown" {
|
|
t.Fatalf("expected input content_type to be preserved, got %q", p.Inputs[0].ContentType)
|
|
}
|
|
})
|
|
|
|
t.Run("valid file-backed prompt", func(t *testing.T) {
|
|
p, err := repo.GetPromptDefinition(ctx, "valid-file-backed", "")
|
|
if err != nil {
|
|
t.Fatalf("expected no error, got %v", err)
|
|
}
|
|
if len(p.Templates) != 2 {
|
|
t.Fatalf("expected 2 messages, got %d", len(p.Templates))
|
|
}
|
|
if !strings.Contains(p.Templates[1].Content, "{{input \"transcript\"}}") {
|
|
t.Fatalf("expected content_file template body to be loaded, got %q", p.Templates[1].Content)
|
|
}
|
|
if p.Templates[1].ContentFile == "" {
|
|
t.Fatal("expected ContentFile source metadata to be preserved")
|
|
}
|
|
if !filepath.IsAbs(p.Templates[1].ContentFile) {
|
|
t.Fatalf("expected resolved content_file path to be absolute, got %q", p.Templates[1].ContentFile)
|
|
}
|
|
})
|
|
|
|
t.Run("valid cache control with ttl", func(t *testing.T) {
|
|
p, err := repo.GetPromptDefinition(ctx, "valid-cache-control-ttl", "")
|
|
if err != nil {
|
|
t.Fatalf("expected no error, got %v", err)
|
|
}
|
|
if len(p.Templates) != 2 {
|
|
t.Fatalf("expected 2 messages, got %d", len(p.Templates))
|
|
}
|
|
assertCacheControl(t, p.Templates[0].CacheControl, domain.CacheControlEphemeral, "1h")
|
|
if p.Templates[1].CacheControl != nil {
|
|
t.Fatalf("expected second message cache control to be nil, got %#v", p.Templates[1].CacheControl)
|
|
}
|
|
})
|
|
|
|
t.Run("valid cache control without ttl", func(t *testing.T) {
|
|
p, err := repo.GetPromptDefinition(ctx, "valid-cache-control-without-ttl", "")
|
|
if err != nil {
|
|
t.Fatalf("expected no error, got %v", err)
|
|
}
|
|
if len(p.Templates) != 2 {
|
|
t.Fatalf("expected 2 messages, got %d", len(p.Templates))
|
|
}
|
|
assertCacheControl(t, p.Templates[0].CacheControl, domain.CacheControlEphemeral, "")
|
|
if p.Templates[1].CacheControl != nil {
|
|
t.Fatalf("expected second message cache control to be nil, got %#v", p.Templates[1].CacheControl)
|
|
}
|
|
})
|
|
|
|
t.Run("valid session id template", func(t *testing.T) {
|
|
p, err := repo.GetPromptDefinition(ctx, "valid-session-id", "")
|
|
if err != nil {
|
|
t.Fatalf("expected no error, got %v", err)
|
|
}
|
|
if p.SessionID != "{{ .session_id }}" {
|
|
t.Fatalf("expected trimmed session_id template, got %q", p.SessionID)
|
|
}
|
|
})
|
|
|
|
t.Run("valid nested file-backed prompt resolves content file relative to nested YAML", func(t *testing.T) {
|
|
nestedDir := filepath.Join(tmpDir, "dnd", "recap")
|
|
if err := os.MkdirAll(nestedDir, 0o755); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
writePromptTestFile(t, filepath.Join(nestedDir, "nested_recap.yaml"), `
|
|
id: nested-recap
|
|
version: "1.0.0"
|
|
messages:
|
|
- role: user
|
|
content_file: ./nested_recap.user.tmpl
|
|
output:
|
|
format: markdown
|
|
validation_mode: basic
|
|
repair_attempts: 0
|
|
`)
|
|
writePromptTestFile(t, filepath.Join(nestedDir, "nested_recap.user.tmpl"), `Nested recap: {{input "transcript"}}`)
|
|
|
|
p, err := repo.GetPromptDefinition(ctx, "nested-recap", "")
|
|
if err != nil {
|
|
t.Fatalf("expected no error, got %v", err)
|
|
}
|
|
if len(p.Templates) != 1 {
|
|
t.Fatalf("expected one template, got %d", len(p.Templates))
|
|
}
|
|
if !strings.Contains(p.Templates[0].Content, "Nested recap") {
|
|
t.Fatalf("expected nested content file body, got %q", p.Templates[0].Content)
|
|
}
|
|
if !strings.Contains(p.Templates[0].ContentFile, filepath.Join("dnd", "recap", "nested_recap.user.tmpl")) {
|
|
t.Fatalf("expected nested content file path, got %q", p.Templates[0].ContentFile)
|
|
}
|
|
})
|
|
|
|
t.Run("prompt with default_profile", func(t *testing.T) {
|
|
p, err := repo.GetPromptDefinition(ctx, "with-default-profile", "")
|
|
if err != nil {
|
|
t.Fatalf("expected no error, got %v", err)
|
|
}
|
|
if p.DefaultProfile != "local-default" {
|
|
t.Fatalf("unexpected default profile: %q", p.DefaultProfile)
|
|
}
|
|
if len(p.Inputs) != 1 {
|
|
t.Fatalf("expected one input, got %d", len(p.Inputs))
|
|
}
|
|
if p.Inputs[0].ContentType != "" {
|
|
t.Fatalf("expected missing content_type to remain empty, got %q", p.Inputs[0].ContentType)
|
|
}
|
|
})
|
|
|
|
t.Run("duplicate prompt IDs fail as ambiguous", func(t *testing.T) {
|
|
writePromptTestFile(t, filepath.Join(tmpDir, "duplicate_a.yaml"), `
|
|
id: duplicate-prompt
|
|
version: "1.0.0"
|
|
messages:
|
|
- role: user
|
|
content: First duplicate.
|
|
output:
|
|
format: markdown
|
|
validation_mode: basic
|
|
repair_attempts: 0
|
|
`)
|
|
nestedDir := filepath.Join(tmpDir, "nested")
|
|
if err := os.MkdirAll(nestedDir, 0o755); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
writePromptTestFile(t, filepath.Join(nestedDir, "duplicate_b.yaml"), `
|
|
id: duplicate-prompt
|
|
version: "2.0.0"
|
|
messages:
|
|
- role: user
|
|
content: Second duplicate.
|
|
output:
|
|
format: markdown
|
|
validation_mode: basic
|
|
repair_attempts: 0
|
|
`)
|
|
|
|
_, err := repo.GetPromptDefinition(ctx, "duplicate-prompt", "")
|
|
if !errors.Is(err, ErrInvalidPromptDefinition) {
|
|
t.Fatalf("expected duplicate prompt to return ErrInvalidPromptDefinition, got %v", err)
|
|
}
|
|
for _, want := range []string{"duplicate prompt definition id", "duplicate_a.yaml", filepath.Join("nested", "duplicate_b.yaml")} {
|
|
if !strings.Contains(err.Error(), want) {
|
|
t.Fatalf("expected error to contain %q, got %v", want, err)
|
|
}
|
|
}
|
|
})
|
|
|
|
t.Run("duplicate prompt ID and requested version fails as ambiguous", func(t *testing.T) {
|
|
writePromptTestFile(t, filepath.Join(tmpDir, "version_duplicate_a.yaml"), `
|
|
id: duplicate-version-prompt
|
|
version: "1.0.0"
|
|
messages:
|
|
- role: user
|
|
content: First duplicate version.
|
|
output:
|
|
format: markdown
|
|
validation_mode: basic
|
|
repair_attempts: 0
|
|
`)
|
|
nestedDir := filepath.Join(tmpDir, "versioned")
|
|
if err := os.MkdirAll(nestedDir, 0o755); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
writePromptTestFile(t, filepath.Join(nestedDir, "version_duplicate_b.yaml"), `
|
|
id: duplicate-version-prompt
|
|
version: "1.0.0"
|
|
messages:
|
|
- role: user
|
|
content: Second duplicate version.
|
|
output:
|
|
format: markdown
|
|
validation_mode: basic
|
|
repair_attempts: 0
|
|
`)
|
|
|
|
_, err := repo.GetPromptDefinition(ctx, "duplicate-version-prompt", "1.0.0")
|
|
if !errors.Is(err, ErrInvalidPromptDefinition) {
|
|
t.Fatalf("expected duplicate prompt version to return ErrInvalidPromptDefinition, got %v", err)
|
|
}
|
|
for _, want := range []string{"duplicate prompt definition id", "version \"1.0.0\"", "version_duplicate_a.yaml", filepath.Join("versioned", "version_duplicate_b.yaml")} {
|
|
if !strings.Contains(err.Error(), want) {
|
|
t.Fatalf("expected error to contain %q, got %v", want, err)
|
|
}
|
|
}
|
|
})
|
|
|
|
t.Run("non-matching malformed nested prompt is ignored for not found lookup", func(t *testing.T) {
|
|
nestedDir := filepath.Join(tmpDir, "broken")
|
|
if err := os.MkdirAll(nestedDir, 0o755); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
writePromptTestFile(t, filepath.Join(nestedDir, "unrelated.yaml"), "id: [")
|
|
|
|
_, err := repo.GetPromptDefinition(ctx, "does-not-exist-even-with-broken-nested-file", "")
|
|
if !errors.Is(err, ErrPromptDefinitionNotFound) {
|
|
t.Fatalf("expected ErrPromptDefinitionNotFound, got %v", err)
|
|
}
|
|
})
|
|
|
|
t.Run("strict decode failure in nested prompt matches by YAML ID", func(t *testing.T) {
|
|
nestedDir := filepath.Join(tmpDir, "strict")
|
|
if err := os.MkdirAll(nestedDir, 0o755); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
writePromptTestFile(t, filepath.Join(nestedDir, "not_named_like_id.yaml"), `
|
|
id: nested-strict-error
|
|
version: "1.0.0"
|
|
unknown_field: true
|
|
messages:
|
|
- role: user
|
|
content: Invalid because of unknown field.
|
|
output:
|
|
format: markdown
|
|
validation_mode: basic
|
|
repair_attempts: 0
|
|
`)
|
|
|
|
_, err := repo.GetPromptDefinition(ctx, "nested-strict-error", "")
|
|
if !errors.Is(err, ErrInvalidYAML) {
|
|
t.Fatalf("expected ErrInvalidYAML, got %v", err)
|
|
}
|
|
if !strings.Contains(err.Error(), filepath.Join("strict", "not_named_like_id.yaml")) {
|
|
t.Fatalf("expected nested path in error, got %v", err)
|
|
}
|
|
})
|
|
|
|
t.Run("version lookup", func(t *testing.T) {
|
|
_, err := repo.GetPromptDefinition(ctx, "valid-inline", "9.9.9")
|
|
if !errors.Is(err, ErrPromptDefinitionNotFound) {
|
|
t.Fatalf("expected ErrPromptDefinitionNotFound, got %v", err)
|
|
}
|
|
})
|
|
|
|
cases := []struct {
|
|
name string
|
|
id string
|
|
targetErr error
|
|
errSubstrs []string
|
|
}{
|
|
{name: "invalid YAML", id: "invalid_yaml", targetErr: ErrInvalidYAML},
|
|
{name: "missing id", id: "missing_id", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"id is required"}},
|
|
{name: "no messages", id: "no_messages", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"at least one message is required"}},
|
|
{name: "both content and content_file", id: "both_content_and_content_file", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"exactly one"}},
|
|
{name: "neither content nor content_file", id: "neither_content_nor_content_file", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"exactly one"}},
|
|
{name: "missing content_file", id: "missing_content_file", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"failed to read content_file"}},
|
|
{name: "duplicate input names", id: "duplicate_input_names", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"duplicate input name"}},
|
|
{name: "invalid validation mode", id: "invalid_validation_mode", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"invalid validation mode"}},
|
|
{name: "json_schema without schema_path", id: "json_schema_without_schema_path", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"schema_path"}},
|
|
{name: "unknown input field", id: "unknown_input_field", targetErr: ErrInvalidYAML, errSubstrs: []string{"field unknown_input_setting not found"}},
|
|
{name: "empty cache control type", id: "empty_cache_control_type", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"cache_control", "type is required"}},
|
|
{name: "unsupported cache control type", id: "unsupported_cache_control_type", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"cache_control", "unsupported type"}},
|
|
{name: "unsupported cache control ttl", id: "unsupported_cache_control_ttl", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"cache_control", "unsupported ttl"}},
|
|
{name: "unknown cache control field", id: "unknown_cache_control_field", targetErr: ErrInvalidYAML, errSubstrs: []string{"field unexpected not found"}},
|
|
}
|
|
|
|
for _, tc := range cases {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
_, err := repo.GetPromptDefinition(ctx, tc.id, "")
|
|
if !errors.Is(err, tc.targetErr) {
|
|
t.Fatalf("expected %v, got %v", tc.targetErr, err)
|
|
}
|
|
for _, sub := range tc.errSubstrs {
|
|
if !strings.Contains(err.Error(), sub) {
|
|
t.Fatalf("expected error to contain %q, got %v", sub, err)
|
|
}
|
|
}
|
|
})
|
|
}
|
|
|
|
t.Run("prompt definition not found", func(t *testing.T) {
|
|
_, err := repo.GetPromptDefinition(ctx, "does-not-exist", "")
|
|
if !errors.Is(err, ErrPromptDefinitionNotFound) {
|
|
t.Fatalf("expected ErrPromptDefinitionNotFound, got %v", err)
|
|
}
|
|
})
|
|
}
|
|
|
|
func TestFSRepositoryGetPromptDefinition(t *testing.T) {
|
|
repo := NewFSRepository(fstest.MapFS{
|
|
"prompts/nested/prompt.yaml": &fstest.MapFile{Data: []byte(`
|
|
id: fs-prompt
|
|
version: "1.0.0"
|
|
inputs:
|
|
- name: transcript
|
|
required: true
|
|
messages:
|
|
- role: user
|
|
content_file: ./messages/user.tmpl
|
|
output:
|
|
format: markdown
|
|
validation_mode: basic
|
|
repair_attempts: 0
|
|
`)},
|
|
"prompts/nested/messages/user.tmpl": &fstest.MapFile{Data: []byte(`Summarize {{input "transcript"}}.`)},
|
|
}, "prompts")
|
|
|
|
got, err := repo.GetPromptDefinition(context.Background(), "fs-prompt", "")
|
|
if err != nil {
|
|
t.Fatalf("expected no error, got %v", err)
|
|
}
|
|
if got.ID != "fs-prompt" {
|
|
t.Fatalf("unexpected prompt id: %q", got.ID)
|
|
}
|
|
if len(got.Templates) != 1 || !strings.Contains(got.Templates[0].Content, `{{input "transcript"}}`) {
|
|
t.Fatalf("expected content_file body to be loaded, got %+v", got.Templates)
|
|
}
|
|
if got.Templates[0].ContentFile != "prompts/nested/messages/user.tmpl" {
|
|
t.Fatalf("unexpected content file path: %q", got.Templates[0].ContentFile)
|
|
}
|
|
}
|
|
|
|
func TestFSRepositoryContentFileContainment(t *testing.T) {
|
|
t.Run("nested prompt can reference file inside root", func(t *testing.T) {
|
|
repo := NewFSRepository(fstest.MapFS{
|
|
"prompts/nested/prompt.yaml": &fstest.MapFile{Data: []byte(`
|
|
id: fs-contained-prompt
|
|
version: "1.0.0"
|
|
messages:
|
|
- role: user
|
|
content_file: ../shared/user.tmpl
|
|
output:
|
|
format: markdown
|
|
validation_mode: basic
|
|
repair_attempts: 0
|
|
`)},
|
|
"prompts/shared/user.tmpl": &fstest.MapFile{Data: []byte(`Inside root.`)},
|
|
}, "prompts")
|
|
|
|
got, err := repo.GetPromptDefinition(context.Background(), "fs-contained-prompt", "")
|
|
if err != nil {
|
|
t.Fatalf("expected no error, got %v", err)
|
|
}
|
|
if len(got.Templates) != 1 || got.Templates[0].Content != "Inside root." {
|
|
t.Fatalf("expected contained content file, got %+v", got.Templates)
|
|
}
|
|
})
|
|
|
|
tests := []struct {
|
|
name string
|
|
contentFile string
|
|
wantErr string
|
|
}{
|
|
{name: "parent escape rejected", contentFile: "../outside.tmpl", wantErr: "escapes source root"},
|
|
{name: "absolute path rejected", contentFile: "/outside.tmpl", wantErr: "must be relative"},
|
|
}
|
|
|
|
for _, tc := range tests {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
fsys := &recordingFS{FS: fstest.MapFS{
|
|
"prompts/prompt.yaml": &fstest.MapFile{Data: []byte(`
|
|
id: fs-escaped-prompt
|
|
version: "1.0.0"
|
|
messages:
|
|
- role: user
|
|
content_file: ` + tc.contentFile + `
|
|
output:
|
|
format: markdown
|
|
validation_mode: basic
|
|
repair_attempts: 0
|
|
`)},
|
|
"outside.tmpl": &fstest.MapFile{Data: []byte(`Outside root.`)},
|
|
}}
|
|
repo := NewFSRepository(fsys, "prompts")
|
|
|
|
_, err := repo.GetPromptDefinition(context.Background(), "fs-escaped-prompt", "")
|
|
if !errors.Is(err, ErrInvalidPromptDefinition) {
|
|
t.Fatalf("expected ErrInvalidPromptDefinition, got %v", err)
|
|
}
|
|
if !strings.Contains(err.Error(), tc.wantErr) {
|
|
t.Fatalf("expected error to contain %q, got %v", tc.wantErr, err)
|
|
}
|
|
if fsys.wasOpened("outside.tmpl") {
|
|
t.Fatal("rejected content path opened the outside file")
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
type recordingFS struct {
|
|
fs.FS
|
|
mu sync.Mutex
|
|
opened []string
|
|
}
|
|
|
|
func (f *recordingFS) Open(name string) (fs.File, error) {
|
|
f.mu.Lock()
|
|
f.opened = append(f.opened, name)
|
|
f.mu.Unlock()
|
|
return f.FS.Open(name)
|
|
}
|
|
|
|
func (f *recordingFS) wasOpened(name string) bool {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
for _, opened := range f.opened {
|
|
if opened == name {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
func TestFSRepositoryRejectsDuplicatePromptIDs(t *testing.T) {
|
|
repo := NewFSRepository(fstest.MapFS{
|
|
"one.yaml": &fstest.MapFile{Data: []byte(`
|
|
id: duplicate-fs-prompt
|
|
version: "1.0.0"
|
|
messages:
|
|
- role: user
|
|
content: First.
|
|
output:
|
|
format: text
|
|
validation_mode: none
|
|
repair_attempts: 0
|
|
`)},
|
|
"nested/two.yaml": &fstest.MapFile{Data: []byte(`
|
|
id: duplicate-fs-prompt
|
|
version: "1.0.0"
|
|
messages:
|
|
- role: user
|
|
content: Second.
|
|
output:
|
|
format: text
|
|
validation_mode: none
|
|
repair_attempts: 0
|
|
`)},
|
|
}, ".")
|
|
|
|
_, err := repo.GetPromptDefinition(context.Background(), "duplicate-fs-prompt", "")
|
|
if !errors.Is(err, ErrInvalidPromptDefinition) {
|
|
t.Fatalf("expected ErrInvalidPromptDefinition, got %v", err)
|
|
}
|
|
if !strings.Contains(err.Error(), "one.yaml") || !strings.Contains(err.Error(), "nested/two.yaml") {
|
|
t.Fatalf("expected duplicate paths in error, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestFSRepositoryRejectsUnknownYAMLFields(t *testing.T) {
|
|
repo := NewFSRepository(fstest.MapFS{
|
|
"not_named_like_id.yaml": &fstest.MapFile{Data: []byte(`
|
|
id: strict-fs-prompt
|
|
version: "1.0.0"
|
|
unknown: true
|
|
messages:
|
|
- role: user
|
|
content: Invalid.
|
|
output:
|
|
format: text
|
|
validation_mode: none
|
|
repair_attempts: 0
|
|
`)},
|
|
}, ".")
|
|
|
|
_, err := repo.GetPromptDefinition(context.Background(), "strict-fs-prompt", "")
|
|
if !errors.Is(err, ErrInvalidYAML) {
|
|
t.Fatalf("expected ErrInvalidYAML, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestPromptRepositoriesApplyOutputContractRules(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
output string
|
|
useFilesystem bool
|
|
wantErr bool
|
|
wantDiagnostic string
|
|
wantSchemaPath string
|
|
}{
|
|
{
|
|
name: "operating-system source rejects unsupported format",
|
|
output: " format: binary\n validation_mode: none\n",
|
|
useFilesystem: true,
|
|
wantErr: true,
|
|
wantDiagnostic: "format",
|
|
},
|
|
{
|
|
name: "fs source rejects negative repair attempts",
|
|
output: " format: text\n validation_mode: none\n repair_attempts: -1\n",
|
|
wantErr: true,
|
|
wantDiagnostic: "repair_attempts",
|
|
},
|
|
{
|
|
name: "source normalization trims a valid schema path",
|
|
output: " format: json\n validation_mode: json_schema\n schema_path: ' schema.json '\n",
|
|
useFilesystem: true,
|
|
wantSchemaPath: "schema.json",
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
data := `id: output-contract
|
|
version: "1"
|
|
messages:
|
|
- role: user
|
|
content: test
|
|
output:
|
|
` + tt.output
|
|
|
|
var repo Repository
|
|
if tt.useFilesystem {
|
|
dir := t.TempDir()
|
|
writePromptTestFile(t, filepath.Join(dir, "output-contract.yaml"), data)
|
|
repo = NewFilesystemRepository(dir)
|
|
} else {
|
|
repo = NewFSRepository(fstest.MapFS{
|
|
"output-contract.yaml": &fstest.MapFile{Data: []byte(data)},
|
|
}, ".")
|
|
}
|
|
|
|
got, err := repo.GetPromptDefinition(context.Background(), "output-contract", "")
|
|
if tt.wantErr {
|
|
if !errors.Is(err, ErrInvalidPromptDefinition) {
|
|
t.Fatalf("expected ErrInvalidPromptDefinition, got %v", err)
|
|
}
|
|
if !strings.Contains(err.Error(), tt.wantDiagnostic) {
|
|
t.Fatalf("expected error containing %q, got %v", tt.wantDiagnostic, err)
|
|
}
|
|
return
|
|
}
|
|
if err != nil {
|
|
t.Fatalf("load prompt definition: %v", err)
|
|
}
|
|
if got.Validation.SchemaPath != tt.wantSchemaPath {
|
|
t.Fatalf("schema path = %q, want %q", got.Validation.SchemaPath, tt.wantSchemaPath)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func assertCacheControl(t *testing.T, got *domain.CacheControl, wantType domain.CacheControlType, wantTTL string) {
|
|
t.Helper()
|
|
if got == nil {
|
|
t.Fatal("expected cache control, got nil")
|
|
}
|
|
if got.Type != wantType {
|
|
t.Fatalf("unexpected cache control type: got %q want %q", got.Type, wantType)
|
|
}
|
|
if got.TTL != wantTTL {
|
|
t.Fatalf("unexpected cache control ttl: got %q want %q", got.TTL, wantTTL)
|
|
}
|
|
}
|
|
|
|
func writePromptTestFile(t *testing.T, path string, content string) {
|
|
t.Helper()
|
|
if err := os.WriteFile(path, []byte(strings.TrimLeft(content, "\n")), 0o644); err != nil {
|
|
t.Fatalf("failed to write prompt test file %q: %v", path, err)
|
|
}
|
|
}
|
|
|
|
func copyTree(src, dst string) error {
|
|
return filepath.WalkDir(src, func(path string, d fs.DirEntry, err error) error {
|
|
if err != nil {
|
|
return err
|
|
}
|
|
rel, err := filepath.Rel(src, path)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if rel == "." {
|
|
return nil
|
|
}
|
|
|
|
target := filepath.Join(dst, rel)
|
|
if d.IsDir() {
|
|
return os.MkdirAll(target, 0o755)
|
|
}
|
|
|
|
data, err := os.ReadFile(path)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return os.WriteFile(target, data, 0o644)
|
|
})
|
|
}
|