diff --git a/litellm/proxy/agent_endpoints/agent_search.py b/litellm/proxy/agent_endpoints/agent_search.py index c249079ccac..4e62e592261 100644 --- a/litellm/proxy/agent_endpoints/agent_search.py +++ b/litellm/proxy/agent_endpoints/agent_search.py @@ -199,17 +199,21 @@ class AgentSearchIndex: if not agents: return AgentSearchHits(hits=()) texts: Final = tuple(agent_search_text(agent) for agent in agents) - cached: Final = self._vectors.get(embedding_model, _NO_VECTORS) - embedded: Final = await _embed_query_and_agents(embed, query, texts, cached) + embedded: Final = await _embed_query_and_agents( + embed, query, texts, self._vectors.get(embedding_model, _NO_VECTORS) + ) if isinstance(embedded, AgentSearchEmbeddingFailed): return embedded if not _same_dimension(embedded.query_vector, embedded.vectors, texts): return AgentSearchEmbeddingFailed( reason=f"embedding model {embedding_model} returned vectors of mixed dimensions" ) - self._vectors = MappingProxyType( - {**self._vectors, embedding_model: MappingProxyType({**cached, **embedded.vectors})} + current: Final = self._vectors.get(embedding_model, _NO_VECTORS) + dim: Final = len(embedded.query_vector) + merged: Final = MappingProxyType( + {text: vector for text, vector in chain(current.items(), embedded.vectors.items()) if len(vector) == dim} ) + self._vectors = MappingProxyType({**self._vectors, embedding_model: merged}) ranked: Final = sorted( ( AgentSearchHit(agent=agent, score=cosine_similarity(embedded.query_vector, embedded.vectors[text])) diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_search.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_search.py index 93f29c4a1d2..5f4bc21a05a 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_search.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_search.py @@ -156,6 +156,19 @@ class TestAgentSearchIndex: ("language translation", *(agent_search_text(agent) for agent in AGENTS)), ] + @pytest.mark.asyncio + async def test_subset_reembed_after_dimension_change_evicts_stale_vectors(self) -> None: + index = AgentSearchIndex() + await index.search("language translation", AGENTS, top_k=5, embed=FakeEmbedder(), embedding_model="m") + narrow = FixedDimensionEmbedder(2) + await index.search("language translation", (TRANSLATOR,), top_k=5, embed=narrow, embedding_model="m") + broader = FixedDimensionEmbedder(2) + outcome = await index.search("language translation", AGENTS, top_k=5, embed=broader, embedding_model="m") + assert isinstance(outcome, AgentSearchHits) + assert broader.calls == [ + ("language translation", agent_search_text(SQL_ANALYST), agent_search_text(TRIP_PLANNER)), + ] + @pytest.mark.asyncio async def test_mixed_dimensions_in_one_batch_become_embedding_failed(self) -> None: async def mixed(texts: Sequence[str]) -> Sequence[Vector]: