mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
fix(a2a): key the agent search vector cache by embedding model and re-embed on dimension changes
This commit is contained in:
parent
e9cc9c9bc3
commit
8e455897a4
2 changed files with 85 additions and 19 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue