diff --git a/kv/internal/resolve/memo.go b/kv/internal/resolve/memo.go new file mode 100644 index 00000000..08df07d0 --- /dev/null +++ b/kv/internal/resolve/memo.go @@ -0,0 +1,50 @@ +package resolve + +import ( + "context" + "sync" +) + +// memoResult is a memoized resolution outcome (value or error). +type memoResult struct { + val string + err error +} + +// memo is the per-ResolveAll-call resolution cache. +type memo struct { + m map[refKey]memoResult + mu sync.Mutex +} + +func (mm *memo) get(k refKey) (memoResult, bool) { + mm.mu.Lock() + defer mm.mu.Unlock() + + mr, ok := mm.m[k] + + return mr, ok +} + +func (mm *memo) set(k refKey, v memoResult) { + mm.mu.Lock() + defer mm.mu.Unlock() + + mm.m[k] = v +} + +type memoCtxKey struct{} + +// withMemo attaches a fresh, per-call resolution memo to ctx. +func withMemo(ctx context.Context) context.Context { + return context.WithValue(ctx, memoCtxKey{}, &memo{m: make(map[refKey]memoResult)}) +} + +// memoFrom returns the per-call memo, or nil when ctx carries none (e.g. a +// direct Resolve call outside ResolveAll), in which case fetches are not +// memoized. +func memoFrom(ctx context.Context) *memo { + //nolint:errcheck + m, _ := ctx.Value(memoCtxKey{}).(*memo) + return m +} diff --git a/kv/internal/resolve/refs.go b/kv/internal/resolve/refs.go new file mode 100644 index 00000000..92f14b89 --- /dev/null +++ b/kv/internal/resolve/refs.go @@ -0,0 +1,119 @@ +package resolve + +import ( + "fmt" + "strings" +) + +// refKey identifies a resolution target. +type refKey struct { + store string + path string + fragment string +} + +// parseWholeValue parses a whole-value reference of the form +// "kv://store/path#frag". +// +// The bool reports whether input is a whole-value reference at all (i.e. has +// the "kv://" prefix). When false, input is an inline/literal string and err is +// nil. When true, err is non-nil if the reference is malformed. +func parseWholeValue(input string) (refKey, bool, error) { + if !strings.HasPrefix(input, "kv://") { + return refKey{}, false, nil + } + + trimmed := strings.TrimPrefix(input, "kv://") + + slashIdx := strings.IndexByte(trimmed, '/') + if slashIdx < 0 { + return refKey{}, true, fmt.Errorf( + "%w: missing path separator in %q", + ErrMalformedReference, + input, + ) + } + + store := trimmed[:slashIdx] + path, fragment, _ := strings.Cut(trimmed[slashIdx+1:], "#") + + if store == "" || path == "" { + return refKey{}, true, fmt.Errorf( + "%w: empty store name or path in %q", + ErrMalformedReference, + input, + ) + } + + return refKey{store: store, path: path, fragment: fragment}, true, nil +} + +// parseInlineToken parses a single inline token. match is the full "$kv{...}" +// text (used verbatim in error messages); its contents are "store:path#fragment". +func parseInlineToken(match string) (refKey, error) { + // strip "$kv{" prefix and "}" suffix + inner := match[len("$kv{") : len(match)-1] + + colonIdx := strings.IndexByte(inner, ':') + if colonIdx < 0 { + return refKey{}, fmt.Errorf( + "%w: missing store separator in %q", + ErrMalformedReference, + match, + ) + } + + store := inner[:colonIdx] + path, fragment, _ := strings.Cut(inner[colonIdx+1:], "#") + + if store == "" || path == "" { + return refKey{}, fmt.Errorf( + "%w: empty store name or path in %q", + ErrMalformedReference, + match, + ) + } + + return refKey{store: store, path: path, fragment: fragment}, nil +} + +// collectRefs walks a decoded JSON document and returns every distinct, +// well-formed reference target. It is best-effort: malformed references are +// skipped here. +func collectRefs(node any) map[refKey]struct{} { + refs := make(map[refKey]struct{}) + collectInto(node, refs) + + return refs +} + +func collectInto(node any, into map[refKey]struct{}) { + switch v := node.(type) { + case string: + collectRefsFromString(v, into) + case map[string]any: + for _, value := range v { + collectInto(value, into) + } + case []any: + for _, value := range v { + collectInto(value, into) + } + } +} + +func collectRefsFromString(input string, into map[refKey]struct{}) { + if ref, ok, err := parseWholeValue(input); ok { + if err == nil { + into[ref] = struct{}{} + } + + return + } + + for _, match := range inlineRe.FindAllString(input, -1) { + if ref, err := parseInlineToken(match); err == nil { + into[ref] = struct{}{} + } + } +} diff --git a/kv/internal/resolve/refs_test.go b/kv/internal/resolve/refs_test.go new file mode 100644 index 00000000..41856ced --- /dev/null +++ b/kv/internal/resolve/refs_test.go @@ -0,0 +1,163 @@ +package resolve + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestParseWholeValue(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + input string + wantOK bool + wantErr bool + want refKey + }{ + { + name: "path only", + input: "kv://vault/secret/data/app", + wantOK: true, + want: refKey{store: "vault", path: "secret/data/app"}, + }, + { + name: "with fragment", + input: "kv://vault/secret/data/app#password", + wantOK: true, + want: refKey{store: "vault", path: "secret/data/app", fragment: "password"}, + }, + { + name: "fragment splits on first hash only", + input: "kv://vault/a/b#x/y#z", + wantOK: true, + want: refKey{store: "vault", path: "a/b", fragment: "x/y#z"}, + }, + { + name: "not a whole-value reference", + input: "prefix-$kv{vault:secret}-suffix", + wantOK: false, + }, + { + name: "plain literal", + input: "just-a-string", + wantOK: false, + }, + { + name: "missing path separator", + input: "kv://vault", + wantOK: true, + wantErr: true, + }, + { + name: "empty store", + input: "kv:///secret/data/app", + wantOK: true, + wantErr: true, + }, + { + name: "empty path", + input: "kv://vault/", + wantOK: true, + wantErr: true, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + ref, ok, err := parseWholeValue(tc.input) + + require.Equal(t, tc.wantOK, ok) + + if tc.wantErr { + require.ErrorIs(t, err, ErrMalformedReference) + return + } + + require.NoError(t, err) + require.Equal(t, tc.want, ref) + }) + } +} + +func TestParseInlineToken(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + match string + wantErr bool + want refKey + }{ + { + name: "path only", + match: "$kv{vault:secret/data/app}", + want: refKey{store: "vault", path: "secret/data/app"}, + }, + { + name: "with fragment", + match: "$kv{vault:secret/data/app#password}", + want: refKey{store: "vault", path: "secret/data/app", fragment: "password"}, + }, + { + name: "missing store separator", + match: "$kv{no-colon-here}", + wantErr: true, + }, + { + name: "empty store", + match: "$kv{:secret}", + wantErr: true, + }, + { + name: "empty path", + match: "$kv{vault:}", + wantErr: true, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + ref, err := parseInlineToken(tc.match) + + if tc.wantErr { + require.ErrorIs(t, err, ErrMalformedReference) + return + } + + require.NoError(t, err) + require.Equal(t, tc.want, ref) + }) + } +} + +func TestCollectRefsMatchesResolverParsing(t *testing.T) { + t.Parallel() + + doc := map[string]any{ + "whole": "kv://vault/secret/data/app#password", + "inline": "postgres://$kv{vault:db/creds#host}:$kv{env:DB_PORT}/prod", + "dupe": "kv://vault/secret/data/app#password", // same target as "whole" + "malformed": "kv://vault", // no path — skipped + "literal": "no references here", + "nested": map[string]any{ + "list": []any{"kv://consul/services/web", "plain"}, + }, + } + + got := collectRefs(doc) + + want := map[refKey]struct{}{ + {store: "vault", path: "secret/data/app", fragment: "password"}: {}, + {store: "vault", path: "db/creds", fragment: "host"}: {}, + {store: "env", path: "DB_PORT"}: {}, + {store: "consul", path: "services/web"}: {}, + } + + require.Equal(t, want, got) +} diff --git a/kv/internal/resolve/resolver.go b/kv/internal/resolve/resolver.go index ba5cde92..14d42600 100644 --- a/kv/internal/resolve/resolver.go +++ b/kv/internal/resolve/resolver.go @@ -9,9 +9,15 @@ import ( "regexp" "strings" + "golang.org/x/sync/errgroup" + "github.com/TykTechnologies/storage/kv" ) +// maxConcurrentResolves bounds how many references are fetched at once during +// the prefetch phase. +const maxConcurrentResolves = 16 + type Resolver struct { registry kv.StoreGetter lenient bool @@ -51,42 +57,43 @@ func NewResolver(registry kv.StoreGetter, opts ...Option) *Resolver { var inlineRe = regexp.MustCompile(`\$kv\{([^}]+)\}`) func (r *Resolver) Resolve(ctx context.Context, input string) (string, error) { - if strings.HasPrefix(input, "kv://") { - trimmed := strings.TrimPrefix(input, "kv://") - - slashIdx := strings.IndexByte(trimmed, '/') - if slashIdx < 0 { - return "", fmt.Errorf( - "%w: missing path separator in %q", - ErrMalformedReference, - input, - ) - } + ref, ok, err := parseWholeValue(input) + if err != nil { + return "", err + } - storeName := trimmed[:slashIdx] - rest := trimmed[slashIdx+1:] - path, fragment, _ := strings.Cut(rest, "#") + if ok { + return r.resolveRef(ctx, ref, input) + } - if storeName == "" || path == "" { - return "", fmt.Errorf( - "%w: empty store name or path in %q", - ErrMalformedReference, - input, - ) - } + return r.resolveInline(ctx, input) +} - res, err := r.fetchAndExtract(ctx, storeName, path, fragment) +// resolveRef resolves a single reference. literal is the original text to emit +// unchanged when lenient mode tolerates a missing store — the whole input for a +// kv:// reference, or the matched token for a $kv{} one. It is the single place +// the lenient store-not-found rule lives. +func (r *Resolver) resolveRef(ctx context.Context, ref refKey, literal string) (string, error) { + val, err := r.fetchAndExtract(ctx, ref.store, ref.path, ref.fragment) + if err != nil { if r.lenient && errors.Is(err, kv.ErrStoreNotFound) { - return input, nil + return literal, nil } - return res, err + return "", err } - // The token regex requires a closing brace, so an unclosed "$kv{" can - // never match — without this check a typo'd reference would silently pass - // through as a literal value. - if idx := unclosedInlineToken(input); idx >= 0 { + return val, nil +} + +// resolveInline replaces every $kv{...} token in input. A malformed or +// unresolvable token is left in place and its error collected; all failures are +// returned joined so the caller sees every problem at once. +func (r *Resolver) resolveInline(ctx context.Context, input string) (string, error) { + // The token regex requires a closing brace, so an unclosed "$kv{" can never + // match — without this check a typo'd reference would silently pass through + // as a literal value. + if unclosedInlineToken(input) >= 0 { return "", fmt.Errorf( "%w: unclosed $kv{ reference in %q", ErrMalformedReference, @@ -94,43 +101,19 @@ func (r *Resolver) Resolve(ctx context.Context, input string) (string, error) { ) } - var resolveErrs []error - result := inlineRe.ReplaceAllStringFunc(input, func(match string) string { - // strip "$kv{" prefix and "}" suffix - inner := match[4 : len(match)-1] - - colonIdx := strings.IndexByte(inner, ':') - if colonIdx < 0 { - resolveErrs = append(resolveErrs, fmt.Errorf( - "%w: missing store separator in %q", - ErrMalformedReference, - match, - )) + var errs []error - return match - } - - storeName := inner[:colonIdx] - rest := inner[colonIdx+1:] - path, fragment, _ := strings.Cut(rest, "#") - - if storeName == "" || path == "" { - resolveErrs = append(resolveErrs, fmt.Errorf( - "%w: empty store name or path in %q", - ErrMalformedReference, - match, - )) + result := inlineRe.ReplaceAllStringFunc(input, func(match string) string { + ref, err := parseInlineToken(match) + if err != nil { + errs = append(errs, err) return match } - val, err := r.fetchAndExtract(ctx, storeName, path, fragment) + val, err := r.resolveRef(ctx, ref, match) if err != nil { - if r.lenient && errors.Is(err, kv.ErrStoreNotFound) { - return match - } - - resolveErrs = append(resolveErrs, err) + errs = append(errs, err) return match } @@ -138,8 +121,8 @@ func (r *Resolver) Resolve(ctx context.Context, input string) (string, error) { return val }) - if len(resolveErrs) > 0 { - return "", errors.Join(resolveErrs...) + if len(errs) > 0 { + return "", errors.Join(errs...) } return result, nil @@ -169,6 +152,12 @@ func (r *Resolver) ResolveAll(ctx context.Context, rawJSON []byte) ([]byte, erro return nil, fmt.Errorf("%w: %w", ErrInvalidJSON, err) } + ctx = withMemo(ctx) + + // Prefetch every distinct reference concurrently into the memo, so the + // sequential substitution walk below reads them without further I/O. + r.prefetch(ctx, doc) + resolved, err := r.walkAndResolve(ctx, doc) if err != nil { return nil, err @@ -187,7 +176,56 @@ func (r *Resolver) ResolveAll(ctx context.Context, rawJSON []byte) ([]byte, erro return bytes.TrimSuffix(buf.Bytes(), []byte("\n")), nil } +// prefetch resolves every distinct reference in doc concurrently, warming the +// per-call memo so the substitution walk reads them without further I/O. +// +// It is best-effort and deliberately non-cancelling. Each goroutine returns nil +// even on failure — to keep one reference's failure from poisoning the others. +func (r *Resolver) prefetch(ctx context.Context, doc any) { + refs := collectRefs(doc) + + if len(refs) <= 1 { + return + } + + g, gctx := errgroup.WithContext(ctx) + g.SetLimit(maxConcurrentResolves) + + for ref := range refs { + g.Go(func() error { + //nolint:errcheck + _, _ = r.fetchAndExtract(gctx, ref.store, ref.path, ref.fragment) + return nil + }) + } + + //nolint:errcheck + _ = g.Wait() +} + +// fetchAndExtract resolves a single target, consulting the per-call memo (when +// ctx carries one) so repeated references — in either syntax form — resolve the +// backend at most once per document. func (r *Resolver) fetchAndExtract(ctx context.Context, storeName, path, fragment string) (string, error) { + mm := memoFrom(ctx) + + key := refKey{store: storeName, path: path, fragment: fragment} + if mm != nil { + if hit, ok := mm.get(key); ok { + return hit.val, hit.err + } + } + + val, err := r.fetch(ctx, storeName, path, fragment) + + if mm != nil { + mm.set(key, memoResult{val: val, err: err}) + } + + return val, err +} + +func (r *Resolver) fetch(ctx context.Context, storeName, path, fragment string) (string, error) { store, err := r.registry.GetStore(storeName) if err != nil { return "", err diff --git a/kv/internal/resolve/resolver_bench_test.go b/kv/internal/resolve/resolver_bench_test.go new file mode 100644 index 00000000..37e9e201 --- /dev/null +++ b/kv/internal/resolve/resolver_bench_test.go @@ -0,0 +1,174 @@ +package resolve_test + +import ( + "context" + "encoding/json" + "fmt" + "strings" + "testing" + "time" + + "github.com/TykTechnologies/storage/kv" + "github.com/TykTechnologies/storage/kv/internal/resolve" + "github.com/stretchr/testify/require" +) + +func BenchmarkResolve_ThreeInlineTokensWithFragment(b *testing.B) { + payload := `{"host":"db.internal","port":"5432"}` + getter := newGetter(map[string]kv.Provider{ + "vault": &mockProvider{value: payload}, + "env": &mockProvider{value: "simple-value"}, + }) + r := resolve.NewResolver(getter) + input := "postgres://$kv{vault:db/creds#host}:$kv{vault:db/creds#port}/$kv{env:DB_NAME}" + ctx := context.Background() + + b.ResetTimer() + + for b.Loop() { + _, err := r.Resolve(ctx, input) + if err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkResolveAll measures the whole-document resolution cost the caller +// pays at startup as a function of reference count. Providers are +// in-memory mocks, so the numbers isolate the library's own overhead +// — parse, walk, substitute, re-serialize — from backend latency. +func BenchmarkResolveAll(b *testing.B) { + getter := newGetter(map[string]kv.Provider{ + "env": &mockProvider{value: "resolved-value"}, + "vault": &mockProvider{value: `{"username":"admin","password":"hunter2"}`}, + }) + r := resolve.NewResolver(getter) + ctx := context.Background() + + for _, n := range []int{20, 50, 100} { + doc := buildBenchDoc(b, n) + + b.Run(fmt.Sprintf("refs=%d", n), func(b *testing.B) { + for b.Loop() { + if _, err := r.ResolveAll(ctx, doc); err != nil { + b.Fatal(err) + } + } + }) + } +} + +const benchLatency = time.Millisecond + +func BenchmarkResolveAll_DuplicateRefs(b *testing.B) { + for _, count := range []int{1, 10, 100, 1000, 10000} { + b.Run(fmt.Sprintf("count=%d", count), func(b *testing.B) { + provider := &latencyProvider{value: "s3cr3t", latency: benchLatency} + r := resolve.NewResolver(newGetter(map[string]kv.Provider{"vault": provider})) + doc := buildDoc("kv://vault/secret/data/app", count) + + ctx := context.Background() + + b.ReportAllocs() + b.ResetTimer() + + for i := 0; i < b.N; i++ { + if _, err := r.ResolveAll(ctx, doc); err != nil { + b.Fatal(err) + } + } + }) + } +} + +func BenchmarkResolveAll_DistinctRefs(b *testing.B) { + for _, count := range []int{1, 16, 64, 256} { + b.Run(fmt.Sprintf("count=%d", count), func(b *testing.B) { + provider := &latencyProvider{value: "v", latency: benchLatency} + r := resolve.NewResolver(newGetter(map[string]kv.Provider{"vault": provider})) + doc := distinctRefsDoc(count) + + ctx := context.Background() + + b.ReportAllocs() + b.ResetTimer() + + for i := 0; i < b.N; i++ { + if _, err := r.ResolveAll(ctx, doc); err != nil { + b.Fatal(err) + } + } + }) + } +} + +func BenchmarkResolveAll_NoRefs(b *testing.B) { + provider := &latencyProvider{value: "unused", latency: benchLatency} + r := resolve.NewResolver(newGetter(map[string]kv.Provider{"vault": provider})) + + fields := make([]string, 1000) + for i := range fields { + fields[i] = fmt.Sprintf(`"api_%d":{"listen_path":"/svc-%d/","target":"http://upstream-%d.internal"}`, i, i, i) + } + + doc := []byte("{" + strings.Join(fields, ",") + "}") + ctx := context.Background() + + b.ReportAllocs() + b.ResetTimer() + + for i := 0; i < b.N; i++ { + if _, err := r.ResolveAll(ctx, doc); err != nil { + b.Fatal(err) + } + } +} + +type latencyProvider struct { + value string + latency time.Duration +} + +func (p *latencyProvider) Get(ctx context.Context, _ string) (string, error) { + select { + case <-time.After(p.latency): + return p.value, nil + case <-ctx.Done(): + return "", ctx.Err() + } +} + +func buildBenchDoc(b *testing.B, n int) []byte { + b.Helper() + + fields := make(map[string]any, n*2) + + for i := 0; i < n; i++ { + key := fmt.Sprintf("field_%d", i) + + switch i % 3 { + case 0: + fields[key] = "kv://env/SOME_KEY" + case 1: + fields[key] = fmt.Sprintf("https://$kv{env:HOST_%d}/v1", i) + case 2: + fields[key] = "kv://vault/db/creds#password" + } + + fields[fmt.Sprintf("plain_%d", i)] = "no reference here" + } + + doc, err := json.Marshal(map[string]any{"config": fields}) + require.NoError(b, err) + + return doc +} + +func distinctRefsDoc(count int) []byte { + fields := make([]string, count) + for i := range fields { + fields[i] = fmt.Sprintf(`"h%d":"kv://vault/secret/%d"`, i, i) + } + + return []byte("{" + strings.Join(fields, ",") + "}") +} diff --git a/kv/internal/resolve/resolver_test.go b/kv/internal/resolve/resolver_test.go index 20bb99b1..198bf84a 100644 --- a/kv/internal/resolve/resolver_test.go +++ b/kv/internal/resolve/resolver_test.go @@ -4,7 +4,11 @@ import ( "context" "encoding/json" "fmt" + "strings" + "sync" + "sync/atomic" "testing" + "time" "github.com/TykTechnologies/storage/kv" "github.com/TykTechnologies/storage/kv/internal/resolve" @@ -12,34 +16,6 @@ import ( "github.com/stretchr/testify/require" ) -type mockProvider struct { - value string - err error - lastKey string -} - -func (m *mockProvider) Get(_ context.Context, key string) (string, error) { - m.lastKey = key - return m.value, m.err -} - -type mockStoreGetter struct { - stores map[string]kv.Provider -} - -func (m *mockStoreGetter) GetStore(name string) (kv.Provider, error) { - p, ok := m.stores[name] - if !ok { - return nil, kv.NewStoreNotFoundError(name) - } - - return p, nil -} - -func newGetter(stores map[string]kv.Provider) kv.StoreGetter { - return &mockStoreGetter{stores: stores} -} - func TestResolve(t *testing.T) { t.Parallel() @@ -629,77 +605,250 @@ func TestResolveAll_NoHTMLEscaping(t *testing.T) { assert.NotContains(t, string(got), `\u003c`) } -// -------------------------------------------------------------------------- -// Benchmarks -// -------------------------------------------------------------------------- +func TestResolveAll_DeduplicatesIdenticalReferences(t *testing.T) { + const n = 100 -func BenchmarkResolve_ThreeInlineTokensWithFragment(b *testing.B) { - payload := `{"host":"db.internal","port":"5432"}` - getter := newGetter(map[string]kv.Provider{ - "vault": &mockProvider{value: payload}, - "env": &mockProvider{value: "simple-value"}, - }) - r := resolve.NewResolver(getter) - input := "postgres://$kv{vault:db/creds#host}:$kv{vault:db/creds#port}/$kv{env:DB_NAME}" - ctx := context.Background() + provider := newCountingProvider("s3cr3t") + r := resolve.NewResolver(newGetter(map[string]kv.Provider{"vault": provider})) + + // Whole-value ref with no #fragment, so the raw value is returned as-is + // (no JSON extraction needed for the mock value). + doc := buildDoc("kv://vault/secret/data/app", n) - b.ResetTimer() + _, err := r.ResolveAll(t.Context(), doc) + require.NoError(t, err) + + require.Equal(t, 1, provider.callsFor("secret/data/app"), + "identical references in one document must hit the backend once, not per-occurrence") + require.EqualValues(t, 1, provider.total.Load(), + "total backend calls must equal the number of UNIQUE references (1), not occurrences (%d)", n) +} + +func TestResolveAll_DistinctReferencesEachResolvedOnce(t *testing.T) { + const repeatEach = 10 - for b.Loop() { - _, err := r.Resolve(ctx, input) - if err != nil { - b.Fatal(err) + provider := newCountingProvider("v") + r := resolve.NewResolver(newGetter(map[string]kv.Provider{"vault": provider})) + + paths := []string{"secret/a", "secret/b", "secret/c"} + + var fields []string + + for i := 0; i < repeatEach; i++ { + for j, p := range paths { + fields = append(fields, fmt.Sprintf(`"h%d_%d":"kv://vault/%s"`, i, j, p)) } } + + doc := []byte("{" + strings.Join(fields, ",") + "}") + + _, err := r.ResolveAll(t.Context(), doc) + require.NoError(t, err) + + for _, p := range paths { + require.Equal(t, 1, provider.callsFor(p), + "each distinct reference must resolve exactly once (path %q)", p) + } + + require.EqualValues(t, len(paths), provider.total.Load(), + "total backend calls must equal the number of unique references") } -// BenchmarkResolveAll measures the whole-document resolution cost the caller -// pays at startup as a function of reference count. Providers are -// in-memory mocks, so the numbers isolate the library's own overhead -// — parse, walk, substitute, re-serialize — from backend latency. -func BenchmarkResolveAll(b *testing.B) { - getter := newGetter(map[string]kv.Provider{ - "env": &mockProvider{value: "resolved-value"}, - "vault": &mockProvider{value: `{"username":"admin","password":"hunter2"}`}, - }) - r := resolve.NewResolver(getter) - ctx := context.Background() +func TestResolveAll_DedupsAcrossSyntaxForms(t *testing.T) { + provider := newCountingProvider("s3cr3t") + r := resolve.NewResolver(newGetter(map[string]kv.Provider{"vault": provider})) - for _, n := range []int{20, 50, 100} { - doc := buildBenchDoc(b, n) + doc := []byte(`{ + "whole": "kv://vault/secret/data/app", + "inline": "prefix-$kv{vault:secret/data/app}-suffix" + }`) - b.Run(fmt.Sprintf("refs=%d", n), func(b *testing.B) { - for b.Loop() { - if _, err := r.ResolveAll(ctx, doc); err != nil { - b.Fatal(err) - } - } - }) + _, err := r.ResolveAll(context.Background(), doc) + require.NoError(t, err) + + require.EqualValues(t, int32(1), provider.total.Load(), + "the same target via different syntax forms must resolve the backend once") +} + +// mutableProvider returns whatever value it currently holds, and counts calls. +// It lets a test change the backing secret between resolutions. +type mutableProvider struct { + mu sync.Mutex + value string + calls int +} + +func (m *mutableProvider) Get(_ context.Context, _ string) (string, error) { + m.mu.Lock() + defer m.mu.Unlock() + + m.calls++ + + return m.value, nil +} + +func (m *mutableProvider) set(v string) { + m.mu.Lock() + defer m.mu.Unlock() + m.value = v +} + +func TestResolveAll_MemoIsPerCallNotPersistent(t *testing.T) { + provider := &mutableProvider{value: "old"} + r := resolve.NewResolver(newGetter(map[string]kv.Provider{"vault": provider})) + + doc := []byte(`{"a":"kv://vault/secret/data/app","b":"kv://vault/secret/data/app"}`) + + first, err := r.ResolveAll(t.Context(), doc) + require.NoError(t, err) + require.JSONEq(t, `{"a":"old","b":"old"}`, string(first)) + require.Equal(t, 1, provider.calls, "duplicate refs collapse to one backend call within a call") + + // Secret rotates between reloads. + provider.set("new") + + second, err := r.ResolveAll(t.Context(), doc) + require.NoError(t, err) + require.JSONEq(t, `{"a":"new","b":"new"}`, string(second), + "a later ResolveAll must see the rotated value, not a memo from the previous call") + require.Equal(t, 2, provider.calls, + "the second call must re-hit the backend — the memo must not persist across calls") +} + +func TestResolveAll_ResolvesDistinctReferencesConcurrently(t *testing.T) { + const n = 8 + + provider := newBarrierProvider(n, "v") + r := resolve.NewResolver(newGetter(map[string]kv.Provider{"vault": provider})) + + // n DISTINCT references so dedup does not collapse them — we want n + // independent fetches that can only complete if run concurrently. + fields := make([]string, n) + for i := range fields { + fields[i] = fmt.Sprintf(`"h%d":"kv://vault/secret/%d"`, i, i) } + + doc := []byte("{" + strings.Join(fields, ",") + "}") + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + _, err := r.ResolveAll(ctx, doc) + require.NoError(t, err, + "distinct references must resolve concurrently; serial resolution never reaches the barrier and times out") + require.Equal(t, n, provider.peakConcurrency(), + "all %d distinct fetches must be in flight simultaneously", n) +} + +type mockProvider struct { + value string + err error + lastKey string } -func buildBenchDoc(b *testing.B, n int) []byte { - b.Helper() +func (m *mockProvider) Get(_ context.Context, key string) (string, error) { + m.lastKey = key + return m.value, m.err +} - fields := make(map[string]any, n*2) +type mockStoreGetter struct { + stores map[string]kv.Provider +} - for i := 0; i < n; i++ { - key := fmt.Sprintf("field_%d", i) +func (m *mockStoreGetter) GetStore(name string) (kv.Provider, error) { + p, ok := m.stores[name] + if !ok { + return nil, kv.NewStoreNotFoundError(name) + } - switch i % 3 { - case 0: - fields[key] = "kv://env/SOME_KEY" - case 1: - fields[key] = fmt.Sprintf("https://$kv{env:HOST_%d}/v1", i) - case 2: - fields[key] = "kv://vault/db/creds#password" - } + return p, nil +} - fields[fmt.Sprintf("plain_%d", i)] = "no reference here" +func newGetter(stores map[string]kv.Provider) kv.StoreGetter { + return &mockStoreGetter{stores: stores} +} + +type countingProvider struct { + value string + mu sync.Mutex + calls map[string]int + total atomic.Int32 +} + +func newCountingProvider(value string) *countingProvider { + return &countingProvider{value: value, calls: map[string]int{}} +} + +func (c *countingProvider) Get(_ context.Context, key string) (string, error) { + c.total.Add(1) + + c.mu.Lock() + c.calls[key]++ + c.mu.Unlock() + + return c.value, nil +} + +func (c *countingProvider) callsFor(key string) int { + c.mu.Lock() + defer c.mu.Unlock() + + return c.calls[key] +} + +// barrierProvider blocks every Get until exactly n calls are simultaneously +// in-flight, then releases them all. +type barrierProvider struct { + n int + value string + + mu sync.Mutex + arrived int + peak int + + gate chan struct{} + once sync.Once +} + +func newBarrierProvider(n int, value string) *barrierProvider { + return &barrierProvider{n: n, value: value, gate: make(chan struct{})} +} + +func (b *barrierProvider) Get(ctx context.Context, _ string) (string, error) { + b.mu.Lock() + b.arrived++ + + if b.arrived > b.peak { + b.peak = b.arrived + } + + reached := b.arrived >= b.n + b.mu.Unlock() + + if reached { + b.once.Do(func() { close(b.gate) }) + } + + select { + case <-b.gate: + return b.value, nil + case <-ctx.Done(): + return "", ctx.Err() } +} - doc, err := json.Marshal(map[string]any{"config": fields}) - require.NoError(b, err) +func (b *barrierProvider) peakConcurrency() int { + b.mu.Lock() + defer b.mu.Unlock() + + return b.peak +} + +func buildDoc(ref string, count int) []byte { + fields := make([]string, count) + for i := range fields { + fields[i] = fmt.Sprintf(`"h%d":%q`, i, ref) + } - return doc + return []byte("{" + strings.Join(fields, ",") + "}") }