From d06fd4b3126b6fb407155ad6c3b1d3159fc01d1e Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 28 Sep 2026 10:41:13 +0000 Subject: [PATCH] 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> --- .../llms/litellm_proxy/skills/skill_search.py | 6 +- .../_experimental/mcp_server/tool_search.py | 6 +- litellm/proxy/agent_endpoints/agent_search.py | 6 +- .../proxy/common_utils/semantic_text_index.py | 95 ++++++++++++--- tests/e2e/coverage_registry/mcp.yaml | 8 ++ tests/e2e/mcp/datadog_mcp.py | 4 +- tests/e2e/mcp/mcp_client.py | 104 +++++++++++++--- tests/e2e/mcp/test_mcp_tool_search_e2e.py | 86 +++++++++++++ tests/e2e/models.py | 1 + .../mcp_server/test_mcp_tool_search.py | 47 +++++-- .../agent_endpoints/test_agent_search.py | 28 +++++ .../common_utils/test_semantic_text_index.py | 115 ++++++++++++++++++ .../litellm_proxy/skills/test_skill_search.py | 26 ++++ 13 files changed, 482 insertions(+), 50 deletions(-) create mode 100644 tests/e2e/mcp/test_mcp_tool_search_e2e.py create mode 100644 tests/test_litellm/proxy/common_utils/test_semantic_text_index.py diff --git a/litellm/llms/litellm_proxy/skills/skill_search.py b/litellm/llms/litellm_proxy/skills/skill_search.py index f975c6c4cab..b3545768cc0 100644 --- a/litellm/llms/litellm_proxy/skills/skill_search.py +++ b/litellm/llms/litellm_proxy/skills/skill_search.py @@ -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, ) diff --git a/litellm/proxy/_experimental/mcp_server/tool_search.py b/litellm/proxy/_experimental/mcp_server/tool_search.py index 9d117a1a1fa..934d0c2bb35 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_search.py +++ b/litellm/proxy/_experimental/mcp_server/tool_search.py @@ -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]) diff --git a/litellm/proxy/agent_endpoints/agent_search.py b/litellm/proxy/agent_endpoints/agent_search.py index 65a89bb2c7a..81ddd8533a4 100644 --- a/litellm/proxy/agent_endpoints/agent_search.py +++ b/litellm/proxy/agent_endpoints/agent_search.py @@ -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, ) diff --git a/litellm/proxy/common_utils/semantic_text_index.py b/litellm/proxy/common_utils/semantic_text_index.py index d3fe68e65f7..eb2e9281ff4 100644 --- a/litellm/proxy/common_utils/semantic_text_index.py +++ b/litellm/proxy/common_utils/semantic_text_index.py @@ -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 + ) diff --git a/tests/e2e/coverage_registry/mcp.yaml b/tests/e2e/coverage_registry/mcp.yaml index 1cdeac7b77f..0cd00a34b40 100644 --- a/tests/e2e/coverage_registry/mcp.yaml +++ b/tests/e2e/coverage_registry/mcp.yaml @@ -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 diff --git a/tests/e2e/mcp/datadog_mcp.py b/tests/e2e/mcp/datadog_mcp.py index 352b4446cfd..199b5be36f9 100644 --- a/tests/e2e/mcp/datadog_mcp.py +++ b/tests/e2e/mcp/datadog_mcp.py @@ -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 diff --git a/tests/e2e/mcp/mcp_client.py b/tests/e2e/mcp/mcp_client.py index 56f7fffba29..d35e39c0186 100644 --- a/tests/e2e/mcp/mcp_client.py +++ b/tests/e2e/mcp/mcp_client.py @@ -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 diff --git a/tests/e2e/mcp/test_mcp_tool_search_e2e.py b/tests/e2e/mcp/test_mcp_tool_search_e2e.py new file mode 100644 index 00000000000..847239e8a66 --- /dev/null +++ b/tests/e2e/mcp/test_mcp_tool_search_e2e.py @@ -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 diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 6e66529ec8f..e835a31c779 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -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): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 4575741aa8b..b74eb3c626d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -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: 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 c02ed1f37e5..750fadadab0 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_search.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_search.py @@ -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() diff --git a/tests/test_litellm/proxy/common_utils/test_semantic_text_index.py b/tests/test_litellm/proxy/common_utils/test_semantic_text_index.py new file mode 100644 index 00000000000..caa843f408e --- /dev/null +++ b/tests/test_litellm/proxy/common_utils/test_semantic_text_index.py @@ -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 diff --git a/tests/unit/llms/litellm_proxy/skills/test_skill_search.py b/tests/unit/llms/litellm_proxy/skills/test_skill_search.py index e28e927cf61..c0754f6f1af 100644 --- a/tests/unit/llms/litellm_proxy/skills/test_skill_search.py +++ b/tests/unit/llms/litellm_proxy/skills/test_skill_search.py @@ -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()