fix(mcp): leave out tool texts too long to embed instead of failing tool search

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-28 10:41:13 +00:00
parent 90e4962c81
commit d06fd4b312
13 changed files with 482 additions and 50 deletions

View file

@ -105,7 +105,11 @@ class SkillSearchIndex:
if isinstance(scores, EmbeddingFailed):
return SkillSearchEmbeddingFailed(reason=scores.reason)
ranked: Final = sorted(
(SkillSearchHit(skill=skill, score=score) for skill, score in zip(skills, scores, strict=True)),
(
SkillSearchHit(skill=skill, score=score)
for skill, score in zip(skills, scores, strict=True)
if score is not None
),
key=lambda hit: hit.score,
reverse=True,
)

View file

@ -183,9 +183,11 @@ def _split_core_tools(tools: Sequence[Tool], core_tools: Sequence[str]) -> tuple
def _top_hits(
tools: Sequence[Tool], scores: Sequence[float], minimum: float, limit: int
tools: Sequence[Tool], scores: Sequence[float | None], minimum: float, limit: int
) -> tuple[tuple[float, Tool], ...]:
hits: Final = ((score, tool) for score, tool in zip(scores, tools, strict=True) if score >= minimum)
hits: Final = (
(score, tool) for score, tool in zip(scores, tools, strict=True) if score is not None and score >= minimum
)
return tuple(sorted(hits, key=lambda hit: hit[0], reverse=True)[:limit])

View file

@ -115,7 +115,11 @@ class AgentSearchIndex:
if isinstance(scores, EmbeddingFailed):
return AgentSearchEmbeddingFailed(reason=scores.reason)
ranked: Final = sorted(
(AgentSearchHit(agent=agent, score=score) for agent, score in zip(agents, scores, strict=True)),
(
AgentSearchHit(agent=agent, score=score)
for agent, score in zip(agents, scores, strict=True)
if score is not None
),
key=lambda hit: hit.score,
reverse=True,
)

View file

