fix(a2a): key the agent search vector cache by embedding model and re-embed on dimension changes

This commit is contained in:
mateo-berri 2026-08-27 19:33:33 -07:00
parent e9cc9c9bc3
commit 8e455897a4
2 changed files with 85 additions and 19 deletions

View file

@ -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
)

View file

@ -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)