diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 2554e46b9b..80ae8217d7 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1058,6 +1058,103 @@ class MCPServerManager: ) ) + async def _call_regular_mcp_tool( + self, + mcp_server: MCPServer, + original_tool_name: str, + arguments: Dict[str, Any], + tasks: List, + mcp_auth_header: Optional[str], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + oauth2_headers: Optional[Dict[str, str]], + raw_headers: Optional[Dict[str, str]], + proxy_logging_obj: Optional[ProxyLogging], + ) -> CallToolResult: + """ + Call a regular MCP tool using the MCP client. + + Args: + mcp_server: The MCP server configuration + original_tool_name: The original tool name (without prefix) + arguments: Tool arguments + tasks: List of async tasks to append to (for during hooks) + mcp_auth_header: MCP auth header (deprecated) + mcp_server_auth_headers: Optional dict of server-specific auth headers + oauth2_headers: Optional OAuth2 headers + raw_headers: Optional raw headers from the request + proxy_logging_obj: Optional ProxyLogging object for hook integration + + Returns: + CallToolResult from the MCP server + + Raises: + BlockedPiiEntityError: If PII is blocked by guardrails + GuardrailRaisedException: If guardrails block the call + HTTPException: If an HTTP error occurs + """ + # Get server-specific auth header if available + server_auth_header: Optional[Union[Dict[str, str], str]] = None + if mcp_server_auth_headers and mcp_server.alias: + server_auth_header = mcp_server_auth_headers.get(mcp_server.alias) + elif mcp_server_auth_headers and mcp_server.server_name: + server_auth_header = mcp_server_auth_headers.get(mcp_server.server_name) + + # Fall back to deprecated mcp_auth_header if no server-specific header found + if server_auth_header is None: + server_auth_header = mcp_auth_header + + # oauth2 headers + extra_headers: Optional[Dict[str, str]] = None + if mcp_server.auth_type == MCPAuth.oauth2: + extra_headers = oauth2_headers + + if mcp_server.extra_headers and raw_headers: + if extra_headers is None: + extra_headers = {} + for header in mcp_server.extra_headers: + if header in raw_headers: + extra_headers[header] = raw_headers[header] + + client = self._create_mcp_client( + server=mcp_server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + ) + + call_tool_params = MCPCallToolRequestParams( + name=original_tool_name, + arguments=arguments, + ) + + async def _call_tool_via_client(client, params): + async with client: + return await client.call_tool(params) + + tasks.append( + asyncio.create_task(_call_tool_via_client(client, call_tool_params)) + ) + + # IMPORTANT: Must await tasks INSIDE the context manager to keep connection alive + try: + mcp_responses = await asyncio.gather(*tasks) + except ( + BlockedPiiEntityError, + GuardrailRaisedException, + HTTPException, + ) as e: + # Re-raise guardrail exceptions to properly fail the MCP call + verbose_logger.error( + f"Guardrail blocked MCP tool call during result check: {str(e)}" + ) + raise e + + # If proxy_logging_obj is None, the tool call result is at index 0 + # If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task) + result_index = 1 if proxy_logging_obj else 0 + result = mcp_responses[result_index] + + return cast(CallToolResult, result) + async def call_tool( self, name: str, @@ -1146,69 +1243,18 @@ class MCPServerManager: ) else: # For regular MCP servers, use the MCP client - # Get server-specific auth header if available - server_auth_header: Optional[Union[Dict[str, str], str]] = None - if mcp_server_auth_headers and mcp_server.alias: - server_auth_header = mcp_server_auth_headers.get(mcp_server.alias) - elif mcp_server_auth_headers and mcp_server.server_name: - server_auth_header = mcp_server_auth_headers.get(mcp_server.server_name) - - # Fall back to deprecated mcp_auth_header if no server-specific header found - if server_auth_header is None: - server_auth_header = mcp_auth_header - - # oauth2 headers - extra_headers: Optional[Dict[str, str]] = None - if mcp_server.auth_type == MCPAuth.oauth2: - extra_headers = oauth2_headers - - if mcp_server.extra_headers and raw_headers: - if extra_headers is None: - extra_headers = {} - for header in mcp_server.extra_headers: - if header in raw_headers: - extra_headers[header] = raw_headers[header] - - client = self._create_mcp_client( - server=mcp_server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - ) - - call_tool_params = MCPCallToolRequestParams( - name=original_tool_name, + return await self._call_regular_mcp_tool( + mcp_server=mcp_server, + original_tool_name=original_tool_name, arguments=arguments, + tasks=tasks, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + proxy_logging_obj=proxy_logging_obj, ) - async def _call_tool_via_client(client, params): - async with client: - return await client.call_tool(params) - - tasks.append( - asyncio.create_task(_call_tool_via_client(client, call_tool_params)) - ) - - # IMPORTANT: Must await tasks INSIDE the context manager to keep connection alive - try: - mcp_responses = await asyncio.gather(*tasks) - except ( - BlockedPiiEntityError, - GuardrailRaisedException, - HTTPException, - ) as e: - # Re-raise guardrail exceptions to properly fail the MCP call - verbose_logger.error( - f"Guardrail blocked MCP tool call during result check: {str(e)}" - ) - raise e - - # If proxy_logging_obj is None, the tool call result is at index 0 - # If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task) - result_index = 1 if proxy_logging_obj else 0 - result = mcp_responses[result_index] - - return cast(CallToolResult, result) - # For OpenAPI tools, await outside the client context try: mcp_responses = await asyncio.gather(*tasks) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 0d02007b8b..94e0fd3d44 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -627,6 +627,54 @@ async def _update_model_table( return _model_id +async def fetch_and_validate_organization( + organization_id: str, + existing_team_row: Any, + llm_router: Optional[Router], + prisma_client: Any, +) -> Any: + """ + Fetch and validate an organization for team update operations. + + Args: + organization_id: The organization ID to fetch + existing_team_row: The existing team row being updated + llm_router: The LLM router instance + prisma_client: The Prisma database client + + Returns: + The organization row from the database + + Raises: + HTTPException: If llm_router is None, organization not found, or validation fails + """ + if llm_router is None: + raise HTTPException( + status_code=500, detail={"error": CommonProxyErrors.no_llm_router.value} + ) + + organization_row = await prisma_client.db.litellm_organizationtable.find_unique( + where={"organization_id": organization_id}, + include={"litellm_budget_table": True, "users": True}, + ) + + if organization_row is None: + raise HTTPException( + status_code=404, + detail={ + "error": f"Organization not found, passed organization_id={organization_id}" + }, + ) + + validate_team_org_change( + team=LiteLLM_TeamTable(**existing_team_row.model_dump()), + organization=LiteLLM_OrganizationTable(**organization_row.model_dump()), + llm_router=llm_router, + ) + + return organization_row + + def validate_team_org_change( team: LiteLLM_TeamTable, organization: LiteLLM_OrganizationTable, llm_router: Router ) -> bool: @@ -817,25 +865,11 @@ async def update_team( if ( data.organization_id is not None and len(data.organization_id) > 0 ): # allow unsetting the organization_id - if llm_router is None: - raise HTTPException( - status_code=500, detail={"error": CommonProxyErrors.no_llm_router.value} - ) - organization_row = await prisma_client.db.litellm_organizationtable.find_unique( - where={"organization_id": data.organization_id}, - include={"litellm_budget_table": True, "users": True}, - ) - if organization_row is None: - raise HTTPException( - status_code=404, - detail={ - "error": f"Organization not found, passed organization_id={data.organization_id}" - }, - ) - validate_team_org_change( - team=LiteLLM_TeamTable(**existing_team_row.model_dump()), - organization=LiteLLM_OrganizationTable(**organization_row.model_dump()), + await fetch_and_validate_organization( + organization_id=data.organization_id, + existing_team_row=existing_team_row, llm_router=llm_router, + prisma_client=prisma_client, ) elif data.organization_id is not None and len(data.organization_id) == 0: # unsetting the organization_id