mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge d06fd4b312 into f285229b51
This commit is contained in:
commit
f13949a101
13 changed files with 482 additions and 50 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
86
tests/e2e/mcp/test_mcp_tool_search_e2e.py
Normal file
86
tests/e2e/mcp/test_mcp_tool_search_e2e.py
Normal 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
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue