package profile import ( "context" "errors" "fmt" "io/fs" "reflect" "strings" "sync" "testing" "testing/fstest" "gitea.maximumdirect.net/eric/promptkit/internal/domain" ) func TestResolvingRepositoryMergesProfileChain(t *testing.T) { repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{ "leaf": { ID: "leaf", BaseProfileID: "middle", BackendID: "leaf-backend", TopP: 0.8, TimeoutSeconds: 45, ReasoningEffort: "high", }, "middle": { ID: "middle", BaseProfileID: "root", Endpoint: "https://middle.example/v1", Model: "middle-model", MaxTokens: 256, APIKeyEnv: "MIDDLE_API_KEY", APIKeyRequired: true, ExtraParams: map[string]any{"middle": map[string]any{"value": "middle"}}, }, "root": { ID: "root", BackendID: "root-backend", Endpoint: "https://root.example/v1", Model: "root-model", Temperature: 0.3, ServiceTier: "priority", ExtraParams: map[string]any{"root": "value"}, }, }} got, err := NewResolvingRepository(repo).GetProfile(context.Background(), "leaf") if err != nil { t.Fatalf("resolve profile: %v", err) } want := &domain.ExecutionProfile{ ID: "leaf", BackendID: "leaf-backend", Endpoint: "https://middle.example/v1", Model: "middle-model", Temperature: 0.3, MaxTokens: 256, TopP: 0.8, TimeoutSeconds: 45, ServiceTier: "priority", ReasoningEffort: "high", APIKeyEnv: "MIDDLE_API_KEY", APIKeyRequired: true, ExtraParams: map[string]any{"middle": map[string]any{"value": "middle"}}, } if !reflect.DeepEqual(got, want) { t.Fatalf("resolved profile:\n got %#v\nwant %#v", got, want) } } func TestResolvingRepositoryRejectsMissingSourceAndProfileID(t *testing.T) { if _, err := NewResolvingRepository(nil).GetProfile(context.Background(), "profile"); !errors.Is(err, ErrInvalidProfile) { t.Fatalf("nil source error = %v, want ErrInvalidProfile", err) } repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{}} if _, err := NewResolvingRepository(repo).GetProfile(context.Background(), " \t "); !errors.Is(err, ErrInvalidProfile) { t.Fatalf("blank id error = %v, want ErrInvalidProfile", err) } if got := repo.callCount(" "); got != 0 { t.Fatalf("blank id looked up source %d times", got) } } func TestResolvingRepositoryCopiesExtraParams(t *testing.T) { baseParams := map[string]any{"nested": map[string]any{"value": "base"}} repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{ "child": {ID: "child", BaseProfileID: "base"}, "base": { ID: "base", Endpoint: "https://base.example/v1", Model: "model", ExtraParams: baseParams, }, }} resolver := NewResolvingRepository(repo) first, err := resolver.GetProfile(context.Background(), "child") if err != nil { t.Fatalf("resolve inherited map: %v", err) } first.ExtraParams["nested"].(map[string]any)["value"] = "mutated" second, err := resolver.GetProfile(context.Background(), "child") if err != nil { t.Fatalf("resolve inherited map again: %v", err) } if got := second.ExtraParams["nested"].(map[string]any)["value"]; got != "base" { t.Fatalf("later result retained mutation: %v", got) } if got := baseParams["nested"].(map[string]any)["value"]; got != "base" { t.Fatalf("source map retained mutation: %v", got) } repo.set("child", &domain.ExecutionProfile{ ID: "child", BaseProfileID: "base", ExtraParams: map[string]any{"child": "replacement"}, }) replaced, err := resolver.GetProfile(context.Background(), "child") if err != nil { t.Fatalf("resolve replacement map: %v", err) } if !reflect.DeepEqual(replaced.ExtraParams, map[string]any{"child": "replacement"}) { t.Fatalf("extra params = %#v, want complete child replacement", replaced.ExtraParams) } } func TestResolvingRepositoryUsesRawOverlayForEachLookup(t *testing.T) { leafSource := NewFSRepository(profileTestFS(map[string]string{ "leaf.yaml": "id: leaf\nbase_profile: base\n", }), ".") fallback := NewFSRepository(profileTestFS(map[string]string{ "base.yaml": "id: base\nendpoint: https://fallback.example/v1\nmodel: fallback-model\n", }), ".") overlay := NewOverlayRepository(leafSource, fallback) resolver := NewResolvingRepository(overlay) got, err := resolver.GetProfile(context.Background(), "leaf") if err != nil { t.Fatalf("resolve fallback base: %v", err) } if got.Model != "fallback-model" { t.Fatalf("fallback base model = %q", got.Model) } shadowing := NewOverlayRepository(NewFSRepository(profileTestFS(map[string]string{ "leaf.yaml": "id: leaf\nbase_profile: base\n", "base.yaml": "id: base\nendpoint: https://primary.example/v1\nmodel: primary-model\n", }), "."), fallback) got, err = NewResolvingRepository(shadowing).GetProfile(context.Background(), "leaf") if err != nil { t.Fatalf("resolve shadowed base: %v", err) } if got.Model != "primary-model" || got.Endpoint != "https://primary.example/v1" { t.Fatalf("shadowed base = %+v", got) } } func TestResolvingRepositoryReportsSafetyAndSourceErrors(t *testing.T) { sourceErr := errors.New("source failure") tests := []struct { name string repo *resolvingTestRepository id string want []error wantNot error contains []string }{ { name: "missing selected profile preserves not found", repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{}}, id: "missing", want: []error{ErrProfileNotFound}, wantNot: ErrInvalidProfile, }, { name: "missing base is invalid but not not found", repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{ "leaf": {ID: "leaf", BaseProfileID: "missing"}, }}, id: "leaf", want: []error{ErrInvalidProfile}, wantNot: ErrProfileNotFound, contains: []string{"missing", "leaf -> missing"}, }, { name: "direct cycle", repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{ "a": {ID: "a", BaseProfileID: "a"}, }}, id: "a", want: []error{ErrInvalidProfile}, contains: []string{"a -> a"}, }, { name: "indirect cycle", repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{ "a": {ID: "a", BaseProfileID: "b"}, "b": {ID: "b", BaseProfileID: "c"}, "c": {ID: "c", BaseProfileID: "a"}, }}, id: "a", want: []error{ErrInvalidProfile}, contains: []string{"a -> b -> c -> a"}, }, { name: "nil result", repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{ "leaf": nil, }}, id: "leaf", want: []error{ErrInvalidProfile}, }, { name: "incomplete resolved profile", repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{ "leaf": {ID: "leaf", BaseProfileID: "base"}, "base": {ID: "base", Model: "model"}, }}, id: "leaf", want: []error{ErrInvalidProfile}, }, { name: "base source error is retained", repo: &resolvingTestRepository{ profiles: map[string]*domain.ExecutionProfile{"leaf": {ID: "leaf", BaseProfileID: "base"}}, errors: map[string]error{"base": sourceErr}, }, id: "leaf", want: []error{ErrInvalidProfile, sourceErr}, }, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { _, err := NewResolvingRepository(tc.repo).GetProfile(context.Background(), tc.id) for _, want := range tc.want { if !errors.Is(err, want) { t.Fatalf("error = %v, want %v", err, want) } } if tc.wantNot != nil && errors.Is(err, tc.wantNot) { t.Fatalf("error = %v, must not match %v", err, tc.wantNot) } for _, fragment := range tc.contains { if !strings.Contains(err.Error(), fragment) { t.Fatalf("error = %v, want %q", err, fragment) } } }) } } func TestResolvingRepositoryEnforcesChainLength(t *testing.T) { for _, count := range []int{maximumProfileChainLength, maximumProfileChainLength + 1} { t.Run(fmt.Sprintf("%d profiles", count), func(t *testing.T) { profiles := make(map[string]*domain.ExecutionProfile, count) for index := 1; index <= count; index++ { id := fmt.Sprintf("profile-%d", index) definition := &domain.ExecutionProfile{ID: id} if index == count { definition.Endpoint = "https://root.example/v1" definition.Model = "model" } else { definition.BaseProfileID = fmt.Sprintf("profile-%d", index+1) } profiles[id] = definition } got, err := NewResolvingRepository(&resolvingTestRepository{profiles: profiles}).GetProfile(context.Background(), "profile-1") if count == maximumProfileChainLength { if err != nil || got == nil { t.Fatalf("profile = %+v, error = %v, want accepted chain", got, err) } return } if !errors.Is(err, ErrInvalidProfile) { t.Fatalf("error = %v, want ErrInvalidProfile", err) } }) } } func TestResolvingRepositoryIsFreshAndCancellationAware(t *testing.T) { repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{ "leaf": {ID: "leaf", BaseProfileID: "base"}, "base": {ID: "base", Endpoint: "https://base.example/v1", Model: "first", ExtraParams: map[string]any{"nested": map[string]any{"value": "first"}}}, }} resolver := NewResolvingRepository(repo) first, err := resolver.GetProfile(context.Background(), "leaf") if err != nil || first.Model != "first" { t.Fatalf("first result=(%+v, %v)", first, err) } repo.set("base", &domain.ExecutionProfile{ID: "base", Endpoint: "https://base.example/v1", Model: "second", ExtraParams: map[string]any{"nested": map[string]any{"value": "second"}}}) second, err := resolver.GetProfile(context.Background(), "leaf") if err != nil || second.Model != "second" { t.Fatalf("second result=(%+v, %v)", second, err) } canceled, cancel := context.WithCancel(context.Background()) cancel() if _, err := resolver.GetProfile(canceled, "leaf"); !errors.Is(err, context.Canceled) { t.Fatalf("canceled lookup error = %v", err) } if got := repo.callCount("leaf"); got != 2 { t.Fatalf("calls after canceled lookup = %d, want 2", got) } duringTraversal, cancelDuringTraversal := context.WithCancel(context.Background()) repo.afterGet = func(id string) { if id == "leaf" { cancelDuringTraversal() } } if _, err := resolver.GetProfile(duringTraversal, "leaf"); !errors.Is(err, context.Canceled) { t.Fatalf("during traversal error = %v", err) } if got := repo.callCount("base"); got != 2 { t.Fatalf("base calls after cancellation = %d, want 2", got) } terminalLookup, cancelTerminalLookup := context.WithCancel(context.Background()) repo.afterGet = func(id string) { if id == "base" { cancelTerminalLookup() } } if _, err := resolver.GetProfile(terminalLookup, "leaf"); !errors.Is(err, context.Canceled) { t.Fatalf("terminal lookup cancellation error = %v", err) } if got := repo.callCount("base"); got != 3 { t.Fatalf("base calls after terminal cancellation = %d, want 3", got) } repo.afterGet = nil var wg sync.WaitGroup errors := make(chan error, 8) for index := 0; index < cap(errors); index++ { wg.Add(1) go func() { defer wg.Done() resolved, err := resolver.GetProfile(context.Background(), "leaf") if err != nil { errors <- err return } resolved.ExtraParams["nested"].(map[string]any)["value"] = "mutated" }() } wg.Wait() close(errors) for err := range errors { t.Errorf("concurrent resolution: %v", err) } latest, err := resolver.GetProfile(context.Background(), "leaf") if err != nil || latest.ExtraParams["nested"].(map[string]any)["value"] != "second" { t.Fatalf("latest result=(%+v, %v)", latest, err) } } type resolvingTestRepository struct { mu sync.Mutex profiles map[string]*domain.ExecutionProfile errors map[string]error calls map[string]int afterGet func(string) } func (r *resolvingTestRepository) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) { if err := ctx.Err(); err != nil { return nil, err } r.mu.Lock() if r.calls == nil { r.calls = make(map[string]int) } r.calls[id]++ err := r.errors[id] profile := r.profiles[id] afterGet := r.afterGet r.mu.Unlock() if afterGet != nil { afterGet(id) } if err != nil { return nil, err } if profile == nil { if _, exists := r.profiles[id]; exists { return nil, nil } return nil, ErrProfileNotFound } copy := *profile return ©, nil } func (r *resolvingTestRepository) set(id string, profile *domain.ExecutionProfile) { r.mu.Lock() defer r.mu.Unlock() r.profiles[id] = profile } func (r *resolvingTestRepository) callCount(id string) int { r.mu.Lock() defer r.mu.Unlock() return r.calls[id] } func profileTestFS(files map[string]string) fs.FS { fsys := make(fstest.MapFS, len(files)) for name, content := range files { fsys[name] = profileMapFile(content) } return fsys }