Skip to content

Commit 2ed2486

Browse files
committed
fix(a2a): drop cached agent vectors whose dimension no longer matches the query
When the configured agent_search_embedding_model or a router fallback switches to a model with a different embedding dimension, the query vector shape stops matching cached agent vectors. cosine_similarity uses zip(..., strict=True), so ranking raised ValueError outside the try/except that treats embedding failures as AgentSearchEmbeddingFailed. Because the cache never dropped the old vectors, every subsequent GET /v1/agents?query= and agent_search MCP call returned 500 until the worker restarted. After embedding the query, filter the cache to only vectors of the same dimension and re-embed any agent texts that got evicted.
1 parent e9cc9c9 commit 2ed2486

2 files changed

Lines changed: 29 additions & 2 deletions

File tree

‎litellm/proxy/agent_endpoints/agent_search.py‎

Lines changed: 16 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -164,10 +164,24 @@ async def search(
164164
return AgentSearchEmbeddingFailed(
165165
reason=f"embedding model returned {len(vectors)} vectors for {len(missing) + 1} inputs"
166166
)
167-
self._vectors = MappingProxyType(dict(chain(self._vectors.items(), zip(missing, vectors[1:], strict=True))))
167+
query_vector: Final = vectors[0]
168+
fresh: Final = dict(zip(missing, vectors[1:], strict=True))
169+
kept: Final = {text: vec for text, vec in self._vectors.items() if len(vec) == len(query_vector)}
170+
stale: Final = tuple(dict.fromkeys(text for text in texts if text not in kept and text not in fresh))
171+
try:
172+
refreshed: Final = await embed(stale) if stale else ()
173+
except (OpenAIError, ValueError, BudgetExceededError) as exc:
174+
return AgentSearchEmbeddingFailed(reason=f"re-embedding stale agent texts failed: {exc}")
175+
if len(refreshed) != len(stale):
176+
return AgentSearchEmbeddingFailed(
177+
reason=f"embedding model returned {len(refreshed)} vectors for {len(stale)} inputs"
178+
)
179+
self._vectors = MappingProxyType(
180+
dict(chain(kept.items(), fresh.items(), zip(stale, refreshed, strict=True)))
181+
)
168182
ranked: Final = sorted(
169183
(
170-
AgentSearchHit(agent=agent, score=cosine_similarity(vectors[0], self._vectors[text]))
184+
AgentSearchHit(agent=agent, score=cosine_similarity(query_vector, self._vectors[text]))
171185
for agent, text in zip(agents, texts, strict=True)
172186
),
173187
key=lambda hit: hit.score,

‎tests/test_litellm/proxy/agent_endpoints/test_agent_search.py‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -147,6 +147,19 @@ async def short(texts: Sequence[str]) -> Sequence[Vector]:
147147
outcome = await AgentSearchIndex().search("q", AGENTS, top_k=5, embed=short)
148148
assert isinstance(outcome, AgentSearchEmbeddingFailed)
149149

150+
@pytest.mark.asyncio
151+
async def test_dimension_change_invalidates_cached_agent_vectors(self) -> None:
152+
index = AgentSearchIndex()
153+
warm = FakeEmbedder()
154+
await index.search("language translation", AGENTS, top_k=5, embed=warm)
155+
156+
async def wider(texts: Sequence[str]) -> Sequence[Vector]:
157+
return tuple((1.0, 0.0, 0.0, 0.0) for _ in texts)
158+
159+
outcome = await index.search("language translation", AGENTS, top_k=5, embed=wider)
160+
assert isinstance(outcome, AgentSearchHits)
161+
assert len(outcome.hits) == len(AGENTS)
162+
150163

151164
class TestSearchAgents:
152165
@pytest.mark.asyncio

0 commit comments

Comments
 (0)