feat(mcp): semantic tool search for the native MCP Gateway (#39404)

The mcp_tool_search virtual tool only did substring token matching, so a native MCP client asking for "FX" could not find a tool described as "foreign exchange rates" even though the same catalog is ranked by embeddings on /responses and /chat/completions.

Adds litellm_settings.mcp_tool_search (embedding_model, top_k, similarity_threshold, core_tools). With an embedding model the caller's authorized catalog from _list_mcp_tools is ranked by cosine similarity of name plus description; configured core tools the caller can reach come first and do not consume top_k. Without an embedding model the keyword fallback keeps the old behavior. Settings are hot-reloadable from the DB, exposed on /get and /update mcp_tool_search_settings, and editable from the Admin UI under MCP Servers > Tool Search. The embedding index is shared with agent_search via a new SemanticTextIndex.

Resolves LIT-6751

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-02 15:09:44 -07:00 • committed by GitHub
parent a701effbad
commit a76cb6feaf
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 1328 additions and 164 deletions

View file

@ -1,6 +1,6 @@
{
"reportAny": {
"limit": 14076
"limit": 14074
},
"reportArgumentType": {
"limit": 2216
@ -24,7 +24,7 @@
"limit": 19
},
"reportExplicitAny": {
"limit": 4128
"limit": 4125
},
"reportFunctionMemberAccess": {
"limit": 7

View file

@ -29,7 +29,7 @@ def _dev_env_hot_reload_enabled() -> bool:
if os.getenv("LITELLM_MODE", "DEV") == "DEV":
_dotenv.load_dotenv(override=_dev_env_hot_reload_enabled())
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from typing import (
Any,
Callable,
@ -490,6 +490,7 @@ public_mcp_hub_strict_whitelist: bool = True
public_model_groups: Optional[List[str]] = None
public_agent_groups: Optional[List[str]] = None
agent_search_embedding_model: Optional[str] = None
mcp_tool_search: Optional[Mapping[str, object]] = None
# Supports both old format (Dict[str, str]) and new format (Dict[str, Dict[str, Any]])
# New format: { "displayName": { "url": "...", "index": 0 } }
# Old format: { "displayName": "url" } (for backward compatibility)

View file

@ -1742,6 +1742,7 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES: Final = [
"anthropic_prompt_caching_ttl",
"max_ui_session_budget",
"budget_rollover",
"mcp_tool_search",
]
SPECIAL_LITELLM_AUTH_TOKEN: Final = ["ui-token"]
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60))

View file

@ -2,20 +2,31 @@ from __future__ import annotations
import json
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, TypedDict, assert_never
from pydantic import ValidationError
from typing_extensions import ReadOnly, Required
import litellm
from litellm.proxy.agent_endpoints.agent_search import DEFAULT_AGENT_SEARCH_TOP_K
from litellm.proxy.common_utils.semantic_text_index import (
Embedder,
EmbeddingFailed,
SemanticTextIndex,
router_embedder,
)
from litellm.types.mcp import MCPToolSearchSettings
if TYPE_CHECKING:
from mcp.types import CallToolResult
from mcp.types import CallToolResult, Tool
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
MCP_TOOL_SEARCH_SETTINGS_KEY: Final[str] = "mcp_tool_search"
MCP_TOOL_SEARCH_TOOL_NAME: Final[str] = "mcp_tool_search"
MCP_TOOL_CALL_TOOL_NAME: Final[str] = "mcp_tool_call"
AGENT_SEARCH_TOOL_NAME: Final[str] = "agent_search"
@ -29,17 +40,91 @@ def coerce_top_k(value: Any, default: int = 5) -> int:
return default
def search_tools(query: str, tools: list[dict[str, Any]], top_k: int = 5) -> list[dict[str, Any]]:
class ToolSearchResult(TypedDict, total=False):
name: Required[ReadOnly[str]]
description: Required[ReadOnly[str]]
inputSchema: Required[ReadOnly[Mapping[str, object]]]
score: ReadOnly[float]
@dataclass(frozen=True, slots=True)
class SemanticToolRanker:
embed: Embedder
embedding_model: str
index: SemanticTextIndex
global_mcp_tool_search_index: Final = SemanticTextIndex()
def mcp_tool_search_settings() -> MCPToolSearchSettings | ValidationError:
try:
return MCPToolSearchSettings.model_validate(litellm.mcp_tool_search or {})
except ValidationError as exc:
return exc
def _tool_result(tool: Tool) -> ToolSearchResult:
return {"name": tool.name, "description": tool.description or "", "inputSchema": tool.inputSchema}
def _scored_result(tool: Tool, score: float) -> ToolSearchResult:
return {"name": tool.name, "description": tool.description or "", "inputSchema": tool.inputSchema, "score": score}
def _tool_text(tool: Tool) -> str:
return "\n".join(part for part in (tool.name, tool.description or "") if part)
def _keyword_score(query: str, tool: Tool) -> float:
haystack: Final = _tool_text(tool).lower()
return float(sum(1 for token in query.lower().split() if token in haystack))
def _split_core_tools(tools: Sequence[Tool], core_tools: Sequence[str]) -> tuple[tuple[Tool, ...], tuple[Tool, ...]]:
by_name: Final = MappingProxyType({tool.name: tool for tool in tools})
core: Final = tuple(by_name[name] for name in dict.fromkeys(core_tools) if name in by_name)
rest: Final = tuple(tool for tool in tools if tool.name not in frozenset(core_tools))
return core, rest
def _top_hits(
tools: Sequence[Tool], scores: Sequence[float], minimum: float, limit: int
) -> tuple[tuple[float, Tool], ...]:
hits: Final = ((score, tool) for score, tool in zip(scores, tools, strict=True) if score >= minimum)
return tuple(sorted(hits, key=lambda hit: hit[0], reverse=True)[:limit])
def search_tools(query: str, tools: Sequence[Tool], top_k: int = 5) -> tuple[ToolSearchResult, ...]:
"""Keyword fallback used when no embedding model is configured: one point per query token found in the tool."""
if not query:
return []
tokens: Final = query.lower().split()
return ()
scores: Final = tuple(_keyword_score(query, tool) for tool in tools)
return tuple(_tool_result(tool) for _, tool in _top_hits(tools, scores, minimum=1.0, limit=top_k))
def _score(tool: dict[str, Any]) -> int:
haystack: Final = (tool.get("name", "") + " " + tool.get("description", "")).lower()
return sum(1 for t in tokens if t in haystack)
scored: Final = ((s, tool) for tool in tools if (s := _score(tool)) > 0)
return [tool for _, tool in sorted(scored, key=lambda x: x[0], reverse=True)[:top_k]]
async def search_mcp_tools(
query: str,
tools: Sequence[Tool],
top_k: int,
settings: MCPToolSearchSettings,
ranker: SemanticToolRanker | None,
) -> tuple[ToolSearchResult, ...] | EmbeddingFailed:
"""Core tools the caller can access come first, then up to `top_k` ranked matches from the remaining tools."""
core, rest = _split_core_tools(tools, settings.core_tools)
limit: Final = min(top_k, settings.top_k)
core_results: Final = tuple(_tool_result(tool) for tool in core)
if ranker is None:
return (*core_results, *search_tools(query, rest, limit))
if not query:
return core_results
scores: Final = await ranker.index.scores(
query, tuple(_tool_text(tool) for tool in rest), ranker.embed, ranker.embedding_model
)
if isinstance(scores, EmbeddingFailed):
return scores
hits: Final = _top_hits(rest, scores, minimum=settings.similarity_threshold, limit=limit)
return (*core_results, *(_scored_result(tool, score) for score, tool in hits))
class _ToolParamSchema(TypedDict, total=False):
@ -66,11 +151,17 @@ def _json_array(*items: str) -> Sequence[str]:
_MCP_TOOL_SEARCH_DEFINITION: Final[VirtualToolDefinition] = {
"name": MCP_TOOL_SEARCH_TOOL_NAME,
"description": "Search for MCP tools by keyword. Returns top matching tools with names, descriptions, and input schemas.",
"description": (
"Search for MCP tools by describing what you need. "
"Returns top matching tools with names, descriptions, and input schemas."
),
"inputSchema": {
"type": "object",
"properties": {
"query": {"type": "string", "description": "Keywords to search for in tool names and descriptions."},
"query": {
"type": "string",
"description": "What the tool should do, matched against names and descriptions.",
},
"top_k": {"type": "integer", "description": "Maximum number of results to return.", "default": 5},
},
"required": _json_array("query"),
@ -165,10 +256,28 @@ async def handle_mcp_tool_search(
oauth2_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
) -> CallToolResult:
from mcp.types import CallToolResult, TextContent
from litellm.proxy._experimental.mcp_server.server import _list_mcp_tools
from litellm.proxy.proxy_server import llm_router
settings: Final = mcp_tool_search_settings()
if isinstance(settings, ValidationError):
return _text_tool_result(
f"litellm_settings.{MCP_TOOL_SEARCH_SETTINGS_KEY} is invalid: {settings}", is_error=True
)
if settings.embedding_model is not None and llm_router is None:
return _text_tool_result(
f"litellm_settings.{MCP_TOOL_SEARCH_SETTINGS_KEY}.embedding_model needs a model_list so it can be called",
is_error=True,
)
ranker: Final = (
SemanticToolRanker(
embed=router_embedder(llm_router, settings.embedding_model, user_api_key_dict),
embedding_model=settings.embedding_model,
index=global_mcp_tool_search_index,
)
if settings.embedding_model is not None and llm_router is not None
else None
)
mcp_listing: Final = await _list_mcp_tools(
user_api_key_auth=user_api_key_dict,
mcp_servers=mcp_servers,
@ -178,17 +287,10 @@ async def handle_mcp_tool_search(
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
mcp_tools: Final = mcp_listing.tools
tools: Final = [
{
"name": t.name,
"description": t.description or "",
"inputSchema": t.inputSchema,
}
for t in mcp_tools
]
results: Final = search_tools(query, tools, top_k)
return CallToolResult(content=[TextContent(type="text", text=json.dumps(results))], isError=False)
results: Final = await search_mcp_tools(query, mcp_listing.tools, top_k, settings, ranker)
if isinstance(results, EmbeddingFailed):
return _text_tool_result(results.reason, is_error=True)
return _text_tool_result(json.dumps(results), is_error=False)
async def handle_mcp_tool_call(

View file

@ -2,17 +2,18 @@
from __future__ import annotations
import math
from collections.abc import Awaitable, Mapping, Sequence
from collections.abc import Sequence
from dataclasses import dataclass
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Protocol, TypeAlias
from typing import TYPE_CHECKING, Final, TypeAlias
from openai import OpenAIError
from pydantic import BaseModel, ConfigDict, ValidationError
from litellm.exceptions import BudgetExceededError
from litellm.proxy.common_utils.semantic_text_index import (
Embedder,
EmbeddingFailed,
SemanticTextIndex,
router_embedder,
)
from litellm.types.agents import AgentResponse
if TYPE_CHECKING:
@ -21,12 +22,6 @@ if TYPE_CHECKING:
DEFAULT_AGENT_SEARCH_TOP_K: Final = 5
Vector: TypeAlias = tuple[float, ...]
class Embedder(Protocol):
def __call__(self, texts: Sequence[str]) -> Awaitable[Sequence[Vector]]: ...
@dataclass(frozen=True, slots=True)
class AgentSearchHit:
@ -67,18 +62,6 @@ class _SearchableCard(BaseModel):
skills: tuple[_SearchableSkill, ...] = ()
class _EmbeddingItem(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
embedding: tuple[float, ...]
class _EmbeddingData(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
data: tuple[_EmbeddingItem, ...]
class AgentSearchResult(BaseModel):
model_config = ConfigDict(frozen=True)
@ -117,110 +100,21 @@ def agent_search_result(hit: AgentSearchHit) -> AgentSearchResult:
)
def cosine_similarity(left: Vector, right: Vector) -> float:
dot: Final = sum(a * b for a, b in zip(left, right, strict=True))
norms: Final = math.sqrt(sum(a * a for a in left)) * math.sqrt(sum(b * b for b in right))
return dot / norms if norms else 0.0
def embedding_spend_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]:
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
return { # mutable-ok: the router mutates the metadata dict it is handed
**LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict),
"user_api_key": user_api_key_dict.api_key,
}
def router_embedder(router: Router, embedding_model: str, user_api_key_dict: UserAPIKeyAuth) -> Embedder:
async def embed(texts: Sequence[str]) -> Sequence[Vector]:
batch: Final = list(texts) # mutable-ok: Router.aembedding accepts only str | list input
response: Final = await router.aembedding(
model=embedding_model, input=batch, metadata=embedding_spend_metadata(user_api_key_dict)
)
return tuple(item.embedding for item in _EmbeddingData.model_validate(response.model_dump()).data)
return embed
_NO_VECTORS: Final[Mapping[str, Vector]] = MappingProxyType({})
async def _embed_all(embed: Embedder, texts: Sequence[str]) -> tuple[Vector, ...] | AgentSearchEmbeddingFailed:
try:
vectors: Final = tuple(await embed(texts))
except (OpenAIError, ValueError, BudgetExceededError) as exc:
return AgentSearchEmbeddingFailed(reason=f"embedding the search query failed: {exc}")
if len(vectors) != len(texts):
return AgentSearchEmbeddingFailed(
reason=f"embedding model returned {len(vectors)} vectors for {len(texts)} inputs"
)
return vectors
@dataclass(frozen=True, slots=True)
class _Embedded:
query_vector: Vector
vectors: Mapping[str, Vector]
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)
async def _embed_query_and_agents(
embed: Embedder, query: str, texts: Sequence[str], cached: Mapping[str, Vector]
) -> _Embedded | AgentSearchEmbeddingFailed:
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, AgentSearchEmbeddingFailed):
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))
if isinstance(reembedded, AgentSearchEmbeddingFailed):
return reembedded
return _Embedded(
query_vector=reembedded[0], vectors=MappingProxyType(dict(zip(unique, reembedded[1:], strict=True)))
)
class AgentSearchIndex:
"""Caches one vector per distinct agent text per embedding model, so repeat searches only embed the query."""
def __init__(self) -> None:
self._vectors: Mapping[str, Mapping[str, Vector]] = MappingProxyType({})
def _merged(self, embedding_model: str, embedded: _Embedded) -> Mapping[str, Vector]:
kept: Final = {
text: vector
for text, vector in self._vectors.get(embedding_model, _NO_VECTORS).items()
if len(vector) == len(embedded.query_vector)
}
return MappingProxyType({**kept, **embedded.vectors})
self._index: Final = SemanticTextIndex()
async def search(
self, query: str, agents: Sequence[AgentResponse], top_k: int, embed: Embedder, embedding_model: str
) -> AgentSearchHits | AgentSearchEmbeddingFailed:
if not agents:
return AgentSearchHits(hits=())
texts: Final = tuple(agent_search_text(agent) for agent in agents)
cached: Final = self._vectors.get(embedding_model, _NO_VECTORS)
embedded: Final = await _embed_query_and_agents(embed, query, texts, cached)
if isinstance(embedded, AgentSearchEmbeddingFailed):
return embedded
if not _same_dimension(embedded.query_vector, embedded.vectors, texts):
return AgentSearchEmbeddingFailed(
reason=f"embedding model {embedding_model} returned vectors of mixed dimensions"
)
self._vectors = MappingProxyType({**self._vectors, embedding_model: self._merged(embedding_model, embedded)})
scores: Final = await self._index.scores(query, texts, embed, embedding_model)
if isinstance(scores, EmbeddingFailed):
return AgentSearchEmbeddingFailed(reason=scores.reason)
ranked: Final = sorted(
(
AgentSearchHit(agent=agent, score=cosine_similarity(embedded.query_vector, embedded.vectors[text]))
for agent, text in zip(agents, texts, strict=True)
),
(AgentSearchHit(agent=agent, score=score) for agent, score in zip(agents, scores, strict=True)),
key=lambda hit: hit.score,
reverse=True,
)

View file

@ -0,0 +1,142 @@
"""Embedding-similarity ranking over short texts with a per-model vector cache, shared by agent search and MCP tool search."""
from __future__ import annotations
import math
from collections.abc import Awaitable, Mapping, Sequence
from dataclasses import dataclass
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Protocol, TypeAlias
from openai import OpenAIError
from pydantic import BaseModel, ConfigDict
from litellm.exceptions import BudgetExceededError
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
from litellm.router import Router
Vector: TypeAlias = tuple[float, ...]
class Embedder(Protocol):
def __call__(self, texts: Sequence[str]) -> Awaitable[Sequence[Vector]]: ...
@dataclass(frozen=True, slots=True)
class EmbeddingFailed:
reason: str
class _EmbeddingItem(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
embedding: tuple[float, ...]
class _EmbeddingData(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
data: tuple[_EmbeddingItem, ...]
def cosine_similarity(left: Vector, right: Vector) -> float:
dot: Final = sum(a * b for a, b in zip(left, right, strict=True))
norms: Final = math.sqrt(sum(a * a for a in left)) * math.sqrt(sum(b * b for b in right))
return dot / norms if norms else 0.0
def embedding_spend_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]: # mutable-ok: router mutates it
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
return { # mutable-ok: the router mutates the metadata dict it is handed
**LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict),
"user_api_key": user_api_key_dict.api_key,
}
def router_embedder(router: Router, embedding_model: str, user_api_key_dict: UserAPIKeyAuth) -> Embedder:
async def embed(texts: Sequence[str]) -> Sequence[Vector]:
batch: Final = list(texts) # mutable-ok: Router.aembedding accepts only str | list input
response: Final = await router.aembedding(
model=embedding_model, input=batch, metadata=embedding_spend_metadata(user_api_key_dict)
)
return tuple(item.embedding for item in _EmbeddingData.model_validate(response.model_dump()).data)
return embed
_NO_VECTORS: Final[Mapping[str, Vector]] = MappingProxyType({})
async def _embed_all(embed: Embedder, texts: Sequence[str]) -> tuple[Vector, ...] | EmbeddingFailed:
try:
vectors: Final = tuple(await embed(texts))
except (OpenAIError, ValueError, BudgetExceededError) as exc:
return EmbeddingFailed(reason=f"embedding the search query failed: {exc}")
if len(vectors) != len(texts):
return EmbeddingFailed(reason=f"embedding model returned {len(vectors)} vectors for {len(texts)} inputs")
return vectors
@dataclass(frozen=True, slots=True)
class _Embedded:
query_vector: Vector
vectors: Mapping[str, Vector]
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)
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))
if isinstance(reembedded, EmbeddingFailed):
return reembedded
return _Embedded(
query_vector=reembedded[0], vectors=MappingProxyType(dict(zip(unique, reembedded[1:], strict=True)))
)
class SemanticTextIndex:
"""Caches one vector per distinct text per embedding model, so repeat searches only embed the query."""
def __init__(self) -> None:
self._vectors: Mapping[str, Mapping[str, Vector]] = MappingProxyType({})
def _merged(self, embedding_model: str, embedded: _Embedded) -> Mapping[str, Vector]:
kept: Final = MappingProxyType(
{
text: vector
for text, vector in self._vectors.get(embedding_model, _NO_VECTORS).items()
if len(vector) == len(embedded.query_vector)
}
)
return MappingProxyType({**kept, **embedded.vectors})
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."""
if not texts:
return ()
cached: Final = self._vectors.get(embedding_model, _NO_VECTORS)
embedded: Final = await _embed_query_and_texts(embed, query, texts, cached)
if isinstance(embedded, EmbeddingFailed):
return embedded
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 = MappingProxyType({**self._vectors, embedding_model: self._merged(embedding_model, embedded)})
return tuple(cosine_similarity(embedded.query_vector, embedded.vectors[text]) for text in texts)

View file

@ -21,6 +21,7 @@ from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys
from litellm.proxy._experimental.mcp_server.tool_search import MCP_TOOL_SEARCH_SETTINGS_KEY
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.config_resolvers.sso import (
@ -38,6 +39,7 @@ from litellm.repositories.table_repositories import (
UISettingsRepository,
)
from litellm.repositories.team_repository import TeamRepository
from litellm.types.mcp import MCPToolSearchSettings
from litellm.types.proxy.management_endpoints.ui_sso import (
DefaultTeamSSOParams,
SSOConfig,
@ -448,6 +450,10 @@ class MCPSemanticFilterSettingsResponse(SettingsResponse):
"""Response model for MCP semantic filter settings"""
class MCPToolSearchSettingsResponse(SettingsResponse):
"""Response model for native MCP tool search settings"""
@router.get(
"/get/allowed_ips",
tags=["Budget & Spend Tracking"],
@ -835,7 +841,7 @@ async def update_default_team_member_budget(teams: list[NewUserRequestTeam], use
async def _update_litellm_setting(
settings: DefaultInternalUserParams | DefaultTeamSSOParams | MCPSemanticFilterSettings,
settings: DefaultInternalUserParams | DefaultTeamSSOParams | MCPSemanticFilterSettings | MCPToolSearchSettings,
settings_key: str,
success_message: str,
user_api_key_dict: UserAPIKeyAuth,
@ -861,7 +867,7 @@ async def _update_litellm_setting(
detail={"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."},
)
in_memory_var: Final = settings.model_dump(exclude_none=True)
in_memory_var: Final = settings.model_dump(mode="json", exclude_none=True)
# Load existing config first, then set in-memory value after,
# because get_config() may overwrite litellm.<key> with stale DB values
@ -1359,6 +1365,59 @@ async def update_mcp_semantic_filter_settings(
return result
@router.get(
"/get/mcp_tool_search_settings",
tags=["Settings"], # mutable-ok: FastAPI's route decorator only accepts a list
dependencies=[Depends(user_api_key_auth)], # mutable-ok: FastAPI's route decorator only accepts a list
response_model=MCPToolSearchSettingsResponse,
)
async def get_mcp_tool_search_settings(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> Mapping[str, object]:
"""
Get the `litellm_settings.mcp_tool_search` configuration used by the native `mcp_tool_search` virtual tool.
"""
from litellm.proxy.proxy_server import prisma_client, proxy_config
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected. Please connect a database.")
config: Final = await proxy_config.get_config()
return await _get_settings_with_schema(
settings_key=MCP_TOOL_SEARCH_SETTINGS_KEY,
settings_class=MCPToolSearchSettings,
config=config,
)
@router.patch(
"/update/mcp_tool_search_settings",
tags=["Settings"], # mutable-ok: FastAPI's route decorator only accepts a list
dependencies=[Depends(user_api_key_auth)], # mutable-ok: FastAPI's route decorator only accepts a list
)
async def update_mcp_tool_search_settings(
settings: MCPToolSearchSettings,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> Mapping[str, object]:
"""
Update `litellm_settings.mcp_tool_search` in the database.
Settings will be picked up by all pods within approximately 10 seconds via background polling.
"""
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(
status_code=403,
detail="Only proxy admins can update MCP tool search settings.",
)
return await _update_litellm_setting(
settings=settings,
settings_key=MCP_TOOL_SEARCH_SETTINGS_KEY,
success_message="MCP tool search settings updated successfully. Changes will be applied across all pods within 10 seconds.",
user_api_key_dict=user_api_key_dict,
)
UI_SETTINGS_CACHE_KEY: Final = "ui_settings:settings_dict"
UI_SETTINGS_CACHE_TTL: Final = 600 # 10 minutes

View file

@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal
from urllib.parse import urlsplit
import httpx
from pydantic import BaseModel
from pydantic import BaseModel, ConfigDict, Field
from typing_extensions import TypedDict
from litellm.types.llms.base import HiddenParams
@ -91,6 +91,33 @@ class MCPPublicServer(BaseModel):
mcp_info: dict[str, Any] | None = None
class MCPToolSearchSettings(BaseModel):
"""`litellm_settings.mcp_tool_search`: how the native `mcp_tool_search` virtual tool ranks the caller's tools."""
model_config = ConfigDict(frozen=True)
embedding_model: str | None = Field(
default=None,
description="Embedding model from model_list used to rank tools by meaning. Unset keeps keyword matching.",
)
top_k: int = Field(
default=5,
ge=1,
le=100,
description="Most ranked tools a search returns. A smaller top_k in the tool call wins. Core tools do not count.",
)
similarity_threshold: float = Field(
default=0.0,
ge=0.0,
le=1.0,
description="Lowest cosine similarity a tool needs to appear in semantic results (0.0 = no cutoff).",
)
core_tools: tuple[str, ...] = Field(
default=(),
description="Tool names always returned first when the caller can access them, e.g. `my_server-get_rates`.",
)
# OAuth 2.0 token-endpoint client authentication method (RFC 6749 section 2.3.1).
MCPTokenEndpointAuthMethod = Literal["client_secret_basic", "client_secret_post"]

