From db02cf81e5e8eddd28fbc9c704d9136fe8dd111f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 27 Aug 2026 19:38:00 -0700 Subject: [PATCH] fix(a2a): re-embed the query with the agents in one call when cached vectors change dimension --- litellm/proxy/agent_endpoints/agent_search.py | 50 +++++++++++++------ .../agent_endpoints/test_agent_search.py | 5 +- 2 files changed, 39 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/agent_endpoints/agent_search.py b/litellm/proxy/agent_endpoints/agent_search.py index 88373e26f51..c249079ccac 100644 --- a/litellm/proxy/agent_endpoints/agent_search.py +++ b/litellm/proxy/agent_endpoints/agent_search.py @@ -158,6 +158,35 @@ async def _embed_all(embed: Embedder, texts: Sequence[str]) -> tuple[Vector, ... return vectors +@dataclass(frozen=True, slots=True) +class _Embedded: + query_vector: Vector + vectors: Mapping[str, Vector] + + +def _same_dimension(query_vector: Vector, vectors: Mapping[str, Vector], texts: Sequence[str]) -> bool: + return all(len(vectors[text]) == len(query_vector) for text in texts) + + +async def _embed_query_and_agents( + embed: Embedder, query: str, texts: Sequence[str], cached: Mapping[str, Vector] +) -> _Embedded | AgentSearchEmbeddingFailed: + 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 + vectors: Final = MappingProxyType(dict(chain(cached.items(), zip(missing, embedded[1:], strict=True)))) + if _same_dimension(embedded[0], vectors, texts): + return _Embedded(query_vector=embedded[0], vectors=vectors) + unique: Final = tuple(dict.fromkeys(texts)) + reembedded: Final = await _embed_all(embed, (query, *unique)) + if isinstance(reembedded, AgentSearchEmbeddingFailed): + return reembedded + return _Embedded( + query_vector=reembedded[0], vectors=MappingProxyType(dict(zip(unique, reembedded[1:], strict=True))) + ) + + class AgentSearchIndex: """Caches one vector per distinct agent text per embedding model, so repeat searches only embed the query.""" @@ -171,28 +200,19 @@ class AgentSearchIndex: return AgentSearchHits(hits=()) texts: Final = tuple(agent_search_text(agent) for agent in agents) 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)) + embedded: Final = await _embed_query_and_agents(embed, query, texts, cached) 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): + 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: vectors}) + self._vectors = MappingProxyType( + {**self._vectors, embedding_model: MappingProxyType({**cached, **embedded.vectors})} + ) ranked: Final = sorted( ( - AgentSearchHit(agent=agent, score=cosine_similarity(query_vector, vectors[text])) + AgentSearchHit(agent=agent, score=cosine_similarity(embedded.query_vector, embedded.vectors[text])) for agent, text in zip(agents, texts, strict=True) ), key=lambda hit: hit.score, 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 4699b45ec94..93f29c4a1d2 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_search.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_search.py @@ -151,7 +151,10 @@ class TestAgentSearchIndex: 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)] + assert fallback.calls == [ + ("language translation",), + ("language translation", *(agent_search_text(agent) for agent in AGENTS)), + ] @pytest.mark.asyncio async def test_mixed_dimensions_in_one_batch_become_embedding_failed(self) -> None: