package promptdef import ( "context" "errors" "fmt" "io/fs" "os" "path/filepath" "strings" "sync" "testing" "testing/fstest" "gitea.maximumdirect.net/eric/promptkit/internal/domain" ) func TestPromptRepositoryDefinitionFixtures(t *testing.T) { sources := []struct { name string newRepository func(string) Repository contentPathsAreFull bool }{ { name: "operating system", newRepository: NewFilesystemRepository, contentPathsAreFull: true, }, { name: "filesystem", newRepository: func(root string) Repository { return NewFSRepository(os.DirFS(root), ".") }, }, } for _, source := range sources { t.Run(source.name, func(t *testing.T) { testPromptRepositoryDefinitionFixtures(t, source.newRepository, source.contentPathsAreFull) }) } } func testPromptRepositoryDefinitionFixtures(t *testing.T, newRepository func(string) Repository, contentPathsAreFull bool) { tmpDir := t.TempDir() if err := copyTree("testdata", tmpDir); err != nil { t.Fatalf("failed to copy testdata: %v", err) } repo := newRepository(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) != contentPathsAreFull { t.Fatalf("unexpected content_file path representation: %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: "unidentifiable invalid YAML is unrelated", id: "invalid-yaml", targetErr: ErrPromptDefinitionNotFound}, {name: "missing id is not selected by filename", id: "missing_id", targetErr: ErrPromptDefinitionNotFound}, {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 TestPromptRepositoryReturnsDefinitionReadFailures(t *testing.T) { readErr := errors.New("definition read failed") fsys := &definitionReadFailureFS{ FS: fstest.MapFS{ "prompts/target.yaml": &fstest.MapFile{Data: []byte("unread")}, }, target: "prompts/target.yaml", err: readErr, } repo := NewFSRepository(fsys, "prompts") definition, err := repo.GetPromptDefinition(context.Background(), "target", "1") if definition != nil || !errors.Is(err, readErr) { t.Fatalf("GetPromptDefinition() = (%#v, %v), want nil and definition read error", definition, err) } if errors.Is(err, ErrPromptDefinitionNotFound) { t.Fatalf("definition read error was classified as absence: %v", err) } if !strings.Contains(err.Error(), "target.yaml") { t.Fatalf("definition read error lacks source context: %v", err) } } type definitionReadFailureFS struct { fs.FS target string err error } func (f *definitionReadFailureFS) Open(name string) (fs.File, error) { if name == f.target { return nil, f.err } return f.FS.Open(name) } 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 { return f.openCount(name) > 0 } func (f *recordingFS) openCount(name string) int { f.mu.Lock() defer f.mu.Unlock() count := 0 for _, opened := range f.opened { if opened == name { count++ } } return count } func TestPromptRepositorySelectionUsesYAMLMetadata(t *testing.T) { const validDefinition = ` id: selected-prompt version: "1" messages: - role: user content: selected output: format: text validation_mode: none ` tests := []struct { name string files map[string]string wantErr error diagnostics []string wantContent string }{ { name: "same-stem strict error with different YAML ID is unrelated", files: map[string]string{ "selected-prompt.yaml": ` id: another-prompt version: "1" unknown: true `, "valid.yaml": validDefinition, }, }, { name: "unidentifiable same-stem YAML is unrelated", files: map[string]string{ "selected-prompt.yaml": "id: [", "valid.yaml": validDefinition, }, }, { name: "same ID invalid different version is unrelated", files: map[string]string{ "invalid-version.yaml": ` id: selected-prompt version: "2" output: format: text validation_mode: none `, "valid.yaml": validDefinition, }, }, { name: "selected content is resolved relative to its definition", files: map[string]string{ "nested/selected.yaml": ` id: selected-prompt version: "1" messages: - role: user content_file: ./content/selected.tmpl output: format: text validation_mode: none `, "nested/content/selected.tmpl": "selected from file", }, wantContent: "selected from file", }, { name: "duplicate selected definitions are ambiguous", files: map[string]string{ "selected-a.yaml": validDefinition, "nested/selected-b.yaml": ` id: selected-prompt version: "1" messages: - role: user content: duplicate output: format: text validation_mode: none `, }, wantErr: ErrInvalidPromptDefinition, diagnostics: []string{"duplicate prompt definition id", "selected-a.yaml", "nested/selected-b.yaml"}, }, { name: "selected strict error is authoritative", files: map[string]string{ "selected-strict.yaml": ` id: selected-prompt version: "1" unknown: true `, }, wantErr: ErrInvalidYAML, diagnostics: []string{"selected-strict.yaml"}, }, { name: "selected semantic error is authoritative", files: map[string]string{ "selected-invalid.yaml": ` id: selected-prompt version: "1" output: format: text validation_mode: none `, }, wantErr: ErrInvalidPromptDefinition, diagnostics: []string{"selected-invalid.yaml", "at least one message"}, }, { name: "selected content error includes definition context", files: map[string]string{ "selected-missing-content.yaml": ` id: selected-prompt version: "1" messages: - role: user content_file: missing.tmpl output: format: text validation_mode: none `, }, wantErr: ErrInvalidPromptDefinition, diagnostics: []string{"selected-missing-content.yaml", "failed to read content_file", "missing.tmpl"}, }, } for _, source := range promptRepositorySources() { for _, tc := range tests { t.Run(source.name+"/"+tc.name, func(t *testing.T) { repo := source.newRepository(t, tc.files) got, err := repo.GetPromptDefinition(context.Background(), "selected-prompt", "1") if tc.wantErr != nil { if !errors.Is(err, tc.wantErr) { t.Fatalf("expected %v, got %v", tc.wantErr, err) } for _, diagnostic := range tc.diagnostics { if !strings.Contains(err.Error(), diagnostic) { t.Fatalf("expected error to contain %q, got %v", diagnostic, err) } } return } if err != nil { t.Fatalf("load selected prompt: %v", err) } wantContent := tc.wantContent if wantContent == "" { wantContent = "selected" } if got.ID != "selected-prompt" || got.Version != "1" || got.Templates[0].Content != wantContent { t.Fatalf("unexpected selected definition: %+v", got) } }) } } } func TestPromptRepositoryHonorsCancellation(t *testing.T) { for _, source := range promptRepositorySources() { t.Run(source.name, func(t *testing.T) { repo := source.newRepository(t, map[string]string{ "definition.yaml": ` id: cancelled-prompt version: "1" messages: - role: user content: selected output: format: text validation_mode: none `, }) ctx, cancel := context.WithCancel(context.Background()) cancel() _, err := repo.GetPromptDefinition(ctx, "cancelled-prompt", "1") if !errors.Is(err, context.Canceled) { t.Fatalf("expected context cancellation, got %v", err) } }) } } func TestPromptRepositoryRequiresOneYAMLDocument(t *testing.T) { const definition = ` id: one-document version: "1" messages: - role: user content: selected output: format: text validation_mode: none ` tests := []struct { name string suffix string wantErr bool }{ {name: "comments and trailing whitespace", suffix: "\n# trailing comment\n\n"}, {name: "second populated document", suffix: "\n---\nid: another\n", wantErr: true}, {name: "second empty document", suffix: "\n---\n", wantErr: true}, {name: "malformed trailing YAML", suffix: "\n---\n[", wantErr: true}, } for _, source := range promptRepositorySources() { for _, tc := range tests { t.Run(source.name+"/"+tc.name, func(t *testing.T) { repo := source.newRepository(t, map[string]string{"definition.yaml": definition + tc.suffix}) _, err := repo.GetPromptDefinition(context.Background(), "one-document", "1") if tc.wantErr { if !errors.Is(err, ErrInvalidYAML) { t.Fatalf("expected ErrInvalidYAML, got %v", err) } if !strings.Contains(err.Error(), "definition.yaml") { t.Fatalf("expected source path in error, got %v", err) } return } if err != nil { t.Fatalf("load one document: %v", err) } }) } } } func TestPromptRepositoryReadsOnlySelectedContent(t *testing.T) { fsys := &recordingFS{FS: fstest.MapFS{ "target.yaml": &fstest.MapFile{Data: []byte(` id: target version: "1" messages: - role: user content_file: target.tmpl output: format: text validation_mode: none `)}, "target.tmpl": &fstest.MapFile{Data: []byte("selected")}, "unrelated.yaml": &fstest.MapFile{Data: []byte(` id: unrelated version: "1" messages: - role: user content_file: unrelated.tmpl output: format: text validation_mode: none `)}, "unrelated.tmpl": &fstest.MapFile{Data: []byte("unrelated")}, "other-version.yaml": &fstest.MapFile{Data: []byte(` id: target version: "2" messages: - role: user content_file: other-version.tmpl output: format: text validation_mode: none `)}, "other-version.tmpl": &fstest.MapFile{Data: []byte("other version")}, }} repo := NewFSRepository(fsys, ".") for lookup := 1; lookup <= 2; lookup++ { got, err := repo.GetPromptDefinition(context.Background(), "target", "1") if err != nil { t.Fatalf("lookup %d: %v", lookup, err) } if got.Templates[0].Content != "selected" { t.Fatalf("lookup %d content = %q", lookup, got.Templates[0].Content) } if count := fsys.openCount("target.tmpl"); count != lookup { t.Fatalf("selected content opens after lookup %d = %d, want %d", lookup, count, lookup) } for _, name := range []string{"unrelated.tmpl", "other-version.tmpl"} { if count := fsys.openCount(name); count != 0 { t.Fatalf("unselected content %q opened %d times", name, count) } } for _, name := range []string{"target.yaml", "unrelated.yaml", "other-version.yaml"} { if count := fsys.openCount(name); count != lookup { t.Fatalf("metadata %q opens after lookup %d = %d, want %d", name, lookup, count, lookup) } } } } func TestPromptDefinitionNormalizationRules(t *testing.T) { tests := []struct { name string definition string wantErr bool wantDiagnostic string wantSchemaPath string }{ { name: "missing version", definition: ` id: normalization-rule messages: - role: user content: test output: format: text validation_mode: none `, wantErr: true, wantDiagnostic: "version", }, { name: "blank input name", definition: ` id: normalization-rule version: "1" inputs: - name: " " messages: - role: user content: test output: format: text validation_mode: none `, wantErr: true, wantDiagnostic: "input 0", }, { name: "blank message role", definition: ` id: normalization-rule version: "1" messages: - role: " " content: test output: format: text validation_mode: none `, wantErr: true, wantDiagnostic: "role", }, { name: "invalid output format", definition: ` id: normalization-rule version: "1" messages: - role: user content: test output: format: binary validation_mode: none `, wantErr: true, wantDiagnostic: "format", }, { name: "negative repair attempts", definition: ` id: normalization-rule version: "1" messages: - role: user content: test output: format: text validation_mode: none repair_attempts: -1 `, wantErr: true, wantDiagnostic: "repair_attempts", }, { name: "repair attempts above maximum", definition: ` id: normalization-rule version: "1" messages: - role: user content: test output: format: text validation_mode: basic repair_attempts: 4 `, wantErr: true, wantDiagnostic: "repair_attempts", }, { name: "none validation with repair attempts", definition: ` id: normalization-rule version: "1" messages: - role: user content: test output: format: text validation_mode: none repair_attempts: 1 `, wantErr: true, wantDiagnostic: "repair_attempts", }, { name: "explicit blank default profile", definition: ` id: normalization-rule version: "1" default_profile: " " messages: - role: user content: test output: format: text validation_mode: none `, wantErr: true, wantDiagnostic: "default_profile", }, { name: "schema path normalization", definition: ` id: normalization-rule version: "1" messages: - role: user content: test output: format: json validation_mode: json_schema schema_path: ' schema.json ' `, wantSchemaPath: "schema.json", }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { repo := NewFSRepository(fstest.MapFS{ "definition.yaml": &fstest.MapFile{Data: []byte(tt.definition)}, }, ".") got, err := repo.GetPromptDefinition(context.Background(), "normalization-rule", "") 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) } }) } } type promptRepositorySource struct { name string newRepository func(t *testing.T, files map[string]string) Repository } func promptRepositorySources() []promptRepositorySource { return []promptRepositorySource{ { name: "operating system", newRepository: func(t *testing.T, files map[string]string) Repository { t.Helper() root := t.TempDir() for name, content := range files { filePath := filepath.Join(root, filepath.FromSlash(name)) if err := os.MkdirAll(filepath.Dir(filePath), 0o755); err != nil { t.Fatalf("create prompt directory: %v", err) } writePromptTestFile(t, filePath, content) } return NewFilesystemRepository(root) }, }, { name: "filesystem", newRepository: func(t *testing.T, files map[string]string) Repository { t.Helper() fsys := make(fstest.MapFS, len(files)) for name, content := range files { fsys[name] = &fstest.MapFile{Data: []byte(strings.TrimLeft(content, "\n"))} } return NewFSRepository(fsys, ".") }, }, } } func BenchmarkPromptRepositoryLookup(b *testing.B) { for _, size := range []int{10, 1000} { b.Run(fmt.Sprintf("catalog-%d", size), func(b *testing.B) { files := fstest.MapFS{ "target.yaml": &fstest.MapFile{Data: []byte(` id: target version: "1" messages: - role: user content_file: target.tmpl output: format: text validation_mode: none `)}, "target.tmpl": &fstest.MapFile{Data: []byte("selected")}, } metadataNames := []string{"target.yaml"} contentNames := make([]string, 0, size-1) for i := 1; i < size; i++ { definitionName := fmt.Sprintf("prompt-%04d.yaml", i) contentName := fmt.Sprintf("prompt-%04d.tmpl", i) files[definitionName] = &fstest.MapFile{Data: []byte(fmt.Sprintf(` id: prompt-%04d version: "1" messages: - role: user content_file: %s output: format: text validation_mode: none `, i, contentName))} files[contentName] = &fstest.MapFile{Data: []byte("unrelated")} metadataNames = append(metadataNames, definitionName) contentNames = append(contentNames, contentName) } fsys := &recordingFS{FS: files} repo := NewFSRepository(fsys, ".") b.ReportAllocs() b.ResetTimer() for i := 0; i < b.N; i++ { if _, err := repo.GetPromptDefinition(context.Background(), "target", "1"); err != nil { b.Fatal(err) } } b.StopTimer() if got := fsys.openCount("target.tmpl"); got != b.N { b.Fatalf("selected content opens = %d, want %d", got, b.N) } for _, name := range contentNames { if got := fsys.openCount(name); got != 0 { b.Fatalf("unrelated content %q opened %d times", name, got) } } for _, name := range metadataNames { if got := fsys.openCount(name); got != b.N { b.Fatalf("metadata %q opens = %d, want %d", name, got, b.N) } } }) } } 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) }) }