View file

@ -10,33 +10,36 @@ Covers:
"""
import json
from collections.abc import Sequence
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from mcp.types import Tool
import litellm
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
from litellm.proxy._experimental.mcp_server.tool_search import (
AGENT_SEARCH_TOOL_NAME,
MCP_TOOL_CALL_TOOL_NAME,
MCP_TOOL_SEARCH_TOOL_NAME,
SemanticToolRanker,
ToolSearchResult,
coerce_top_k,
get_virtual_tool_definitions,
search_mcp_tools,
search_tools,
)
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.common_utils.semantic_text_index import EmbeddingFailed, SemanticTextIndex, Vector
from litellm.types.mcp import MCPToolSearchSettings
def _make_tools(specs: list[tuple[str, str]]) -> list[dict[str, Any]]:
return [
{
"name": name,
"description": desc,
"inputSchema": {"type": "object", "properties": {}},
}
for name, desc in specs
]
def _make_tools(specs: list[tuple[str, str]]) -> tuple[Tool, ...]:
return tuple(
Tool(name=name, description=desc, inputSchema={"type": "object", "properties": {}}) for name, desc in specs
)
def _make_perm(**kwargs: Any) -> LiteLLM_ObjectPermissionTable:
@ -54,6 +57,160 @@ SAMPLE_TOOLS = _make_tools(
)
FX_TOOL = Tool(
name="treasury-get_rates",
description="Get foreign exchange rates for a currency pair",
inputSchema={"type": "object", "properties": {}},
)
WEATHER_TOOL = Tool(
name="weather-forecast",
description="Get the weather forecast for a city",
inputSchema={"type": "object", "properties": {}},
)
CALENDAR_TOOL = Tool(
name="calendar-create_event",
description="Create a calendar event",
inputSchema={"type": "object", "properties": {}},
)
CATALOG = (FX_TOOL, WEATHER_TOOL, CALENDAR_TOOL)
# A stand-in embedding space: "FX" sits next to the foreign-exchange tool and far from the rest.
FAKE_VECTORS: dict[str, Vector] = {
"FX": (1.0, 0.0),
f"{FX_TOOL.name}\n{FX_TOOL.description}": (0.9, 0.1),
f"{WEATHER_TOOL.name}\n{WEATHER_TOOL.description}": (0.3, 1.0),
f"{CALENDAR_TOOL.name}\n{CALENDAR_TOOL.description}": (0.0, 1.0),
}
class RecordingEmbedder:
def __init__(self) -> None:
self.calls: list[tuple[str, ...]] = []
async def __call__(self, texts: Sequence[str]) -> Sequence[Vector]:
self.calls.append(tuple(texts))
return tuple(FAKE_VECTORS[text] for text in texts)
def _ranker(embedder: RecordingEmbedder | None = None) -> SemanticToolRanker:
return SemanticToolRanker(embed=embedder or RecordingEmbedder(), embedding_model="emb", index=SemanticTextIndex())
def _names(results: Sequence[ToolSearchResult] | EmbeddingFailed) -> list[str]:
assert not isinstance(results, EmbeddingFailed)
return [tool["name"] for tool in results]
class TestSearchMcpTools:
@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)
assert _names(keyword_only) == []
results = await search_mcp_tools("FX", CATALOG, 5, MCPToolSearchSettings(embedding_model="emb"), _ranker())
assert _names(results) == [FX_TOOL.name, WEATHER_TOOL.name, CALENDAR_TOOL.name]
assert not isinstance(results, EmbeddingFailed)
assert results[0]["score"] > results[1]["score"] > results[2]["score"]
assert results[0]["inputSchema"] == FX_TOOL.inputSchema
@pytest.mark.asyncio
async def test_similarity_threshold_drops_weak_matches(self) -> None:
settings = MCPToolSearchSettings(embedding_model="emb", similarity_threshold=0.5)
results = await search_mcp_tools("FX", CATALOG, 5, settings, _ranker())
assert _names(results) == [FX_TOOL.name]
@pytest.mark.asyncio
async def test_request_top_k_limits_semantic_results(self) -> None:
results = await search_mcp_tools("FX", CATALOG, 2, MCPToolSearchSettings(embedding_model="emb"), _ranker())
assert _names(results) == [FX_TOOL.name, WEATHER_TOOL.name]
@pytest.mark.asyncio
async def test_configured_top_k_caps_request_top_k(self) -> None:
settings = MCPToolSearchSettings(embedding_model="emb", top_k=1)
assert _names(await search_mcp_tools("FX", CATALOG, 50, settings, _ranker())) == [FX_TOOL.name]
assert _names(await search_mcp_tools("weather", CATALOG, 50, MCPToolSearchSettings(top_k=1), None)) == [
WEATHER_TOOL.name
]
@pytest.mark.asyncio
async def test_core_tools_lead_and_do_not_consume_top_k(self) -> None:
settings = MCPToolSearchSettings(embedding_model="emb", top_k=1, core_tools=(CALENDAR_TOOL.name,))
results = await search_mcp_tools("FX", CATALOG, 1, settings, _ranker())
assert _names(results) == [CALENDAR_TOOL.name, FX_TOOL.name]
assert not isinstance(results, EmbeddingFailed)
assert "score" not in results[0]
@pytest.mark.asyncio
async def test_core_tools_apply_in_keyword_mode_too(self) -> None:
settings = MCPToolSearchSettings(core_tools=(CALENDAR_TOOL.name,))
assert _names(await search_mcp_tools("weather", CATALOG, 5, settings, None)) == [
CALENDAR_TOOL.name,
WEATHER_TOOL.name,
]
@pytest.mark.asyncio
async def test_core_tools_outside_the_callers_catalog_are_not_returned(self) -> None:
settings = MCPToolSearchSettings(embedding_model="emb", core_tools=("payroll-run", CALENDAR_TOOL.name))
results = await search_mcp_tools("FX", (FX_TOOL, WEATHER_TOOL), 5, settings, _ranker())
assert _names(results) == [FX_TOOL.name, WEATHER_TOOL.name]
@pytest.mark.asyncio
async def test_core_tools_are_listed_once_and_never_embedded(self) -> None:
embedder = RecordingEmbedder()
settings = MCPToolSearchSettings(embedding_model="emb", core_tools=(FX_TOOL.name, FX_TOOL.name))
results = await search_mcp_tools("FX", CATALOG, 5, settings, _ranker(embedder))
assert _names(results) == [FX_TOOL.name, WEATHER_TOOL.name, CALENDAR_TOOL.name]
assert all(FX_TOOL.description not in text for call in embedder.calls for text in call)
@pytest.mark.asyncio
async def test_empty_query_returns_only_core_tools_without_embedding(self) -> None:
embedder = RecordingEmbedder()
settings = MCPToolSearchSettings(embedding_model="emb", core_tools=(CALENDAR_TOOL.name,))
assert _names(await search_mcp_tools("", CATALOG, 5, settings, _ranker(embedder))) == [CALENDAR_TOOL.name]
assert embedder.calls == []
@pytest.mark.asyncio
async def test_repeat_searches_only_embed_the_query(self) -> None:
embedder = RecordingEmbedder()
ranker = _ranker(embedder)
settings = MCPToolSearchSettings(embedding_model="emb")
await search_mcp_tools("FX", CATALOG, 5, settings, ranker)
await search_mcp_tools("FX", CATALOG, 5, settings, ranker)
assert [len(call) for call in embedder.calls] == [4, 1]
@pytest.mark.asyncio
async def test_embedding_failure_is_reported_not_raised(self) -> None:
async def failing(texts: Sequence[str]) -> Sequence[Vector]:
raise ValueError("embedding model is down")
ranker = SemanticToolRanker(embed=failing, embedding_model="emb", index=SemanticTextIndex())
result = await search_mcp_tools("FX", CATALOG, 5, MCPToolSearchSettings(embedding_model="emb"), ranker)
assert isinstance(result, EmbeddingFailed)
assert "embedding model is down" in result.reason
class TestMcpToolSearchSettings:
def test_rejects_out_of_range_values(self) -> None:
from pydantic import ValidationError
with pytest.raises(ValidationError):
MCPToolSearchSettings(top_k=0)
with pytest.raises(ValidationError):
MCPToolSearchSettings(similarity_threshold=1.5)
def test_yaml_shape_round_trips(self) -> None:
settings = MCPToolSearchSettings.model_validate(
{"embedding_model": "emb", "top_k": 3, "similarity_threshold": 0.2, "core_tools": ["a", "b"]}
)
assert settings.core_tools == ("a", "b")
assert settings.model_dump() == {
"embedding_model": "emb",
"top_k": 3,
"similarity_threshold": 0.2,
"core_tools": ("a", "b"),
}
class TestCoerceTopK:
def test_int_passthrough(self) -> None:
assert coerce_top_k(3) == 3
@ -92,10 +249,10 @@ class TestSearchTools:
assert len(results) <= 2
def test_empty_query_returns_empty(self) -> None:
assert search_tools("", SAMPLE_TOOLS) == []
assert search_tools("", SAMPLE_TOOLS) == ()
def test_no_match_returns_empty(self) -> None:
assert search_tools("xyzzy_nonexistent_zzz", SAMPLE_TOOLS) == []
assert search_tools("xyzzy_nonexistent_zzz", SAMPLE_TOOLS) == ()
def test_matches_description_not_just_name(self) -> None:
results = search_tools("channel", SAMPLE_TOOLS)
@ -603,6 +760,63 @@ class TestCallToolRestApiVirtualTools:
assert result.isError is True
assert result.content[0].text == "set agent_search_embedding_model"
def _semantic_request(self, query: str = "FX") -> MagicMock:
return self._make_request({"name": MCP_TOOL_SEARCH_TOOL_NAME, "arguments": {"query": query}})
@pytest.mark.asyncio
async def test_mcp_tool_search_ranks_the_callers_catalog_with_the_configured_embedding_model(
self, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(litellm, "mcp_tool_search", {"embedding_model": "emb", "similarity_threshold": 0.5})
user_api_key_dict = UserAPIKeyAuth(
api_key="k", team_id="team-1", object_permission=_make_perm(mcp_tool_search_enabled=True)
)
async def fake_aembedding(model: str, input: list[str], metadata: dict[str, Any]) -> MagicMock:
assert model == "emb"
assert metadata["user_api_key"] == "k"
assert metadata["user_api_key_team_id"] == "team-1"
response = MagicMock()
response.model_dump.return_value = {"data": [{"embedding": list(FAKE_VECTORS[t])} for t in input]}
return response
router = MagicMock()
router.aembedding = AsyncMock(side_effect=fake_aembedding)
with (
patch( # test-quality-ok: the proxy's router is a module global; the handler reaches it the way production does
"litellm.proxy.proxy_server.llm_router", router
),
patch( # test-quality-ok: the authorized catalog is the seam every virtual tool shares; the ranking under test stays real
"litellm.proxy._experimental.mcp_server.server._list_mcp_tools",
new_callable=AsyncMock,
return_value=AggregateToolListing(tools=list(CATALOG), outcomes={}),
) as mock_list,
):
result = await self._get_call_fn()(request=self._semantic_request(), user_api_key_dict=user_api_key_dict)
assert mock_list.await_args.kwargs["user_api_key_auth"] is user_api_key_dict
assert result.isError is False
assert [t["name"] for t in json.loads(result.content[0].text)] == [FX_TOOL.name]
@pytest.mark.asyncio
async def test_mcp_tool_search_reports_missing_router_as_tool_error(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "mcp_tool_search", {"embedding_model": "emb"})
user_api_key_dict = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True))
with patch( # test-quality-ok: the proxy's router is a module global; the handler reaches it the way production does
"litellm.proxy.proxy_server.llm_router", None
):
result = await self._get_call_fn()(request=self._semantic_request(), user_api_key_dict=user_api_key_dict)
assert result.isError is True
assert "mcp_tool_search.embedding_model" in result.content[0].text
@pytest.mark.asyncio
async def test_mcp_tool_search_reports_invalid_settings_as_tool_error(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "mcp_tool_search", {"top_k": 0})
user_api_key_dict = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True))
result = await self._get_call_fn()(request=self._semantic_request(), user_api_key_dict=user_api_key_dict)
assert result.isError is True
assert "top_k" in result.content[0].text
@pytest.mark.asyncio
async def test_agent_search_requires_flag_enabled(self) -> None:
from fastapi import HTTPException

View file

@ -16,13 +16,12 @@ from litellm.proxy.agent_endpoints.agent_search import (
AgentSearchHits,
AgentSearchIndex,
AgentSearchNotConfigured,
Vector,
agent_search_text,
cosine_similarity,
search_agents,
)
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import RestrictedAgentAccess
from litellm.proxy.agent_endpoints.endpoints import router, user_api_key_auth
from litellm.proxy.common_utils.semantic_text_index import Vector, cosine_similarity
from litellm.types.agents import AgentResponse
CALLER: Final = UserAPIKeyAuth(api_key="hashed-caller-key", team_id="team-1", user_id="user-1")

View file

@ -3006,6 +3006,78 @@ def test_update_mcp_semantic_filter_settings_requires_proxy_admin(monkeypatch):
app.dependency_overrides.pop(user_api_key_auth, None)
class TestMcpToolSearchSettingsEndpoints:
"""`litellm_settings.mcp_tool_search` drives the native `mcp_tool_search` virtual tool, so the UI must round-trip it."""
@staticmethod
def _override_auth(role: LitellmUserRoles):
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="u", api_key="hashed", user_role=role
)
def test_get_returns_stored_values_and_field_schema(self, mock_proxy_config, mock_auth, monkeypatch):
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object())
mock_proxy_config["config"]["litellm_settings"]["mcp_tool_search"] = {
"embedding_model": "text-embedding-3-small",
"core_tools": ["treasury-get_rates"],
}
resp = client.get("/get/mcp_tool_search_settings")
assert resp.status_code == 200, resp.text
assert resp.json()["values"] == {
"embedding_model": "text-embedding-3-small",
"top_k": 5,
"similarity_threshold": 0.0,
"core_tools": ["treasury-get_rates"],
}
assert resp.json()["field_schema"]["properties"]["core_tools"]["type"] == "array"
def test_update_requires_proxy_admin(self, monkeypatch):
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
self._override_auth(LitellmUserRoles.INTERNAL_USER)
try:
resp = client.patch("/update/mcp_tool_search_settings", json={"top_k": 3})
finally:
app.dependency_overrides.clear()
assert resp.status_code == 403
def test_update_persists_and_applies_in_memory(self, mock_proxy_config, monkeypatch):
import litellm
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.setattr(litellm, "mcp_tool_search", None)
self._override_auth(LitellmUserRoles.PROXY_ADMIN)
payload = {
"embedding_model": "text-embedding-3-small",
"top_k": 3,
"similarity_threshold": 0.25,
"core_tools": ["treasury-get_rates"],
}
try:
resp = client.patch("/update/mcp_tool_search_settings", json=payload)
finally:
app.dependency_overrides.clear()
assert resp.status_code == 200, resp.text
assert mock_proxy_config["save_call_count"]() == 1
assert litellm.mcp_tool_search == payload
assert mock_proxy_config["config"]["litellm_settings"]["mcp_tool_search"] == payload
def test_update_rejects_out_of_range_top_k(self, mock_proxy_config, monkeypatch):
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
self._override_auth(LitellmUserRoles.PROXY_ADMIN)
try:
resp = client.patch("/update/mcp_tool_search_settings", json={"top_k": 0})
finally:
app.dependency_overrides.clear()
assert resp.status_code == 422
assert mock_proxy_config["save_call_count"]() == 0
def test_upload_logo_requires_proxy_admin(monkeypatch):
"""Any authenticated key could previously write a file to the server's disk here."""
from litellm.proxy._types import UserAPIKeyAuth

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 22364
"limit": 22358
},
"LIT002": {
"limit": 26777
"limit": 26774
},
"LIT003": {
"limit": 269

View file

@ -0,0 +1,47 @@
import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query";
import { apiClient } from "@/components/networking";
import type { components } from "@/lib/http/schema";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { createQueryKeys } from "../common/queryKeysFactory";
export type MCPToolSearchSettings = components["schemas"]["MCPToolSearchSettings"];
export type MCPToolSearchSettingsResponse = components["schemas"]["MCPToolSearchSettingsResponse"];
const GET_PATH = "/get/mcp_tool_search_settings";
const UPDATE_PATH = "/update/mcp_tool_search_settings";
const mcpToolSearchSettingsKeys = createQueryKeys("mcpToolSearchSettings");
export const getMCPToolSearchSettings = (accessToken: string): Promise<MCPToolSearchSettingsResponse> =>
apiClient.get<MCPToolSearchSettingsResponse>(GET_PATH, { accessToken });
export const updateMCPToolSearchSettings = (
accessToken: string,
settings: MCPToolSearchSettings,
): Promise<MCPToolSearchSettings> =>
apiClient.patch<MCPToolSearchSettings>(UPDATE_PATH, { accessToken, body: settings });
export const useMCPToolSearchSettings = () => {
const { accessToken } = useAuthorized();
return useQuery<MCPToolSearchSettingsResponse>({
queryKey: mcpToolSearchSettingsKeys.list({}),
queryFn: () => getMCPToolSearchSettings(accessToken),
enabled: !!accessToken,
});
};
export const useUpdateMCPToolSearchSettings = () => {
const { accessToken } = useAuthorized();
const queryClient = useQueryClient();
return useMutation<MCPToolSearchSettings, Error, MCPToolSearchSettings>({
mutationFn: (settings) => {
if (!accessToken) {
throw new Error("Access token is required");
}
return updateMCPToolSearchSettings(accessToken, settings);
},
onSuccess: () => {
queryClient.invalidateQueries({ queryKey: mcpToolSearchSettingsKeys.all });
},
});
};

View file

@ -36,6 +36,7 @@ import type {
Team,
} from "@/components/mcp_tools/types";
import MCPSemanticFilterSettings from "@/components/Settings/AdminSettings/MCPSemanticFilterSettings/MCPSemanticFilterSettings";
import MCPToolSearchSettings from "@/components/Settings/AdminSettings/MCPToolSearchSettings/MCPToolSearchSettings";
import MCPNetworkSettings from "./MCPNetworkSettings";
import MCPDiscovery from "./mcp_discovery";
import { ByokCredentialModal } from "@/components/mcp_tools/ByokCredentialModal";
@ -544,6 +545,11 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
Semantic Filter
</TabsTrigger>
)}
{isAdminRole(userRole) && (
<TabsTrigger value="tool-search" className="flex-none rounded-none px-4 py-2">
Tool Search
</TabsTrigger>
)}
{isAdminRole(userRole) && (
<TabsTrigger value="network-settings" className="flex-none rounded-none px-4 py-2">
Network Settings
@ -726,6 +732,11 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
<MCPSemanticFilterSettings accessToken={accessToken} />
</TabsContent>
)}
{isAdminRole(userRole) && (
<TabsContent value="tool-search" keepMounted>
<MCPToolSearchSettings accessToken={accessToken} />
</TabsContent>
)}
{isAdminRole(userRole) && (
<TabsContent value="network-settings" keepMounted>
<MCPNetworkSettings accessToken={accessToken} />

View file

@ -0,0 +1,96 @@
import React from "react";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { render, screen, act, fireEvent } from "@testing-library/react";
import MCPToolSearchSettings from "./MCPToolSearchSettings";
import {
useMCPToolSearchSettings,
useUpdateMCPToolSearchSettings,
} from "@/app/(dashboard)/hooks/mcpToolSearchSettings/useMCPToolSearchSettings";
vi.mock("@/app/(dashboard)/hooks/mcpToolSearchSettings/useMCPToolSearchSettings", () => ({
useMCPToolSearchSettings: vi.fn(),
useUpdateMCPToolSearchSettings: vi.fn(),
}));
vi.mock("@/components/llm_calls/fetch_models", () => ({
fetchAvailableModels: vi.fn().mockResolvedValue([{ model_group: "text-embedding-3-small", mode: "embedding" }]),
}));
vi.mock("@/lib/toast", () => ({ toast: { success: vi.fn(), fromError: vi.fn() } }));
const mockMutate = vi.fn();
const EDITED_PAYLOAD = {
embedding_model: "text-embedding-3-small",
top_k: 8,
similarity_threshold: 0.25,
core_tools: ["treasury-get_rates", "weather-forecast"],
};
const STORED = {
field_schema: {},
values: {
embedding_model: "text-embedding-3-small",
top_k: 3,
similarity_threshold: 0.25,
core_tools: ["treasury-get_rates"],
},
};
type SettingsQuery = ReturnType<typeof useMCPToolSearchSettings>;
type SettingsMutation = ReturnType<typeof useUpdateMCPToolSearchSettings>;
const settled = (data: typeof STORED | undefined, overrides: Partial<SettingsQuery> = {}) =>
({ data, isLoading: false, isError: false, error: null, ...overrides }) as SettingsQuery;
async function renderSettings(accessToken: string | null = "token") {
const result = render(<MCPToolSearchSettings accessToken={accessToken} />);
await act(async () => {});
return result;
}
describe("MCPToolSearchSettings", () => {
beforeEach(() => {
vi.clearAllMocks();
vi.mocked(useMCPToolSearchSettings).mockReturnValue(settled(STORED));
vi.mocked(useUpdateMCPToolSearchSettings).mockReturnValue({
mutate: mockMutate,
isPending: false,
} as unknown as SettingsMutation);
});
it("shows the stored settings and keeps Save disabled until something changes", async () => {
await renderSettings();
expect(screen.getByLabelText(/top k results/i)).toHaveValue(3);
expect(screen.getByLabelText(/always returned first/i)).toHaveValue("treasury-get_rates");
expect(screen.getByRole("slider", { hidden: true })).toHaveAttribute("aria-valuenow", "0.25");
expect(screen.getByRole("button", { name: /save settings/i })).toBeDisabled();
});
it("sends the edited settings as the proxy's PATCH payload", async () => {
await renderSettings();
fireEvent.change(screen.getByLabelText(/top k results/i), { target: { value: "8" } });
fireEvent.change(screen.getByLabelText(/always returned first/i), {
target: { value: "treasury-get_rates\nweather-forecast" },
});
await act(async () => {
fireEvent.click(screen.getByRole("button", { name: /save settings/i }));
});
expect(mockMutate).toHaveBeenCalledTimes(1);
expect(mockMutate.mock.calls[0][0]).toEqual(EDITED_PAYLOAD);
});
it("asks the user to log in without a token and surfaces load errors", async () => {
await renderSettings(null);
expect(screen.getByText(/please log in/i)).toBeInTheDocument();
vi.mocked(useMCPToolSearchSettings).mockReturnValue(
settled(undefined, { isError: true, error: new Error("Database not connected") }),
);
await renderSettings();
expect(screen.getByText("Database not connected")).toBeInTheDocument();
});
});

View file

@ -0,0 +1,251 @@
"use client";
import {
useMCPToolSearchSettings,
useUpdateMCPToolSearchSettings,
} from "@/app/(dashboard)/hooks/mcpToolSearchSettings/useMCPToolSearchSettings";
import { toast } from "@/lib/toast";
import { Skeleton } from "@/components/ui/skeleton";
import { CircleHelp, Info, Save } from "lucide-react";
import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert";
import { useEffect, useState } from "react";
import { useForm } from "react-hook-form";
import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models";
import { FieldGroup } from "@/components/ui/field";
import { FormField } from "@/components/shared/form/FormField";
import { SearchSelect } from "@/components/shared/SearchSelect";
import { Button } from "@/components/ui/button";
import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
import { Input } from "@/components/ui/input";
import { Slider } from "@/components/ui/slider";
import { Textarea } from "@/components/ui/textarea";
import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import {
DEFAULT_FORM_VALUES,
TOP_K_MAX,
TOP_K_MIN,
clampTopK,
formToPayload,
storedValuesToForm,
ToolSearchFormValues,
} from "./toolSearchForm";
interface MCPToolSearchSettingsProps {
accessToken: string | null;
}
const SIMILARITY_THRESHOLD_MARKS = [0, 0.3, 0.5, 0.7, 1];
const labelWithHint = (label: string, hint: string): React.ReactNode => (
<>
{label}
<Tooltip>
<TooltipTrigger render={<CircleHelp className="size-3.5 shrink-0 cursor-help text-muted-foreground" />} />
<TooltipContent>{hint}</TooltipContent>
</Tooltip>
</>
);
export default function MCPToolSearchSettings({ accessToken }: MCPToolSearchSettingsProps) {
const { data, isLoading, isError, error } = useMCPToolSearchSettings();
const { mutate: updateSettings, isPending: isUpdating } = useUpdateMCPToolSearchSettings();
const form = useForm<ToolSearchFormValues>({ defaultValues: DEFAULT_FORM_VALUES });
const isDirty = form.formState.isDirty;
const [embeddingModels, setEmbeddingModels] = useState<ModelGroup[]>([]);
const [loadingModels, setLoadingModels] = useState(true);
const storedValues = data?.values;
useEffect(() => {
if (!accessToken) return;
fetchAvailableModels(accessToken)
.then((models) => setEmbeddingModels(models.filter((model) => model.mode === "embedding")))
.catch((fetchError: unknown) => console.error("Error fetching embedding models:", fetchError))
.finally(() => setLoadingModels(false));
}, [accessToken]);
useEffect(() => {
if (!storedValues) return;
form.reset(storedValuesToForm(storedValues));
}, [storedValues, form]);
const handleSave = (formValues: ToolSearchFormValues) => {
updateSettings(formToPayload(formValues), {
onSuccess: () => {
form.reset(formValues);
toast.success("Settings updated successfully. Changes will be applied across all pods within 10 seconds.");
},
onError: (saveError) => toast.fromError(saveError),
});
};
if (!accessToken) {
return <div className="p-6 text-center text-muted-foreground">Please log in to configure tool search.</div>;
}
if (isLoading) {
return (
<div className="flex flex-col gap-3">
<Skeleton className="h-4 w-2/5" />
<Skeleton className="h-4 w-full" />
<Skeleton className="h-4 w-3/5" />
</div>
);
}
if (isError) {
return (
<Alert variant="error" className="mb-6">
<AlertTitle>Could not load MCP tool search settings</AlertTitle>
{error instanceof Error && <AlertDescription>{error.message}</AlertDescription>}
</Alert>
);
}
return (
<div className="w-full">
<Alert variant="info" className="mb-6">
<Info />
<AlertTitle>Native MCP Tool Search</AlertTitle>
<AlertDescription>
Controls the <code>mcp_tool_search</code> virtual tool that native MCP clients call to discover tools. With an
embedding model set, tools are ranked by the meaning of their name and description, so a query like
&quot;FX&quot; finds a &quot;foreign exchange rates&quot; tool. Without one, keyword matching is used. Callers
only ever see tools their key, team and server permissions already allow.
</AlertDescription>
</Alert>
<TooltipProvider>
<form onSubmit={(event) => event.preventDefault()} noValidate>
<Card className="mb-4">
<CardHeader className="border-b">
<CardTitle>Ranking</CardTitle>
</CardHeader>
<CardContent>
<FieldGroup>
<FormField
control={form.control}
name="embedding_model"
label={labelWithHint(
"Embedding Model",
"Embedding model from your model list used to rank tools by meaning. Clear it to fall back to keyword matching.",
)}
>
{({ value, onChange, id }) => (
<SearchSelect
inputId={id}
options={embeddingModels.map((model) => ({ label: model.model_group, value: model.model_group }))}
value={value}
onValueChange={onChange}
allowClear
placeholder={loadingModels ? "Loading models..." : "Keyword matching (no embedding model)"}
emptyText={loadingModels ? "Loading..." : "No embedding models available"}
disabled={isUpdating || loadingModels}
/>
)}
</FormField>
<FormField
control={form.control}
name="top_k"
label={labelWithHint(
"Top K Results",
"Most ranked tools a search returns. A smaller top_k in the tool call wins. Core tools do not count.",
)}
>
{({ ref, value, onChange, onBlur, id }) => (
<Input
id={id}
ref={ref}
type="number"
min={TOP_K_MIN}
max={TOP_K_MAX}
value={value}
onChange={(event) => onChange(event.target.valueAsNumber)}
onBlur={() => {
onChange(Number.isNaN(value) ? DEFAULT_FORM_VALUES.top_k : clampTopK(value));
onBlur();
}}
disabled={isUpdating}
/>
)}
</FormField>
<FormField
control={form.control}
name="similarity_threshold"
label={labelWithHint(
"Similarity Threshold",
"Lowest cosine similarity a tool needs to appear in semantic results. 0 means no cutoff.",
)}
>
{({ value, onChange, id }) => (
<div className="w-full">
<Slider
id={id}
min={0}
max={1}
step={0.05}
value={[value]}
onValueChange={(next) => onChange(Array.isArray(next) ? next[0] : next)}
disabled={isUpdating}
/>
<div className="relative mt-2 h-4 text-xs text-muted-foreground">
{SIMILARITY_THRESHOLD_MARKS.map((mark) => (
<span key={mark} className="absolute -translate-x-1/2" style={{ left: `${mark * 100}%` }}>
{mark.toFixed(1)}
</span>
))}
</div>
</div>
)}
</FormField>
</FieldGroup>
</CardContent>
</Card>
<Card className="mb-4">
<CardHeader className="border-b">
<CardTitle>Core Tools</CardTitle>
</CardHeader>
<CardContent>
<FieldGroup>
<FormField
control={form.control}
name="core_tools_text"
label={labelWithHint(
"Always Returned First",
"One tool name per line, e.g. my_server-get_rates. Listed before ranked results whenever the caller is allowed to use them.",
)}
>
{({ ref, value, onChange, onBlur, id }) => (
<Textarea
id={id}
ref={ref}
value={value}
placeholder={"my_server-get_rates\nmy_server-list_accounts"}
onChange={(event) => onChange(event.target.value)}
onBlur={onBlur}
disabled={isUpdating}
/>
)}
</FormField>
</FieldGroup>
</CardContent>
</Card>
<div className="flex justify-end gap-2">
<Button
type="button"
onClick={() => void form.handleSubmit(handleSave)()}
disabled={!isDirty || isUpdating}
>
{isUpdating ? <UiLoadingSpinner className="size-4" /> : <Save />}
Save Settings
</Button>
</div>
</form>
</TooltipProvider>
</div>
);
}

View file

@ -0,0 +1,60 @@
import { describe, expect, it } from "vitest";
import { DEFAULT_FORM_VALUES, formToPayload, parseCoreTools, storedValuesToForm } from "./toolSearchForm";
const STORED_SEMANTIC = {
embedding_model: "text-embedding-3-small",
top_k: 3,
similarity_threshold: 0.25,
core_tools: ["treasury-get_rates", "treasury-list_accounts"],
};
const SEMANTIC_FORM = {
embedding_model: "text-embedding-3-small",
top_k: 3,
similarity_threshold: 0.25,
core_tools_text: "treasury-get_rates\ntreasury-list_accounts",
};
const KEYWORD_PAYLOAD = { embedding_model: null, top_k: 5, similarity_threshold: 0, core_tools: [] };
const OVERSIZED_FORM = {
embedding_model: "emb",
top_k: 400,
similarity_threshold: 0.5,
core_tools_text: "treasury-get_rates",
};
const CLAMPED_PAYLOAD = {
embedding_model: "emb",
top_k: 100,
similarity_threshold: 0.5,
core_tools: ["treasury-get_rates"],
};
describe("storedValuesToForm", () => {
it("maps stored settings onto the form, joining core tools one per line", () => {
expect(storedValuesToForm(STORED_SEMANTIC)).toEqual(SEMANTIC_FORM);
});
it("falls back to keyword defaults when nothing usable is stored", () => {
expect(storedValuesToForm({})).toEqual(DEFAULT_FORM_VALUES);
expect(storedValuesToForm({ embedding_model: null, top_k: "7" })).toEqual(DEFAULT_FORM_VALUES);
});
});
describe("parseCoreTools", () => {
it("splits on newlines and commas, trims, drops blanks and duplicates in order", () => {
expect(parseCoreTools(" a-x \n\nb-y, a-x ,c-z\n")).toEqual(["a-x", "b-y", "c-z"]);
});
});
describe("formToPayload", () => {
it("sends null for a cleared embedding model so the proxy returns to keyword matching", () => {
expect(formToPayload({ ...DEFAULT_FORM_VALUES, embedding_model: " " })).toEqual(KEYWORD_PAYLOAD);
});
it("clamps top_k into the range the proxy accepts and lists core tools", () => {
expect(formToPayload(OVERSIZED_FORM)).toEqual(CLAMPED_PAYLOAD);
expect(formToPayload({ ...DEFAULT_FORM_VALUES, top_k: 0 }).top_k).toBe(1);
});
});

View file

@ -0,0 +1,49 @@
import type { MCPToolSearchSettings } from "@/app/(dashboard)/hooks/mcpToolSearchSettings/useMCPToolSearchSettings";
export interface ToolSearchFormValues {
embedding_model: string;
top_k: number;
similarity_threshold: number;
core_tools_text: string;
}
export const TOP_K_MIN = 1;
export const TOP_K_MAX = 100;
export const DEFAULT_FORM_VALUES: ToolSearchFormValues = {
embedding_model: "",
top_k: 5,
similarity_threshold: 0,
core_tools_text: "",
};
const isString = (value: unknown): value is string => typeof value === "string";
const isNumber = (value: unknown): value is number => typeof value === "number" && Number.isFinite(value);
export const parseCoreTools = (text: string): string[] =>
Array.from(
new Set(
text
.split(/[\n,]/)
.map((name) => name.trim())
.filter((name) => name.length > 0),
),
);
export const clampTopK = (value: number): number => Math.min(TOP_K_MAX, Math.max(TOP_K_MIN, Math.round(value)));
export const storedValuesToForm = (values: Record<string, unknown>): ToolSearchFormValues => ({
embedding_model: isString(values.embedding_model) ? values.embedding_model : DEFAULT_FORM_VALUES.embedding_model,
top_k: isNumber(values.top_k) ? values.top_k : DEFAULT_FORM_VALUES.top_k,
similarity_threshold: isNumber(values.similarity_threshold)
? values.similarity_threshold
: DEFAULT_FORM_VALUES.similarity_threshold,
core_tools_text: Array.isArray(values.core_tools) ? values.core_tools.filter(isString).join("\n") : "",
});
export const formToPayload = (form: ToolSearchFormValues): MCPToolSearchSettings => ({
embedding_model: form.embedding_model.trim() === "" ? null : form.embedding_model.trim(),
top_k: clampTopK(form.top_k),
similarity_threshold: form.similarity_threshold,
core_tools: parseCoreTools(form.core_tools_text),
});

View file

@ -5110,6 +5110,26 @@ export interface paths {
patch?: never;
trace?: never;
};
"/get/mcp_tool_search_settings": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
/**
* Get Mcp Tool Search Settings
* @description Get the `litellm_settings.mcp_tool_search` configuration used by the native `mcp_tool_search` virtual tool.
*/
get: operations["get_mcp_tool_search_settings_get_mcp_tool_search_settings_get"];
put?: never;
post?: never;
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/get/sso_settings": {
parameters: {
query?: never;
@ -16126,6 +16146,27 @@ export interface paths {
patch: operations["update_mcp_semantic_filter_settings_update_mcp_semantic_filter_settings_patch"];
trace?: never;
};
"/update/mcp_tool_search_settings": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
get?: never;
put?: never;
post?: never;
delete?: never;
options?: never;
head?: never;
/**
* Update Mcp Tool Search Settings
* @description Update `litellm_settings.mcp_tool_search` in the database.
* Settings will be picked up by all pods within approximately 10 seconds via background polling.
*/
patch: operations["update_mcp_tool_search_settings_update_mcp_tool_search_settings_patch"];
trace?: never;
};
"/update/sso_settings": {
parameters: {
query?: never;
@ -31113,6 +31154,49 @@ export interface components {
/** Total */
total: number;
};
/**
* MCPToolSearchSettings
* @description `litellm_settings.mcp_tool_search`: how the native `mcp_tool_search` virtual tool ranks the caller's tools.
*/
MCPToolSearchSettings: {
/**
* Core Tools
* @description Tool names always returned first when the caller can access them, e.g. `my_server-get_rates`.
* @default []
*/
core_tools: string[];
/**
* Embedding Model
* @description Embedding model from model_list used to rank tools by meaning. Unset keeps keyword matching.
*/
embedding_model?: string | null;
/**
* Similarity Threshold
* @description Lowest cosine similarity a tool needs to appear in semantic results (0.0 = no cutoff).
* @default 0
*/
similarity_threshold: number;
/**
* Top K
* @description Most ranked tools a search returns. A smaller top_k in the tool call wins. Core tools do not count.
* @default 5
*/
top_k: number;
};
/**
* MCPToolSearchSettingsResponse
* @description Response model for native MCP tool search settings
*/
MCPToolSearchSettingsResponse: {
/** Field Schema */
field_schema: {
[key: string]: unknown;
};
/** Values */
values: {
[key: string]: unknown;
};
};
/** MCPToolsetTool */
MCPToolsetTool: {
/** Server Id */
@ -46645,6 +46729,26 @@ export interface operations {
};
};
};
get_mcp_tool_search_settings_get_mcp_tool_search_settings_get: {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["MCPToolSearchSettingsResponse"];
};
};
};
};
get_sso_settings_get_sso_settings_get: {
parameters: {
query?: never;
@ -58976,6 +59080,41 @@ export interface operations {
};
};
};
update_mcp_tool_search_settings_update_mcp_tool_search_settings_patch: {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
requestBody: {
content: {
"application/json": components["schemas"]["MCPToolSearchSettings"];
};
};
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": {
[key: string]: unknown;
};
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
update_sso_settings_update_sso_settings_patch: {
parameters: {
query?: never;