From 8e455897a41fd16b462c8f43f108c5463bf45dca Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 27 Aug 2026 19:33:33 -0700 Subject: [PATCH] fix(a2a): key the agent search vector cache by embedding model and re-embed on dimension changes --- litellm/proxy/agent_endpoints/agent_search.py | 53 ++++++++++++++----- .../agent_endpoints/test_agent_search.py | 51 +++++++++++++++--- 2 files changed, 85 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/agent_endpoints/agent_search.py b/litellm/proxy/agent_endpoints/agent_search.py index 896a22fbf9d..88373e26f51 100644 --- a/litellm/proxy/agent_endpoints/agent_search.py +++ b/litellm/proxy/agent_endpoints/agent_search.py @@ -143,31 +143,56 @@ def router_embedder(router: Router, embedding_model: str, user_api_key_dict: Use return embed +_NO_VECTORS: Final[Mapping[str, Vector]] = MappingProxyType({}) + + +async def _embed_all(embed: Embedder, texts: Sequence[str]) -> tuple[Vector, ...] | AgentSearchEmbeddingFailed: + try: + vectors: Final = tuple(await embed(texts)) + except (OpenAIError, ValueError, BudgetExceededError) as exc: + return AgentSearchEmbeddingFailed(reason=f"embedding the search query failed: {exc}") + if len(vectors) != len(texts): + return AgentSearchEmbeddingFailed( + reason=f"embedding model returned {len(vectors)} vectors for {len(texts)} inputs" + ) + return vectors + + class AgentSearchIndex: - """Caches one vector per distinct agent text, so repeat searches only embed the query.""" + """Caches one vector per distinct agent text per embedding model, so repeat searches only embed the query.""" def __init__(self) -> None: - self._vectors: Mapping[str, Vector] = MappingProxyType({}) + self._vectors: Mapping[str, Mapping[str, Vector]] = MappingProxyType({}) async def search( - self, query: str, agents: Sequence[AgentResponse], top_k: int, embed: Embedder + self, query: str, agents: Sequence[AgentResponse], top_k: int, embed: Embedder, embedding_model: str ) -> AgentSearchHits | AgentSearchEmbeddingFailed: if not agents: return AgentSearchHits(hits=()) texts: Final = tuple(agent_search_text(agent) for agent in agents) - missing: Final = tuple(dict.fromkeys(text for text in texts if text not in self._vectors)) - try: - vectors: Final = await embed((query, *missing)) - except (OpenAIError, ValueError, BudgetExceededError) as exc: - return AgentSearchEmbeddingFailed(reason=f"embedding the search query failed: {exc}") - if len(vectors) != len(missing) + 1: + cached: Final = self._vectors.get(embedding_model, _NO_VECTORS) + missing: Final = tuple(dict.fromkeys(text for text in texts if text not in cached)) + embedded: Final = await _embed_all(embed, (query, *missing)) + if isinstance(embedded, AgentSearchEmbeddingFailed): + return embedded + query_vector: Final = embedded[0] + stale: Final = tuple( + dict.fromkeys(text for text in texts if text in cached and len(cached[text]) != len(query_vector)) + ) + refreshed: Final = await _embed_all(embed, stale) if stale else () + if isinstance(refreshed, AgentSearchEmbeddingFailed): + return refreshed + vectors: Final = MappingProxyType( + dict(chain(cached.items(), zip(missing, embedded[1:], strict=True), zip(stale, refreshed, strict=True))) + ) + if any(len(vectors[text]) != len(query_vector) for text in texts): return AgentSearchEmbeddingFailed( - reason=f"embedding model returned {len(vectors)} vectors for {len(missing) + 1} inputs" + reason=f"embedding model {embedding_model} returned vectors of mixed dimensions" ) - self._vectors = MappingProxyType(dict(chain(self._vectors.items(), zip(missing, vectors[1:], strict=True)))) + self._vectors = MappingProxyType({**self._vectors, embedding_model: vectors}) ranked: Final = sorted( ( - AgentSearchHit(agent=agent, score=cosine_similarity(vectors[0], self._vectors[text])) + AgentSearchHit(agent=agent, score=cosine_similarity(query_vector, vectors[text])) for agent, text in zip(agents, texts, strict=True) ), key=lambda hit: hit.score, @@ -194,4 +219,6 @@ async def search_agents( ) if router is None: return AgentSearchNotConfigured(reason="agent search needs a model_list so the embedding model can be called") - return await index.search(query, agents, top_k, router_embedder(router, embedding_model, user_api_key_dict)) + return await index.search( + query, agents, top_k, router_embedder(router, embedding_model, user_api_key_dict), embedding_model + ) 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 4b674eb142b..4699b45ec94 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_search.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_search.py @@ -77,6 +77,16 @@ class FakeEmbedder: return tuple(VECTORS[text] for text in texts) +class FixedDimensionEmbedder: + def __init__(self, dimensions: int) -> None: + self.dimensions: Final = dimensions + self.calls: list[tuple[str, ...]] = [] # mutable-ok: test spy recording embed inputs + + async def __call__(self, texts: Sequence[str]) -> Sequence[Vector]: + self.calls.append(tuple(texts)) + return tuple((1.0,) * self.dimensions for _ in texts) + + class TestAgentSearchText: def test_joins_name_description_and_skills_with_tags(self) -> None: assert agent_search_text(TRANSLATOR) == ( @@ -109,7 +119,9 @@ class TestCosineSimilarity: class TestAgentSearchIndex: @pytest.mark.asyncio async def test_ranks_by_similarity_and_truncates_to_top_k(self) -> None: - outcome = await AgentSearchIndex().search("language translation", AGENTS, top_k=2, embed=FakeEmbedder()) + outcome = await AgentSearchIndex().search( + "language translation", AGENTS, top_k=2, embed=FakeEmbedder(), embedding_model="m" + ) assert isinstance(outcome, AgentSearchHits) assert [hit.agent.agent_id for hit in outcome.hits] == ["translator", "trip"] assert outcome.hits[0].score > outcome.hits[1].score @@ -118,15 +130,42 @@ class TestAgentSearchIndex: async def test_second_search_only_embeds_the_query(self) -> None: index = AgentSearchIndex() embedder = FakeEmbedder() - await index.search("language translation", AGENTS, top_k=5, embed=embedder) - await index.search("language translation", AGENTS, top_k=5, embed=embedder) + await index.search("language translation", AGENTS, top_k=5, embed=embedder, embedding_model="m") + await index.search("language translation", AGENTS, top_k=5, embed=embedder, embedding_model="m") assert len(embedder.calls[0]) == 1 + len(AGENTS) assert embedder.calls[1] == ("language translation",) + @pytest.mark.asyncio + async def test_switching_embedding_models_does_not_reuse_cached_vectors(self) -> None: + index = AgentSearchIndex() + await index.search("language translation", AGENTS, top_k=5, embed=FakeEmbedder(), embedding_model="small") + wide = FixedDimensionEmbedder(2) + outcome = await index.search("language translation", AGENTS, top_k=5, embed=wide, embedding_model="wide") + assert isinstance(outcome, AgentSearchHits) + assert len(wide.calls[0]) == 1 + len(AGENTS) + + @pytest.mark.asyncio + async def test_cached_vectors_of_another_dimension_are_re_embedded(self) -> None: + index = AgentSearchIndex() + await index.search("language translation", AGENTS, top_k=5, embed=FakeEmbedder(), embedding_model="m") + fallback = FixedDimensionEmbedder(2) + outcome = await index.search("language translation", AGENTS, top_k=5, embed=fallback, embedding_model="m") + assert isinstance(outcome, AgentSearchHits) + assert fallback.calls == [("language translation",), tuple(agent_search_text(agent) for agent in AGENTS)] + + @pytest.mark.asyncio + async def test_mixed_dimensions_in_one_batch_become_embedding_failed(self) -> None: + async def mixed(texts: Sequence[str]) -> Sequence[Vector]: + return ((1.0, 0.0), *((1.0, 0.0, 0.0) for _ in texts[1:])) + + outcome = await AgentSearchIndex().search("q", AGENTS, top_k=5, embed=mixed, embedding_model="m") + assert isinstance(outcome, AgentSearchEmbeddingFailed) + assert "mixed dimensions" in outcome.reason + @pytest.mark.asyncio async def test_empty_registry_returns_no_hits_without_embedding(self) -> None: embedder = FakeEmbedder() - outcome = await AgentSearchIndex().search("anything", (), top_k=5, embed=embedder) + outcome = await AgentSearchIndex().search("anything", (), top_k=5, embed=embedder, embedding_model="m") assert outcome == AgentSearchHits(hits=()) assert embedder.calls == [] @@ -135,7 +174,7 @@ class TestAgentSearchIndex: async def failing(texts: Sequence[str]) -> Sequence[Vector]: raise APIConnectionError(request=MagicMock()) - outcome = await AgentSearchIndex().search("q", AGENTS, top_k=5, embed=failing) + outcome = await AgentSearchIndex().search("q", AGENTS, top_k=5, embed=failing, embedding_model="m") assert isinstance(outcome, AgentSearchEmbeddingFailed) assert "embedding the search query failed" in outcome.reason @@ -144,7 +183,7 @@ class TestAgentSearchIndex: async def short(texts: Sequence[str]) -> Sequence[Vector]: return ((1.0, 0.0, 0.0),) - outcome = await AgentSearchIndex().search("q", AGENTS, top_k=5, embed=short) + outcome = await AgentSearchIndex().search("q", AGENTS, top_k=5, embed=short, embedding_model="m") assert isinstance(outcome, AgentSearchEmbeddingFailed)