diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index f23dd75a753..6b00650035d 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -452,20 +452,33 @@ if MCP_AVAILABLE: dependencies=[Depends(user_api_key_auth)], ) async def get_mcp_access_groups( + team_id: Optional[str] = Query( + None, + description="When provided, return only access groups the team is permitted to use.", + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ - Get MCP access groups available to the user. Non-admins only see groups - they have access to via their teams; admins see all groups. + Get MCP access groups available to the user. When team_id is provided, returns only + the access groups the team is configured to use. Non-admins without a team_id only + see groups accessible via their own teams. """ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) from litellm.proxy.management_helpers.object_permission_utils import ( get_allowed_mcp_access_groups_for_user, + get_team_mcp_permissions, ) from litellm.proxy.proxy_server import prisma_client + # When a specific team is requested, return only that team's configured groups + if team_id is not None: + team_perms = await get_team_mcp_permissions(team_id, prisma_client) + if team_perms is not None: + return {"access_groups": sorted(team_perms["mcp_access_groups"])} + # Team has no restrictions — fall through to return all groups below + access_groups = set() # Get from config-loaded servers @@ -569,6 +582,10 @@ if MCP_AVAILABLE: response_model=List[LiteLLM_MCPServerTable], ) async def fetch_all_mcp_servers( + team_id: Optional[str] = Query( + None, + description="When provided, return only MCP servers the team is permitted to use.", + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -578,6 +595,29 @@ if MCP_AVAILABLE: --header 'Authorization: Bearer your_api_key_here' ``` """ + from litellm.proxy.proxy_server import prisma_client + + # When a specific team is requested, filter to that team's allowed servers + if team_id is not None: + from litellm.proxy.management_helpers.object_permission_utils import ( + get_team_mcp_permissions, + ) + + team_perms = await get_team_mcp_permissions(team_id, prisma_client) + if team_perms is not None: + # Team has explicit restrictions — return only its allowed servers + allowed_ids = set(team_perms["mcp_servers"]) + allow_all_ids = set(global_mcp_server_manager.get_allow_all_keys_server_ids()) + all_servers = await global_mcp_server_manager.get_all_mcp_servers_unfiltered() + filtered = [ + s for s in all_servers + if s.server_id in allowed_ids or s.server_id in allow_all_ids + ] + return _redact_mcp_credentials_list(filtered) + else: + # Team has no MCP restrictions — return all servers + all_servers = await global_mcp_server_manager.get_all_mcp_servers_unfiltered() + return _redact_mcp_credentials_list(all_servers) user_mcp_management_mode = _get_user_mcp_management_mode() is_restricted_virtual_key = _is_restricted_virtual_key_request( diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index d49f62183da..8cca5e220e0 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -240,6 +240,71 @@ async def get_allowed_mcp_access_groups_for_user( return allowed_groups +async def get_team_mcp_permissions( + team_id: str, + prisma_client: Optional[PrismaClient], +) -> Optional[Dict[str, Any]]: + """ + Return the MCP servers and access groups that a team is explicitly permitted to use. + + Returns a dict: + { + "mcp_servers": List[str], # expanded server IDs (direct + access-group-resolved + tool perm keys) + "mcp_access_groups": List[str], # team's raw access group names + } + + Returns None when the team has no object_permission (meaning no restrictions configured). + Callers treat None as "no restriction" (show/allow all servers). + """ + try: + from litellm.proxy.auth.auth_checks import get_team_object + from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache + except ImportError: + return None + + if prisma_client is None or user_api_key_cache is None: + return None + + try: + team_obj = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + except HTTPException: + return None + + if team_obj is None or team_obj.object_permission is None: + return None + + obj_perm = team_obj.object_permission + server_ids: List[str] = list(obj_perm.mcp_servers or []) + + if obj_perm.mcp_access_groups: + try: + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + + resolved = await MCPRequestHandler._get_mcp_servers_from_access_groups( + obj_perm.mcp_access_groups + ) + server_ids.extend(resolved) + except Exception as e: + verbose_proxy_logger.warning( + f"get_team_mcp_permissions: failed to resolve access groups: {e}" + ) + + if obj_perm.mcp_tool_permissions: + server_ids.extend(obj_perm.mcp_tool_permissions.keys()) + + return { + "mcp_servers": list(set(server_ids)), + "mcp_access_groups": list(obj_perm.mcp_access_groups or []), + } + + async def validate_key_mcp_servers_against_team( object_permission: Optional[Union[Dict, Any]], team_obj: Optional[LiteLLM_TeamTableCachedObj], @@ -247,11 +312,14 @@ async def validate_key_mcp_servers_against_team( """ Validate that a key's requested MCP servers/access groups are allowed by its team. - Mirrors the runtime intersection logic: only restricts when the team has restrictions - configured (non-empty). allow_all_keys servers always pass validation. + Security model: + - Team has explicit MCP config → key must be a subset of team's servers/groups + - Team has no MCP config (object_permission is None) → only allow_all_keys servers pass; + all other MCP server assignments are denied (deny-by-default) + - allow_all_keys servers always pass regardless of team config - Raises HTTPException(403) if the key requests MCP servers or access groups that - the team does not allow. + Raises HTTPException(403) if the key requests MCP servers or access groups the team + does not allow. """ if object_permission is None or team_obj is None: return @@ -267,16 +335,49 @@ async def validate_key_mcp_servers_against_team( if not key_mcp_servers and not key_mcp_access_groups: return + # Get allow_all_keys server IDs - these bypass all per-key restrictions + allow_all_server_ids: Set[str] = set() + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + allow_all_server_ids = set(global_mcp_server_manager.get_allow_all_keys_server_ids()) + except Exception as e: + verbose_proxy_logger.warning( + f"validate_key_mcp_servers_against_team: failed to get allow_all_keys servers: {e}" + ) + team_object_permission = team_obj.object_permission if team_object_permission is None: - # Team has no MCP config - no restriction + # Team has no MCP config: deny-by-default — only allow_all_keys servers are permitted + disallowed = [s for s in key_mcp_servers if s not in allow_all_server_ids] + if disallowed: + raise HTTPException( + status_code=403, + detail={ + "error": ( + f"MCP servers not allowed: {disallowed}. " + "The key's team has no MCP servers configured. " + "Ask an admin to add MCP servers to the team." + ) + }, + ) + if key_mcp_access_groups: + raise HTTPException( + status_code=403, + detail={ + "error": ( + f"MCP access groups not allowed: {key_mcp_access_groups}. " + "The key's team has no MCP access groups configured. " + "Ask an admin to add MCP access groups to the team." + ) + }, + ) return - # Build the team's allowed server set from direct servers + access group resolution + tool permissions - team_allowed_servers: List[str] = [] - - if team_object_permission.mcp_servers: - team_allowed_servers.extend(team_object_permission.mcp_servers) + # Team has MCP config — build its allowed server set + team_allowed_servers: List[str] = list(team_object_permission.mcp_servers or []) if team_object_permission.mcp_access_groups: try: @@ -296,21 +397,8 @@ async def validate_key_mcp_servers_against_team( if team_object_permission.mcp_tool_permissions: team_allowed_servers.extend(team_object_permission.mcp_tool_permissions.keys()) - # Get allow_all_keys server IDs - these bypass per-key restrictions - allow_all_server_ids: set = set() - try: - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) - - allow_all_server_ids = set(global_mcp_server_manager.get_allow_all_keys_server_ids()) - except Exception as e: - verbose_proxy_logger.warning( - f"validate_key_mcp_servers_against_team: failed to get allow_all_keys servers: {e}" - ) - - # Validate key's mcp_servers only when the team has server restrictions configured - if key_mcp_servers and team_allowed_servers: + # Validate mcp_servers + if key_mcp_servers: team_allowed_set = set(team_allowed_servers) disallowed = [ s for s in key_mcp_servers @@ -327,9 +415,9 @@ async def validate_key_mcp_servers_against_team( }, ) - # Validate key's mcp_access_groups only when the team has access group restrictions configured - if key_mcp_access_groups and team_object_permission.mcp_access_groups: - team_access_group_set = set(team_object_permission.mcp_access_groups) + # Validate mcp_access_groups + if key_mcp_access_groups: + team_access_group_set = set(team_object_permission.mcp_access_groups or []) disallowed_groups = [ g for g in key_mcp_access_groups if g not in team_access_group_set ] diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index a7e8c0cd7ef..46d63b9f99f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -6515,8 +6515,10 @@ async def test_mcp_validation_key_creation_allows_permitted_server(): @pytest.mark.asyncio -async def test_mcp_validation_no_restriction_when_team_has_no_mcp_config(): - """When the team has no object_permission, any MCP server is allowed.""" +async def test_mcp_validation_deny_by_default_when_team_has_no_mcp_config(): + """When team has no object_permission, non-allow_all_keys servers are denied (deny-by-default).""" + from fastapi import HTTPException + from litellm.proxy.management_helpers.object_permission_utils import ( validate_key_mcp_servers_against_team, ) @@ -6524,8 +6526,11 @@ async def test_mcp_validation_no_restriction_when_team_has_no_mcp_config(): team_obj = _make_team_obj_with_mcp_servers() # no object_permission object_permission = {"mcp_servers": ["any-server"]} - # Should not raise - await validate_key_mcp_servers_against_team(object_permission, team_obj) + with pytest.raises(HTTPException) as exc_info: + await validate_key_mcp_servers_against_team(object_permission, team_obj) + + assert exc_info.value.status_code == 403 + assert "team has no MCP servers configured" in str(exc_info.value.detail) @pytest.mark.asyncio diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPAccessGroups.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPAccessGroups.ts index 0e88b62b0f3..40d258cfe61 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPAccessGroups.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPAccessGroups.ts @@ -2,13 +2,14 @@ import { useQuery } from "@tanstack/react-query"; import { createQueryKeys } from "../common/queryKeysFactory"; import { fetchMCPAccessGroups } from "@/components/networking"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; + const mcpAccessGroupsKeys = createQueryKeys("mcpAccessGroups"); -export const useMCPAccessGroups = () => { +export const useMCPAccessGroups = (teamId?: string) => { const { accessToken } = useAuthorized(); return useQuery({ - queryKey: mcpAccessGroupsKeys.list({}), - queryFn: async () => await fetchMCPAccessGroups(accessToken!), + queryKey: mcpAccessGroupsKeys.list({ teamId }), + queryFn: async () => await fetchMCPAccessGroups(accessToken!, teamId), enabled: Boolean(accessToken), }); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServers.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServers.ts index 8746baae148..9aff24f13b4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServers.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServers.ts @@ -6,11 +6,11 @@ import useAuthorized from "../useAuthorized"; const mcpServersKeys = createQueryKeys("mcpServers"); -export const useMCPServers = () => { +export const useMCPServers = (teamId?: string) => { const { accessToken } = useAuthorized(); return useQuery({ - queryKey: mcpServersKeys.list({}), - queryFn: async () => await fetchMCPServers(accessToken!), + queryKey: mcpServersKeys.list({ teamId }), + queryFn: async () => await fetchMCPServers(accessToken!, teamId), enabled: !!accessToken, }); }; diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx index d94a80e502c..63666847b02 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx @@ -13,6 +13,7 @@ interface MCPServerSelectorProps { accessToken: string; placeholder?: string; disabled?: boolean; + teamId?: string; } const MCPServerSelector: React.FC = ({ @@ -22,9 +23,10 @@ const MCPServerSelector: React.FC = ({ accessToken, placeholder = "Select MCP servers", disabled = false, + teamId, }) => { - const { data: mcpServers = [], isLoading: serversLoading } = useMCPServers(); - const { data: accessGroups = [], isLoading: groupsLoading } = useMCPAccessGroups(); + const { data: mcpServers = [], isLoading: serversLoading } = useMCPServers(teamId); + const { data: accessGroups = [], isLoading: groupsLoading } = useMCPAccessGroups(teamId); const loading = serversLoading || groupsLoading; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 3d941cf638f..15be3ce8e60 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -6973,10 +6973,13 @@ export const fetchDiscoverableMCPServers = async (accessToken: string) => { } }; -export const fetchMCPServers = async (accessToken: string) => { +export const fetchMCPServers = async (accessToken: string, teamId?: string) => { try { // Construct base URL - const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/mcp/server` : `/v1/mcp/server`; + const baseUrl = proxyBaseUrl ? `${proxyBaseUrl}/v1/mcp/server` : `/v1/mcp/server`; + const url = teamId + ? `${baseUrl}?${new URLSearchParams({ team_id: teamId }).toString()}` + : baseUrl; console.log("Fetching MCP servers from:", url); @@ -7042,10 +7045,13 @@ export const fetchMCPServerHealth = async (accessToken: string, serverIds?: stri } }; -export const fetchMCPAccessGroups = async (accessToken: string) => { +export const fetchMCPAccessGroups = async (accessToken: string, teamId?: string) => { try { // Construct base URL - const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/mcp/access_groups` : `/v1/mcp/access_groups`; + const baseUrl = proxyBaseUrl ? `${proxyBaseUrl}/v1/mcp/access_groups` : `/v1/mcp/access_groups`; + const url = teamId + ? `${baseUrl}?${new URLSearchParams({ team_id: teamId }).toString()}` + : baseUrl; console.log("Fetching MCP access groups from:", url); diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx index 8a8a018f267..d1bfcf4efed 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx @@ -85,6 +85,7 @@ export function KeyEditView({ premiumUser = false, }: KeyEditViewProps) { const [form] = Form.useForm(); + const selectedTeamId = Form.useWatch("team_id", form) as string | undefined; const [promptsList, setPromptsList] = useState([]); const [tagsList, setTagsList] = useState>({}); const team = teams?.find((team) => team.team_id === keyData.team_id); @@ -567,6 +568,7 @@ export function KeyEditView({ value={form.getFieldValue("mcp_servers_and_groups")} accessToken={accessToken || ""} placeholder="Select MCP servers or access groups (optional)" + teamId={selectedTeamId} />