diff --git a/litellm/constants.py b/litellm/constants.py index 7423d9b2211..3804f1b93af 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -89,6 +89,7 @@ DEFAULT_MCP_SEMANTIC_FILTER_TOP_K = int(os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_T DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD = float( os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD", 0.3) ) +DEFAULT_EMBEDDING_ENCODER_BATCH_SIZE = int(os.getenv("DEFAULT_EMBEDDING_ENCODER_BATCH_SIZE", 1024)) MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH = int(os.getenv("MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH", 150)) # Semantic Guard Defaults diff --git a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py index f24d5715e83..d75abc1ffdc 100644 --- a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py +++ b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py @@ -7,6 +7,7 @@ Filters MCP tools semantically for /chat/completions and /responses endpoints. from typing import TYPE_CHECKING, Any, Dict, List, Optional from litellm._logging import verbose_logger +from litellm.constants import DEFAULT_EMBEDDING_ENCODER_BATCH_SIZE from litellm.proxy._experimental.mcp_server.utils import MCP_TOOL_PREFIX_SEPARATOR if TYPE_CHECKING: @@ -25,6 +26,7 @@ class SemanticMCPToolFilter: top_k: int = 10, similarity_threshold: float = 0.3, enabled: bool = True, + embedding_batch_size: int = DEFAULT_EMBEDDING_ENCODER_BATCH_SIZE, ): """ Initialize the semantic tool filter. @@ -35,11 +37,15 @@ class SemanticMCPToolFilter: top_k: Maximum number of tools to return similarity_threshold: Minimum similarity score for filtering enabled: Whether filtering is enabled + embedding_batch_size: Maximum number of tool descriptions embedded per + provider request when building the router index. Prevents exceeding + the embedding provider's input-array limit for large tool catalogs. """ self.enabled = enabled self.top_k = top_k self.similarity_threshold = similarity_threshold self.embedding_model = embedding_model + self.embedding_batch_size = embedding_batch_size self.router_instance = litellm_router_instance self.tool_router: Optional["SemanticRouter"] = None self._tool_map: Dict[str, Any] = {} # MCPTool objects or OpenAI function dicts @@ -134,6 +140,7 @@ class SemanticMCPToolFilter: litellm_router_instance=self.router_instance, model_name=self.embedding_model, score_threshold=self.similarity_threshold, + embedding_batch_size=self.embedding_batch_size, ), auto_sync="local", ) diff --git a/litellm/proxy/hooks/mcp_semantic_filter/hook.py b/litellm/proxy/hooks/mcp_semantic_filter/hook.py index 7379096bf9b..da77e2738d9 100644 --- a/litellm/proxy/hooks/mcp_semantic_filter/hook.py +++ b/litellm/proxy/hooks/mcp_semantic_filter/hook.py @@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from litellm._logging import verbose_proxy_logger from litellm.constants import ( + DEFAULT_EMBEDDING_ENCODER_BATCH_SIZE, DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL, DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD, DEFAULT_MCP_SEMANTIC_FILTER_TOP_K, @@ -445,6 +446,7 @@ class SemanticToolFilterHook(CustomLogger): embedding_model = config.get("embedding_model", DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL) top_k = config.get("top_k", DEFAULT_MCP_SEMANTIC_FILTER_TOP_K) similarity_threshold = config.get("similarity_threshold", DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD) + embedding_batch_size = config.get("embedding_batch_size", DEFAULT_EMBEDDING_ENCODER_BATCH_SIZE) semantic_filter = SemanticMCPToolFilter( embedding_model=embedding_model, @@ -452,6 +454,7 @@ class SemanticToolFilterHook(CustomLogger): top_k=top_k, similarity_threshold=similarity_threshold, enabled=True, + embedding_batch_size=embedding_batch_size, ) # Build router from MCP registry on startup @@ -462,7 +465,8 @@ class SemanticToolFilterHook(CustomLogger): verbose_proxy_logger.info( f"✅ MCP Semantic Tool Filter enabled: " f"embedding_model={embedding_model}, top_k={top_k}, " - f"similarity_threshold={similarity_threshold}" + f"similarity_threshold={similarity_threshold}, " + f"embedding_batch_size={embedding_batch_size}" ) return hook diff --git a/litellm/router_strategy/auto_router/litellm_encoder.py b/litellm/router_strategy/auto_router/litellm_encoder.py index 7e163ba16a6..27bab26069f 100644 --- a/litellm/router_strategy/auto_router/litellm_encoder.py +++ b/litellm/router_strategy/auto_router/litellm_encoder.py @@ -5,6 +5,7 @@ from semantic_router.encoders import DenseEncoder from semantic_router.encoders.base import AsymmetricDenseMixin import litellm +from litellm.constants import DEFAULT_EMBEDDING_ENCODER_BATCH_SIZE if TYPE_CHECKING: from litellm.router import Router @@ -50,6 +51,7 @@ class LiteLLMRouterEncoder(CustomDenseEncoder, AsymmetricDenseMixin): litellm_router_instance: "Router", model_name: str, score_threshold: Union[float, None] = None, + embedding_batch_size: int = DEFAULT_EMBEDDING_ENCODER_BATCH_SIZE, ): """Initialize the LiteLLMEncoder. @@ -60,6 +62,11 @@ class LiteLLMRouterEncoder(CustomDenseEncoder, AsymmetricDenseMixin): :type model_name: str :param score_threshold: The score threshold for the embeddings. :type score_threshold: float + :param embedding_batch_size: Maximum number of inputs sent to the embedding + provider per request. Documents beyond this size are embedded across + multiple calls and concatenated, preserving order. Guards against + provider input-array limits (e.g. OpenAI 2048, DeepInfra 1024). + :type embedding_batch_size: int """ super().__init__( name=model_name, @@ -67,6 +74,16 @@ class LiteLLMRouterEncoder(CustomDenseEncoder, AsymmetricDenseMixin): ) self.model_name = model_name self.litellm_router_instance = litellm_router_instance + self.embedding_batch_size = ( + embedding_batch_size + if embedding_batch_size and embedding_batch_size > 0 + else DEFAULT_EMBEDDING_ENCODER_BATCH_SIZE + ) + + def _batches(self, docs: list[str]) -> tuple[list[str], ...]: + """Split docs into ordered batches no larger than embedding_batch_size.""" + size = self.embedding_batch_size + return tuple(docs[i : i + size] for i in range(0, len(docs), size)) def __call__(self, docs: list[Any], **kwargs) -> list[list[float]]: """Encode a list of text documents into embeddings using LiteLLM. @@ -86,34 +103,32 @@ class LiteLLMRouterEncoder(CustomDenseEncoder, AsymmetricDenseMixin): if self.litellm_router_instance is None: raise ValueError("litellm_router_instance is not set") try: - embeds = self.litellm_router_instance.embedding(input=docs, model=self.model_name, **kwargs) - return litellm_to_list(embeds) + return [ + embedding + for batch in self._batches(docs) + for embedding in litellm_to_list( + self.litellm_router_instance.embedding(input=batch, model=self.model_name, **kwargs) + ) + ] except Exception as e: raise ValueError(f"{self.type.capitalize()} API call failed. Error: {e}") from e def encode_documents(self, docs: list[str], **kwargs) -> list[list[float]]: - if self.litellm_router_instance is None: - raise ValueError("litellm_router_instance is not set") - try: - embeds = self.litellm_router_instance.embedding(input=docs, model=self.model_name, **kwargs) - return litellm_to_list(embeds) - except Exception as e: - raise ValueError(f"{self.type.capitalize()} API call failed. Error: {e}") from e + return self.encode_queries(docs, **kwargs) async def aencode_queries(self, docs: list[str], **kwargs) -> list[list[float]]: if self.litellm_router_instance is None: raise ValueError("litellm_router_instance is not set") try: - embeds = await self.litellm_router_instance.aembedding(input=docs, model=self.model_name, **kwargs) - return litellm_to_list(embeds) + return [ + embedding + for batch in self._batches(docs) + for embedding in litellm_to_list( + await self.litellm_router_instance.aembedding(input=batch, model=self.model_name, **kwargs) + ) + ] except Exception as e: raise ValueError(f"{self.type.capitalize()} API call failed. Error: {e}") from e async def aencode_documents(self, docs: list[str], **kwargs) -> list[list[float]]: - if self.litellm_router_instance is None: - raise ValueError("litellm_router_instance is not set") - try: - embeds = await self.litellm_router_instance.aembedding(input=docs, model=self.model_name, **kwargs) - return litellm_to_list(embeds) - except Exception as e: - raise ValueError(f"{self.type.capitalize()} API call failed. Error: {e}") from e + return await self.aencode_queries(docs, **kwargs) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py index 82c0aa3ccda..a7c7bd47dd2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py @@ -1356,3 +1356,107 @@ def test_truncate_csv_at_tool_name_boundary_edges(): assert _truncate_csv_at_tool_name_boundary(tool_names_csv="ab,cd,ef", max_length=5) == "ab,cd" assert _truncate_csv_at_tool_name_boundary(tool_names_csv="ab,cd,ef", max_length=4) == "ab" assert _truncate_csv_at_tool_name_boundary(tool_names_csv="single_name_longer_than_cap", max_length=10) == "" + + +def _batching_mock_router(): + """Router mock that returns one embedding per input item and records the + size of every embedding request, so batching can be asserted.""" + from litellm.types.utils import Embedding, EmbeddingResponse + + call_input_sizes: list[int] = [] + + def mock_embedding_sync(*args, **kwargs): + inputs = kwargs["input"] + call_input_sizes.append(len(inputs)) + return EmbeddingResponse( + data=[Embedding(embedding=[0.1] * 8, index=i, object="embedding") for i in range(len(inputs))], + model="text-embedding-3-small", + object="list", + usage={"prompt_tokens": 10, "total_tokens": 10}, + ) + + async def mock_embedding_async(*args, **kwargs): + return mock_embedding_sync(*args, **kwargs) + + mock_router = Mock() + mock_router.embedding = mock_embedding_sync + mock_router.aembedding = mock_embedding_async + return mock_router, call_input_sizes + + +def test_build_router_batches_embeddings_when_tools_exceed_batch_size(): + """ + Regression for https://github.com/BerriAI/litellm/issues/32528 + + With more tools than the embedding provider's input-array limit, building + the semantic router must split embedding requests into batches no larger + than embedding_batch_size instead of sending every tool in one oversized + call (which hangs / errors against the provider). + """ + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + + tools = [ + MCPTool(name=f"tool_{i}", description=f"Tool number {i}", inputSchema={"type": "object"}) + for i in range(2500) + ] + + mock_router, call_input_sizes = _batching_mock_router() + + filter_instance = SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=mock_router, + top_k=5, + similarity_threshold=0.3, + enabled=True, + embedding_batch_size=1024, + ) + + filter_instance._build_router(tools) + + assert filter_instance.tool_router is not None + assert len(filter_instance._tool_map) == len(tools) + assert call_input_sizes, "Expected embedding to be called while building the router" + # The core fix: no single embedding request exceeds the configured limit. + assert max(call_input_sizes) <= 1024, f"A batch exceeded the limit: {max(call_input_sizes)}" + # 2500 utterances at batch size 1024 require ceil(2500/1024) == 3 batched calls; + # a non-batching encoder would embed all 2500 in a single oversized request. + assert sum(1 for size in call_input_sizes if size > 1) >= 3 + assert sum(call_input_sizes) >= len(tools) + + +def test_semantic_filter_forwards_embedding_batch_size_to_encoder(): + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + + mock_router, _ = _batching_mock_router() + + filter_instance = SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=mock_router, + embedding_batch_size=256, + ) + + captured = {} + + real_encoder_cls = __import__( + "litellm.router_strategy.auto_router.litellm_encoder", + fromlist=["LiteLLMRouterEncoder"], + ).LiteLLMRouterEncoder + + class _SpyEncoder(real_encoder_cls): # type: ignore[misc, valid-type] + def __init__(self, *args, **kwargs): + captured["embedding_batch_size"] = kwargs.get("embedding_batch_size") + super().__init__(*args, **kwargs) + + with patch( + "litellm.router_strategy.auto_router.litellm_encoder.LiteLLMRouterEncoder", + _SpyEncoder, + ): + filter_instance._build_router( + [MCPTool(name="t", description="d", inputSchema={"type": "object"})] + ) + + assert captured["embedding_batch_size"] == 256 diff --git a/tests/test_litellm/router_strategy/auto_router/test_litellm_encoder.py b/tests/test_litellm/router_strategy/auto_router/test_litellm_encoder.py new file mode 100644 index 00000000000..d7b467dd33d --- /dev/null +++ b/tests/test_litellm/router_strategy/auto_router/test_litellm_encoder.py @@ -0,0 +1,110 @@ +""" +Regression tests for LiteLLMRouterEncoder embedding batching. + +Large MCP tool catalogs (semantic tool filter) and auto-router route sets can +produce more documents than an embedding provider accepts in a single request +(e.g. OpenAI 2048, DeepInfra 1024). The encoder must split the input into +batches no larger than ``embedding_batch_size`` and concatenate the results in +order instead of sending everything in one oversized call. +""" + +from typing import List, Optional + +import pytest + +import litellm +from litellm.router_strategy.auto_router.litellm_encoder import LiteLLMRouterEncoder + + +class _RecordingRouter: + """Fake Router that records embedding call inputs and echoes a deterministic + embedding per document so ordering can be asserted.""" + + def __init__(self) -> None: + self.sync_batches: List[List[str]] = [] + self.async_batches: List[List[str]] = [] + + @staticmethod + def _response(docs: List[str]) -> litellm.EmbeddingResponse: + return litellm.EmbeddingResponse( + model="m", + data=[ + {"embedding": [float(len(doc))], "index": i, "object": "embedding"} + for i, doc in enumerate(docs) + ], + ) + + def embedding(self, input: List[str], model: str, **kwargs) -> litellm.EmbeddingResponse: + self.sync_batches.append(list(input)) + return self._response(input) + + async def aembedding(self, input: List[str], model: str, **kwargs) -> litellm.EmbeddingResponse: + self.async_batches.append(list(input)) + return self._response(input) + + +def _make_encoder(router: _RecordingRouter, batch_size: Optional[int] = None) -> LiteLLMRouterEncoder: + kwargs = {} if batch_size is None else {"embedding_batch_size": batch_size} + return LiteLLMRouterEncoder( + litellm_router_instance=router, # type: ignore[arg-type] + model_name="openai/text-embedding-3-small", + **kwargs, + ) + + +def test_encode_documents_splits_into_batches(): + router = _RecordingRouter() + encoder = _make_encoder(router, batch_size=3) + docs = [f"doc-{i}" for i in range(7)] + + result = encoder.encode_documents(docs) + + # 7 docs, batch size 3 -> batches of 3, 3, 1 + assert [len(b) for b in router.sync_batches] == [3, 3, 1] + assert all(len(b) <= 3 for b in router.sync_batches) + # concatenation preserves order: one embedding per doc, in input order + assert result == [[float(len(d))] for d in docs] + + +def test_encode_documents_single_batch_when_under_limit(): + router = _RecordingRouter() + encoder = _make_encoder(router, batch_size=100) + docs = [f"doc-{i}" for i in range(10)] + + encoder.encode_documents(docs) + + assert len(router.sync_batches) == 1 + assert router.sync_batches[0] == docs + + +@pytest.mark.asyncio +async def test_aencode_documents_splits_into_batches(): + router = _RecordingRouter() + encoder = _make_encoder(router, batch_size=2) + docs = [f"doc-{i}" for i in range(5)] + + result = await encoder.aencode_documents(docs) + + assert [len(b) for b in router.async_batches] == [2, 2, 1] + assert result == [[float(len(d))] for d in docs] + + +def test_call_uses_batching(): + router = _RecordingRouter() + encoder = _make_encoder(router, batch_size=4) + docs = [f"doc-{i}" for i in range(9)] + + result = encoder(docs) + + assert all(len(b) <= 4 for b in router.sync_batches) + assert sum(len(b) for b in router.sync_batches) == 9 + assert result == [[float(len(d))] for d in docs] + + +def test_invalid_batch_size_falls_back_to_default(): + from litellm.constants import DEFAULT_EMBEDDING_ENCODER_BATCH_SIZE + + router = _RecordingRouter() + encoder = _make_encoder(router, batch_size=0) + + assert encoder.embedding_batch_size == DEFAULT_EMBEDDING_ENCODER_BATCH_SIZE