mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
a701effbad
commit
a76cb6feaf
19 changed files with 1328 additions and 164 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
142
litellm/proxy/common_utils/semantic_text_index.py
Normal file
142
litellm/proxy/common_utils/semantic_text_index.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22364
|
||||
"limit": 22358
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26777
|
||||
"limit": 26774
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 269
|
||||
|
|
|
|||
|
|
@ -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 });
|
||||
},
|
||||
});
|
||||
};
|
||||
|
|
@ -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} />
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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
|
||||
"FX" finds a "foreign exchange rates" 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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
@ -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),
|
||||
});
|
||||
139
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
139
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue