From 75396918e876667e157836d53a28678711b7497f Mon Sep 17 00:00:00 2001 From: TANTIOPE Date: Sat, 25 Jul 2026 21:24:24 +0200 Subject: [PATCH 1/2] Reuse tool embeddings across sessions The optimizer re-embeds the whole tool set on every client session, so tools/list blocks on a full index build each time a client connects. At 140 aggregated tools against a CPU embedding backend that is 16-19s per connect, and under concurrency the redundant rebuilds queue until sessions fail. THV-0022 describes the store as a regenerable cache whose cold-start cost falls on the first session after a pod restart. This restores that: an embedding is reused when the tool's embedded text and the embedding backend are both unchanged, keyed on a hash of the text plus the backend identity. Because vectors now outlive a single build, two things follow. Stored vectors whose width differs from the current backend's are skipped in search rather than compared, since cosine distance indexes both slices positionally. And a fixed probe string is re-embedded on each build and compared with the stored one, because neither the content hash nor the vector width can observe a same-width model swap behind an unchanged service URL. Fixes #5847 Signed-off-by: TANTIOPE --- .../optimizer/internal/toolstore/schema.sql | 36 +- .../internal/toolstore/sqlite_store.go | 406 +++++++- .../toolstore/sqlite_store_cache_test.go | 969 ++++++++++++++++++ .../toolstore/sqlite_store_livemodel_test.go | 253 +++++ pkg/vmcp/server/serve_optimizer_live_test.go | 174 ++++ 5 files changed, 1818 insertions(+), 20 deletions(-) create mode 100644 pkg/vmcp/optimizer/internal/toolstore/sqlite_store_cache_test.go create mode 100644 pkg/vmcp/optimizer/internal/toolstore/sqlite_store_livemodel_test.go create mode 100644 pkg/vmcp/server/serve_optimizer_live_test.go diff --git a/pkg/vmcp/optimizer/internal/toolstore/schema.sql b/pkg/vmcp/optimizer/internal/toolstore/schema.sql index 94e19aeb92..48029efb37 100644 --- a/pkg/vmcp/optimizer/internal/toolstore/schema.sql +++ b/pkg/vmcp/optimizer/internal/toolstore/schema.sql @@ -1,11 +1,43 @@ -- SPDX-FileCopyrightText: Copyright 2025 Stacklok, Inc. -- SPDX-License-Identifier: Apache-2.0 --- Capabilities table stores tool/resource/prompt metadata +-- Capabilities table stores tool/resource/prompt metadata. +-- +-- content_hash identifies the exact input the stored embedding was produced +-- from, covering both the embedded text and the embedding backend (see +-- embeddingCacheKey), so UpsertTools can reuse a vector when the hash still +-- matches. NULL means no embedding (FTS5-only mode). +-- +-- The database is recreated in memory on every process start, so this column +-- needs no migration. That changes if the store ever becomes file-backed. CREATE TABLE IF NOT EXISTS llm_capabilities ( name TEXT PRIMARY KEY, description TEXT NOT NULL DEFAULT '', - embedding BLOB + embedding BLOB, + content_hash TEXT +); + +-- The reuse probe selects by content_hash across the whole table, not by the +-- name primary key. +CREATE INDEX IF NOT EXISTS llm_capabilities_content_hash_idx + ON llm_capabilities (content_hash); + +-- Embedding of a fixed probe string, used to detect that the embedding backend +-- started returning different vectors. +-- +-- The content hash covers the configured provider, endpoint and model, but the +-- TEI model is fixed by the running container rather than by config, so +-- swapping it behind an unchanged service URL is invisible to the hash. When +-- the replacement has the same vector width the dimension check cannot see it +-- either, and the stored vectors would be reused across a change of semantic +-- space. Comparing a re-embedded probe against the stored one detects that +-- regardless of what the configuration says. +-- +-- Single row by construction: the probe describes the one backend this store +-- talks to. +CREATE TABLE IF NOT EXISTS embedding_canary ( + id INTEGER PRIMARY KEY CHECK (id = 1), + embedding BLOB NOT NULL ); -- FTS5 virtual table for full-text search with BM25 ranking. diff --git a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store.go b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store.go index acb8a154a4..2cbd8f5807 100644 --- a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store.go +++ b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store.go @@ -8,9 +8,11 @@ package toolstore import ( "context" + "crypto/sha256" "database/sql" _ "embed" "encoding/binary" + "encoding/hex" "encoding/json" "errors" "fmt" @@ -18,6 +20,8 @@ import ( "math" "sort" "strings" + "sync" + "sync/atomic" "golang.org/x/sync/errgroup" _ "modernc.org/sqlite" // registers the "sqlite" database/sql driver @@ -77,6 +81,30 @@ type sqliteToolStore struct { maxToolsToReturn int hybridSemanticRatio float64 semanticDistanceThreshold float64 + + // embeddingIdentity describes the backend that produces embeddings. It is + // mixed into every content hash so vectors are never reused across a + // provider, endpoint, or model change. Immutable after construction. + embeddingIdentity string + + // embeddingDim is the vector width most recently observed from the + // embedding backend, or 0 before any embedding call has succeeded. It + // bounds which stored vectors may be reused, catching a model swap that + // embeddingIdentity cannot see (the TEI model is fixed by the running + // container, not by config). Held by pointer because the store is used by + // value. + embeddingDim *atomic.Int64 + + // canary serializes the backend probe so concurrent builds cannot race on + // the stored probe row. Held by pointer because the store is used by value. + canary *canaryState +} + +// canaryState serializes the backend probe and counts completed probes, so a +// build that waited through one can tell that its result already applies. +type canaryState struct { + mu sync.Mutex + generation atomic.Uint64 } // NewSQLiteToolStore creates a new ToolStore backed by a shared in-memory @@ -136,6 +164,9 @@ func newSQLiteToolStore( maxToolsToReturn: maxTools, hybridSemanticRatio: hybridRatio, semanticDistanceThreshold: semanticThreshold, + embeddingIdentity: embeddingIdentity(cfg), + embeddingDim: &atomic.Int64{}, + canary: &canaryState{}, } slog.Debug("optimizer tool store created", @@ -150,6 +181,17 @@ func newSQLiteToolStore( // UpsertTools adds or updates tools in the store. func (s sqliteToolStore) UpsertTools(ctx context.Context, tools []server.ServerTool) (retErr error) { + // Resolve embeddings before opening the write transaction, so no lock is + // held across the multi-second embedding call. SQLite in shared-cache mode + // takes table-level locks: a read inside the transaction would pin a read + // lock on llm_capabilities, and concurrent builds would then fail to take + // the write lock with SQLITE_LOCKED, which the busy handler does not retry. + // See TestSQLiteToolStore_UpsertTools_ConcurrentBuilds. + embBlobs, hashes, err := s.resolveEmbeddings(ctx, tools) + if err != nil { + return err + } + tx, err := s.db.BeginTx(ctx, nil) if err != nil { return fmt.Errorf("failed to begin transaction: %w", err) @@ -160,19 +202,15 @@ func (s sqliteToolStore) UpsertTools(ctx context.Context, tools []server.ServerT } }() - embBlobs, err := s.generateEmbeddings(ctx, tools) - if err != nil { - return err - } - - stmt, err := tx.PrepareContext(ctx, "INSERT OR REPLACE INTO llm_capabilities (name, description, embedding) VALUES (?, ?, ?)") + stmt, err := tx.PrepareContext(ctx, + "INSERT OR REPLACE INTO llm_capabilities (name, description, embedding, content_hash) VALUES (?, ?, ?, ?)") if err != nil { return fmt.Errorf("failed to prepare statement: %w", err) } defer func() { _ = stmt.Close() }() for i, tool := range tools { - if _, err := stmt.ExecContext(ctx, tool.Tool.Name, tool.Tool.Description, embBlobs[i]); err != nil { + if _, err := stmt.ExecContext(ctx, tool.Tool.Name, tool.Tool.Description, embBlobs[i], hashes[i]); err != nil { return fmt.Errorf("failed to upsert tool %s: %w", tool.Tool.Name, err) } } @@ -182,30 +220,342 @@ func (s sqliteToolStore) UpsertTools(ctx context.Context, tools []server.ServerT return tx.Commit() } -// generateEmbeddings produces encoded embedding blobs for each tool. -// If no embedding client is configured, it returns a slice of nil byte slices. -func (s sqliteToolStore) generateEmbeddings(ctx context.Context, tools []server.ServerTool) ([][]byte, error) { +// resolveEmbeddings returns an encoded embedding blob and a content hash for +// each tool, embedding only the tools whose hash is not already stored. +// +// The Serve path rebuilds a per-session optimizer on every session registration +// and every cross-pod rehydration, each upserting the session's whole tool set. +// Embedding all of it every time costs O(tools x sessions) on the client's +// initialize round-trip; reuse makes it O(tools whose text changed). See +// stacklok/toolhive#5847. +// +// With no embedding client it returns nil blobs and hashes (FTS5-only mode). +func (s sqliteToolStore) resolveEmbeddings( + ctx context.Context, tools []server.ServerTool, +) ([][]byte, []sql.NullString, error) { blobs := make([][]byte, len(tools)) + hashes := make([]sql.NullString, len(tools)) if s.embeddingClient == nil { - return blobs, nil + return blobs, hashes, nil } texts := make([]string, len(tools)) + keys := make([]string, len(tools)) for i, tool := range tools { - texts[i] = fmt.Sprintf("name: %s description: %s", tool.Tool.Name, tool.Tool.Description) + texts[i] = embeddedText(tool.Tool.Name, tool.Tool.Description) + keys[i] = embeddingCacheKey(s.embeddingIdentity, texts[i]) + hashes[i] = sql.NullString{String: keys[i], Valid: true} + } + + // Discards stored vectors first if the backend has changed under us, so the + // lookup below simply finds nothing to reuse for them. + s.syncBackendProbe(ctx) + + cached, err := s.cachedEmbeddings(ctx, keys) + if err != nil { + return nil, nil, err + } + + // Deduplicated by key: the same tool may appear twice in one batch. + missIndexByKey := make(map[string]int, len(keys)) + var missTexts []string + reused := 0 + for i, key := range keys { + if blob, ok := cached[key]; ok { + blobs[i] = blob + reused++ + continue + } + if _, seen := missIndexByKey[key]; !seen { + missIndexByKey[key] = len(missTexts) + missTexts = append(missTexts, texts[i]) + } + } + + slog.Debug("resolved tool embeddings", + "tools", len(tools), "reused", reused, "embedded", len(missTexts)) + + if len(missTexts) == 0 { + return blobs, hashes, nil + } + + embeddings, err := s.embeddingClient.EmbedBatch(ctx, missTexts) + if err != nil { + return nil, nil, fmt.Errorf("failed to generate embeddings: %w", err) + } + if len(embeddings) != len(missTexts) { + return nil, nil, fmt.Errorf("embedding client returned %d embeddings for %d inputs", + len(embeddings), len(missTexts)) + } + if n := len(embeddings[0]); n > 0 { + s.embeddingDim.Store(int64(n)) + } + + for i, key := range keys { + if blobs[i] != nil { + continue + } + idx, ok := missIndexByKey[key] + if !ok { + // Unreachable, but a missing key would otherwise index 0 and store + // another tool's vector under this name. + return nil, nil, fmt.Errorf("no embedding resolved for tool %s", tools[i].Tool.Name) + } + blobs[i] = encodeEmbedding(embeddings[idx]) + } + + return blobs, hashes, nil +} + +// canaryText is the fixed probe embedded to detect a change of embedding +// backend. Its content is arbitrary but must never change: an edit would make +// every stored canary incomparable and force one needless full re-embed. +const canaryText = "toolhive optimizer embedding canary v1" + +// canaryMaxDistance is the cosine distance below which two probe embeddings are +// considered to come from the same backend. +// +// The threshold is not zero because a backend is not required to be +// deterministic: reduction order can differ across hardware and runtimes, so the +// same model may return slightly different vectors on different deployments. +// It is small because the signal it must not miss is large — two different +// models of equal width place the same text roughly a full unit apart, i.e. +// effectively orthogonal, so this leaves two orders of magnitude of margin. +const canaryMaxDistance = 0.01 + +// syncBackendProbe discards stored embeddings when the embedding backend +// has started returning different vectors. +// +// It runs on every build, not once per store. The store lives for the whole +// process and the embedding service is addressed by a stable URL, so the backend +// can be replaced underneath a running server — redeploying the embedding +// service with a different model of the same width changes neither the URL nor +// the vector length, leaving it invisible to both the content hash and the +// dimension check. Probing once at startup would never see it, and reuse would +// then serve vectors from the old model for the rest of the process's life. +// Before embedding reuse existed the next session simply re-embedded everything +// and the condition healed on its own. +// +// Cost is one embedding per build, against the ~140 it saves. +// +// Every failure path leaves the stored embeddings reusable. A probe that cannot +// be taken says the backend is unreachable, not that it changed — and an +// unreachable backend cannot re-embed the catalogue either, so refusing reuse +// would turn a working build into a failed one. Serving possibly-stale vectors +// is bounded: the next build with a reachable backend re-probes and discards +// them. A build that still works beats a client with no tools. +func (s sqliteToolStore) syncBackendProbe(ctx context.Context) { + // Every build must be ordered after any in-flight probe, because a probe may + // be about to discard the very vectors this build is about to read and write + // back. Skipping the lock instead of waiting for it lets a build read + // pre-discard rows and re-insert them — restoring both the stale vector and + // its content hash — after which the freshly written probe certifies the + // store as current and nothing ever re-checks. That is permanent, not + // bounded. See TestSQLiteToolStore_ConcurrentBuilds_OrderedAfterProbe. + // + // Waiting is still cheap: a build that waited through someone else's probe + // sees the generation move and skips the network call, so a burst of + // concurrent builds costs one embedding between them, not one each. + gen := s.canary.generation.Load() + s.canary.mu.Lock() + defer s.canary.mu.Unlock() + if s.canary.generation.Load() != gen { + return + } + + // Embedded outside any transaction so no database lock is held across the + // network call (see the note in UpsertTools). + probe, err := s.embeddingClient.Embed(ctx, canaryText) + if err != nil { + slog.Warn("could not probe the embedding backend; reusing stored embeddings unverified", "error", err) + return + } + if len(probe) == 0 { + slog.Warn("embedding backend returned an empty probe; reusing stored embeddings unverified") + return + } + s.embeddingDim.Store(int64(len(probe))) + + changed, err := s.reconcileCanary(ctx, probe) + if err != nil { + slog.Warn("could not reconcile the embedding probe; reusing stored embeddings unverified", "error", err) + return + } + + s.canary.generation.Add(1) + if changed { + slog.Warn("embedding backend changed; stored embeddings were discarded and will be recomputed") + } +} + +// reconcileCanary compares probe against the stored canary, discarding every +// stored embedding when they differ, and records probe as the new canary. +// Reports whether stored embeddings were discarded. +// +// The discard and the new canary are written in one transaction so a failure +// cannot leave a canary that claims vectors are current when they are not. +func (s sqliteToolStore) reconcileCanary(ctx context.Context, probe []float32) (discarded bool, retErr error) { + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return false, fmt.Errorf("failed to begin canary transaction: %w", err) + } + defer func() { + if retErr != nil { + _ = tx.Rollback() + } + }() + + var storedBlob []byte + err = tx.QueryRowContext(ctx, "SELECT embedding FROM embedding_canary WHERE id = 1").Scan(&storedBlob) + switch { + case errors.Is(err, sql.ErrNoRows): + // No probe recorded. Any embeddings already present came from a backend + // this store cannot vouch for, so they are not reusable. + var existing int + if err := tx.QueryRowContext(ctx, + "SELECT COUNT(*) FROM llm_capabilities WHERE embedding IS NOT NULL").Scan(&existing); err != nil { + return false, fmt.Errorf("failed to count stored embeddings: %w", err) + } + discarded = existing > 0 + case err != nil: + return false, fmt.Errorf("failed to read the stored canary: %w", err) + default: + stored := decodeEmbedding(storedBlob) + // Compare widths first: cosine distance indexes both slices positionally + // and would panic on a shorter stored vector. + discarded = len(stored) != len(probe) || + similarity.CosineDistance(stored, probe) > canaryMaxDistance + } + + if discarded { + // Clear the vectors but keep the rows: they also back the external-content + // FTS5 index, so deleting them would break keyword search for tools + // outside the current session's set. + if _, err := tx.ExecContext(ctx, + "UPDATE llm_capabilities SET embedding = NULL, content_hash = NULL"); err != nil { + return false, fmt.Errorf("failed to discard stale embeddings: %w", err) + } } - embeddings, err := s.embeddingClient.EmbedBatch(ctx, texts) + if _, err := tx.ExecContext(ctx, + "INSERT OR REPLACE INTO embedding_canary (id, embedding) VALUES (1, ?)", + encodeEmbedding(probe)); err != nil { + return false, fmt.Errorf("failed to record the canary: %w", err) + } + + return discarded, tx.Commit() +} + +// cachedEmbeddings returns the reusable stored embeddings among the given +// content hashes, keyed by hash. Hashes with no usable vector are absent. +// +// Matching on content_hash rather than tool name lets a renamed tool keep its +// vector. Runs outside any transaction — see the lock note in UpsertTools. +// +// A stale-width vector is treated as a miss rather than reused: reuse would be +// permanent, since it is handed back and re-stored on every rebuild while +// search discards it, silently dropping the tool from semantic results. +func (s sqliteToolStore) cachedEmbeddings(ctx context.Context, keys []string) (map[string][]byte, error) { + keysJSON, err := json.Marshal(keys) if err != nil { - return nil, fmt.Errorf("failed to generate embeddings: %w", err) + return nil, fmt.Errorf("failed to marshal content hashes: %w", err) + } + + queryStr := `SELECT content_hash, embedding + FROM llm_capabilities + WHERE embedding IS NOT NULL + AND content_hash IN (SELECT value FROM json_each(?))` + + rows, err := s.db.QueryContext(ctx, queryStr, string(keysJSON)) + if err != nil { + return nil, fmt.Errorf("embedding cache lookup failed: %w", err) + } + defer func() { _ = rows.Close() }() + + wantBytes := int(s.embeddingDim.Load()) * 4 + cached := make(map[string][]byte, len(keys)) + var unusable int + for rows.Next() { + var hash string + var blob []byte + if err := rows.Scan(&hash, &blob); err != nil { + return nil, fmt.Errorf("failed to scan cached embedding: %w", err) + } + if len(blob) == 0 || (wantBytes > 0 && len(blob) != wantBytes) { + unusable++ + continue + } + cached[hash] = blob + } + if err := rows.Err(); err != nil { + return nil, err } - for i, emb := range embeddings { - blobs[i] = encodeEmbedding(emb) + if unusable > 0 { + slog.Warn("ignoring unusable stored embeddings, they will be recomputed", + "count", unusable, "expected_bytes", wantBytes) } - return blobs, nil + return cached, nil +} + +// embeddedText builds the exact string sent to the embedding backend. The +// content hash is only meaningful if it covers precisely what was embedded, so +// callers must not construct this string themselves. +func embeddedText(name, description string) string { + return fmt.Sprintf("name: %s description: %s", name, description) +} + +// cacheKeyVersion invalidates every stored key if the embedded-text format changes. +const cacheKeyVersion = "v1" + +// embeddingCacheKey hashes the embedded text together with the identity of the +// backend that produced it. +// +// The backend identity is what makes reuse safe: an embedding is interchangeable +// only with one from the same provider, endpoint, and model. Without it, +// repointing the service would silently serve vectors from a different semantic +// space. +func embeddingCacheKey(identity, text string) string { + return hashParts(cacheKeyVersion, identity, text) +} + +// hashParts hashes an ordered list of strings so that distinct lists can never +// produce the same digest. +// +// Each part is length-prefixed rather than delimiter-joined. A delimiter alone +// is ambiguous — moving the separator between two adjacent parts yields an +// identical byte stream, so ("a\x00b", "c") and ("a", "b\x00c") would collide. +// That matters because tool names and descriptions reach this function from +// aggregated backend servers, so a backend could otherwise craft a description +// that collides with another tool's key and take over its stored embedding. +func hashParts(parts ...string) string { + h := sha256.New() + var length [8]byte + for _, part := range parts { + binary.LittleEndian.PutUint64(length[:], uint64(len(part))) + h.Write(length[:]) + h.Write([]byte(part)) + } + return hex.EncodeToString(h.Sum(nil)) +} + +// embeddingIdentity derives the backend identity mixed into every content hash. +// +// Known limitation: for the TEI provider the model is fixed by the running +// container rather than by config, so swapping the model behind an unchanged +// service URL is not detected here. Search tolerates the resulting stale +// vectors (see searchSemantic) but they remain semantically stale until the +// process restarts. Reading the model id from the TEI /info endpoint would +// close this gap. +func embeddingIdentity(cfg *types.OptimizerConfig) string { + if cfg == nil { + return "" + } + // Digested rather than joined so the identity is fixed-length and cannot + // shift the field boundaries of the cache key it is folded into. + return hashParts(cfg.EmbeddingProvider, cfg.EmbeddingService, cfg.EmbeddingModel) } // Search finds tools matching q using FTS5 full-text search and optional @@ -390,6 +740,11 @@ func (s sqliteToolStore) searchSemantic( if err != nil { return nil, fmt.Errorf("failed to embed query: %w", err) } + // A build whose tools are all cache hits never embeds anything, so the query + // round-trip is the only place a model swap is still observable. + if len(queryVec) > 0 { + s.embeddingDim.Store(int64(len(queryVec))) + } allowedJSON, err := json.Marshal(allowedTools) if err != nil { @@ -414,7 +769,7 @@ func (s sqliteToolStore) searchSemantic( } var ranked []rankedMatch - var candidatesEvaluated int + var candidatesEvaluated, dimensionMismatches int for rows.Next() { var name, description string var embBlob []byte @@ -424,6 +779,15 @@ func (s sqliteToolStore) searchSemantic( candidatesEvaluated++ emb := decodeEmbedding(embBlob) + + // Cosine distance indexes both slices positionally: a shorter stored + // vector panics, a longer one silently ignores its tail. A mismatch + // means the vector survived a model change (see embeddingIdentity). + if len(emb) != len(queryVec) { + dimensionMismatches++ + continue + } + dist := similarity.CosineDistance(queryVec, emb) // Filter by semantic distance threshold. @@ -461,10 +825,16 @@ func (s sqliteToolStore) searchSemantic( } } + if dimensionMismatches > 0 { + slog.Warn("skipped stored embeddings with a mismatched dimension", + "count", dimensionMismatches, "query_dimensions", len(queryVec)) + } + slog.Debug("semantic search completed", "allowed_tools", len(allowedTools), "limit", limit, "candidates_evaluated", candidatesEvaluated, + "dimension_mismatches", dimensionMismatches, "results", len(matches), "matched_tools", matchNames(matches), ) diff --git a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_cache_test.go b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_cache_test.go new file mode 100644 index 0000000000..0cfba89f4e --- /dev/null +++ b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_cache_test.go @@ -0,0 +1,969 @@ +// SPDX-FileCopyrightText: Copyright 2025 Stacklok, Inc. +// SPDX-License-Identifier: Apache-2.0 + +package toolstore + +import ( + "context" + "fmt" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/stacklok/toolhive-core/mcpcompat/mcp" + "github.com/stacklok/toolhive-core/mcpcompat/server" + "github.com/stacklok/toolhive/pkg/vmcp/optimizer/internal/types" +) + +// countingEmbeddingClient wraps fakeEmbeddingClient and records how many texts +// were sent to the embedding backend, so tests can assert on embedding work +// avoided rather than on wall-clock time. +type countingEmbeddingClient struct { + *fakeEmbeddingClient + texts atomic.Int64 + calls atomic.Int64 + + mu sync.Mutex + embedded []string +} + +// countingClientDim is the embedding width used by the counting client. Tests +// that need to vary the width exercise dimension handling directly with +// newFakeEmbeddingClient instead. +const countingClientDim = 384 + +func newCountingEmbeddingClient() *countingEmbeddingClient { + return &countingEmbeddingClient{fakeEmbeddingClient: newFakeEmbeddingClient(countingClientDim)} +} + +func (c *countingEmbeddingClient) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) { + c.calls.Add(1) + c.texts.Add(int64(len(texts))) + c.mu.Lock() + c.embedded = append(c.embedded, texts...) + c.mu.Unlock() + return c.fakeEmbeddingClient.EmbedBatch(ctx, texts) +} + +// textsEmbedded returns the total number of texts sent to the backend. +func (c *countingEmbeddingClient) textsEmbedded() int { return int(c.texts.Load()) } + +// batchCalls returns the number of EmbedBatch round-trips made to the backend. +func (c *countingEmbeddingClient) batchCalls() int { return int(c.calls.Load()) } + +// embeddedTexts returns a copy of every text sent to the backend, in order. +func (c *countingEmbeddingClient) embeddedTexts() []string { + c.mu.Lock() + defer c.mu.Unlock() + return append([]string(nil), c.embedded...) +} + +// slowEmbeddingClient holds each EmbedBatch open long enough for concurrent +// builds to overlap, modelling a real embedding backend where a batch takes +// seconds. Without the delay a lock conflict between concurrent builds is a +// race that usually resolves before it can be observed. +type slowEmbeddingClient struct { + *countingEmbeddingClient + delay time.Duration +} + +func (c *slowEmbeddingClient) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) { + select { + case <-time.After(c.delay): + case <-ctx.Done(): + return nil, ctx.Err() + } + return c.countingEmbeddingClient.EmbedBatch(ctx, texts) +} + +// TestSQLiteToolStore_UpsertTools_ConcurrentBuilds covers the production shape +// the embedding cache targets: one shared store, several sessions registering +// at once, each building over its own tool set against a slow backend. +// +// The existing concurrency tests pass a nil embedding client, so they never +// reach the embedding path at all. +func TestSQLiteToolStore_UpsertTools_ConcurrentBuilds(t *testing.T) { + t.Parallel() + + client := &slowEmbeddingClient{countingEmbeddingClient: newCountingEmbeddingClient(), delay: 250 * time.Millisecond} + store := newTestStore(t, client, nil) + ctx := context.Background() + + const sessions = 4 + errs := make(chan error, sessions) + var wg sync.WaitGroup + for i := range sessions { + wg.Add(1) + go func(idx int) { + defer wg.Done() + tools := makeTools(mcp.NewTool( + fmt.Sprintf("session_%d_tool", idx), + mcp.WithDescription(fmt.Sprintf("Tool contributed by session %d", idx)), + )) + errs <- store.UpsertTools(ctx, tools) + }(i) + } + waitOrFail(t, &wg, "concurrent session builds") + close(errs) + + for err := range errs { + require.NoError(t, err, "concurrent session builds must not contend on the store") + } +} + +// waitOrFail blocks until wg completes, failing instead of hanging if it does +// not. A store deadlock — the failure these concurrency tests exist to catch — +// would otherwise stall the run until the global go test timeout, with no +// indication of which test was stuck. +func waitOrFail(t *testing.T, wg *sync.WaitGroup, what string) { + t.Helper() + done := make(chan struct{}) + go func() { + wg.Wait() + close(done) + }() + select { + case <-done: + case <-time.After(30 * time.Second): + t.Fatalf("timeout waiting for %s to finish", what) + } +} + +// newTestStoreDSN builds a store over an explicit database name so several +// stores can share one database, as successive processes would. +func newTestStoreDSN(t *testing.T, dsn string, client types.EmbeddingClient) sqliteToolStore { + t.Helper() + store, err := newSQLiteToolStore(dsn, client, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = store.Close() }) + return store +} + +func catalog(n int) []server.ServerTool { + tools := make([]server.ServerTool, n) + for i := range n { + tools[i] = server.ServerTool{ + Tool: mcp.Tool{ + Name: fmt.Sprintf("tool_%04d", i), + Description: fmt.Sprintf("Tool number %d performing operation %d", i, i%20), + }, + } + } + return tools +} + +// TestSQLiteToolStore_UpsertTools_ReusesCachedEmbeddings locks in the core +// guarantee behind stacklok/toolhive#5847: re-upserting an unchanged tool set +// must not re-embed it. +// +// The Serve path builds a per-session optimizer, and every build calls +// UpsertTools over the session's whole tool set. Without embedding reuse the +// cost is O(tools x sessions) and lands on the client's initialize/tools/list +// round-trip; with reuse it is O(tools whose text changed). +func TestSQLiteToolStore_UpsertTools_ReusesCachedEmbeddings(t *testing.T) { + t.Parallel() + + client := newCountingEmbeddingClient() + store := newTestStore(t, client, nil) + ctx := context.Background() + + tools := catalog(10) + + require.NoError(t, store.UpsertTools(ctx, tools), "first build") + require.Equal(t, 10, client.textsEmbedded(), "the first build must embed every tool") + + // Three more sessions registering against an unchanged catalog. + for i := range 3 { + require.NoError(t, store.UpsertTools(ctx, tools), "rebuild %d", i) + } + + assert.Equal(t, 10, client.textsEmbedded(), + "an unchanged tool set must not be re-embedded on subsequent builds") + assert.Equal(t, 1, client.batchCalls(), + "rebuilds over an unchanged tool set must not reach the embedding backend at all") +} + +// TestSQLiteToolStore_UpsertTools_ReEmbedsOnlyChangedTools asserts the cache is +// keyed on tool content, not merely on presence: a tool whose description +// changed must be re-embedded, and only that tool. +func TestSQLiteToolStore_UpsertTools_ReEmbedsOnlyChangedTools(t *testing.T) { + t.Parallel() + + client := newCountingEmbeddingClient() + store := newTestStore(t, client, nil) + ctx := context.Background() + + tools := catalog(10) + require.NoError(t, store.UpsertTools(ctx, tools)) + require.Equal(t, 10, client.textsEmbedded()) + + // One backend redeploys with a reworded description. + changed := catalog(10) + changed[4].Tool.Description = "Reworded description for tool 4" + require.NoError(t, store.UpsertTools(ctx, changed)) + + assert.Equal(t, 11, client.textsEmbedded(), + "only the tool whose embedded text changed may be re-embedded") + + embedded := client.embeddedTexts() + assert.Contains(t, embedded[len(embedded)-1], "Reworded description for tool 4", + "the re-embedded text must be the changed tool") +} + +// TestSQLiteToolStore_UpsertTools_EmbeddingIdentityInvalidatesCache asserts +// that stored vectors are never reused across a change of embedding backend. +// +// An embedding is only interchangeable with another produced by the same +// provider, endpoint, and model. Keying the cache on the tool text alone would +// serve vectors from a different semantic space after an operator repoints the +// embedding service — silently degrading find_tool rather than failing. +func TestSQLiteToolStore_UpsertTools_EmbeddingIdentityInvalidatesCache(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + next *types.OptimizerConfig + // reuse is true when the second store must reuse the first store's vectors. + reuse bool + }{ + { + name: "same backend reuses", + next: &types.OptimizerConfig{EmbeddingProvider: "tei", EmbeddingService: "http://tei:8080"}, + reuse: true, + }, + { + name: "different endpoint re-embeds", + next: &types.OptimizerConfig{EmbeddingProvider: "tei", EmbeddingService: "http://other:8080"}, + reuse: false, + }, + { + name: "different provider re-embeds", + next: &types.OptimizerConfig{EmbeddingProvider: "openai", EmbeddingService: "http://tei:8080"}, + reuse: false, + }, + { + name: "different model re-embeds", + next: &types.OptimizerConfig{ + EmbeddingProvider: "tei", EmbeddingService: "http://tei:8080", EmbeddingModel: "bge-large", + }, + reuse: false, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + ctx := context.Background() + dsn := fmt.Sprintf("file:identitydb_%d?mode=memory&cache=shared", testDBCounter.Add(1)) + tools := catalog(5) + + first := &types.OptimizerConfig{EmbeddingProvider: "tei", EmbeddingService: "http://tei:8080"} + store, err := newSQLiteToolStore(dsn, newCountingEmbeddingClient(), first) + require.NoError(t, err) + t.Cleanup(func() { _ = store.Close() }) + require.NoError(t, store.UpsertTools(ctx, tools)) + + // A second store over the same database, as after a config change. + nextClient := newCountingEmbeddingClient() + nextStore, err := newSQLiteToolStore(dsn, nextClient, tc.next) + require.NoError(t, err) + t.Cleanup(func() { _ = nextStore.Close() }) + require.NoError(t, nextStore.UpsertTools(ctx, tools)) + + if tc.reuse { + assert.Zero(t, nextClient.textsEmbedded(), "an unchanged backend must reuse stored vectors") + return + } + assert.Equal(t, len(tools), nextClient.textsEmbedded(), + "a changed embedding backend must re-embed every tool") + }) + } +} + +// TestSQLiteToolStore_UpsertTools_RepairsStaleDimension asserts that a stored +// vector left behind by a previous model is recomputed rather than reused +// forever. +// +// Reuse would otherwise make the damage permanent: the stale vector is handed +// back and re-stored on every rebuild while search discards it as +// incomparable, so the tool disappears from semantic results for good. Before +// embedding reuse the next session re-embedded everything and self-healed. +func TestSQLiteToolStore_UpsertTools_RepairsStaleDimension(t *testing.T) { + t.Parallel() + + ctx := context.Background() + dsn := fmt.Sprintf("file:repairdb_%d?mode=memory&cache=shared", testDBCounter.Add(1)) + tools := makeTools(mcp.NewTool("archive_file", mcp.WithDescription("Archive a file to cold storage"))) + + // Indexed by the outgoing model. + old, err := newSQLiteToolStore(dsn, newFakeEmbeddingClient(384), nil) + require.NoError(t, err) + t.Cleanup(func() { _ = old.Close() }) + require.NoError(t, old.UpsertTools(ctx, tools)) + + // The embedding container is replaced by a model of a different width, with + // no config change: same provider, same service URL. + store, err := newSQLiteToolStore(dsn, newFakeEmbeddingClient(768), nil) + require.NoError(t, err) + t.Cleanup(func() { _ = store.Close() }) + + allowed := []string{"archive_file"} + + // The first search cannot rank the stale vector, but it does reveal the + // backend's current width to the store. + before, err := store.searchSemantic(ctx, "store a file", allowed, DefaultMaxToolsToReturn) + require.NoError(t, err) + require.Empty(t, before, "a vector of the previous width must not be ranked") + + require.NoError(t, store.UpsertTools(ctx, tools), "rebuild after the model change") + + after, err := store.searchSemantic(ctx, "store a file", allowed, DefaultMaxToolsToReturn) + require.NoError(t, err) + assert.Equal(t, []string{"archive_file"}, matchNames(after), + "the rebuild must recompute the stale vector so the tool returns to semantic search") +} + +// TestSQLiteToolStore_UpsertTools_IgnoresEmptyStoredEmbedding asserts that a +// zero-length stored vector is recomputed and, above all, that no tool is ever +// handed a different tool's embedding. +// +// An empty blob is not reusable but still matches on content_hash. Treated as a +// hit it would leave a gap that a positional lookup could fill with an unrelated +// vector, silently corrupting search for that tool. +// +// The rebuild deliberately runs on a store that has not yet observed the +// backend's vector width, which is the only state where the width check cannot +// mask an empty blob — a fresh process whose first build is all cache hits. +func TestSQLiteToolStore_UpsertTools_IgnoresEmptyStoredEmbedding(t *testing.T) { + t.Parallel() + + ctx := context.Background() + dsn := fmt.Sprintf("file:emptydb_%d?mode=memory&cache=shared", testDBCounter.Add(1)) + + tools := makeTools( + mcp.NewTool("alpha", mcp.WithDescription("First tool")), + mcp.NewTool("bravo", mcp.WithDescription("Second tool")), + mcp.NewTool("charlie", mcp.WithDescription("Third tool")), + ) + + seed, err := newSQLiteToolStore(dsn, newCountingEmbeddingClient(), nil) + require.NoError(t, err) + t.Cleanup(func() { _ = seed.Close() }) + require.NoError(t, seed.UpsertTools(ctx, tools)) + + // Corrupt one row the way a backend returning an empty vector would. + _, err = seed.db.ExecContext(ctx, "UPDATE llm_capabilities SET embedding = ? WHERE name = ?", []byte{}, "alpha") + require.NoError(t, err) + + // A restarted process: same database, no observed vector width yet. + client := newCountingEmbeddingClient() + store, err := newSQLiteToolStore(dsn, client, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = store.Close() }) + require.Zero(t, store.embeddingDim.Load(), "the width must still be unknown for this test to be meaningful") + + require.NoError(t, store.UpsertTools(ctx, tools), "rebuild over the corrupted row") + + var alpha, charlie []byte + require.NoError(t, store.db.QueryRowContext(ctx, + "SELECT embedding FROM llm_capabilities WHERE name = ?", "alpha").Scan(&alpha)) + require.NoError(t, store.db.QueryRowContext(ctx, + "SELECT embedding FROM llm_capabilities WHERE name = ?", "charlie").Scan(&charlie)) + + assert.NotEmpty(t, alpha, "an unusable stored vector must be recomputed") + assert.NotEqual(t, charlie, alpha, "a tool must never be given another tool's embedding") + + // The expected text is spelled out rather than built with embeddedText: the + // cache key is only meaningful if that format is exactly this, so deriving + // the expectation from the code under test would hide a format change. + want, err := client.Embed(ctx, "name: alpha description: First tool") + require.NoError(t, err) + assert.Equal(t, encodeEmbedding(want), alpha, "the recomputed vector must be alpha's own") +} + +// shiftedEmbeddingClient produces vectors of the configured width that differ +// from fakeEmbeddingClient's for the same text, standing in for a replacement +// model of identical width — the swap neither the content hash nor the +// dimension check can see. +type shiftedEmbeddingClient struct { + *fakeEmbeddingClient + calls atomic.Int64 +} + +func newShiftedEmbeddingClient() *shiftedEmbeddingClient { + return &shiftedEmbeddingClient{fakeEmbeddingClient: newFakeEmbeddingClient(countingClientDim)} +} + +func (c *shiftedEmbeddingClient) Embed(ctx context.Context, text string) ([]float32, error) { + vec, err := c.fakeEmbeddingClient.Embed(ctx, text) + if err != nil { + return nil, err + } + for i := range vec { + vec[i] = -vec[i] + } + return vec, nil +} + +func (c *shiftedEmbeddingClient) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) { + c.calls.Add(1) + out := make([][]float32, len(texts)) + for i, text := range texts { + vec, err := c.Embed(ctx, text) + if err != nil { + return nil, err + } + out[i] = vec + } + return out, nil +} + +// swappableEmbeddingClient returns one model's vectors until swap() is called, +// then a different model's — same width, same client, same URL. +// +// This is what redeploying the embedding service with a different model looks +// like to a running vmcp: the Service URL is unchanged, so the store keeps the +// same client and never learns the backend moved. +type swappableEmbeddingClient struct { + *fakeEmbeddingClient + swapped atomic.Bool + texts atomic.Int64 +} + +func newSwappableEmbeddingClient() *swappableEmbeddingClient { + return &swappableEmbeddingClient{fakeEmbeddingClient: newFakeEmbeddingClient(countingClientDim)} +} + +func (c *swappableEmbeddingClient) swap() { c.swapped.Store(true) } + +// Embed counts every text, including the backend probe, which is issued here +// rather than through EmbedBatch. +func (c *swappableEmbeddingClient) Embed(ctx context.Context, text string) ([]float32, error) { + c.texts.Add(1) + vec, err := c.fakeEmbeddingClient.Embed(ctx, text) + if err != nil { + return nil, err + } + if c.swapped.Load() { + for i := range vec { + vec[i] = -vec[i] + } + } + return vec, nil +} + +func (c *swappableEmbeddingClient) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) { + out := make([][]float32, len(texts)) + for i, text := range texts { + vec, err := c.Embed(ctx, text) + if err != nil { + return nil, err + } + out[i] = vec + } + return out, nil +} + +func (c *swappableEmbeddingClient) embedded() int { return int(c.texts.Load()) } + +// TestSQLiteToolStore_BackendChange_DiscardsStaleEmbeddings is the regression +// test for the failure this cache would otherwise introduce. +// +// A replacement model of the same width is invisible to the content hash (the +// config is unchanged) and to the dimension check (the width is unchanged), so +// without the backend probe the stale vectors would be reused and re-stored on +// every later build — permanently, where re-embedding every session used to heal +// it on its own. +// +// The topology matters: production creates ONE store per process +// (NewOptimizerFactory) and every session build calls UpsertTools on it, so the +// swap must be exercised as successive builds on a single store. Two stores over +// one DSN would test a configuration that does not exist. +func TestSQLiteToolStore_BackendChange_DiscardsStaleEmbeddings(t *testing.T) { + t.Parallel() + + ctx := context.Background() + client := newSwappableEmbeddingClient() + store := newTestStore(t, client, nil) + tools := catalog(5) + + require.NoError(t, store.UpsertTools(ctx, tools), "first build") + afterCold := client.embedded() + require.Equal(t, len(tools)+1, afterCold, "cold build embeds every tool plus the probe") + + require.NoError(t, store.UpsertTools(ctx, tools), "second build, backend unchanged") + afterWarm := client.embedded() + require.Equal(t, afterCold+1, afterWarm, + "an unchanged backend re-embeds only the probe") + + // The embedding service is redeployed with a different model. Same URL, same + // client, same store — only the vectors change. + client.swap() + require.NoError(t, store.UpsertTools(ctx, tools), "build after the model swap") + + assert.Equal(t, afterWarm+1+len(tools), client.embedded(), + "a changed backend must re-embed every tool, not just the probe") + + var stored []byte + require.NoError(t, store.db.QueryRowContext(ctx, + "SELECT embedding FROM llm_capabilities WHERE name = ?", tools[0].Tool.Name).Scan(&stored)) + want, err := client.Embed(ctx, embeddedText(tools[0].Tool.Name, tools[0].Tool.Description)) + require.NoError(t, err) + assert.Equal(t, encodeEmbedding(want), stored, + "stored vectors must come from the current backend, not the previous one") +} + +// TestSQLiteToolStore_BackendUnreachable_StillServesTools pins the availability +// property this cache buys: once warm, a build survives the embedding backend +// being completely down. +// +// Measured live on 2026-07-25 — all four TEI pods deleted, a session built +// normally from cache. It is asserted here because that scenario cannot run in +// CI, and because it is fragile: making the probe refuse reuse on failure +// silently destroys it, turning a working build into a client with no tools. +// That is the original stacklok/toolhive#5847 symptom. +func TestSQLiteToolStore_BackendUnreachable_StillServesTools(t *testing.T) { + t.Parallel() + + ctx := context.Background() + client := newFlakyEmbeddingClient() + store := newTestStore(t, client, nil) + tools := catalog(20) + + require.NoError(t, store.UpsertTools(ctx, tools), "warm the cache while the backend is up") + + // The embedding backend goes away entirely. + client.down.Store(true) + before := client.embedded() + + require.NoError(t, store.UpsertTools(ctx, tools), + "a warm build must survive the embedding backend being unreachable") + assert.Equal(t, before, client.embedded(), + "nothing may be re-embedded while the backend is down") + + var stored int + require.NoError(t, store.db.QueryRowContext(ctx, + "SELECT COUNT(*) FROM llm_capabilities WHERE embedding IS NOT NULL").Scan(&stored)) + assert.Equal(t, len(tools), stored, + "an unverifiable backend must not cause stored vectors to be discarded") +} + +// flakyEmbeddingClient fails every call once down is set, modelling the +// embedding service being unreachable. +type flakyEmbeddingClient struct { + *fakeEmbeddingClient + down atomic.Bool + texts atomic.Int64 +} + +func newFlakyEmbeddingClient() *flakyEmbeddingClient { + return &flakyEmbeddingClient{fakeEmbeddingClient: newFakeEmbeddingClient(countingClientDim)} +} + +func (c *flakyEmbeddingClient) Embed(ctx context.Context, text string) ([]float32, error) { + if c.down.Load() { + return nil, fmt.Errorf("embedding backend unreachable") + } + c.texts.Add(1) + return c.fakeEmbeddingClient.Embed(ctx, text) +} + +func (c *flakyEmbeddingClient) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) { + out := make([][]float32, len(texts)) + for i, text := range texts { + vec, err := c.Embed(ctx, text) + if err != nil { + return nil, err + } + out[i] = vec + } + return out, nil +} + +func (c *flakyEmbeddingClient) embedded() int { return int(c.texts.Load()) } + +// blockingProbeClient holds the probe open until released, so a test can +// observe what a concurrent build does while a probe is in flight. +type blockingProbeClient struct { + *fakeEmbeddingClient + entered chan struct{} + release chan struct{} +} + +func newBlockingProbeClient() *blockingProbeClient { + return &blockingProbeClient{ + fakeEmbeddingClient: newFakeEmbeddingClient(countingClientDim), + entered: make(chan struct{}, 1), + release: make(chan struct{}), + } +} + +func (c *blockingProbeClient) Embed(ctx context.Context, text string) ([]float32, error) { + if text == canaryText { + select { + case c.entered <- struct{}{}: + default: + } + select { + case <-c.release: + case <-ctx.Done(): + return nil, ctx.Err() + } + } + return c.fakeEmbeddingClient.Embed(ctx, text) +} + +func (c *blockingProbeClient) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) { + out := make([][]float32, len(texts)) + for i, text := range texts { + vec, err := c.Embed(ctx, text) + if err != nil { + return nil, err + } + out[i] = vec + } + return out, nil +} + +// TestSQLiteToolStore_ConcurrentBuilds_OrderedAfterProbe asserts no build reads +// the cache while a probe is in flight. +// +// A probe may be about to discard the vectors a concurrent build is reading. If +// that build is allowed to proceed, its INSERT OR REPLACE writes the stale +// vector AND its content hash back after the discard commits. The identity is +// unchanged on a same-width backend swap, so later builds recompute the same +// hash, find the restored row, and reuse it — while the freshly written probe +// certifies the store as current, so nothing re-checks. Permanently stale, with +// nothing logged. +// +// Found by cross-model review; the earlier TryLock implementation had exactly +// this hole. +func TestSQLiteToolStore_ConcurrentBuilds_OrderedAfterProbe(t *testing.T) { + t.Parallel() + + ctx := context.Background() + client := newBlockingProbeClient() + store := newTestStore(t, client, nil) + + go func() { _ = store.UpsertTools(ctx, catalog(3)) }() + <-client.entered // a probe is now in flight and holding the lock + + done := make(chan struct{}) + go func() { + _ = store.UpsertTools(ctx, makeTools( + mcp.NewTool("concurrent", mcp.WithDescription("Indexed while a probe is in flight")))) + close(done) + }() + + select { + case <-done: + close(client.release) + t.Fatal("a build completed while a probe was in flight; it can write pre-discard state back") + case <-time.After(600 * time.Millisecond): + // Correct: the second build is ordered behind the probe. + } + close(client.release) + + select { + case <-done: + case <-time.After(30 * time.Second): + t.Fatal("timeout waiting for the queued build to finish") + } +} + +// TestSQLiteToolStore_BackendUnchanged_KeepsReuse asserts the probe does not +// invalidate the cache when the backend is the same, which would silently undo +// the reuse this cache exists for. +func TestSQLiteToolStore_BackendUnchanged_KeepsReuse(t *testing.T) { + t.Parallel() + + ctx := context.Background() + dsn := fmt.Sprintf("file:canarysame_%d?mode=memory&cache=shared", testDBCounter.Add(1)) + tools := catalog(5) + + first := newTestStoreDSN(t, dsn, newCountingEmbeddingClient()) + require.NoError(t, first.UpsertTools(ctx, tools)) + + second := newCountingEmbeddingClient() + rebuilt := newTestStoreDSN(t, dsn, second) + require.NoError(t, rebuilt.UpsertTools(ctx, tools)) + + assert.Zero(t, second.textsEmbedded(), + "an unchanged backend must keep every stored vector reusable") +} + +// TestSQLiteToolStore_BackendChange_PreservesKeywordSearch asserts the stale +// vectors are cleared without removing the rows, which also back the +// external-content FTS5 index. +func TestSQLiteToolStore_BackendChange_PreservesKeywordSearch(t *testing.T) { + t.Parallel() + + ctx := context.Background() + dsn := fmt.Sprintf("file:canaryfts_%d?mode=memory&cache=shared", testDBCounter.Add(1)) + indexed := newTestStoreDSN(t, dsn, newCountingEmbeddingClient()) + require.NoError(t, indexed.UpsertTools(ctx, makeTools( + mcp.NewTool("archive_file", mcp.WithDescription("Archive a file to cold storage"))))) + + // A different backend, and a build whose tool set does not include the + // previously indexed tool. + swapped := newTestStoreDSN(t, dsn, newShiftedEmbeddingClient()) + require.NoError(t, swapped.UpsertTools(ctx, makeTools( + mcp.NewTool("unrelated_tool", mcp.WithDescription("Something else entirely"))))) + + found, err := swapped.searchFTS5(ctx, "archive", []string{"archive_file"}, DefaultMaxToolsToReturn) + require.NoError(t, err) + assert.Equal(t, []string{"archive_file"}, matchNames(found), + "discarding stale vectors must not remove the rows backing keyword search") +} + +// probeCountingClient counts probe embeds and holds each one open, so a test +// can tell whether concurrent builds queue behind the probe or skip it. +type probeCountingClient struct { + *fakeEmbeddingClient + probes atomic.Int64 + delay time.Duration +} + +func (c *probeCountingClient) Embed(ctx context.Context, text string) ([]float32, error) { + if text == canaryText { + c.probes.Add(1) + select { + case <-time.After(c.delay): + case <-ctx.Done(): + return nil, ctx.Err() + } + } + return c.fakeEmbeddingClient.Embed(ctx, text) +} + +func (c *probeCountingClient) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) { + out := make([][]float32, len(texts)) + for i, text := range texts { + vec, err := c.Embed(ctx, text) + if err != nil { + return nil, err + } + out[i] = vec + } + return out, nil +} + +// TestSQLiteToolStore_ConcurrentBuilds_ShareOneProbe asserts that concurrent +// builds skip the probe when one is already in flight, rather than queueing +// behind it. +// +// The probe runs on every build and costs a network round-trip. Serializing it +// would put one round-trip per concurrent session on the tools/list path: +// measured at 4 builds x a 200ms probe = 805ms before this, 201ms after. The +// probe is idempotent, so its result applies to every build racing it. +func TestSQLiteToolStore_ConcurrentBuilds_ShareOneProbe(t *testing.T) { + t.Parallel() + + ctx := context.Background() + client := &probeCountingClient{ + fakeEmbeddingClient: newFakeEmbeddingClient(countingClientDim), + delay: 200 * time.Millisecond, + } + store := newTestStore(t, client, nil) + + // Warm first so the concurrent builds do the probe and nothing else. + require.NoError(t, store.UpsertTools(ctx, catalog(5))) + client.probes.Store(0) + + const builds = 4 + var wg sync.WaitGroup + errs := make(chan error, builds) + for i := range builds { + wg.Add(1) + go func(idx int) { + defer wg.Done() + errs <- store.UpsertTools(ctx, makeTools(mcp.NewTool( + fmt.Sprintf("concurrent_%d", idx), mcp.WithDescription("A concurrently indexed tool")))) + }(i) + } + waitOrFail(t, &wg, "concurrent probing builds") + close(errs) + for err := range errs { + require.NoError(t, err) + } + + assert.Less(t, int(client.probes.Load()), builds, + "concurrent builds must skip an in-flight probe, not queue behind it") +} + +// TestSQLiteToolStore_BackendProbe_RunsOnce asserts concurrent first builds +// probe the backend once rather than once each. +func TestSQLiteToolStore_BackendProbe_RunsOnce(t *testing.T) { + t.Parallel() + + ctx := context.Background() + client := newCountingEmbeddingClient() + store := newTestStore(t, client, nil) + + var wg sync.WaitGroup + errs := make(chan error, 4) + for i := range 4 { + wg.Add(1) + go func(idx int) { + defer wg.Done() + errs <- store.UpsertTools(ctx, makeTools(mcp.NewTool( + fmt.Sprintf("tool_%d", idx), mcp.WithDescription("A concurrently indexed tool")))) + }(i) + } + waitOrFail(t, &wg, "concurrent first builds") + close(errs) + for err := range errs { + require.NoError(t, err) + } + + var canaries int + require.NoError(t, store.db.QueryRowContext(ctx, + "SELECT COUNT(*) FROM embedding_canary").Scan(&canaries)) + assert.Equal(t, 1, canaries, "the probe must record exactly one canary") +} + +// TestEmbeddedText pins the exact string sent to the embedding backend. Every +// stored cache key is a hash of it, so changing the format silently invalidates +// every entry — it must be a deliberate edit here, with a cacheKeyVersion bump. +func TestEmbeddedText(t *testing.T) { + t.Parallel() + assert.Equal(t, "name: read_file description: Read a file", embeddedText("read_file", "Read a file")) +} + +// TestEmbeddingCacheKey_Injective asserts that distinct inputs never share a +// cache key. +// +// The parts are length-prefixed rather than delimiter-joined precisely because +// a delimiter is ambiguous: with "v1\x00"+identity+"\x00"+text, moving a NUL +// across the boundary produces an identical byte stream. Tool names and +// descriptions come from aggregated backend servers, so a backend could craft a +// description that collides with another tool's key and take over its embedding. +func TestEmbeddingCacheKey_Injective(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + aID, aText string + bID, bText string + wantSameAsFirst bool + }{ + { + name: "identical inputs agree", + aID: "tei", aText: "name: A description: B", + bID: "tei", bText: "name: A description: B", + wantSameAsFirst: true, + }, + { + name: "NUL shifted across the identity/text boundary", + aID: "tei\x00http://x\x00", aText: "name: A description: B", + bID: "tei\x00http://x", bText: "\x00name: A description: B", + }, + { + name: "content moved between fields", + aID: "abc", aText: "def", + bID: "ab", bText: "cdef", + }, + { + name: "different identity, same text", + aID: "tei", aText: "name: A description: B", + bID: "openai", bText: "name: A description: B", + }, + { + name: "same identity, different text", + aID: "tei", aText: "name: A description: B", + bID: "tei", bText: "name: A description: C", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + a := embeddingCacheKey(tc.aID, tc.aText) + b := embeddingCacheKey(tc.bID, tc.bText) + if tc.wantSameAsFirst { + assert.Equal(t, a, b, "identical inputs must produce one key") + return + } + assert.NotEqual(t, a, b, "distinct inputs must not share a cache key") + }) + } +} + +// FuzzEmbeddingCacheKey checks the same injectivity property over arbitrary +// inputs: two key derivations agree only when their inputs do. +func FuzzEmbeddingCacheKey(f *testing.F) { + f.Add("tei", "name: A description: B", "tei", "name: A description: B") + f.Add("tei\x00svc\x00", "name: A", "tei\x00svc", "\x00name: A") + f.Add("", "", "", "") + + f.Fuzz(func(t *testing.T, aID, aText, bID, bText string) { + sameKey := embeddingCacheKey(aID, aText) == embeddingCacheKey(bID, bText) + sameInput := aID == bID && aText == bText + if sameKey != sameInput { + t.Fatalf("key equality %v but input equality %v for (%q,%q) vs (%q,%q)", + sameKey, sameInput, aID, aText, bID, bText) + } + }) +} + +// TestSQLiteToolStore_UpsertTools_FTS5OnlyStoresNoHash covers the branch taken +// when no embedding client is configured: rows must carry neither an embedding +// nor a content hash, so nothing can later be mistaken for a reusable vector. +func TestSQLiteToolStore_UpsertTools_FTS5OnlyStoresNoHash(t *testing.T) { + t.Parallel() + + ctx := context.Background() + store := newTestStore(t, nil, nil) + require.NoError(t, store.UpsertTools(ctx, makeTools( + mcp.NewTool("read_file", mcp.WithDescription("Read a file"))))) + + var embedding, hash any + require.NoError(t, store.db.QueryRowContext(ctx, + "SELECT embedding, content_hash FROM llm_capabilities WHERE name = ?", "read_file"). + Scan(&embedding, &hash)) + + assert.Nil(t, embedding, "FTS5-only mode must store no embedding") + assert.Nil(t, hash, "FTS5-only mode must store no content hash") +} + +// TestSQLiteToolStore_SearchSemantic_SkipsMismatchedDimensions asserts that a +// stored vector whose dimension differs from the query's is dropped from the +// ranking instead of crashing the search. +// +// Cosine distance indexes both slices positionally, so a shorter stored vector +// panics. Before embedding reuse every vector was rewritten on each session and +// dimensions could not diverge; with reuse, a vector can outlive the model that +// produced it. +func TestSQLiteToolStore_SearchSemantic_SkipsMismatchedDimensions(t *testing.T) { + t.Parallel() + + ctx := context.Background() + dsn := fmt.Sprintf("file:dimdb_%d?mode=memory&cache=shared", testDBCounter.Add(1)) + + stale := makeTools(mcp.NewTool("stale_tool", mcp.WithDescription("Indexed by the previous model"))) + store, err := newSQLiteToolStore(dsn, newFakeEmbeddingClient(384), nil) + require.NoError(t, err) + t.Cleanup(func() { _ = store.Close() }) + require.NoError(t, store.UpsertTools(ctx, stale)) + + // The model is replaced by one with a different output dimension. + wider, err := newSQLiteToolStore(dsn, newFakeEmbeddingClient(768), nil) + require.NoError(t, err) + t.Cleanup(func() { _ = wider.Close() }) + + fresh := makeTools(mcp.NewTool("fresh_tool", mcp.WithDescription("Indexed by the current model"))) + require.NoError(t, wider.UpsertTools(ctx, fresh)) + + allowed := []string{"stale_tool", "fresh_tool"} + require.NotPanics(t, func() { + results, searchErr := wider.searchSemantic(ctx, "indexed by a model", allowed, DefaultMaxToolsToReturn) + require.NoError(t, searchErr) + assert.Equal(t, []string{"fresh_tool"}, matchNames(results), + "only vectors comparable with the query may be ranked") + }) +} diff --git a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_livemodel_test.go b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_livemodel_test.go new file mode 100644 index 0000000000..3dcd2725cc --- /dev/null +++ b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_livemodel_test.go @@ -0,0 +1,253 @@ +// SPDX-FileCopyrightText: Copyright 2025 Stacklok, Inc. +// SPDX-License-Identifier: Apache-2.0 + +package toolstore + +import ( + "cmp" + "context" + "fmt" + "os" + "sort" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/stacklok/toolhive-core/mcpcompat/mcp" + "github.com/stacklok/toolhive/pkg/vmcp/optimizer/internal/similarity" + "github.com/stacklok/toolhive/pkg/vmcp/optimizer/internal/types" +) + +// Environment variables gating the live model-swap tests. Two models are +// required: one whose embedding width differs from the baseline, and one whose +// width is identical, since the two swaps are handled by different mechanisms. +const ( + liveEmbedURLEnv = "VMCP_LIVE_EMBEDDING_URL" + liveModelBaseEnv = "VMCP_LIVE_MODEL_BASE" + liveModelSameWidthEnv = "VMCP_LIVE_MODEL_SAME_WIDTH" + liveModelDiffWidthEnv = "VMCP_LIVE_MODEL_DIFF_WIDTH" +) + +// liveModelStore builds a store against a real OpenAI-compatible embedding +// endpoint, sharing dsn so successive stores see the same rows. +func liveModelStore(t *testing.T, dsn, endpoint, model string) sqliteToolStore { + t.Helper() + cfg := &types.OptimizerConfig{ + EmbeddingService: endpoint, + EmbeddingProvider: types.EmbeddingProviderOpenAI, + EmbeddingModel: model, + EmbeddingServiceTimeout: time.Minute, + } + client, err := similarity.NewEmbeddingClient(cfg) + require.NoError(t, err) + store, err := newSQLiteToolStore(dsn, client, cfg) + require.NoError(t, err) + t.Cleanup(func() { _ = store.Close() }) + return store +} + +// liveModelEndpoint returns the configured endpoint, or skips the test. +func liveModelEndpoint(t *testing.T) string { + t.Helper() + endpoint := os.Getenv(liveEmbedURLEnv) + if endpoint == "" { + t.Skipf("%s not set; skipping live embedding model test", liveEmbedURLEnv) + } + return endpoint +} + +// TestLiveModelSwap_DifferentWidth verifies against real models that swapping +// the embedding model to one of a different width does not corrupt search: the +// stale vector must be discarded, then repaired on the next build. +// +// The unit tests cover this with a fake client, which cannot show that real +// vectors of two widths coexist in one column without panicking cosine distance. +func TestLiveModelSwap_DifferentWidth(t *testing.T) { + t.Parallel() + + endpoint := liveModelEndpoint(t) + base := cmp.Or(os.Getenv(liveModelBaseEnv), "bge-m3") + other := cmp.Or(os.Getenv(liveModelDiffWidthEnv), "all-minilm") + + ctx := context.Background() + dsn := fmt.Sprintf("file:livediff_%d?mode=memory&cache=shared", testDBCounter.Add(1)) + tools := makeTools(mcp.NewTool("archive_file", mcp.WithDescription("Archive a file to cold storage"))) + + indexed := liveModelStore(t, dsn, endpoint, base) + require.NoError(t, indexed.UpsertTools(ctx, tools)) + + swapped := liveModelStore(t, dsn, endpoint, other) + allowed := []string{"archive_file"} + + before, err := swapped.searchSemantic(ctx, "store a file", allowed, DefaultMaxToolsToReturn) + require.NoError(t, err, "a stale-width vector must not panic cosine distance") + assert.Empty(t, before, "a vector of the previous width must not be ranked") + + require.NoError(t, swapped.UpsertTools(ctx, tools), "rebuild after the model change") + + after, err := swapped.searchSemantic(ctx, "store a file", allowed, DefaultMaxToolsToReturn) + require.NoError(t, err) + assert.Equal(t, []string{"archive_file"}, matchNames(after), + "the rebuild must recompute the stale vector so the tool returns to semantic search") +} + +// TestLiveModelSwap_SameWidth verifies against two real models of identical +// width that a swap invisible to both the content hash and the dimension check +// is still caught, and the stale vectors recomputed. +// +// For the TEI provider the model is fixed by the running container rather than +// by config, so this swap changes nothing the cache key can see; equal widths +// hide it from the dimension check too. Only re-embedding the probe detects it. +// +// It also measures how far apart the two spaces are, which is what makes the +// probe's tolerance safe: the same model repeats bit-identically, while two +// different models put the same text roughly a full unit apart. +func TestLiveModelSwap_SameWidth(t *testing.T) { + t.Parallel() + + endpoint := liveModelEndpoint(t) + base := cmp.Or(os.Getenv(liveModelBaseEnv), "bge-m3") + other := os.Getenv(liveModelSameWidthEnv) + if other == "" { + t.Skipf("%s not set; skipping same-width model swap test", liveModelSameWidthEnv) + } + + ctx := context.Background() + text := embeddedText("archive_file", "Archive a file to cold storage") + + // Confirm the premise: both models must produce the same width, or this test + // is silently measuring the different-width case instead. + baseVec := liveEmbed(ctx, t, endpoint, base, text) + otherVec := liveEmbed(ctx, t, endpoint, other, text) + require.Equal(t, len(baseVec), len(otherVec), + "this test requires two models of equal width (%s=%d, %s=%d)", base, len(baseVec), other, len(otherVec)) + + // The same text embedded by two different models: how different are the spaces? + dist := similarity.CosineDistance(baseVec, otherVec) + t.Logf("same text under %s vs %s: cosine distance %.4f (width %d)", base, other, dist, len(baseVec)) + + dsn := fmt.Sprintf("file:livesame_%d?mode=memory&cache=shared", testDBCounter.Add(1)) + tools := makeTools(mcp.NewTool("archive_file", mcp.WithDescription("Archive a file to cold storage"))) + + indexed := liveModelStore(t, dsn, endpoint, base) + require.NoError(t, indexed.UpsertTools(ctx, tools)) + + // A TEI-style swap: the model changes behind an unchanged service URL, so + // the identity the cache key is derived from is byte-identical. The store is + // configured with `base` while the client actually embeds with `other`. + swapped := liveModelStore(t, dsn, endpoint, other) + swapped.embeddingIdentity = indexed.embeddingIdentity + require.NoError(t, swapped.UpsertTools(ctx, tools)) + + var stored []byte + require.NoError(t, swapped.db.QueryRowContext(ctx, + "SELECT embedding FROM llm_capabilities WHERE name = ?", "archive_file").Scan(&stored)) + + assert.Equal(t, encodeEmbedding(otherVec), stored, + "a same-width model swap must be detected and the stale vector recomputed") + assert.NotEqual(t, encodeEmbedding(baseVec), stored, + "the previous model's vector must not survive the swap") +} + +// TestLiveModelSwap_StaleDistanceDistribution measures how many stale vectors +// survive the semantic distance filter after an undetectable same-width model +// swap, across a realistic catalogue rather than a single pair. +// +// A single aggregate distance is misleading here. Two unrelated embedding spaces +// give a cosine similarity centred on zero, so per-tool distances scatter around +// 1.0 — and DefaultSemanticDistanceThreshold is exactly 1.0. Whether the filter +// is a real backstop or a coin flip therefore depends on the spread, which only +// a distribution can show. +func TestLiveModelSwap_StaleDistanceDistribution(t *testing.T) { + t.Parallel() + + endpoint := liveModelEndpoint(t) + base := cmp.Or(os.Getenv(liveModelBaseEnv), "bge-m3") + other := os.Getenv(liveModelSameWidthEnv) + if other == "" { + t.Skipf("%s not set; skipping stale-distance distribution measurement", liveModelSameWidthEnv) + } + + ctx := context.Background() + backends := []string{"grafana", "datadog", "argocd", "k8s", "gitnexus", "dbhub", "firecrawl", "context7"} + const toolCount = 64 + + texts := make([]string, toolCount) + for i := range toolCount { + b := backends[i%len(backends)] + texts[i] = embeddedText( + fmt.Sprintf("%s_operation_%03d", b, i), + fmt.Sprintf("Perform operation %d against the %s backend, applying the requested "+ + "filters and returning the matching records with their metadata.", i, b), + ) + } + + // Tools indexed by the outgoing model; queries embedded by the incoming one. + staleVecs := liveEmbedBatch(ctx, t, endpoint, base, texts) + queries := []string{ + "search dashboards", "list kubernetes pods", "run a SQL query", + "fetch a web page", "look up library documentation", "check deployment sync status", + } + queryVecs := liveEmbedBatch(ctx, t, endpoint, other, queries) + require.Equal(t, len(staleVecs[0]), len(queryVecs[0]), "the two models must share a width") + + var dists []float64 + for _, q := range queryVecs { + for _, s := range staleVecs { + dists = append(dists, similarity.CosineDistance(q, s)) + } + } + sort.Float64s(dists) + + belowThreshold := 0 + for _, d := range dists { + if d <= DefaultSemanticDistanceThreshold { + belowThreshold++ + } + } + pct := 100 * float64(belowThreshold) / float64(len(dists)) + + t.Logf("stale-vector distances after a same-width swap (%s indexed, %s querying), n=%d:", + base, other, len(dists)) + t.Logf(" min %.4f | p50 %.4f | max %.4f", dists[0], dists[len(dists)/2], dists[len(dists)-1]) + t.Logf(" %d/%d (%.1f%%) fall at or below the %.1f distance threshold and would be ranked", + belowThreshold, len(dists), pct, DefaultSemanticDistanceThreshold) + + // No pass/fail on the fraction: it is a property of the model pair, not of + // this code. The assertion is only that the measurement is meaningful. + require.NotEmpty(t, dists) +} + +// liveEmbedBatch returns embeddings for several texts from one model. +func liveEmbedBatch(ctx context.Context, t *testing.T, endpoint, model string, texts []string) [][]float32 { + t.Helper() + client, err := similarity.NewEmbeddingClient(&types.OptimizerConfig{ + EmbeddingService: endpoint, + EmbeddingProvider: types.EmbeddingProviderOpenAI, + EmbeddingModel: model, + EmbeddingServiceTimeout: 5 * time.Minute, + }) + require.NoError(t, err) + t.Cleanup(func() { _ = client.Close() }) + vecs, err := client.EmbedBatch(ctx, texts) + require.NoError(t, err) + return vecs +} + +// liveEmbed returns one embedding from the live endpoint for the given model. +func liveEmbed(ctx context.Context, t *testing.T, endpoint, model, text string) []float32 { + t.Helper() + client, err := similarity.NewEmbeddingClient(&types.OptimizerConfig{ + EmbeddingService: endpoint, + EmbeddingProvider: types.EmbeddingProviderOpenAI, + EmbeddingModel: model, + EmbeddingServiceTimeout: time.Minute, + }) + require.NoError(t, err) + t.Cleanup(func() { _ = client.Close() }) + vec, err := client.Embed(ctx, text) + require.NoError(t, err) + return vec +} diff --git a/pkg/vmcp/server/serve_optimizer_live_test.go b/pkg/vmcp/server/serve_optimizer_live_test.go new file mode 100644 index 0000000000..4f79b9cc77 --- /dev/null +++ b/pkg/vmcp/server/serve_optimizer_live_test.go @@ -0,0 +1,174 @@ +// SPDX-FileCopyrightText: Copyright 2025 Stacklok, Inc. +// SPDX-License-Identifier: Apache-2.0 + +package server + +import ( + "cmp" + "context" + "fmt" + "net/http" + "net/http/httptest" + "os" + "strconv" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" + + "github.com/stacklok/toolhive-core/mcpcompat/server" + "github.com/stacklok/toolhive/pkg/vmcp" + "github.com/stacklok/toolhive/pkg/vmcp/optimizer" + "github.com/stacklok/toolhive/pkg/vmcp/server/sessionmanager" + "github.com/stacklok/toolhive/pkg/vmcp/session/optimizerdec" +) + +// Environment variables gating the live Serve-path optimizer test. +const ( + liveEmbeddingURLEnv = "VMCP_LIVE_EMBEDDING_URL" + liveEmbeddingModelEnv = "VMCP_LIVE_EMBEDDING_MODEL" + liveToolCountEnv = "VMCP_LIVE_TOOL_COUNT" +) + +// TestServeOptimizerLive_SecondSessionServedWarm measures what a client actually +// experiences on the Serve path — the wait from `initialize` until find_tool is +// advertised — across consecutive sessions, against a real embedding backend. +// +// This is the end-to-end counterpart to the store-level embedding-reuse tests: +// it exercises the real optimizer factory, the real SQLite store, and a real +// embedding round-trip through the session-registration path that +// stacklok/toolhive#5847 describes, rather than asserting on a mock. +// +// Skipped unless VMCP_LIVE_EMBEDDING_URL points at an OpenAI-compatible +// /embeddings endpoint, so the default `task test` run stays green. Any such +// endpoint works; a local Ollama needs no API key: +// +// VMCP_LIVE_EMBEDDING_URL=http://127.0.0.1:11434/v1 \ +// VMCP_LIVE_EMBEDDING_MODEL=bge-m3 \ +// VMCP_LIVE_TOOL_COUNT=140 \ +// go test ./pkg/vmcp/server/ -run TestServeOptimizerLive -v +func TestServeOptimizerLive_SecondSessionServedWarm(t *testing.T) { + // Safe to run alongside other tests: every assertion is relative to this + // run's own cold measurement, so a loaded machine moves both numbers. + t.Parallel() + + endpoint := os.Getenv(liveEmbeddingURLEnv) + if endpoint == "" { + t.Skipf("%s not set; skipping live Serve-path optimizer test", liveEmbeddingURLEnv) + } + model := cmp.Or(os.Getenv(liveEmbeddingModelEnv), "bge-m3") + + toolCount := 140 + if raw := os.Getenv(liveToolCountEnv); raw != "" { + parsed, err := strconv.Atoi(raw) + require.NoErrorf(t, err, "%s must be an integer", liveToolCountEnv) + require.Positive(t, parsed, "%s must be positive", liveToolCountEnv) + toolCount = parsed + } + + optFactory, cleanup, err := optimizer.NewOptimizerFactory(&optimizer.Config{ + EmbeddingService: endpoint, + EmbeddingProvider: "openai", + EmbeddingModel: model, + EmbeddingServiceTimeout: 2 * time.Minute, + }) + require.NoError(t, err) + t.Cleanup(func() { _ = cleanup(context.Background()) }) + + fc := &fakeCore{tools: liveTools(toolCount)} + baseURL := serveWithOptimizerFactory(t, fc, optFactory) + + const sessions = 3 + elapsed := make([]time.Duration, sessions) + for i := range sessions { + elapsed[i] = timeToolsAvailable(t, baseURL) + t.Logf("session %d: find_tool advertised after %s (%d tools)", i+1, elapsed[i].Round(time.Millisecond), toolCount) + } + + // The first session pays for embedding the catalog; every later session must + // be served from the stored vectors. The bound is deliberately loose — this + // asserts the cold build is gone, not a particular backend's speed. + warmBudget := max(elapsed[0]/4, 500*time.Millisecond) + for i := 1; i < sessions; i++ { + assert.Lessf(t, elapsed[i], warmBudget, + "session %d must be served from stored embeddings, not rebuilt (cold was %s)", i+1, elapsed[0]) + } +} + +// liveTools builds a synthetic catalog roughly the shape of an aggregated vMCP: +// several backends, each contributing tools with prose descriptions long enough +// to be representative of real embedding input. +func liveTools(n int) []vmcp.Tool { + backends := []string{"grafana", "datadog", "argocd", "k8s", "gitnexus", "dbhub", "firecrawl", "context7"} + tools := make([]vmcp.Tool, n) + for i := range n { + backend := backends[i%len(backends)] + tools[i] = vmcp.Tool{ + Name: fmt.Sprintf("%s_operation_%03d", backend, i), + Description: fmt.Sprintf( + "Perform operation %d against the %s backend. This tool queries the %s subsystem, "+ + "applies the requested filters, and returns the matching records together with their "+ + "metadata so the caller can decide how to proceed.", i, backend, backend), + } + } + return tools +} + +// serveWithOptimizerFactory starts a Serve-path test server wired to the given +// optimizer factory and returns its base URL. It mirrors +// registerServeOptimizerSession but takes a real factory and registers no +// session, so each session can be timed individually. +func serveWithOptimizerFactory( + t *testing.T, vmcpCore *fakeCore, optFactory func(context.Context, []server.ServerTool) (optimizer.Optimizer, error), +) string { + t.Helper() + ctrl := gomock.NewController(t) + factory, _ := newToolSessionFactory(t, ctrl, vmcpCore.tools) + + srv, err := Serve(context.Background(), vmcpCore, &ServerConfig{ + SessionTTL: time.Minute, + SessionManagerConfig: &sessionmanager.FactoryConfig{ + Base: factory, + OptimizerFactory: optFactory, + AdvertiseFromCore: true, + }, + BackendRegistry: vmcp.NewImmutableRegistry([]vmcp.Backend{}), + }) + require.NoError(t, err) + t.Cleanup(func() { _ = srv.Stop(context.Background()) }) + + streamable := server.NewStreamableHTTPServer( + srv.mcpServer, + server.WithEndpointPath("/mcp"), + server.WithSessionIdManager(srv.vmcpSessionMgr), + ) + ts := httptest.NewServer(streamable) + t.Cleanup(ts.Close) + return ts.URL +} + +// timeToolsAvailable opens a fresh session and returns how long it takes before +// find_tool is advertised — the moment a client can actually use the server. +func timeToolsAvailable(t *testing.T, baseURL string) time.Duration { + t.Helper() + start := time.Now() + + initResp := postServeMCP(t, baseURL, initBody, "") + defer initResp.Body.Close() + require.Equal(t, http.StatusOK, initResp.StatusCode) + sessionID := initResp.Header.Get("Mcp-Session-Id") + require.NotEmpty(t, sessionID) + + require.Eventually(t, func() bool { + for _, name := range serveToolNames(t, baseURL, sessionID) { + if name == optimizerdec.FindToolName { + return true + } + } + return false + }, 3*time.Minute, 20*time.Millisecond, "find_tool should eventually be advertised") + + return time.Since(start) +} From 303b1f5edd169a418091aacd821779a9aa3ee8d3 Mon Sep 17 00:00:00 2001 From: TANTIOPE Date: Sun, 9 Aug 2026 10:53:30 +0200 Subject: [PATCH 2/2] Derive the cache identity from the backend model id, dropping the canary MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The embedding cache key now folds in the model id read from the backend on every build: TEI reports it from /info, the OpenAI client knows it from configuration. A model swap changes the keys, so stale vectors stop being found instead of needing to be detected and discarded — which makes the canary probe, its table, its ordering lock and its distance threshold unnecessary. A failed id read falls back to the last id seen, keeping keys stable across transient failures. The identity is re-read after each embedding batch; a build whose batch spanned a swap is discarded and re-run under the new identity rather than committing vectors under keys naming the wrong model. The dimension guard moves into CosineSimilarity/CosineDistance, which now refuse mismatched widths instead of documenting the requirement. Signed-off-by: TANTIOPE --- .../optimizer/internal/similarity/cosine.go | 30 +- .../internal/similarity/cosine_bench_test.go | 4 +- .../internal/similarity/cosine_test.go | 36 +- .../internal/similarity/openai_client.go | 7 + .../internal/similarity/openai_client_test.go | 12 + .../internal/similarity/tei_client.go | 49 +- .../internal/similarity/tei_client_test.go | 84 +++ .../optimizer/internal/toolstore/schema.sql | 20 +- .../internal/toolstore/sqlite_store.go | 385 ++++++------- .../toolstore/sqlite_store_cache_test.go | 519 +++++++++++------- .../toolstore/sqlite_store_livemodel_test.go | 113 +--- .../internal/toolstore/sqlite_store_test.go | 8 + .../internal/types/mocks/mock_types.go | 15 + pkg/vmcp/optimizer/internal/types/types.go | 9 + 14 files changed, 748 insertions(+), 543 deletions(-) diff --git a/pkg/vmcp/optimizer/internal/similarity/cosine.go b/pkg/vmcp/optimizer/internal/similarity/cosine.go index fdf1463217..ef6f882e92 100644 --- a/pkg/vmcp/optimizer/internal/similarity/cosine.go +++ b/pkg/vmcp/optimizer/internal/similarity/cosine.go @@ -4,13 +4,24 @@ // Package similarity provides vector distance functions for semantic search. package similarity -import "math" +import ( + "fmt" + "math" +) // CosineSimilarity computes the cosine similarity between two vectors. // Returns a value in [-1, 1] where 1 means identical direction, // 0 means orthogonal, and -1 means opposite direction. -// Both vectors must have the same length. -func CosineSimilarity(a, b []float32) float64 { +// +// Vectors of different lengths return an error. The guard lives here rather +// than at the call sites because the loop below indexes both slices +// positionally: a shorter b would panic and a longer b would silently ignore +// its tail, and a caller that forgets its own check gets one or the other. +func CosineSimilarity(a, b []float32) (float64, error) { + if len(a) != len(b) { + return 0, fmt.Errorf("vectors have different dimensions: %d and %d", len(a), len(b)) + } + var dot, normA, normB float64 for i := range a { ai := float64(a[i]) @@ -22,14 +33,19 @@ func CosineSimilarity(a, b []float32) float64 { denom := math.Sqrt(normA) * math.Sqrt(normB) if denom == 0 { - return 0 + return 0, nil } - return dot / denom + return dot / denom, nil } // CosineDistance computes the cosine distance between two vectors. // Returns a value in [0, 2] where 0 means identical direction and 2 means // opposite direction. Lower values indicate more similar vectors. -func CosineDistance(a, b []float32) float64 { - return 1 - CosineSimilarity(a, b) +// Vectors of different lengths return an error (see CosineSimilarity). +func CosineDistance(a, b []float32) (float64, error) { + sim, err := CosineSimilarity(a, b) + if err != nil { + return 0, err + } + return 1 - sim, nil } diff --git a/pkg/vmcp/optimizer/internal/similarity/cosine_bench_test.go b/pkg/vmcp/optimizer/internal/similarity/cosine_bench_test.go index 3d3f27ddb5..90f1e0ed70 100644 --- a/pkg/vmcp/optimizer/internal/similarity/cosine_bench_test.go +++ b/pkg/vmcp/optimizer/internal/similarity/cosine_bench_test.go @@ -23,7 +23,7 @@ func BenchmarkCosineDistance_384(b *testing.B) { b.ResetTimer() b.ReportAllocs() for b.Loop() { - CosineDistance(a, v) + _, _ = CosineDistance(a, v) } } @@ -33,6 +33,6 @@ func BenchmarkCosineDistance_768(b *testing.B) { b.ResetTimer() b.ReportAllocs() for b.Loop() { - CosineDistance(a, v) + _, _ = CosineDistance(a, v) } } diff --git a/pkg/vmcp/optimizer/internal/similarity/cosine_test.go b/pkg/vmcp/optimizer/internal/similarity/cosine_test.go index f967985e7e..5b8c4920d3 100644 --- a/pkg/vmcp/optimizer/internal/similarity/cosine_test.go +++ b/pkg/vmcp/optimizer/internal/similarity/cosine_test.go @@ -28,7 +28,37 @@ func TestCosineSimilarity(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { t.Parallel() - require.InDelta(t, tc.want, CosineSimilarity(tc.a, tc.b), 1e-7) + got, err := CosineSimilarity(tc.a, tc.b) + require.NoError(t, err) + require.InDelta(t, tc.want, got, 1e-7) + }) + } +} + +// TestCosineSimilarity_DimensionMismatch asserts mismatched widths are refused +// rather than computed. Without the guard, a shorter b panics on indexing and +// a longer b silently ignores its tail — a wrong answer, not an error. +func TestCosineSimilarity_DimensionMismatch(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + a, b []float32 + }{ + {name: "b shorter would panic", a: []float32{1, 2, 3}, b: []float32{1, 2}}, + {name: "b longer would be silently truncated", a: []float32{1, 2}, b: []float32{1, 2, 3}}, + {name: "empty against non-empty", a: nil, b: []float32{1}}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + _, err := CosineSimilarity(tc.a, tc.b) + require.Error(t, err) + + _, err = CosineDistance(tc.a, tc.b) + require.Error(t, err) }) } } @@ -49,7 +79,9 @@ func TestCosineDistance(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { t.Parallel() - require.InDelta(t, tc.want, CosineDistance(tc.a, tc.b), 1e-7) + got, err := CosineDistance(tc.a, tc.b) + require.NoError(t, err) + require.InDelta(t, tc.want, got, 1e-7) }) } } diff --git a/pkg/vmcp/optimizer/internal/similarity/openai_client.go b/pkg/vmcp/optimizer/internal/similarity/openai_client.go index 907164d16c..f3b4c46e1f 100644 --- a/pkg/vmcp/optimizer/internal/similarity/openai_client.go +++ b/pkg/vmcp/optimizer/internal/similarity/openai_client.go @@ -178,6 +178,13 @@ func (c *openAIClient) embedChunk(ctx context.Context, texts []string) ([][]floa return embeddings, nil } +// ModelID returns the configured model name. The OpenAI /embeddings API +// selects the model per request from this same field, so unlike TEI there is +// no server-side state that could drift from it. +func (c *openAIClient) ModelID(context.Context) (string, error) { + return c.model, nil +} + // Close is a no-op for the OpenAI client. func (*openAIClient) Close() error { return nil diff --git a/pkg/vmcp/optimizer/internal/similarity/openai_client_test.go b/pkg/vmcp/optimizer/internal/similarity/openai_client_test.go index dd8a2b6879..a340471285 100644 --- a/pkg/vmcp/optimizer/internal/similarity/openai_client_test.go +++ b/pkg/vmcp/optimizer/internal/similarity/openai_client_test.go @@ -59,6 +59,18 @@ func Test_newOpenAIClient(t *testing.T) { }) } +func TestOpenAIClient_ModelID(t *testing.T) { + t.Parallel() + + client, err := newOpenAIClient("http://embeddings:8080/v1", "text-embedding-3-small", "key", nil, 0) + require.NoError(t, err) + + id, err := client.ModelID(context.Background()) + require.NoError(t, err) + require.Equal(t, "text-embedding-3-small", id, + "the OpenAI client's model is fixed by configuration and sent per request") +} + func TestOpenAIClient_Embed(t *testing.T) { t.Parallel() diff --git a/pkg/vmcp/optimizer/internal/similarity/tei_client.go b/pkg/vmcp/optimizer/internal/similarity/tei_client.go index 035f3e0f9d..e2e0d94234 100644 --- a/pkg/vmcp/optimizer/internal/similarity/tei_client.go +++ b/pkg/vmcp/optimizer/internal/similarity/tei_client.go @@ -68,24 +68,41 @@ func newTEIClient(baseURL string, timeout time.Duration) (*teiClient, error) { // teiInfoResponse is a subset of the TEI /info endpoint response. type teiInfoResponse struct { - MaxClientBatchSize int `json:"max_client_batch_size"` + ModelID string `json:"model_id"` + MaxClientBatchSize int `json:"max_client_batch_size"` } -// fetchMaxBatchSize queries the TEI /info endpoint and returns the max client batch size. -func fetchMaxBatchSize(baseURL string, httpClient *http.Client) (int, error) { - resp, err := httpClient.Get(baseURL + infoPath) // #nosec G107 -- URL is built from the configured TEI base URL +// fetchInfo queries the TEI /info endpoint. +func fetchInfo(ctx context.Context, baseURL string, httpClient *http.Client) (teiInfoResponse, error) { + var info teiInfoResponse + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, baseURL+infoPath, nil) + if err != nil { + return info, fmt.Errorf("failed to create TEI /info request: %w", err) + } + + resp, err := httpClient.Do(req) // #nosec G704 -- URL is built from the configured TEI base URL if err != nil { - return 0, fmt.Errorf("TEI /info request failed: %w", err) + return info, fmt.Errorf("TEI /info request failed: %w", err) } defer func() { _ = resp.Body.Close() }() if resp.StatusCode != http.StatusOK { - return 0, fmt.Errorf("TEI /info returned status %d", resp.StatusCode) + return info, fmt.Errorf("TEI /info returned status %d", resp.StatusCode) } - var info teiInfoResponse if err := json.NewDecoder(resp.Body).Decode(&info); err != nil { - return 0, fmt.Errorf("failed to decode TEI /info response: %w", err) + return info, fmt.Errorf("failed to decode TEI /info response: %w", err) + } + + return info, nil +} + +// fetchMaxBatchSize queries the TEI /info endpoint and returns the max client batch size. +func fetchMaxBatchSize(baseURL string, httpClient *http.Client) (int, error) { + info, err := fetchInfo(context.Background(), baseURL, httpClient) + if err != nil { + return 0, err } if info.MaxClientBatchSize <= 0 { @@ -95,6 +112,22 @@ func fetchMaxBatchSize(baseURL string, httpClient *http.Client) (int, error) { return info.MaxClientBatchSize, nil } +// ModelID returns the id of the model the TEI server is currently running, +// read from /info on every call. The model is a property of the running +// container, not of this client's configuration, so it is deliberately not +// cached: the point is letting a caller observe a redeploy that swapped the +// model behind an unchanged URL. +func (c *teiClient) ModelID(ctx context.Context) (string, error) { + info, err := fetchInfo(ctx, c.baseURL, c.httpClient) + if err != nil { + return "", err + } + if info.ModelID == "" { + return "", fmt.Errorf("TEI /info reported no model_id") + } + return info.ModelID, nil +} + // embedRequest is the JSON body sent to the TEI /embed endpoint. type embedRequest struct { Inputs []string `json:"inputs"` diff --git a/pkg/vmcp/optimizer/internal/similarity/tei_client_test.go b/pkg/vmcp/optimizer/internal/similarity/tei_client_test.go index 484ae35dfb..2272b8d3bb 100644 --- a/pkg/vmcp/optimizer/internal/similarity/tei_client_test.go +++ b/pkg/vmcp/optimizer/internal/similarity/tei_client_test.go @@ -9,6 +9,7 @@ import ( "fmt" "net/http" "net/http/httptest" + "sync/atomic" "testing" "time" @@ -69,6 +70,89 @@ func Test_newTEIClient(t *testing.T) { }) } +func TestTEIClient_ModelID(t *testing.T) { + t.Parallel() + + t.Run("returns the model id from /info", func(t *testing.T) { + t.Parallel() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, infoPath, r.URL.Path) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"model_id": "BAAI/bge-small-en-v1.5", "max_client_batch_size": 16}`)) + })) + defer srv.Close() + + client, err := newTEIClient(srv.URL, 0) + require.NoError(t, err) + + id, err := client.ModelID(context.Background()) + require.NoError(t, err) + require.Equal(t, "BAAI/bge-small-en-v1.5", id) + }) + + t.Run("reads per call so a swap is observable", func(t *testing.T) { + t.Parallel() + var calls atomic.Int64 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + if calls.Add(1) > 2 { // the constructor's own /info read is call 1 + _, _ = w.Write([]byte(`{"model_id": "model-b", "max_client_batch_size": 16}`)) + return + } + _, _ = w.Write([]byte(`{"model_id": "model-a", "max_client_batch_size": 16}`)) + })) + defer srv.Close() + + client, err := newTEIClient(srv.URL, 0) + require.NoError(t, err) + + first, err := client.ModelID(context.Background()) + require.NoError(t, err) + second, err := client.ModelID(context.Background()) + require.NoError(t, err) + + require.Equal(t, "model-a", first) + require.Equal(t, "model-b", second, + "the id must be read live per call, not cached at construction") + }) + + t.Run("missing model_id is an error", func(t *testing.T) { + t.Parallel() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"max_client_batch_size": 16}`)) + })) + defer srv.Close() + + client, err := newTEIClient(srv.URL, 0) + require.NoError(t, err) + + _, err = client.ModelID(context.Background()) + require.ErrorContains(t, err, "no model_id") + }) + + t.Run("non-200 is an error", func(t *testing.T) { + t.Parallel() + var constructed atomic.Bool + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + if !constructed.Load() { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"model_id": "model-a", "max_client_batch_size": 16}`)) + return + } + w.WriteHeader(http.StatusServiceUnavailable) + })) + defer srv.Close() + + client, err := newTEIClient(srv.URL, 0) + require.NoError(t, err) + constructed.Store(true) + + _, err = client.ModelID(context.Background()) + require.ErrorContains(t, err, "status 503") + }) +} + func TestTEIClient_Embed(t *testing.T) { t.Parallel() diff --git a/pkg/vmcp/optimizer/internal/toolstore/schema.sql b/pkg/vmcp/optimizer/internal/toolstore/schema.sql index 48029efb37..29c5ebbd71 100644 --- a/pkg/vmcp/optimizer/internal/toolstore/schema.sql +++ b/pkg/vmcp/optimizer/internal/toolstore/schema.sql @@ -17,29 +17,11 @@ CREATE TABLE IF NOT EXISTS llm_capabilities ( content_hash TEXT ); --- The reuse probe selects by content_hash across the whole table, not by the +-- The reuse lookup selects by content_hash across the whole table, not by the -- name primary key. CREATE INDEX IF NOT EXISTS llm_capabilities_content_hash_idx ON llm_capabilities (content_hash); --- Embedding of a fixed probe string, used to detect that the embedding backend --- started returning different vectors. --- --- The content hash covers the configured provider, endpoint and model, but the --- TEI model is fixed by the running container rather than by config, so --- swapping it behind an unchanged service URL is invisible to the hash. When --- the replacement has the same vector width the dimension check cannot see it --- either, and the stored vectors would be reused across a change of semantic --- space. Comparing a re-embedded probe against the stored one detects that --- regardless of what the configuration says. --- --- Single row by construction: the probe describes the one backend this store --- talks to. -CREATE TABLE IF NOT EXISTS embedding_canary ( - id INTEGER PRIMARY KEY CHECK (id = 1), - embedding BLOB NOT NULL -); - -- FTS5 virtual table for full-text search with BM25 ranking. -- tokenize='porter' uses the Porter stemming algorithm so that morphological -- variants of a word (e.g. "running", "runs", "ran") match the root form "run". diff --git a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store.go b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store.go index 2cbd8f5807..01c67a7093 100644 --- a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store.go +++ b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store.go @@ -20,7 +20,6 @@ import ( "math" "sort" "strings" - "sync" "sync/atomic" "golang.org/x/sync/errgroup" @@ -82,29 +81,30 @@ type sqliteToolStore struct { hybridSemanticRatio float64 semanticDistanceThreshold float64 - // embeddingIdentity describes the backend that produces embeddings. It is - // mixed into every content hash so vectors are never reused across a - // provider, endpoint, or model change. Immutable after construction. - embeddingIdentity string + // configIdentity digests the configured embedding provider, endpoint and + // model. The model id read live from the backend on every build is folded + // in next to it (see embeddingIdentity), so vectors are not reused across + // any provider, endpoint, or model change the identity can see — including + // a TEI container redeployed with a different model behind an unchanged + // URL, which the config alone cannot. When the id cannot be read the + // identity falls back and reuse proceeds unverified, within the bounds + // documented on embeddingIdentity. Immutable after construction. + configIdentity string + + // lastModelID is the most recent model id successfully read from the + // embedding backend, or nil before the first successful read. It is the + // fallback for a failed read: keeping the last-seen id keeps cache keys + // stable across a transient failure, where switching to a different + // identity would force a full re-embed and a second one on recovery. + // Held by pointer because the store is used by value. + lastModelID *atomic.Pointer[string] // embeddingDim is the vector width most recently observed from the // embedding backend, or 0 before any embedding call has succeeded. It - // bounds which stored vectors may be reused, catching a model swap that - // embeddingIdentity cannot see (the TEI model is fixed by the running - // container, not by config). Held by pointer because the store is used by - // value. + // bounds which stored vectors may be reused, as defense in depth behind + // the model id in the cache key. Held by pointer because the store is + // used by value. embeddingDim *atomic.Int64 - - // canary serializes the backend probe so concurrent builds cannot race on - // the stored probe row. Held by pointer because the store is used by value. - canary *canaryState -} - -// canaryState serializes the backend probe and counts completed probes, so a -// build that waited through one can tell that its result already applies. -type canaryState struct { - mu sync.Mutex - generation atomic.Uint64 } // NewSQLiteToolStore creates a new ToolStore backed by a shared in-memory @@ -164,9 +164,9 @@ func newSQLiteToolStore( maxToolsToReturn: maxTools, hybridSemanticRatio: hybridRatio, semanticDistanceThreshold: semanticThreshold, - embeddingIdentity: embeddingIdentity(cfg), + configIdentity: configIdentity(cfg), + lastModelID: &atomic.Pointer[string]{}, embeddingDim: &atomic.Int64{}, - canary: &canaryState{}, } slog.Debug("optimizer tool store created", @@ -220,6 +220,12 @@ func (s sqliteToolStore) UpsertTools(ctx context.Context, tools []server.ServerT return tx.Commit() } +// maxResolveAttempts bounds how many times one build may restart because the +// embedding model changed under it. Two is enough for the legitimate case of +// a single redeploy; a backend that swaps models on consecutive builds is an +// operational problem no retry count fixes. +const maxResolveAttempts = 2 + // resolveEmbeddings returns an encoded embedding blob and a content hash for // each tool, embedding only the tools whose hash is not already stored. // @@ -229,6 +235,19 @@ func (s sqliteToolStore) UpsertTools(ctx context.Context, tools []server.ServerT // initialize round-trip; reuse makes it O(tools whose text changed). See // stacklok/toolhive#5847. // +// The backend identity is read at the start of an attempt and re-read after +// the embedding batch. If the two differ, the batch spanned a model swap and +// none of its vectors is attributable to either model — each chunk may have +// hit either side of the swap — so the attempt is discarded and re-run under +// the new identity, where the stale rows simply miss by key. Without the +// re-read, vectors produced by the new model would be committed under keys +// naming the old one: wrong vectors under valid-looking keys, which no later +// build would ever re-check. Reused blobs need no such guard: a blob is only +// found under the identity it was committed with, and only identities that +// were verified when the blob was committed ever carry a hash — vectors +// embedded while the id could not be read are committed hashless, searchable +// but never reusable (see the note at the fill loop below). +// // With no embedding client it returns nil blobs and hashes (FTS5-only mode). func (s sqliteToolStore) resolveEmbeddings( ctx context.Context, tools []server.ServerTool, @@ -241,57 +260,77 @@ func (s sqliteToolStore) resolveEmbeddings( } texts := make([]string, len(tools)) - keys := make([]string, len(tools)) for i, tool := range tools { texts[i] = embeddedText(tool.Tool.Name, tool.Tool.Description) - keys[i] = embeddingCacheKey(s.embeddingIdentity, texts[i]) - hashes[i] = sql.NullString{String: keys[i], Valid: true} } - // Discards stored vectors first if the backend has changed under us, so the - // lookup below simply finds nothing to reuse for them. - s.syncBackendProbe(ctx) + keys := make([]string, len(tools)) + for attempt := 1; ; attempt++ { + identity, verified := s.embeddingIdentity(ctx) - cached, err := s.cachedEmbeddings(ctx, keys) - if err != nil { - return nil, nil, err - } + for i := range tools { + keys[i] = embeddingCacheKey(identity, texts[i]) + hashes[i] = sql.NullString{String: keys[i], Valid: true} + } - // Deduplicated by key: the same tool may appear twice in one batch. - missIndexByKey := make(map[string]int, len(keys)) - var missTexts []string - reused := 0 - for i, key := range keys { - if blob, ok := cached[key]; ok { - blobs[i] = blob - reused++ + cached, err := s.cachedEmbeddings(ctx, keys) + if err != nil { + return nil, nil, err + } + + missIndexByKey, missTexts, reused := partitionMisses(blobs, keys, texts, cached) + + slog.Debug("resolved tool embeddings", + "tools", len(tools), "reused", reused, "embedded", len(missTexts)) + + // All cache hits: nothing was embedded, so there is no batch to have + // spanned a swap. The reused vectors are correct for the keys they + // carry by construction. + if len(missTexts) == 0 { + return blobs, hashes, nil + } + + embeddings, err := s.embedTexts(ctx, missTexts) + if err != nil { + return nil, nil, err + } + + after, afterVerified := s.embeddingIdentity(ctx) + if after != identity { + if attempt >= maxResolveAttempts { + return nil, nil, fmt.Errorf( + "embedding model changed during %d consecutive build attempts; giving up", attempt) + } + slog.Warn("embedding model changed during the batch; discarding it and re-embedding under the new identity") continue } - if _, seen := missIndexByKey[key]; !seen { - missIndexByKey[key] = len(missTexts) - missTexts = append(missTexts, texts[i]) + verified = verified && afterVerified + + // Published only for a batch that will commit: a discarded attempt's + // width must not make concurrent builds reject their own valid rows. + if n := len(embeddings[0]); n > 0 { + s.embeddingDim.Store(int64(n)) } - } - slog.Debug("resolved tool embeddings", - "tools", len(tools), "reused", reused, "embedded", len(missTexts)) + if err := fillMisses(blobs, hashes, keys, tools, missIndexByKey, embeddings, verified); err != nil { + return nil, nil, err + } - if len(missTexts) == 0 { return blobs, hashes, nil } +} - embeddings, err := s.embeddingClient.EmbedBatch(ctx, missTexts) - if err != nil { - return nil, nil, fmt.Errorf("failed to generate embeddings: %w", err) - } - if len(embeddings) != len(missTexts) { - return nil, nil, fmt.Errorf("embedding client returned %d embeddings for %d inputs", - len(embeddings), len(missTexts)) - } - if n := len(embeddings[0]); n > 0 { - s.embeddingDim.Store(int64(n)) - } - +// fillMisses encodes the freshly embedded vectors into their tools' slots. +// When the attempt's identity was not verified on both sides of the batch, +// the fresh rows are committed hashless — searchable but never reusable. A +// hashed row here would be permanent poison if the backend is later rolled +// back to the fallback model: the next verified build would derive the very +// identity these keys name, cache-hit the mislabelled vectors, and re-key +// nothing. +func fillMisses( + blobs [][]byte, hashes []sql.NullString, keys []string, + tools []server.ServerTool, missIndexByKey map[string]int, embeddings [][]float32, verified bool, +) error { for i, key := range keys { if blobs[i] != nil { continue @@ -300,158 +339,100 @@ func (s sqliteToolStore) resolveEmbeddings( if !ok { // Unreachable, but a missing key would otherwise index 0 and store // another tool's vector under this name. - return nil, nil, fmt.Errorf("no embedding resolved for tool %s", tools[i].Tool.Name) + return fmt.Errorf("no embedding resolved for tool %s", tools[i].Tool.Name) } blobs[i] = encodeEmbedding(embeddings[idx]) + if !verified { + hashes[i] = sql.NullString{} + } } - - return blobs, hashes, nil + return nil } -// canaryText is the fixed probe embedded to detect a change of embedding -// backend. Its content is arbitrary but must never change: an edit would make -// every stored canary incomparable and force one needless full re-embed. -const canaryText = "toolhive optimizer embedding canary v1" - -// canaryMaxDistance is the cosine distance below which two probe embeddings are -// considered to come from the same backend. -// -// The threshold is not zero because a backend is not required to be -// deterministic: reduction order can differ across hardware and runtimes, so the -// same model may return slightly different vectors on different deployments. -// It is small because the signal it must not miss is large — two different -// models of equal width place the same text roughly a full unit apart, i.e. -// effectively orthogonal, so this leaves two orders of magnitude of margin. -const canaryMaxDistance = 0.01 - -// syncBackendProbe discards stored embeddings when the embedding backend -// has started returning different vectors. -// -// It runs on every build, not once per store. The store lives for the whole -// process and the embedding service is addressed by a stable URL, so the backend -// can be replaced underneath a running server — redeploying the embedding -// service with a different model of the same width changes neither the URL nor -// the vector length, leaving it invisible to both the content hash and the -// dimension check. Probing once at startup would never see it, and reuse would -// then serve vectors from the old model for the rest of the process's life. -// Before embedding reuse existed the next session simply re-embedded everything -// and the condition healed on its own. -// -// Cost is one embedding per build, against the ~140 it saves. -// -// Every failure path leaves the stored embeddings reusable. A probe that cannot -// be taken says the backend is unreachable, not that it changed — and an -// unreachable backend cannot re-embed the catalogue either, so refusing reuse -// would turn a working build into a failed one. Serving possibly-stale vectors -// is bounded: the next build with a reachable backend re-probes and discards -// them. A build that still works beats a client with no tools. -func (s sqliteToolStore) syncBackendProbe(ctx context.Context) { - // Every build must be ordered after any in-flight probe, because a probe may - // be about to discard the very vectors this build is about to read and write - // back. Skipping the lock instead of waiting for it lets a build read - // pre-discard rows and re-insert them — restoring both the stale vector and - // its content hash — after which the freshly written probe certifies the - // store as current and nothing ever re-checks. That is permanent, not - // bounded. See TestSQLiteToolStore_ConcurrentBuilds_OrderedAfterProbe. - // - // Waiting is still cheap: a build that waited through someone else's probe - // sees the generation move and skips the network call, so a burst of - // concurrent builds costs one embedding between them, not one each. - gen := s.canary.generation.Load() - s.canary.mu.Lock() - defer s.canary.mu.Unlock() - if s.canary.generation.Load() != gen { - return - } - - // Embedded outside any transaction so no database lock is held across the - // network call (see the note in UpsertTools). - probe, err := s.embeddingClient.Embed(ctx, canaryText) - if err != nil { - slog.Warn("could not probe the embedding backend; reusing stored embeddings unverified", "error", err) - return - } - if len(probe) == 0 { - slog.Warn("embedding backend returned an empty probe; reusing stored embeddings unverified") - return +// partitionMisses fills blobs with the cached vectors and returns the texts +// still to embed, deduplicated by key — the same tool may appear twice in one +// batch — along with the number of reused entries. Entries with no cached +// vector are reset to nil: a retry after a model swap must not carry blobs +// reused under the previous attempt's identity. +func partitionMisses( + blobs [][]byte, keys, texts []string, cached map[string][]byte, +) (map[string]int, []string, int) { + missIndexByKey := make(map[string]int, len(keys)) + var missTexts []string + reused := 0 + for i, key := range keys { + if blob, ok := cached[key]; ok { + blobs[i] = blob + reused++ + continue + } + blobs[i] = nil + if _, seen := missIndexByKey[key]; !seen { + missIndexByKey[key] = len(missTexts) + missTexts = append(missTexts, texts[i]) + } } - s.embeddingDim.Store(int64(len(probe))) + return missIndexByKey, missTexts, reused +} - changed, err := s.reconcileCanary(ctx, probe) +// embedTexts runs one embedding batch, returning exactly one vector per text. +func (s sqliteToolStore) embedTexts(ctx context.Context, texts []string) ([][]float32, error) { + embeddings, err := s.embeddingClient.EmbedBatch(ctx, texts) if err != nil { - slog.Warn("could not reconcile the embedding probe; reusing stored embeddings unverified", "error", err) - return + return nil, fmt.Errorf("failed to generate embeddings: %w", err) } - - s.canary.generation.Add(1) - if changed { - slog.Warn("embedding backend changed; stored embeddings were discarded and will be recomputed") + if len(embeddings) != len(texts) { + return nil, fmt.Errorf("embedding client returned %d embeddings for %d inputs", + len(embeddings), len(texts)) } + return embeddings, nil } -// reconcileCanary compares probe against the stored canary, discarding every -// stored embedding when they differ, and records probe as the new canary. -// Reports whether stored embeddings were discarded. +// embeddingIdentity derives the backend identity mixed into every cache key +// for one build, folding the live model id into the configured identity. // -// The discard and the new canary are written in one transaction so a failure -// cannot leave a canary that claims vectors are current when they are not. -func (s sqliteToolStore) reconcileCanary(ctx context.Context, probe []float32) (discarded bool, retErr error) { - tx, err := s.db.BeginTx(ctx, nil) - if err != nil { - return false, fmt.Errorf("failed to begin canary transaction: %w", err) - } - defer func() { - if retErr != nil { - _ = tx.Rollback() - } - }() - - var storedBlob []byte - err = tx.QueryRowContext(ctx, "SELECT embedding FROM embedding_canary WHERE id = 1").Scan(&storedBlob) - switch { - case errors.Is(err, sql.ErrNoRows): - // No probe recorded. Any embeddings already present came from a backend - // this store cannot vouch for, so they are not reusable. - var existing int - if err := tx.QueryRowContext(ctx, - "SELECT COUNT(*) FROM llm_capabilities WHERE embedding IS NOT NULL").Scan(&existing); err != nil { - return false, fmt.Errorf("failed to count stored embeddings: %w", err) - } - discarded = existing > 0 - case err != nil: - return false, fmt.Errorf("failed to read the stored canary: %w", err) - default: - stored := decodeEmbedding(storedBlob) - // Compare widths first: cosine distance indexes both slices positionally - // and would panic on a shorter stored vector. - discarded = len(stored) != len(probe) || - similarity.CosineDistance(stored, probe) > canaryMaxDistance - } - - if discarded { - // Clear the vectors but keep the rows: they also back the external-content - // FTS5 index, so deleting them would break keyword search for tools - // outside the current session's set. - if _, err := tx.ExecContext(ctx, - "UPDATE llm_capabilities SET embedding = NULL, content_hash = NULL"); err != nil { - return false, fmt.Errorf("failed to discard stale embeddings: %w", err) +// The id is read from the backend on every build rather than once at +// construction: the store lives for the whole process and the embedding +// service is addressed by a stable URL, so a TEI container can be redeployed +// with a different model behind an unchanged URL — invisible to the config, +// and to the dimension check when the widths match. A live id turns that +// swap into different cache keys, so stale rows simply stop being found; +// there is nothing to detect and nothing to discard. +// +// Every failure path keeps the previous identity and reports it unverified. +// A read that fails says the backend is unreachable, not that it changed — +// and an unreachable backend cannot re-embed the catalogue either, so +// flipping the identity would turn a working build into a full re-embed now +// and a second one on recovery. Before any successful read it degrades to +// the configured identity alone. Previously verified rows stay reusable +// under an unverified identity; what an unverified identity must never do is +// attribute NEW vectors (see resolveEmbeddings). +func (s sqliteToolStore) embeddingIdentity(ctx context.Context) (identity string, verified bool) { + modelID, err := s.embeddingClient.ModelID(ctx) + verified = err == nil && modelID != "" + if !verified { + modelID = "" + if last := s.lastModelID.Load(); last != nil { + modelID = *last } + slog.Warn("could not read the embedding model id; proceeding with the last known identity unverified", + "error", err, "assumed_model_id", modelID) + } else { + s.lastModelID.Store(&modelID) } - - if _, err := tx.ExecContext(ctx, - "INSERT OR REPLACE INTO embedding_canary (id, embedding) VALUES (1, ?)", - encodeEmbedding(probe)); err != nil { - return false, fmt.Errorf("failed to record the canary: %w", err) - } - - return discarded, tx.Commit() + // Digested rather than joined so the identity is fixed-length and cannot + // shift the field boundaries of the cache key it is folded into. + return hashParts(s.configIdentity, modelID), verified } // cachedEmbeddings returns the reusable stored embeddings among the given // content hashes, keyed by hash. Hashes with no usable vector are absent. // -// Matching on content_hash rather than tool name lets a renamed tool keep its -// vector. Runs outside any transaction — see the lock note in UpsertTools. +// Matching on content_hash rather than tool name confirms in one lookup that +// the embedded text and the producing backend are both unchanged; a changed +// description or a repointed backend cannot quietly reuse a vector. A rename +// changes the text, so it re-embeds. Runs outside any transaction — see the +// lock note in UpsertTools. // // A stale-width vector is treated as a miss rather than reused: reuse would be // permanent, since it is handed back and re-stored on every rebuild while @@ -541,15 +522,9 @@ func hashParts(parts ...string) string { return hex.EncodeToString(h.Sum(nil)) } -// embeddingIdentity derives the backend identity mixed into every content hash. -// -// Known limitation: for the TEI provider the model is fixed by the running -// container rather than by config, so swapping the model behind an unchanged -// service URL is not detected here. Search tolerates the resulting stale -// vectors (see searchSemantic) but they remain semantically stale until the -// process restarts. Reading the model id from the TEI /info endpoint would -// close this gap. -func embeddingIdentity(cfg *types.OptimizerConfig) string { +// configIdentity digests the configured half of the backend identity; the +// model id read live per build completes it (see embeddingIdentity). +func configIdentity(cfg *types.OptimizerConfig) string { if cfg == nil { return "" } @@ -780,16 +755,14 @@ func (s sqliteToolStore) searchSemantic( candidatesEvaluated++ emb := decodeEmbedding(embBlob) - // Cosine distance indexes both slices positionally: a shorter stored - // vector panics, a longer one silently ignores its tail. A mismatch - // means the vector survived a model change (see embeddingIdentity). - if len(emb) != len(queryVec) { + // A width mismatch means the vector survived a model change (see + // embeddingIdentity); skip it rather than fail the whole search. + dist, err := similarity.CosineDistance(queryVec, emb) + if err != nil { dimensionMismatches++ continue } - dist := similarity.CosineDistance(queryVec, emb) - // Filter by semantic distance threshold. // This is meaningful only for cosine distance (semantic search). // FTS5 ranks are normalized BM25 scores, not true distance measures. diff --git a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_cache_test.go b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_cache_test.go index 0cfba89f4e..d1c83d3ae7 100644 --- a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_cache_test.go +++ b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_cache_test.go @@ -24,8 +24,9 @@ import ( // avoided rather than on wall-clock time. type countingEmbeddingClient struct { *fakeEmbeddingClient - texts atomic.Int64 - calls atomic.Int64 + texts atomic.Int64 + calls atomic.Int64 + modelIDGets atomic.Int64 mu sync.Mutex embedded []string @@ -49,6 +50,11 @@ func (c *countingEmbeddingClient) EmbedBatch(ctx context.Context, texts []string return c.fakeEmbeddingClient.EmbedBatch(ctx, texts) } +func (c *countingEmbeddingClient) ModelID(ctx context.Context) (string, error) { + c.modelIDGets.Add(1) + return c.fakeEmbeddingClient.ModelID(ctx) +} + // textsEmbedded returns the total number of texts sent to the backend. func (c *countingEmbeddingClient) textsEmbedded() int { return int(c.texts.Load()) } @@ -185,6 +191,12 @@ func TestSQLiteToolStore_UpsertTools_ReusesCachedEmbeddings(t *testing.T) { "an unchanged tool set must not be re-embedded on subsequent builds") assert.Equal(t, 1, client.batchCalls(), "rebuilds over an unchanged tool set must not reach the embedding backend at all") + + // The cold build reads the id twice (once per side of its batch); a warm + // build reads it once. This bounds the backend round-trips on the warm + // path, which the probe this design replaced paid one embedding for. + assert.Equal(t, int64(2+3), client.modelIDGets.Load(), + "a warm build must cost exactly one model id read and nothing else") } // TestSQLiteToolStore_UpsertTools_ReEmbedsOnlyChangedTools asserts the cache is @@ -388,13 +400,17 @@ func TestSQLiteToolStore_UpsertTools_IgnoresEmptyStoredEmbedding(t *testing.T) { // shiftedEmbeddingClient produces vectors of the configured width that differ // from fakeEmbeddingClient's for the same text, standing in for a replacement -// model of identical width — the swap neither the content hash nor the -// dimension check can see. +// model of identical width. It reports its own model id, which is how the +// replacement becomes visible to the cache key. type shiftedEmbeddingClient struct { *fakeEmbeddingClient calls atomic.Int64 } +func (*shiftedEmbeddingClient) ModelID(context.Context) (string, error) { + return "shifted-model", nil +} + func newShiftedEmbeddingClient() *shiftedEmbeddingClient { return &shiftedEmbeddingClient{fakeEmbeddingClient: newFakeEmbeddingClient(countingClientDim)} } @@ -428,7 +444,8 @@ func (c *shiftedEmbeddingClient) EmbedBatch(ctx context.Context, texts []string) // // This is what redeploying the embedding service with a different model looks // like to a running vmcp: the Service URL is unchanged, so the store keeps the -// same client and never learns the backend moved. +// same client. The swap is observable only through ModelID, which reports the +// model currently serving — exactly what the real TEI client reads from /info. type swappableEmbeddingClient struct { *fakeEmbeddingClient swapped atomic.Bool @@ -441,8 +458,13 @@ func newSwappableEmbeddingClient() *swappableEmbeddingClient { func (c *swappableEmbeddingClient) swap() { c.swapped.Store(true) } -// Embed counts every text, including the backend probe, which is issued here -// rather than through EmbedBatch. +func (c *swappableEmbeddingClient) ModelID(context.Context) (string, error) { + if c.swapped.Load() { + return "swapped-model", nil + } + return fakeModelID, nil +} + func (c *swappableEmbeddingClient) Embed(ctx context.Context, text string) ([]float32, error) { c.texts.Add(1) vec, err := c.fakeEmbeddingClient.Embed(ctx, text) @@ -474,11 +496,12 @@ func (c *swappableEmbeddingClient) embedded() int { return int(c.texts.Load()) } // TestSQLiteToolStore_BackendChange_DiscardsStaleEmbeddings is the regression // test for the failure this cache would otherwise introduce. // -// A replacement model of the same width is invisible to the content hash (the -// config is unchanged) and to the dimension check (the width is unchanged), so -// without the backend probe the stale vectors would be reused and re-stored on -// every later build — permanently, where re-embedding every session used to heal -// it on its own. +// A replacement model of the same width is invisible to the configured +// identity (the config is unchanged) and to the dimension check (the width is +// unchanged). The model id read per build is what makes it visible: a +// different id produces different cache keys, so every stale row misses and +// the catalogue is re-embedded — where keying on text alone would reuse and +// re-store the stale vectors permanently. // // The topology matters: production creates ONE store per process // (NewOptimizerFactory) and every session build calls UpsertTools on it, so the @@ -494,20 +517,19 @@ func TestSQLiteToolStore_BackendChange_DiscardsStaleEmbeddings(t *testing.T) { require.NoError(t, store.UpsertTools(ctx, tools), "first build") afterCold := client.embedded() - require.Equal(t, len(tools)+1, afterCold, "cold build embeds every tool plus the probe") + require.Equal(t, len(tools), afterCold, "cold build embeds every tool") require.NoError(t, store.UpsertTools(ctx, tools), "second build, backend unchanged") - afterWarm := client.embedded() - require.Equal(t, afterCold+1, afterWarm, - "an unchanged backend re-embeds only the probe") + require.Equal(t, afterCold, client.embedded(), + "an unchanged backend embeds nothing on a rebuild") // The embedding service is redeployed with a different model. Same URL, same - // client, same store — only the vectors change. + // client, same store — only the vectors and the reported model id change. client.swap() require.NoError(t, store.UpsertTools(ctx, tools), "build after the model swap") - assert.Equal(t, afterWarm+1+len(tools), client.embedded(), - "a changed backend must re-embed every tool, not just the probe") + assert.Equal(t, afterCold+len(tools), client.embedded(), + "a changed model id must re-embed every tool") var stored []byte require.NoError(t, store.db.QueryRowContext(ctx, @@ -520,13 +542,14 @@ func TestSQLiteToolStore_BackendChange_DiscardsStaleEmbeddings(t *testing.T) { // TestSQLiteToolStore_BackendUnreachable_StillServesTools pins the availability // property this cache buys: once warm, a build survives the embedding backend -// being completely down. +// being completely down — the model id read fails along with everything else, +// and the identity falls back to the last id seen. // // Measured live on 2026-07-25 — all four TEI pods deleted, a session built // normally from cache. It is asserted here because that scenario cannot run in -// CI, and because it is fragile: making the probe refuse reuse on failure -// silently destroys it, turning a working build into a client with no tools. -// That is the original stacklok/toolhive#5847 symptom. +// CI, and because it is fragile: flipping the identity when the id cannot be +// read silently destroys it, turning a working build into a client with no +// tools. That is the original stacklok/toolhive#5847 symptom. func TestSQLiteToolStore_BackendUnreachable_StillServesTools(t *testing.T) { t.Parallel() @@ -554,7 +577,7 @@ func TestSQLiteToolStore_BackendUnreachable_StillServesTools(t *testing.T) { } // flakyEmbeddingClient fails every call once down is set, modelling the -// embedding service being unreachable. +// embedding service being unreachable — including its model id endpoint. type flakyEmbeddingClient struct { *fakeEmbeddingClient down atomic.Bool @@ -565,6 +588,13 @@ func newFlakyEmbeddingClient() *flakyEmbeddingClient { return &flakyEmbeddingClient{fakeEmbeddingClient: newFakeEmbeddingClient(countingClientDim)} } +func (c *flakyEmbeddingClient) ModelID(ctx context.Context) (string, error) { + if c.down.Load() { + return "", fmt.Errorf("embedding backend unreachable") + } + return c.fakeEmbeddingClient.ModelID(ctx) +} + func (c *flakyEmbeddingClient) Embed(ctx context.Context, text string) ([]float32, error) { if c.down.Load() { return nil, fmt.Errorf("embedding backend unreachable") @@ -587,244 +617,331 @@ func (c *flakyEmbeddingClient) EmbedBatch(ctx context.Context, texts []string) ( func (c *flakyEmbeddingClient) embedded() int { return int(c.texts.Load()) } -// blockingProbeClient holds the probe open until released, so a test can -// observe what a concurrent build does while a probe is in flight. -type blockingProbeClient struct { - *fakeEmbeddingClient - entered chan struct{} - release chan struct{} +// midBatchSwapClient swaps the model at the moment the swapOnBatch-th +// embedding batch begins, so that batch's vectors come from the new model +// while the build's identity was read under the old one — the narrowest form +// of the swap window. With infoDownAfterSwap, the model id also becomes +// unreadable from the swap on, modelling a redeploy that takes /info away in +// the same instant it changes the model. +type midBatchSwapClient struct { + *swappableEmbeddingClient + batches atomic.Int64 + swapOnBatch int64 + infoDownAfterSwap bool } -func newBlockingProbeClient() *blockingProbeClient { - return &blockingProbeClient{ - fakeEmbeddingClient: newFakeEmbeddingClient(countingClientDim), - entered: make(chan struct{}, 1), - release: make(chan struct{}), - } +func newMidBatchSwapClient() *midBatchSwapClient { + return &midBatchSwapClient{swappableEmbeddingClient: newSwappableEmbeddingClient(), swapOnBatch: 1} } -func (c *blockingProbeClient) Embed(ctx context.Context, text string) ([]float32, error) { - if text == canaryText { - select { - case c.entered <- struct{}{}: - default: - } - select { - case <-c.release: - case <-ctx.Done(): - return nil, ctx.Err() - } +func (c *midBatchSwapClient) ModelID(ctx context.Context) (string, error) { + if c.infoDownAfterSwap && c.swapped.Load() { + return "", fmt.Errorf("info endpoint unreachable") } - return c.fakeEmbeddingClient.Embed(ctx, text) + return c.swappableEmbeddingClient.ModelID(ctx) } -func (c *blockingProbeClient) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) { - out := make([][]float32, len(texts)) - for i, text := range texts { - vec, err := c.Embed(ctx, text) - if err != nil { - return nil, err - } - out[i] = vec +func (c *midBatchSwapClient) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) { + if c.batches.Add(1) == c.swapOnBatch { + c.swap() } - return out, nil + return c.swappableEmbeddingClient.EmbedBatch(ctx, texts) } -// TestSQLiteToolStore_ConcurrentBuilds_OrderedAfterProbe asserts no build reads -// the cache while a probe is in flight. -// -// A probe may be about to discard the vectors a concurrent build is reading. If -// that build is allowed to proceed, its INSERT OR REPLACE writes the stale -// vector AND its content hash back after the discard commits. The identity is -// unchanged on a same-width backend swap, so later builds recompute the same -// hash, find the restored row, and reuse it — while the freshly written probe -// certifies the store as current, so nothing re-checks. Permanently stale, with -// nothing logged. +// TestSQLiteToolStore_ModelSwapDuringBatch_ReembedsUnderNewIdentity covers the +// window the post-batch identity re-read exists for: the model is replaced +// while a build sits in EmbedBatch, which takes seconds against a real backend. // -// Found by cross-model review; the earlier TryLock implementation had exactly -// this hole. -func TestSQLiteToolStore_ConcurrentBuilds_OrderedAfterProbe(t *testing.T) { +// Committing that batch would store the new model's vectors under keys naming +// the old model. A later build under the new id re-embeds and heals — but a +// swap back to the old model recomputes exactly those keys and reuses the +// mislabelled vectors as the old model's, permanently. And a batch the swap +// landed in the middle of is a mixture no single identity describes. The +// re-read discards the attempt and re-runs it under the new identity instead. +func TestSQLiteToolStore_ModelSwapDuringBatch_ReembedsUnderNewIdentity(t *testing.T) { t.Parallel() ctx := context.Background() - client := newBlockingProbeClient() + client := newMidBatchSwapClient() store := newTestStore(t, client, nil) + tools := catalog(5) - go func() { _ = store.UpsertTools(ctx, catalog(3)) }() - <-client.entered // a probe is now in flight and holding the lock + require.NoError(t, store.UpsertTools(ctx, tools), "build spanning the swap must still succeed") + require.Equal(t, int64(2), client.batches.Load(), + "the batch that spanned the swap must be discarded and re-run") - done := make(chan struct{}) - go func() { - _ = store.UpsertTools(ctx, makeTools( - mcp.NewTool("concurrent", mcp.WithDescription("Indexed while a probe is in flight")))) - close(done) - }() + // The committed vectors must be the new model's, stored under the new + // model's keys: a rebuild under the (unchanged) new id reuses everything. + before := client.embedded() + require.NoError(t, store.UpsertTools(ctx, tools), "rebuild under the new model") + assert.Equal(t, before, client.embedded(), + "a rebuild after the swap settled must embed nothing") - select { - case <-done: - close(client.release) - t.Fatal("a build completed while a probe was in flight; it can write pre-discard state back") - case <-time.After(600 * time.Millisecond): - // Correct: the second build is ordered behind the probe. - } - close(client.release) + var stored []byte + require.NoError(t, store.db.QueryRowContext(ctx, + "SELECT embedding FROM llm_capabilities WHERE name = ?", tools[0].Tool.Name).Scan(&stored)) + want, err := client.Embed(ctx, embeddedText(tools[0].Tool.Name, tools[0].Tool.Description)) + require.NoError(t, err) + assert.Equal(t, encodeEmbedding(want), stored, + "the committed vector must come from the post-swap model") +} - select { - case <-done: - case <-time.After(30 * time.Second): - t.Fatal("timeout waiting for the queued build to finish") +// TestSQLiteToolStore_ModelSwapDuringBatch_DropsBlobsReusedUnderOldIdentity +// covers the retry from a WARM store: attempt one reuses most of the catalogue +// under the old identity and embeds only the changed tools; the swap lands in +// that batch, so the retry must drop the reused blobs too — every vector it +// commits, including the unchanged tools', must come from the new model. This +// is the path where forgetting to reset carried-over blobs commits old-model +// vectors under new-model keys. +func TestSQLiteToolStore_ModelSwapDuringBatch_DropsBlobsReusedUnderOldIdentity(t *testing.T) { + t.Parallel() + + ctx := context.Background() + client := newMidBatchSwapClient() + client.swapOnBatch = 2 // the warm build below is batch 1 + store := newTestStore(t, client, nil) + + require.NoError(t, store.UpsertTools(ctx, catalog(5)), "warm build under the old model") + + changed := catalog(5) + changed[4].Tool.Description = "Reworded while the model swaps" + require.NoError(t, store.UpsertTools(ctx, changed), "build spanning the swap") + require.Equal(t, int64(3), client.batches.Load(), + "the spanning batch must be discarded and the whole set re-embedded") + + // An unchanged tool is the telling one: attempt one reused its old-model + // blob, and that blob must not have survived into the commit. + var stored []byte + require.NoError(t, store.db.QueryRowContext(ctx, + "SELECT embedding FROM llm_capabilities WHERE name = ?", changed[0].Tool.Name).Scan(&stored)) + want, err := client.Embed(ctx, embeddedText(changed[0].Tool.Name, changed[0].Tool.Description)) + require.NoError(t, err) + assert.Equal(t, encodeEmbedding(want), stored, + "a blob reused under the pre-swap identity must be re-embedded, not committed under the new keys") +} + +// flappingModelIDClient reports a different model id on every read, modelling +// a backend that cannot settle — e.g. two replicas running different models +// behind one Service. +type flappingModelIDClient struct { + *fakeEmbeddingClient + reads atomic.Int64 +} + +func (c *flappingModelIDClient) ModelID(context.Context) (string, error) { + if c.reads.Add(1)%2 == 1 { + return "model-a", nil } + return "model-b", nil } -// TestSQLiteToolStore_BackendUnchanged_KeepsReuse asserts the probe does not -// invalidate the cache when the backend is the same, which would silently undo -// the reuse this cache exists for. -func TestSQLiteToolStore_BackendUnchanged_KeepsReuse(t *testing.T) { +// TestSQLiteToolStore_ModelIDFlapping_FailsTheBuild pins the give-up bound: +// a backend whose identity moves on every read must fail the build with a +// clear error, not retry forever or commit under an arbitrary identity. +func TestSQLiteToolStore_ModelIDFlapping_FailsTheBuild(t *testing.T) { + t.Parallel() + + client := &flappingModelIDClient{fakeEmbeddingClient: newFakeEmbeddingClient(countingClientDim)} + store := newTestStore(t, client, nil) + + err := store.UpsertTools(context.Background(), catalog(3)) + require.ErrorContains(t, err, "model changed during", + "an identity that moves on every read must fail the build after the retry budget") +} + +// TestSQLiteToolStore_SwapWithUnreadableID_CommitsFailOpen pins the documented +// fail-open bound: when the swap lands mid-batch AND the id becomes unreadable +// in the same instant, the post-batch re-read falls back to the pre-batch id, +// the check passes vacuously, and the batch commits — searchable, but never +// reusable (see TestSQLiteToolStore_RollbackAfterUnverifiedCommit_NeverReusesPoison +// for why the unverified rows must not carry a hash). Failing the build +// instead would turn every transient /info outage into an outage of its own. +func TestSQLiteToolStore_SwapWithUnreadableID_CommitsFailOpen(t *testing.T) { t.Parallel() ctx := context.Background() - dsn := fmt.Sprintf("file:canarysame_%d?mode=memory&cache=shared", testDBCounter.Add(1)) + client := newMidBatchSwapClient() + client.infoDownAfterSwap = true + store := newTestStore(t, client, nil) tools := catalog(5) - first := newTestStoreDSN(t, dsn, newCountingEmbeddingClient()) - require.NoError(t, first.UpsertTools(ctx, tools)) - - second := newCountingEmbeddingClient() - rebuilt := newTestStoreDSN(t, dsn, second) - require.NoError(t, rebuilt.UpsertTools(ctx, tools)) + require.NoError(t, store.UpsertTools(ctx, tools), + "a swap the store cannot observe must not fail the build") + require.Equal(t, int64(1), client.batches.Load(), + "the unverifiable batch commits; there is nothing to compare a retry against") - assert.Zero(t, second.textsEmbedded(), - "an unchanged backend must keep every stored vector reusable") + // The id becomes readable again: the next build re-keys under the real + // post-swap identity and re-embeds — the bounded self-healing. + client.infoDownAfterSwap = false + before := client.embedded() + require.NoError(t, store.UpsertTools(ctx, tools), "first build after /info recovers") + assert.Equal(t, before+len(tools), client.embedded(), + "recovering the id must re-key and re-embed the mislabelled rows") } -// TestSQLiteToolStore_BackendChange_PreservesKeywordSearch asserts the stale -// vectors are cleared without removing the rows, which also back the -// external-content FTS5 index. -func TestSQLiteToolStore_BackendChange_PreservesKeywordSearch(t *testing.T) { +// TestSQLiteToolStore_RollbackAfterUnverifiedCommit_NeverReusesPoison covers +// the case where "the next readable build re-keys everything" does NOT heal: +// vectors embedded by model B but committed under model A's identity (swap +// with the id unreadable — the fail-open window) are poison if the backend is +// then rolled back to A. The next readable build derives A's identity again, +// so the keys never change and a hashed poison row would be cache-hit forever. +// The guard: rows committed under an unverified identity carry no content +// hash — searchable, never reusable — so the rollback build re-embeds them. +func TestSQLiteToolStore_RollbackAfterUnverifiedCommit_NeverReusesPoison(t *testing.T) { t.Parallel() ctx := context.Background() - dsn := fmt.Sprintf("file:canaryfts_%d?mode=memory&cache=shared", testDBCounter.Add(1)) - indexed := newTestStoreDSN(t, dsn, newCountingEmbeddingClient()) - require.NoError(t, indexed.UpsertTools(ctx, makeTools( - mcp.NewTool("archive_file", mcp.WithDescription("Archive a file to cold storage"))))) + client := newMidBatchSwapClient() + client.infoDownAfterSwap = true + store := newTestStore(t, client, nil) + tools := catalog(5) - // A different backend, and a build whose tool set does not include the - // previously indexed tool. - swapped := newTestStoreDSN(t, dsn, newShiftedEmbeddingClient()) - require.NoError(t, swapped.UpsertTools(ctx, makeTools( - mcp.NewTool("unrelated_tool", mcp.WithDescription("Something else entirely"))))) + // Swap lands mid-batch with /info down: model-b vectors, identity model-a. + require.NoError(t, store.UpsertTools(ctx, tools), "unverified build must still succeed") - found, err := swapped.searchFTS5(ctx, "archive", []string{"archive_file"}, DefaultMaxToolsToReturn) + // The operator rolls back to model-a and /info recovers. The identity the + // next build derives is the very one the poison was committed under. + client.swapped.Store(false) + client.infoDownAfterSwap = false + + before := client.embedded() + require.NoError(t, store.UpsertTools(ctx, tools), "build after the rollback") + require.Equal(t, before+len(tools), client.embedded(), + "rows committed under an unverified identity must be re-embedded, never reused") + + var stored []byte + require.NoError(t, store.db.QueryRowContext(ctx, + "SELECT embedding FROM llm_capabilities WHERE name = ?", tools[0].Tool.Name).Scan(&stored)) + want, err := client.Embed(ctx, embeddedText(tools[0].Tool.Name, tools[0].Tool.Description)) require.NoError(t, err) - assert.Equal(t, []string{"archive_file"}, matchNames(found), - "discarding stale vectors must not remove the rows backing keyword search") + assert.Equal(t, encodeEmbedding(want), stored, + "after the rollback the stored vector must be model-a's, not the mislabelled model-b one") } -// probeCountingClient counts probe embeds and holds each one open, so a test -// can tell whether concurrent builds queue behind the probe or skip it. -type probeCountingClient struct { - *fakeEmbeddingClient - probes atomic.Int64 - delay time.Duration +// infoDownClient keeps embedding while its model id read fails, modelling a +// backend whose /info route is broken or filtered while /embed still works. +type infoDownClient struct { + *countingEmbeddingClient + infoDown atomic.Bool } -func (c *probeCountingClient) Embed(ctx context.Context, text string) ([]float32, error) { - if text == canaryText { - c.probes.Add(1) - select { - case <-time.After(c.delay): - case <-ctx.Done(): - return nil, ctx.Err() - } - } - return c.fakeEmbeddingClient.Embed(ctx, text) +func newInfoDownClient() *infoDownClient { + return &infoDownClient{countingEmbeddingClient: newCountingEmbeddingClient()} } -func (c *probeCountingClient) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) { - out := make([][]float32, len(texts)) - for i, text := range texts { - vec, err := c.Embed(ctx, text) - if err != nil { - return nil, err - } - out[i] = vec +func (c *infoDownClient) ModelID(ctx context.Context) (string, error) { + if c.infoDown.Load() { + return "", fmt.Errorf("info endpoint unreachable") } - return out, nil + return c.countingEmbeddingClient.ModelID(ctx) } -// TestSQLiteToolStore_ConcurrentBuilds_ShareOneProbe asserts that concurrent -// builds skip the probe when one is already in flight, rather than queueing -// behind it. +// TestSQLiteToolStore_ModelIDUnreadable_KeepsKeysStable asserts a failed model +// id read falls back to the last id seen instead of changing the identity. // -// The probe runs on every build and costs a network round-trip. Serializing it -// would put one round-trip per concurrent session on the tools/list path: -// measured at 4 builds x a 200ms probe = 805ms before this, 201ms after. The -// probe is idempotent, so its result applies to every build racing it. -func TestSQLiteToolStore_ConcurrentBuilds_ShareOneProbe(t *testing.T) { +// Flipping the identity on a transient read failure would invalidate every +// key, force a full re-embed, and force a second one when the read recovers — +// cache churn in both directions, caused by the very mechanism that exists to +// avoid re-embedding. +func TestSQLiteToolStore_ModelIDUnreadable_KeepsKeysStable(t *testing.T) { t.Parallel() ctx := context.Background() - client := &probeCountingClient{ - fakeEmbeddingClient: newFakeEmbeddingClient(countingClientDim), - delay: 200 * time.Millisecond, - } + client := newInfoDownClient() store := newTestStore(t, client, nil) - // Warm first so the concurrent builds do the probe and nothing else. - require.NoError(t, store.UpsertTools(ctx, catalog(5))) - client.probes.Store(0) + require.NoError(t, store.UpsertTools(ctx, catalog(5)), "warm build with the id readable") + require.Equal(t, 5, client.textsEmbedded()) - const builds = 4 - var wg sync.WaitGroup - errs := make(chan error, builds) - for i := range builds { - wg.Add(1) - go func(idx int) { - defer wg.Done() - errs <- store.UpsertTools(ctx, makeTools(mcp.NewTool( - fmt.Sprintf("concurrent_%d", idx), mcp.WithDescription("A concurrently indexed tool")))) - }(i) - } - waitOrFail(t, &wg, "concurrent probing builds") - close(errs) - for err := range errs { - require.NoError(t, err) - } + // The id becomes unreadable; the catalogue gains one tool. + client.infoDown.Store(true) + grown := append(catalog(5), makeTools( + mcp.NewTool("late_arrival", mcp.WithDescription("Registered while /info is down")))...) - assert.Less(t, int(client.probes.Load()), builds, - "concurrent builds must skip an in-flight probe, not queue behind it") + require.NoError(t, store.UpsertTools(ctx, grown), "build with the id unreadable") + assert.Equal(t, 6, client.textsEmbedded(), + "only the new tool may be embedded: reuse must survive an unreadable model id") } -// TestSQLiteToolStore_BackendProbe_RunsOnce asserts concurrent first builds -// probe the backend once rather than once each. -func TestSQLiteToolStore_BackendProbe_RunsOnce(t *testing.T) { +// TestSQLiteToolStore_ModelIDNeverRead_ServesWithoutCaching covers the +// remaining fallback step: a store whose model id has never been readable +// still builds and serves, but cannot attribute what it embeds, so nothing +// it embeds is cached — every unverified build re-embeds. (Caching under the +// config-only identity instead would be rollback poison: see +// TestSQLiteToolStore_RollbackAfterUnverifiedCommit_NeverReusesPoison.) +// Once the id becomes readable, one verified build seeds the cache and +// reuse begins. +func TestSQLiteToolStore_ModelIDNeverRead_ServesWithoutCaching(t *testing.T) { t.Parallel() ctx := context.Background() - client := newCountingEmbeddingClient() + client := newInfoDownClient() + client.infoDown.Store(true) store := newTestStore(t, client, nil) + tools := catalog(5) - var wg sync.WaitGroup - errs := make(chan error, 4) - for i := range 4 { - wg.Add(1) - go func(idx int) { - defer wg.Done() - errs <- store.UpsertTools(ctx, makeTools(mcp.NewTool( - fmt.Sprintf("tool_%d", idx), mcp.WithDescription("A concurrently indexed tool")))) - }(i) - } - waitOrFail(t, &wg, "concurrent first builds") - close(errs) - for err := range errs { - require.NoError(t, err) - } + require.NoError(t, store.UpsertTools(ctx, tools), "first build with the id unreadable") + require.Equal(t, 5, client.textsEmbedded()) - var canaries int - require.NoError(t, store.db.QueryRowContext(ctx, - "SELECT COUNT(*) FROM embedding_canary").Scan(&canaries)) - assert.Equal(t, 1, canaries, "the probe must record exactly one canary") + require.NoError(t, store.UpsertTools(ctx, tools), "rebuild with the id still unreadable") + require.Equal(t, 10, client.textsEmbedded(), + "vectors embedded under an unverifiable identity must not be reused") + + // The id becomes readable: the first verified build embeds once more and + // its rows, now attributable, seed the cache. + client.infoDown.Store(false) + require.NoError(t, store.UpsertTools(ctx, tools), "first build after the id recovers") + assert.Equal(t, 15, client.textsEmbedded(), + "the first verified build re-embeds and seeds the cache") + + require.NoError(t, store.UpsertTools(ctx, tools), "steady state after recovery") + assert.Equal(t, 15, client.textsEmbedded(), + "reuse must resume once the identity is verified") +} + +// TestSQLiteToolStore_BackendUnchanged_KeepsReuse asserts an unchanged backend +// derives an unchanged identity across processes, so a restart does not +// silently undo the reuse this cache exists for. +func TestSQLiteToolStore_BackendUnchanged_KeepsReuse(t *testing.T) { + t.Parallel() + + ctx := context.Background() + dsn := fmt.Sprintf("file:identsame_%d?mode=memory&cache=shared", testDBCounter.Add(1)) + tools := catalog(5) + + first := newTestStoreDSN(t, dsn, newCountingEmbeddingClient()) + require.NoError(t, first.UpsertTools(ctx, tools)) + + second := newCountingEmbeddingClient() + rebuilt := newTestStoreDSN(t, dsn, second) + require.NoError(t, rebuilt.UpsertTools(ctx, tools)) + + assert.Zero(t, second.textsEmbedded(), + "an unchanged backend must keep every stored vector reusable") +} + +// TestSQLiteToolStore_BackendChange_PreservesKeywordSearch asserts a backend +// change leaves the rows backing the external-content FTS5 index in place: +// stale vectors become unreachable by key, they are never deleted. +func TestSQLiteToolStore_BackendChange_PreservesKeywordSearch(t *testing.T) { + t.Parallel() + + ctx := context.Background() + dsn := fmt.Sprintf("file:identfts_%d?mode=memory&cache=shared", testDBCounter.Add(1)) + indexed := newTestStoreDSN(t, dsn, newCountingEmbeddingClient()) + require.NoError(t, indexed.UpsertTools(ctx, makeTools( + mcp.NewTool("archive_file", mcp.WithDescription("Archive a file to cold storage"))))) + + // A different backend, and a build whose tool set does not include the + // previously indexed tool. + swapped := newTestStoreDSN(t, dsn, newShiftedEmbeddingClient()) + require.NoError(t, swapped.UpsertTools(ctx, makeTools( + mcp.NewTool("unrelated_tool", mcp.WithDescription("Something else entirely"))))) + + found, err := swapped.searchFTS5(ctx, "archive", []string{"archive_file"}, DefaultMaxToolsToReturn) + require.NoError(t, err) + assert.Equal(t, []string{"archive_file"}, matchNames(found), + "discarding stale vectors must not remove the rows backing keyword search") } // TestEmbeddedText pins the exact string sent to the embedding backend. Every diff --git a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_livemodel_test.go b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_livemodel_test.go index 3dcd2725cc..a882ef634f 100644 --- a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_livemodel_test.go +++ b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_livemodel_test.go @@ -8,7 +8,6 @@ import ( "context" "fmt" "os" - "sort" "testing" "time" @@ -94,16 +93,18 @@ func TestLiveModelSwap_DifferentWidth(t *testing.T) { } // TestLiveModelSwap_SameWidth verifies against two real models of identical -// width that a swap invisible to both the content hash and the dimension check -// is still caught, and the stale vectors recomputed. +// width that the swap is caught through the model id in the cache key, and the +// stale vectors recomputed. // -// For the TEI provider the model is fixed by the running container rather than -// by config, so this swap changes nothing the cache key can see; equal widths -// hide it from the dimension check too. Only re-embedding the probe detects it. +// For the TEI provider the model is a property of the running container rather +// than of the config, so this swap changes nothing the configured identity can +// see, and equal widths hide it from the dimension check too. The configured +// halves of the two identities are forced equal below to reproduce that; only +// the live model id separates them. // -// It also measures how far apart the two spaces are, which is what makes the -// probe's tolerance safe: the same model repeats bit-identically, while two -// different models put the same text roughly a full unit apart. +// It also measures how far apart the two spaces are — the same text under two +// models lands roughly a full unit apart — which is the size of the corruption +// a missed swap would serve. func TestLiveModelSwap_SameWidth(t *testing.T) { t.Parallel() @@ -125,7 +126,8 @@ func TestLiveModelSwap_SameWidth(t *testing.T) { "this test requires two models of equal width (%s=%d, %s=%d)", base, len(baseVec), other, len(otherVec)) // The same text embedded by two different models: how different are the spaces? - dist := similarity.CosineDistance(baseVec, otherVec) + dist, err := similarity.CosineDistance(baseVec, otherVec) + require.NoError(t, err) t.Logf("same text under %s vs %s: cosine distance %.4f (width %d)", base, other, dist, len(baseVec)) dsn := fmt.Sprintf("file:livesame_%d?mode=memory&cache=shared", testDBCounter.Add(1)) @@ -135,10 +137,10 @@ func TestLiveModelSwap_SameWidth(t *testing.T) { require.NoError(t, indexed.UpsertTools(ctx, tools)) // A TEI-style swap: the model changes behind an unchanged service URL, so - // the identity the cache key is derived from is byte-identical. The store is - // configured with `base` while the client actually embeds with `other`. + // the configured half of the identity is byte-identical. Only the model id + // the client reports separates the two stores. swapped := liveModelStore(t, dsn, endpoint, other) - swapped.embeddingIdentity = indexed.embeddingIdentity + swapped.configIdentity = indexed.configIdentity require.NoError(t, swapped.UpsertTools(ctx, tools)) var stored []byte @@ -151,91 +153,6 @@ func TestLiveModelSwap_SameWidth(t *testing.T) { "the previous model's vector must not survive the swap") } -// TestLiveModelSwap_StaleDistanceDistribution measures how many stale vectors -// survive the semantic distance filter after an undetectable same-width model -// swap, across a realistic catalogue rather than a single pair. -// -// A single aggregate distance is misleading here. Two unrelated embedding spaces -// give a cosine similarity centred on zero, so per-tool distances scatter around -// 1.0 — and DefaultSemanticDistanceThreshold is exactly 1.0. Whether the filter -// is a real backstop or a coin flip therefore depends on the spread, which only -// a distribution can show. -func TestLiveModelSwap_StaleDistanceDistribution(t *testing.T) { - t.Parallel() - - endpoint := liveModelEndpoint(t) - base := cmp.Or(os.Getenv(liveModelBaseEnv), "bge-m3") - other := os.Getenv(liveModelSameWidthEnv) - if other == "" { - t.Skipf("%s not set; skipping stale-distance distribution measurement", liveModelSameWidthEnv) - } - - ctx := context.Background() - backends := []string{"grafana", "datadog", "argocd", "k8s", "gitnexus", "dbhub", "firecrawl", "context7"} - const toolCount = 64 - - texts := make([]string, toolCount) - for i := range toolCount { - b := backends[i%len(backends)] - texts[i] = embeddedText( - fmt.Sprintf("%s_operation_%03d", b, i), - fmt.Sprintf("Perform operation %d against the %s backend, applying the requested "+ - "filters and returning the matching records with their metadata.", i, b), - ) - } - - // Tools indexed by the outgoing model; queries embedded by the incoming one. - staleVecs := liveEmbedBatch(ctx, t, endpoint, base, texts) - queries := []string{ - "search dashboards", "list kubernetes pods", "run a SQL query", - "fetch a web page", "look up library documentation", "check deployment sync status", - } - queryVecs := liveEmbedBatch(ctx, t, endpoint, other, queries) - require.Equal(t, len(staleVecs[0]), len(queryVecs[0]), "the two models must share a width") - - var dists []float64 - for _, q := range queryVecs { - for _, s := range staleVecs { - dists = append(dists, similarity.CosineDistance(q, s)) - } - } - sort.Float64s(dists) - - belowThreshold := 0 - for _, d := range dists { - if d <= DefaultSemanticDistanceThreshold { - belowThreshold++ - } - } - pct := 100 * float64(belowThreshold) / float64(len(dists)) - - t.Logf("stale-vector distances after a same-width swap (%s indexed, %s querying), n=%d:", - base, other, len(dists)) - t.Logf(" min %.4f | p50 %.4f | max %.4f", dists[0], dists[len(dists)/2], dists[len(dists)-1]) - t.Logf(" %d/%d (%.1f%%) fall at or below the %.1f distance threshold and would be ranked", - belowThreshold, len(dists), pct, DefaultSemanticDistanceThreshold) - - // No pass/fail on the fraction: it is a property of the model pair, not of - // this code. The assertion is only that the measurement is meaningful. - require.NotEmpty(t, dists) -} - -// liveEmbedBatch returns embeddings for several texts from one model. -func liveEmbedBatch(ctx context.Context, t *testing.T, endpoint, model string, texts []string) [][]float32 { - t.Helper() - client, err := similarity.NewEmbeddingClient(&types.OptimizerConfig{ - EmbeddingService: endpoint, - EmbeddingProvider: types.EmbeddingProviderOpenAI, - EmbeddingModel: model, - EmbeddingServiceTimeout: 5 * time.Minute, - }) - require.NoError(t, err) - t.Cleanup(func() { _ = client.Close() }) - vecs, err := client.EmbedBatch(ctx, texts) - require.NoError(t, err) - return vecs -} - // liveEmbed returns one embedding from the live endpoint for the given model. func liveEmbed(ctx context.Context, t *testing.T, endpoint, model, text string) []float32 { t.Helper() diff --git a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_test.go b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_test.go index e2a8e2caa4..7b5c6c5576 100644 --- a/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_test.go +++ b/pkg/vmcp/optimizer/internal/toolstore/sqlite_store_test.go @@ -911,6 +911,14 @@ func (f *fakeEmbeddingClient) EmbedBatch(_ context.Context, texts []string) ([][ return result, nil } +// fakeModelID is the model id every fakeEmbeddingClient reports. Fixtures +// standing in for a replaced or swapped model override ModelID instead. +const fakeModelID = "fake-model" + +func (*fakeEmbeddingClient) ModelID(context.Context) (string, error) { + return fakeModelID, nil +} + // embedVector produces the deterministic vector without recording the text, // keeping embeddedTexts limited to the read path. func (f *fakeEmbeddingClient) embedVector(text string) []float32 { diff --git a/pkg/vmcp/optimizer/internal/types/mocks/mock_types.go b/pkg/vmcp/optimizer/internal/types/mocks/mock_types.go index 6ac9e8e9e4..4622fbfd68 100644 --- a/pkg/vmcp/optimizer/internal/types/mocks/mock_types.go +++ b/pkg/vmcp/optimizer/internal/types/mocks/mock_types.go @@ -153,3 +153,18 @@ func (mr *MockEmbeddingClientMockRecorder) EmbedBatch(ctx, texts any) *gomock.Ca mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "EmbedBatch", reflect.TypeOf((*MockEmbeddingClient)(nil).EmbedBatch), ctx, texts) } + +// ModelID mocks base method. +func (m *MockEmbeddingClient) ModelID(ctx context.Context) (string, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ModelID", ctx) + ret0, _ := ret[0].(string) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// ModelID indicates an expected call of ModelID. +func (mr *MockEmbeddingClientMockRecorder) ModelID(ctx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ModelID", reflect.TypeOf((*MockEmbeddingClient)(nil).ModelID), ctx) +} diff --git a/pkg/vmcp/optimizer/internal/types/types.go b/pkg/vmcp/optimizer/internal/types/types.go index 3b064d5c8d..1d7b15daae 100644 --- a/pkg/vmcp/optimizer/internal/types/types.go +++ b/pkg/vmcp/optimizer/internal/types/types.go @@ -72,6 +72,15 @@ type EmbeddingClient interface { // EmbedBatch returns vector embeddings for multiple texts. EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) + // ModelID returns the identity of the model currently serving requests. + // Implementations that fix the model by configuration return the + // configured name; implementations whose backend can be replaced + // underneath a running process (e.g. TEI, where the model is a property + // of the container rather than of the client) must query it live, so two + // calls can observe a swap. Callers own any caching or fallback policy + // for the error case. + ModelID(ctx context.Context) (string, error) + // Close releases any resources held by the client. Close() error }