@ -2,6 +2,7 @@
from __future__ import annotations
import asyncio
import math
from collections.abc import Awaitable, Mapping, Sequence
from dataclasses import dataclass
@ -13,7 +14,7 @@ from fastapi import HTTPException
from openai import OpenAIError
from pydantic import BaseModel, ConfigDict
from litellm.exceptions import BudgetExceededError
from litellm.exceptions import BudgetExceededError, ContextWindowExceededError
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
@ -35,6 +36,11 @@ class EmbeddingFailed:
reason: str
@dataclass(frozen=True, slots=True)
class _InputTooLong:
reason: str
class _EmbeddingItem(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
@ -101,11 +107,13 @@ def router_embedder(
_CacheKey: TypeAlias = tuple[str, str]
async def _embed_all(embed: Embedder, texts: Sequence[str]) -> tuple[Vector, ...] | EmbeddingFailed:
async def _embed_all(embed: Embedder, texts: Sequence[str]) -> tuple[Vector, ...] | EmbeddingFailed | _InputTooLong:
try:
vectors: Final = tuple(await embed(texts))
except HTTPException:
raise
except ContextWindowExceededError as exc:
return _InputTooLong(reason=f"embedding the search query failed: {exc}")
except (OpenAIError, ValueError, BudgetExceededError) as exc:
return EmbeddingFailed(reason=f"embedding the search query failed: {exc}")
if len(vectors) != len(texts):
@ -113,6 +121,52 @@ async def _embed_all(embed: Embedder, texts: Sequence[str]) -> tuple[Vector, ...
return vectors
async def _embed_texts(embed: Embedder, texts: Sequence[str]) -> Mapping[str, Vector] | EmbeddingFailed:
if not texts:
return MappingProxyType({})
embedded: Final = await _embed_all(embed, texts)
if isinstance(embedded, EmbeddingFailed):
return embedded
if not isinstance(embedded, _InputTooLong):
return MappingProxyType(dict(zip(texts, embedded, strict=True)))
if len(texts) == 1:
return MappingProxyType({})
middle: Final = len(texts) // 2
left, right = await asyncio.gather(_embed_texts(embed, texts[:middle]), _embed_texts(embed, texts[middle:]))
if isinstance(left, EmbeddingFailed):
return left
if isinstance(right, EmbeddingFailed):
return right
return MappingProxyType({**left, **right})
async def _embed_query_with(
embed: Embedder, query: str, texts: Sequence[str]
) -> tuple[Vector, Mapping[str, Vector]] | EmbeddingFailed:
if not texts:
embedded_query: Final = await _embed_all(embed, (query,))
if isinstance(embedded_query, EmbeddingFailed):
return embedded_query
if isinstance(embedded_query, _InputTooLong):
return EmbeddingFailed(reason=embedded_query.reason)
return embedded_query[0], MappingProxyType({})
embedded: Final = await _embed_all(embed, (query, *texts))
if isinstance(embedded, EmbeddingFailed):
return embedded
if not isinstance(embedded, _InputTooLong):
return embedded[0], MappingProxyType(dict(zip(texts, embedded[1:], strict=True)))
query_embedding: Final = await _embed_all(embed, (query,))
if isinstance(query_embedding, EmbeddingFailed):
return query_embedding
if isinstance(query_embedding, _InputTooLong):
return EmbeddingFailed(reason=query_embedding.reason)
vectors: Final = await _embed_texts(embed, texts)
if isinstance(vectors, EmbeddingFailed):
return vectors
return query_embedding[0], vectors
@dataclass(frozen=True, slots=True)
class _Embedded:
query_vector: Vector
@ -120,26 +174,26 @@ class _Embedded:
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)
return all(len(vectors[text]) == len(query_vector) for text in texts if text in vectors)
async def _embed_query_and_texts(
embed: Embedder, query: str, texts: Sequence[str], cached: Mapping[str, Vector]
) -> _Embedded | EmbeddingFailed:
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, EmbeddingFailed):
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))
query_embedding: Final = await _embed_query_with(embed, query, missing)
if isinstance(query_embedding, EmbeddingFailed):
return query_embedding
query_vector, newly_embedded = query_embedding
vectors: Final = MappingProxyType(dict(chain(cached.items(), newly_embedded.items())))
if _same_dimension(query_vector, vectors, texts):
return _Embedded(query_vector=query_vector, vectors=vectors)
unique: Final = tuple(dict.fromkeys(text for text in texts if text in vectors))
reembedded: Final = await _embed_query_with(embed, query, unique)
if isinstance(reembedded, EmbeddingFailed):
return reembedded
return _Embedded(
query_vector=reembedded[0], vectors=MappingProxyType(dict(zip(unique, reembedded[1:], strict=True)))
)
reembedded_query, reembedded_vectors = reembedded
return _Embedded(query_vector=reembedded_query, vectors=reembedded_vectors)
class SemanticTextIndex:
@ -158,7 +212,9 @@ class SemanticTextIndex:
def _merged(self, embedding_model: str, embedded: _Embedded, texts: Sequence[str]) -> Mapping[_CacheKey, Vector]:
dimension: Final = len(embedded.query_vector)
touched: Final = MappingProxyType({(embedding_model, text): embedded.vectors[text] for text in texts})
touched: Final = MappingProxyType(
{(embedding_model, text): embedded.vectors[text] for text in texts if text in embedded.vectors}
)
untouched: Final = MappingProxyType(
{
key: vector
@ -174,8 +230,8 @@ class SemanticTextIndex:
async def scores(
self, query: str, texts: Sequence[str], embed: Embedder, embedding_model: str
) -> tuple[float, ...] | EmbeddingFailed:
"""Cosine similarity of `query` to each entry of `texts`, in the same order."""
) -> tuple[float | None, ...] | EmbeddingFailed:
"""Cosine similarity for embeddable entries of `texts`, with None for omitted entries."""
if not texts:
return ()
embedded: Final = await _embed_query_and_texts(embed, query, texts, self._cached(embedding_model))
@ -184,4 +240,7 @@ class SemanticTextIndex:
if not _same_dimension(embedded.query_vector, embedded.vectors, texts):
return EmbeddingFailed(reason=f"embedding model {embedding_model} returned vectors of mixed dimensions")
self._vectors = self._merged(embedding_model, embedded, texts)
return tuple(cosine_similarity(embedded.query_vector, embedded.vectors[text]) for text in texts)
return tuple(
cosine_similarity(embedded.query_vector, embedded.vectors[text]) if text in embedded.vectors else None
for text in texts
)

View file

@ -39,6 +39,14 @@
assertions: [denied_without_permission]
source: "rest_endpoints.py:305-386"
rationale: Tool-level permission guard; multi-tenant safety
- id: mcp.tool_search.api_key.oversized_description_omitted
module: mcp
tier: P0
operation: tool_search
auth_family: api_key
assertions: [oversized_description_omitted]
source: "test_mcp_tool_search_e2e.py"
rationale: Native semantic search omits tools whose descriptions exceed the embedding input limit
- id: mcp.list_tools.bearer.succeeds
module: mcp
tier: P1

View file

@ -3,7 +3,7 @@
from __future__ import annotations
import os
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from e2e_config import datadog_mcp_url, unique_marker
from lifecycle import ResourceManager
@ -37,6 +37,7 @@ def register_datadog_mcp(
*,
mcp_access_groups: list[str] | None = None,
allowed_tools: Sequence[str] | None = (SEARCH_LOGS_TOOL,),
tool_name_to_description: Mapping[str, str] | None = None,
) -> str:
"""Register the core Datadog toolset with its credentials from the env. By default
the server exposes only `search_datadog_logs`; pass `allowed_tools=None` to expose
@ -54,6 +55,7 @@ def register_datadog_mcp(
},
allowed_tools=None if allowed_tools is None else list(allowed_tools),
mcp_access_groups=mcp_access_groups,
tool_name_to_description=tool_name_to_description,
)
resources.defer(lambda: client.delete_server(server_id))
return server_id

View file

@ -41,6 +41,7 @@ class McpServerNewBody(BaseModel):
static_headers: dict[str, str] | None = None
allowed_tools: list[str] | None = None
mcp_access_groups: list[str] | None = None
tool_name_to_description: dict[str, str] | None = None
class McpServerNewResponse(BaseModel):
@ -78,9 +79,7 @@ class McpToolsListResponse(BaseModel):
def tool_names_for_server(self, server_id: str) -> frozenset[str]:
return frozenset(
tool.name
for tool in self.tools
if tool.mcp_info is not None and tool.mcp_info.server_id == server_id
tool.name for tool in self.tools if tool.mcp_info is not None and tool.mcp_info.server_id == server_id
)
def tool_name_containing(self, server_id: str, needle: str) -> str | None:
@ -130,6 +129,34 @@ class McpCallToolBody(BaseModel):
server_id: str
class McpVirtualToolCallBody(BaseModel):
name: str
arguments: dict[str, McpToolArg]
class McpToolsListParams(BaseModel):
server_id: str | None = None
class McpToolSearchSettings(BaseModel):
model_config = ConfigDict(frozen=True)
embedding_model: str | None = None
top_k: int = 5
similarity_threshold: float = 0.0
core_tools: tuple[str, ...] = ()
class McpToolSearchSettingsResponse(BaseModel):
values: McpToolSearchSettings
class McpToolSearchSettingsUpdateResponse(BaseModel):
message: str
status: str
settings: McpToolSearchSettings
class McpCallContent(BaseModel):
type: str | None = None
text: str | None = None
@ -164,6 +191,7 @@ class McpClient:
static_headers: dict[str, str] | None = None,
allowed_tools: list[str] | None = None,
mcp_access_groups: list[str] | None = None,
tool_name_to_description: Mapping[str, str] | None = None,
) -> str:
return unwrap(
self.proxy.transport.post(
@ -178,6 +206,9 @@ class McpClient:
static_headers=static_headers,
allowed_tools=allowed_tools,
mcp_access_groups=mcp_access_groups,
tool_name_to_description=(
dict(tool_name_to_description) if tool_name_to_description is not None else None
),
),
response_type=McpServerNewResponse,
)
@ -224,9 +255,7 @@ class McpClient:
McpServerListResponse,
settled=lambda response: any(row.server_id == server_id for row in response.root),
)
return next(
row for response in registered.values() for row in response.root if row.server_id == server_id
)
return next(row for response in registered.values() for row in response.root if row.server_id == server_id)
def generate_key(
self,
@ -236,14 +265,21 @@ class McpClient:
mcp_access_groups: list[str] | None = None,
mcp_toolsets: list[str] | None = None,
models: list[str] | None = None,
mcp_tool_search_enabled: bool | None = None,
) -> str:
object_permission = (
ObjectPermission(
mcp_servers=mcp_servers,
mcp_access_groups=mcp_access_groups,
mcp_toolsets=mcp_toolsets,
mcp_tool_search_enabled=mcp_tool_search_enabled,
)
if (
mcp_servers is not None
or mcp_access_groups is not None
or mcp_toolsets is not None
or mcp_tool_search_enabled is not None
)
if mcp_servers is not None or mcp_access_groups is not None or mcp_toolsets is not None
else None
)
return self.proxy.generate_key(
@ -254,14 +290,35 @@ class McpClient:
)
)
def list_tools(self, key: str) -> Result[McpToolsListResponse]:
def list_tools(self, key: str, server_id: str | None = None) -> Result[McpToolsListResponse]:
return self.proxy.transport.get(
"/mcp-rest/tools/list",
headers=ApiKeyHeaders(x_litellm_api_key=key),
params=NoBody(),
params=McpToolsListParams(server_id=server_id),
response_type=McpToolsListResponse,
)
def get_mcp_tool_search_settings(self) -> McpToolSearchSettings:
return unwrap(
self.proxy.transport.get(
"/get/mcp_tool_search_settings",
headers=self.proxy.transport.master,
params=NoBody(),
response_type=McpToolSearchSettingsResponse,
)
).values
def update_mcp_tool_search_settings(self, settings: McpToolSearchSettings) -> None:
_ = unwrap(
self.proxy.transport.patch(
"/update/mcp_tool_search_settings",
headers=self.proxy.transport.master,
json=settings,
response_type=McpToolSearchSettingsUpdateResponse,
)
)
settle_propagation(time.monotonic())
def await_tool(self, key: str, server_id: str, needle: str) -> str:
"""Poll tools/list until `server_id` serves a tool matching `needle`, and
return its fully-qualified name. Fails at poll_timeout.
@ -345,13 +402,10 @@ class McpClient:
if isinstance(last, UnknownApiError) and last.status_code == 403:
return last
if not _is_mcp_not_synced(last, tool_name=name):
raise AssertionError(
f"ungranted key's tools/call was not 403 access_denied: {last}"
)
raise AssertionError(f"ungranted key's tools/call was not 403 access_denied: {last}")
if time.monotonic() >= deadline:
raise AssertionError(
f"ungranted key never got 403 for {name!r} within {self.proxy.poll_timeout}s; "
f"last result: {last}"
f"ungranted key never got 403 for {name!r} within {self.proxy.poll_timeout}s; last result: {last}"
)
time.sleep(self.proxy.poll_interval)
@ -397,9 +451,21 @@ class McpClient:
return self.proxy.transport.post(
"/mcp-rest/tools/call",
headers=ApiKeyHeaders(x_litellm_api_key=key),
json=McpCallToolBody(
name=name, arguments=dict(arguments), server_id=server_id
),
json=McpCallToolBody(name=name, arguments=dict(arguments), server_id=server_id),
response_type=McpCallToolResponse,
)
def call_virtual_tool(
self,
key: str,
*,
name: str,
arguments: McpToolArguments,
) -> Result[McpCallToolResponse]:
return self.proxy.transport.post(
"/mcp-rest/tools/call",
headers=ApiKeyHeaders(x_litellm_api_key=key),
json=McpVirtualToolCallBody(name=name, arguments=dict(arguments)),
response_type=McpCallToolResponse,
)
@ -432,9 +498,7 @@ def _is_mcp_not_synced(
# Gateway: "Tool search_datadog_logs not found" (optionally inside a longer message)
if tool_name is not None:
return (
re.search(rf"\btool\s+{re.escape(tool_name)}\s+not found\b", body_l) is not None
)
return re.search(rf"\btool\s+{re.escape(tool_name)}\s+not found\b", body_l) is not None
return re.search(r"\btool\s+\S+\s+not found\b", body_l) is not None

View file

@ -0,0 +1,86 @@
from __future__ import annotations
from typing import Final
import pytest
from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp
from e2e_config import provider_edge_base, unique_marker
from e2e_http import unwrap
from lifecycle import ResourceManager
from mcp_client import McpClient, McpToolSearchSettings
from models import LiteLLMParamsBody
from proxy_client import ProxyClient
pytestmark = pytest.mark.e2e
MCP_TOOL_SEARCH_TOOL_NAME: Final = "mcp_tool_search"
OVERSIZED_DESCRIPTION: Final = f"OVERSIZED {'weather ' * 12_000}"
class TestMcpToolSearch:
@pytest.mark.covers("mcp.tool_search.api_key.oversized_description_omitted")
def test_native_tool_search_omits_oversized_description(
self, proxy: ProxyClient, client: McpClient, resources: ResourceManager
) -> None:
"""Exercise shared ranking through the native virtual tool because proxy mode does not apply description overrides."""
model_name: Final = f"e2e-mcp-tool-search-{unique_marker()}"
embedding_api_base: Final = provider_edge_base("openai")
model_id: Final = proxy.create_model(
model_name,
LiteLLMParamsBody(
model="openai/text-embedding-3-large",
api_key="os.environ/OPENAI_API_KEY",
api_base=None if embedding_api_base is None else f"{embedding_api_base}/v1",
),
)
resources.defer(lambda: proxy.delete_model(model_id))
previous_settings: Final = client.get_mcp_tool_search_settings()
resources.defer(lambda: client.update_mcp_tool_search_settings(previous_settings))
client.update_mcp_tool_search_settings(
McpToolSearchSettings(
embedding_model=model_name,
top_k=100,
similarity_threshold=0.0,
core_tools=previous_settings.core_tools,
)
)
server_id: Final = register_datadog_mcp(
client,
resources,
allowed_tools=None,
tool_name_to_description={SEARCH_LOGS_TOOL: OVERSIZED_DESCRIPTION},
)
client.await_registered(server_id)
key: Final = client.generate_key(
user_id=f"e2e-mcp-tool-search-{unique_marker()}",
mcp_servers=[server_id],
models=[model_name],
mcp_tool_search_enabled=True,
)
resources.defer(lambda: proxy.delete_key(key))
tools: Final = unwrap(client.list_tools(key, server_id=server_id))
assert len(tools.tool_names_for_server(server_id)) >= 2
oversized_tool_name: Final = tools.tool_name_containing(server_id, SEARCH_LOGS_TOOL)
assert oversized_tool_name is not None
oversized_tool: Final = next((tool for tool in tools.tools if tool.name == oversized_tool_name), None)
assert oversized_tool is not None
assert oversized_tool.description == OVERSIZED_DESCRIPTION
normal_tool_name: Final = next(
(name for name in tools.tool_names_for_server(server_id) if name != oversized_tool_name), None
)
assert normal_tool_name is not None
result: Final = unwrap(
client.call_virtual_tool(
key,
name=MCP_TOOL_SEARCH_TOOL_NAME,
arguments={"query": normal_tool_name, "top_k": 100},
)
)
assert result.is_error is not True, result.all_text
assert normal_tool_name in result.all_text
assert oversized_tool_name not in result.all_text

View file

@ -67,6 +67,7 @@ class ObjectPermission(BaseModel):
mcp_servers: list[str] | None = None
mcp_access_groups: list[str] | None = None
mcp_toolsets: list[str] | None = None
mcp_tool_search_enabled: bool | None = None
class KeyGenerateBody(BaseModel):

View file

@ -1,4 +1,5 @@
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
"""
Tests for MCP tool search feature.
@ -13,7 +14,7 @@ Covers:
import json
from collections.abc import Sequence
from types import SimpleNamespace
from typing import Any
from typing import Any, Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -31,6 +32,7 @@ from litellm.proxy._experimental.mcp_server.tool_search import (
ToolSearchResult,
coerce_top_k,
get_virtual_tool_definitions,
rank_mcp_tools,
search_mcp_tools,
search_tools,
)
@ -86,13 +88,12 @@ FAKE_VECTORS: dict[str, Vector] = {
}
def _paged_params():
from mcp.types import PaginatedRequestParams
return PaginatedRequestParams()
class RecordingEmbedder:
def __init__(self) -> None:
self.calls: list[tuple[str, ...]] = []
@ -102,6 +103,19 @@ class RecordingEmbedder:
return tuple(FAKE_VECTORS[text] for text in texts)
class OversizedRecordingEmbedder:
def __init__(self) -> None:
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))
if any("OVERSIZED" in text for text in texts):
raise litellm.ContextWindowExceededError(
message="input is too long", model="text-embedding-3-large", llm_provider="openai"
)
return tuple((1.0, 0.0) if "weather" in text.lower() else (0.0, 1.0) for text in texts)
def _ranker(embedder: RecordingEmbedder | None = None) -> SemanticToolRanker:
return SemanticToolRanker(embed=embedder or RecordingEmbedder(), embedding_model="emb", index=SemanticTextIndex())
@ -112,6 +126,26 @@ def _names(results: Sequence[ToolSearchResult] | EmbeddingFailed) -> list[str]:
class TestSearchMcpTools:
@pytest.mark.asyncio
async def test_rank_mcp_tools_omits_oversized_description(self) -> None:
tools: Final = _make_tools(
[
("get_weather", "Get the current weather forecast for a city"),
("send_email", "Send an email to a recipient"),
("list_files", "List files in a directory"),
("huge_tool", "OVERSIZED description"),
]
)
embedder = OversizedRecordingEmbedder()
ranker = SemanticToolRanker(embed=embedder, embedding_model="emb", index=SemanticTextIndex())
hits: Final = await rank_mcp_tools("weather", tools, 5, MCPToolSearchSettings(embedding_model="emb"), ranker)
assert not isinstance(hits, EmbeddingFailed)
names: Final = [hit.tool.name for hit in hits]
assert names[0] == "get_weather"
assert "huge_tool" not in names
@pytest.mark.asyncio
async def test_semantic_mode_finds_foreign_exchange_tool_for_fx(self) -> None:
keyword_only = await search_mcp_tools("FX", CATALOG, 5, MCPToolSearchSettings(), ranker=None)
@ -1028,7 +1062,8 @@ class TestDispatchVirtualMcpTool:
sentinel_logging_obj = object()
with (
patch.object(
mcp_operations, "_build_virtual_call_logging_obj",
mcp_operations,
"_build_virtual_call_logging_obj",
new_callable=AsyncMock,
return_value=sentinel_logging_obj,
) as mock_build,
@ -1171,9 +1206,7 @@ class TestCaptureHostProgressCallback:
callback = _capture_host_progress_callback(_mcp_request_ctx(meta=params.meta, session=session))
assert callback is not None
await callback(0.5, 1.0)
session.send_progress_notification.assert_awaited_once_with(
progress_token=token, progress=0.5, total=1.0
)
session.send_progress_notification.assert_awaited_once_with(progress_token=token, progress=0.5, total=1.0)
class TestHandleListToolsVirtual:

View file

@ -88,6 +88,19 @@ class FixedDimensionEmbedder:
return tuple((1.0,) * self.dimensions for _ in texts)
class OversizedAwareEmbedder:
def __init__(self) -> None:
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))
if any("OVERSIZED" in text for text in texts):
raise litellm.ContextWindowExceededError(
message="input is too long", model="text-embedding-3-large", llm_provider="openai"
)
return tuple(VECTORS[text] for text in texts)
class TestAgentSearchText:
def test_joins_name_description_and_skills_with_tags(self) -> None:
assert agent_search_text(TRANSLATOR) == (
@ -127,6 +140,21 @@ class TestAgentSearchIndex:
assert [hit.agent.agent_id for hit in outcome.hits] == ["translator", "trip"]
assert outcome.hits[0].score > outcome.hits[1].score
@pytest.mark.asyncio
async def test_omits_agents_with_oversized_search_text(self) -> None:
oversized: Final = AgentResponse(
agent_id="oversized",
agent_name="oversized-agent",
agent_card_params={"description": "OVERSIZED description"},
)
outcome: Final = await AgentSearchIndex().search(
"language translation", (*AGENTS, oversized), top_k=5, embed=OversizedAwareEmbedder(), embedding_model="m"
)
assert isinstance(outcome, AgentSearchHits)
assert {hit.agent.agent_id for hit in outcome.hits} == {agent.agent_id for agent in AGENTS}
@pytest.mark.asyncio
async def test_second_search_only_embeds_the_query(self) -> None:
index = AgentSearchIndex()

View file

@ -0,0 +1,115 @@
from __future__ import annotations
from collections.abc import Sequence
from typing import Final
import litellm
import pytest
from litellm.proxy.common_utils.semantic_text_index import EmbeddingFailed, SemanticTextIndex, Vector
QUERY: Final = "weather query"
TEXTS: Final = (
"valid one",
"OVERSIZED first",
"valid two",
"valid three",
"OVERSIZED second",
"valid four",
"valid five",
)
class ContextWindowEmbedder:
def __init__(self, dimensions: int = 2) -> None:
self.calls: list[tuple[str, ...]] = [] # mutable-ok: test spy recording embed inputs
self.vector: Final = (1.0,) * dimensions
async def __call__(self, texts: Sequence[str]) -> Sequence[Vector]:
self.calls.append(tuple(texts))
if any("OVERSIZED" in text for text in texts):
raise litellm.ContextWindowExceededError(
message="input is too long", model="text-embedding-3-large", llm_provider="openai"
)
return tuple(self.vector for _ in texts)
class RateLimitedEmbedder:
def __init__(self) -> None:
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))
raise litellm.RateLimitError(message="rate limited", model="text-embedding-3-large", llm_provider="openai")
class FixedDimensionEmbedder:
def __init__(self, dimensions: int) -> None:
self.vector: Final = (1.0,) * dimensions
async def __call__(self, texts: Sequence[str]) -> Sequence[Vector]:
return tuple(self.vector for _ in texts)
class TestSemanticTextIndex:
@pytest.mark.asyncio
async def test_omits_multiple_oversized_texts_and_preserves_score_order(self) -> None:
embedder = ContextWindowEmbedder()
scores: Final = await SemanticTextIndex().scores(QUERY, TEXTS, embedder, "emb")
assert not isinstance(scores, EmbeddingFailed)
assert len(scores) == len(TEXTS)
assert scores[0] == pytest.approx(1.0)
assert scores[1] is None
assert scores[2] == pytest.approx(1.0)
assert scores[3] == pytest.approx(1.0)
assert scores[4] is None
assert scores[5] == pytest.approx(1.0)
assert scores[6] == pytest.approx(1.0)
@pytest.mark.asyncio
async def test_repeat_search_only_reembeds_the_query_and_oversized_texts(self) -> None:
index = SemanticTextIndex()
embedder = ContextWindowEmbedder()
first: Final = await index.scores(QUERY, TEXTS, embedder, "emb")
assert not isinstance(first, EmbeddingFailed)
first_call_count: Final = len(embedder.calls)
second: Final = await index.scores(QUERY, TEXTS, embedder, "emb")
assert not isinstance(second, EmbeddingFailed)
second_batches: Final = embedder.calls[first_call_count:]
assert second_batches[0] == (QUERY, "OVERSIZED first", "OVERSIZED second")
assert all(text in (QUERY, "OVERSIZED first", "OVERSIZED second") for batch in second_batches for text in batch)
@pytest.mark.asyncio
async def test_oversized_query_returns_embedding_failed(self) -> None:
result: Final = await SemanticTextIndex().scores(
"OVERSIZED query", ("valid text",), ContextWindowEmbedder(), "emb"
)
assert isinstance(result, EmbeddingFailed)
@pytest.mark.asyncio
async def test_non_context_provider_error_fails_after_one_embed_call(self) -> None:
embedder = RateLimitedEmbedder()
result: Final = await SemanticTextIndex().scores(QUERY, ("valid text",), embedder, "emb")
assert isinstance(result, EmbeddingFailed)
assert len(embedder.calls) == 1
@pytest.mark.asyncio
async def test_mixed_dimension_reembedding_omits_oversized_text(self) -> None:
index = SemanticTextIndex()
seeded: Final = await index.scores(QUERY, ("cached text", "OVERSIZED text"), FixedDimensionEmbedder(3), "emb")
assert not isinstance(seeded, EmbeddingFailed)
embedder = ContextWindowEmbedder(2)
scores: Final = await index.scores(QUERY, ("cached text", "OVERSIZED text"), embedder, "emb")
assert not isinstance(scores, EmbeddingFailed)
assert scores[0] == pytest.approx(1.0)
assert scores[1] is None
assert (QUERY, "cached text", "OVERSIZED text") in embedder.calls
assert ("OVERSIZED text",) in embedder.calls

View file

@ -89,6 +89,19 @@ class FixedDimensionEmbedder:
return tuple((1.0,) * self.dimensions for _ in texts)
class OversizedAwareEmbedder:
def __init__(self) -> None:
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))
if any("OVERSIZED" in text for text in texts):
raise litellm.ContextWindowExceededError(
message="input is too long", model="text-embedding-3-large", llm_provider="openai"
)
return tuple(VECTORS[text] for text in texts)
class TestSkillSearchText:
def test_joins_title_description_and_instructions(self) -> None:
assert skill_search_text(TRANSLATOR) == (
@ -133,6 +146,19 @@ class TestSkillSearchIndex:
assert [hit.skill.skill_id for hit in outcome.hits] == ["translate-file", "trip-planner"]
assert outcome.hits[0].score > outcome.hits[1].score
@pytest.mark.asyncio
async def test_omits_skills_with_oversized_search_text(self) -> None:
oversized: Final = LiteLLM_SkillsTable(
skill_id="oversized", display_title="Oversized", instructions="OVERSIZED instructions"
)
outcome: Final = await SkillSearchIndex().search(
"language translation", (*SKILLS, oversized), top_k=5, embed=OversizedAwareEmbedder(), embedding_model="m"
)
assert isinstance(outcome, SkillSearchHits)
assert {hit.skill.skill_id for hit in outcome.hits} == {skill.skill_id for skill in SKILLS}
@pytest.mark.asyncio
async def test_second_search_only_embeds_the_query(self) -> None:
index = SkillSearchIndex()