mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(mcp): batch embedding requests when building semantic tool filter index
Large MCP tool catalogs sent every tool description to the embedding provider in a single request, exceeding provider input-array limits (OpenAI 2048, DeepInfra 1024) and causing tool listing to hang or fail. LiteLLMRouterEncoder now splits inputs into batches no larger than a configurable embedding_batch_size (default 1024) and concatenates the results in order. Exposed via mcp_semantic_tool_filter.embedding_batch_size. Fixes #32528
This commit is contained in:
parent
86a9871ae9
commit
005a9c47b3
6 changed files with 260 additions and 19 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue