From a135bfc376fd51094bb0637bbd5d9a5b6a8e8ff3 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 12 Jan 2026 17:57:22 -0800 Subject: [PATCH 01/29] Add create key test --- .../e2e_tests/tests/keys/createKey.spec.ts | 23 +++++++++++++++++++ 1 file changed, 23 insertions(+) create mode 100644 ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts b/ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts new file mode 100644 index 00000000000..ac15d606256 --- /dev/null +++ b/ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts @@ -0,0 +1,23 @@ +import { test, expect } from "@playwright/test"; +import { ADMIN_STORAGE_PATH } from "../../constants"; +import { Page } from "../../fixtures/pages"; +import { navigateToPage } from "../../helpers/navigation"; + +test.describe.only("Create Key", () => { + test.use({ storageState: ADMIN_STORAGE_PATH }); + + test("Able to create a key with all team models", async ({ page }) => { + await navigateToPage(page, Page.ApiKeys); + await expect(page.getByRole("button", { name: "Next" })).toBeVisible(); + await page.getByRole("button", { name: "+ Create New Key" }).click(); + await page.getByTestId("base-input").click(); + await page.getByTestId("base-input").fill("e2eUITestingCreateKeyAllTeamModels"); + await page.locator(".ant-select-selection-overflow").click(); + await page.getByText("All Team Models").click(); + await page.getByRole("combobox", { name: "* Models info-circle :" }).press("Escape"); + await page.getByRole("button", { name: "Create Key" }).click(); + await expect(page.getByText("Virtual Key Created")).toBeVisible(); + await page.keyboard.press("Escape"); + await expect(page.getByText("e2eUITestingCreateKeyAllTeamModels")).toBeVisible(); + }); +}); From 81b1becd9588777f00127196eaa35f378fd0e91c Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 12 Jan 2026 18:06:41 -0800 Subject: [PATCH 02/29] remove only --- ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts b/ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts index ac15d606256..633673b3297 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts +++ b/ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts @@ -3,7 +3,7 @@ import { ADMIN_STORAGE_PATH } from "../../constants"; import { Page } from "../../fixtures/pages"; import { navigateToPage } from "../../helpers/navigation"; -test.describe.only("Create Key", () => { +test.describe("Create Key", () => { test.use({ storageState: ADMIN_STORAGE_PATH }); test("Able to create a key with all team models", async ({ page }) => { From 9875ef97201f0c3b906b232a5f2b68d502e5d91d Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 12 Jan 2026 18:34:06 -0800 Subject: [PATCH 03/29] Remove Notification Check --- ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts | 1 - 1 file changed, 1 deletion(-) diff --git a/ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts b/ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts index 633673b3297..4343063b305 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts +++ b/ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts @@ -16,7 +16,6 @@ test.describe("Create Key", () => { await page.getByText("All Team Models").click(); await page.getByRole("combobox", { name: "* Models info-circle :" }).press("Escape"); await page.getByRole("button", { name: "Create Key" }).click(); - await expect(page.getByText("Virtual Key Created")).toBeVisible(); await page.keyboard.press("Escape"); await expect(page.getByText("e2eUITestingCreateKeyAllTeamModels")).toBeVisible(); }); From aad92c0b25e517ebaf15c24b55a3486f3da41d1c Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Thu, 15 Jan 2026 15:50:43 -0800 Subject: [PATCH 04/29] Merge pull request #19116 from BerriAI/litellm_org_admin_escalte [Fix] /user/new Privilege Escalation --- .../test_internal_user_endpoints.py | 83 +++++++++++++++++++ 1 file changed, 83 insertions(+) diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 33f2a75fac6..397a6af556f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -12,6 +12,7 @@ sys.path.insert( from litellm.proxy._types import ( LiteLLM_UserTableFiltered, + LitellmUserRoles, NewUserRequest, ProxyException, UpdateUserRequest, @@ -306,6 +307,88 @@ async def test_new_user_license_over_limit(mocker): mock_license_check.is_over_limit.assert_called_once_with(total_users=1000) +@pytest.mark.asyncio +async def test_new_user_non_admin_cannot_create_admin(mocker): + """ + Test that non-admin users cannot create administrative users (PROXY_ADMIN or PROXY_ADMIN_VIEW_ONLY). + This prevents privilege escalation vulnerabilities. + """ + from litellm.proxy.management_endpoints.internal_user_endpoints import new_user + + # Mock the prisma client + mock_prisma_client = mocker.MagicMock() + + # Setup the mock count response (under license limit) + async def mock_count(*args, **kwargs): + return 5 # Low user count, under limit + + mock_prisma_client.db.litellm_usertable.count = mock_count + + # Mock duplicate checks to pass + async def mock_check_duplicate_user_email(*args, **kwargs): + return None # No duplicate found + + async def mock_check_duplicate_user_id(*args, **kwargs): + return None # No duplicate found + + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_email", + mock_check_duplicate_user_email, + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_id", + mock_check_duplicate_user_id, + ) + + # Mock the license check to return False (under limit) + mock_license_check = mocker.MagicMock() + mock_license_check.is_over_limit.return_value = False + + # Patch the imports in the endpoint + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch("litellm.proxy.proxy_server._license_check", mock_license_check) + + # Test Case 1: INTERNAL_USER trying to create PROXY_ADMIN + user_request = NewUserRequest( + user_email="admin@example.com", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + # Mock user_api_key_dict with non-admin role + mock_user_api_key_dict = UserAPIKeyAuth( + user_id="test_internal_user", user_role=LitellmUserRoles.INTERNAL_USER + ) + + # Call new_user function and expect ProxyException + with pytest.raises(ProxyException) as exc_info: + await new_user(data=user_request, user_api_key_dict=mock_user_api_key_dict) + + # Verify the exception details + assert exc_info.value.code == 403 or exc_info.value.code == "403" + assert "Only proxy admins can create administrative users" in str(exc_info.value.message) + assert "proxy_admin" in str(exc_info.value.message) + assert "proxy_admin_viewer" in str(exc_info.value.message) + assert str(LitellmUserRoles.PROXY_ADMIN) in str(exc_info.value.message) + assert str(LitellmUserRoles.INTERNAL_USER) in str(exc_info.value.message) + + # Test Case 2: INTERNAL_USER trying to create PROXY_ADMIN_VIEW_ONLY + user_request_viewer = NewUserRequest( + user_email="admin_viewer@example.com", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + ) + + with pytest.raises(ProxyException) as exc_info2: + await new_user( + data=user_request_viewer, user_api_key_dict=mock_user_api_key_dict + ) + + # Verify the exception details + assert exc_info2.value.code == 403 or exc_info2.value.code == "403" + assert "Only proxy admins can create administrative users" in str( + exc_info2.value.message + ) + assert str(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) in str(exc_info2.value.message) + + @pytest.mark.asyncio async def test_user_info_url_encoding_plus_character(mocker): """ From 6267f1689b96725bc00f28386c85eccbb42d7fb6 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Wed, 21 Jan 2026 12:00:22 +0900 Subject: [PATCH 05/29] feat: save mcp fail log --- .../proxy/_experimental/mcp_server/server.py | 138 ++++++++++-------- 1 file changed, 74 insertions(+), 64 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 76a28344856..e32ae85f9ad 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -6,6 +6,7 @@ LiteLLM MCP Server Routes import asyncio import contextlib from datetime import datetime +import traceback from typing import Any, AsyncIterator, Dict, List, Optional, Tuple, Union, cast from fastapi import FastAPI, HTTPException @@ -13,6 +14,7 @@ from pydantic import AnyUrl, ConfigDict from starlette.types import Receive, Scope, Send from litellm._logging import verbose_logger +from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, @@ -25,7 +27,7 @@ from litellm.proxy._experimental.mcp_server.utils import ( from litellm.proxy._types import UserAPIKeyAuth from litellm.types.mcp import MCPAuth from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer -from litellm.types.utils import StandardLoggingMCPToolCall +from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall from litellm.utils import client # Check if MCP is available @@ -1320,33 +1322,6 @@ if MCP_AVAILABLE: content=cast(Any, local_content), isError=False ) - ######################################################### - # Post MCP Tool Call Hook - # Allow modifying the MCP tool call response before it is returned to the user - ######################################################### - if litellm_logging_obj: - litellm_logging_obj.post_call(original_response=response) - end_time = datetime.now() - await litellm_logging_obj.async_post_mcp_tool_call_hook( - kwargs=litellm_logging_obj.model_call_details, - response_obj=response, - start_time=start_time, - end_time=end_time, - ) - # Set call_type to call_mcp_tool so cost calculator recognizes it - from litellm.types.utils import CallTypes - - litellm_logging_obj.call_type = CallTypes.call_mcp_tool.value - # Trigger success logging to build standard_logging_object and call callbacks - # async_success_handler will: - # 1. Call _success_handler_helper_fn which recognizes call_mcp_tool - # 2. Call _process_hidden_params_and_response_cost which: - # - Calculates cost via _response_cost_calculator -> MCPCostCalculator - # - Builds standard_logging_object - # 3. Call async_log_success_event on all callbacks - await litellm_logging_obj.async_success_handler( - result=response, start_time=start_time, end_time=end_time - ) return response @client @@ -1365,49 +1340,84 @@ if MCP_AVAILABLE: Call a specific tool with the provided arguments (handles prefixed tool names). """ start_time = datetime.now() - if arguments is None: - raise HTTPException( - status_code=400, detail="Request arguments are required" + litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get( + "litellm_logging_obj", None + ) + + try: + if arguments is None: + raise HTTPException( + status_code=400, detail="Request arguments are required" + ) + + ## CHECK IF USER IS ALLOWED TO CALL THIS TOOL + allowed_mcp_server_ids = ( + await global_mcp_server_manager.get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + ) ) - ## CHECK IF USER IS ALLOWED TO CALL THIS TOOL - allowed_mcp_server_ids = ( - await global_mcp_server_manager.get_allowed_mcp_servers( + allowed_mcp_servers: List[MCPServer] = [] + for allowed_mcp_server_id in allowed_mcp_server_ids: + allowed_server = global_mcp_server_manager.get_mcp_server_by_id( + allowed_mcp_server_id + ) + if allowed_server is not None: + allowed_mcp_servers.append(allowed_server) + + allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=mcp_servers, + allowed_mcp_servers=allowed_mcp_servers, + ) + if not allowed_mcp_servers: + raise HTTPException( + status_code=403, + detail="User not allowed to call this tool.", + ) + + # Delegate to execute_mcp_tool for execution + response = await execute_mcp_tool( + name=name, + arguments=arguments, + allowed_mcp_servers=allowed_mcp_servers, + start_time=start_time, user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + **kwargs, ) - ) - - allowed_mcp_servers: List[MCPServer] = [] - for allowed_mcp_server_id in allowed_mcp_server_ids: - allowed_server = global_mcp_server_manager.get_mcp_server_by_id( - allowed_mcp_server_id + except Exception as e: + traceback_str = traceback.format_exc( + limit=MAXIMUM_TRACEBACK_LINES_TO_LOG ) - if allowed_server is not None: - allowed_mcp_servers.append(allowed_server) + from litellm.proxy.proxy_server import proxy_logging_obj + + if proxy_logging_obj and user_api_key_auth: + await proxy_logging_obj.post_call_failure_hook( + request_data=kwargs, + original_exception=e, + user_api_key_dict=user_api_key_auth, + route="/mcp/call_tool", + traceback_str=traceback_str, + ) + raise - allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( - mcp_servers=mcp_servers, - allowed_mcp_servers=allowed_mcp_servers, - ) - if not allowed_mcp_servers: - raise HTTPException( - status_code=403, - detail="User not allowed to call this tool.", + if litellm_logging_obj: + litellm_logging_obj.post_call(original_response=response) + end_time = datetime.now() + await litellm_logging_obj.async_post_mcp_tool_call_hook( + kwargs=litellm_logging_obj.model_call_details, + response_obj=response, + start_time=start_time, + end_time=end_time, ) - - # Delegate to execute_mcp_tool for execution - return await execute_mcp_tool( - name=name, - arguments=arguments, - allowed_mcp_servers=allowed_mcp_servers, - start_time=start_time, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - **kwargs, - ) + litellm_logging_obj.call_type = CallTypes.call_mcp_tool.value + await litellm_logging_obj.async_success_handler( + result=response, start_time=start_time, end_time=end_time + ) + return response async def mcp_get_prompt( name: str, From ae4d92ad509ebadbb3e72dc790f6621ca2e13de3 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Wed, 21 Jan 2026 13:54:24 +0900 Subject: [PATCH 06/29] feat: save mcp call log via responses --- litellm/responses/main.py | 3 + .../responses/mcp/chat_completions_handler.py | 2 + .../mcp/litellm_proxy_mcp_handler.py | 190 +++++++++++++++++- .../responses/mcp/mcp_streaming_iterator.py | 4 + 4 files changed, 196 insertions(+), 3 deletions(-) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 71d94287e82..78d358a1e39 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -301,6 +301,8 @@ async def aresponses_api_with_mcp( mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers_from_request, + litellm_call_id=kwargs.get("litellm_call_id"), + litellm_trace_id=kwargs.get("litellm_trace_id"), ) if tool_results: @@ -349,6 +351,7 @@ async def aresponses_api_with_mcp( tool_server_map=tool_server_map, base_iterator=final_response, mcp_events=tool_execution_events, + user_api_key_auth=user_api_key_auth, ) # Add custom output elements to the final response (for non-streaming) diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 6ce59e3e67f..26853b30596 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -142,6 +142,8 @@ async def acompletion_with_mcp( mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + litellm_call_id=kwargs.get("litellm_call_id"), + litellm_trace_id=kwargs.get("litellm_trace_id"), ) if not tool_results: diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 9cdcd3894e0..f5757e4d52e 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -1,3 +1,5 @@ +import traceback +from datetime import datetime from typing import ( TYPE_CHECKING, Any, @@ -11,14 +13,18 @@ from typing import ( ) from litellm._logging import verbose_logger +from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._experimental.mcp_server.utils import split_server_prefix_from_name from litellm.responses.main import aresponses from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator from litellm.types.llms.openai import ResponsesAPIResponse, ToolParam -from litellm.types.utils import Choices, ModelResponse +from litellm.types.utils import CallTypes, Choices, ModelResponse, StandardLoggingMCPToolCall +from litellm.utils import Rules, function_setup if TYPE_CHECKING: from mcp.types import Tool as MCPTool + from litellm.proxy.utils import ProxyLogging else: MCPTool = Any @@ -470,6 +476,8 @@ class LiteLLM_Proxy_MCP_Handler: mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, oauth2_headers: Optional[Dict[str, str]] = None, raw_headers: Optional[Dict[str, str]] = None, + litellm_call_id: Optional[str] = None, + litellm_trace_id: Optional[str] = None, ) -> List[Dict[str, Any]]: """Execute tool calls and return results.""" from fastapi import HTTPException @@ -478,10 +486,16 @@ class LiteLLM_Proxy_MCP_Handler: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) + from litellm.proxy.proxy_server import proxy_logging_obj + + from litellm._uuid import uuid tool_results = [] tool_call_id: Optional[str] = None + rules_obj = Rules() for tool_call in tool_calls: + logging_request_data: Dict[str, Any] = {} + tool_name: str = "" try: ( tool_name, @@ -514,6 +528,101 @@ class LiteLLM_Proxy_MCP_Handler: ): sanitized_tool_name = unprefixed_name + start_time = datetime.now() + logging_input = [ + { + "role": "tool", + "content": { + "tool_name": sanitized_tool_name, + "arguments": parsed_arguments, + }, + } + ] + tool_logging_call_id = litellm_call_id or str(uuid.uuid4()) + logging_request_data: Dict[str, Any] = { + "model": f"MCP: {tool_name}", + "metadata": { + "tool_call_id": tool_call_id, + "tool_name": sanitized_tool_name, + "server_name": server_name, + }, + "input": logging_input, + "call_type": CallTypes.call_mcp_tool.value, + "litellm_call_id": tool_logging_call_id, + } + if litellm_trace_id: + logging_request_data["litellm_trace_id"] = litellm_trace_id + user_identifier = None + if user_api_key_auth is not None: + user_api_key = getattr(user_api_key_auth, "api_key", None) + if user_api_key: + logging_request_data["metadata"]["user_api_key"] = user_api_key + + user_identifier = getattr(user_api_key_auth, "end_user_id", None) or getattr( + user_api_key_auth, "user_id", None + ) + if user_identifier: + logging_request_data["user"] = user_identifier + + litellm_logging_obj: Optional[LiteLLMLoggingObj] = None + try: + litellm_logging_obj, _ = function_setup( + original_function="call_mcp_tool", + rules_obj=rules_obj, + start_time=start_time, + **logging_request_data, + ) + except Exception as logging_error: + verbose_logger.debug( + "Failed to initialize logging for MCP tool call %s: %s", + tool_name, + logging_error, + ) + litellm_logging_obj = None + + logging_request_data["litellm_logging_obj"] = litellm_logging_obj + logging_request_data["arguments"] = parsed_arguments + + if litellm_logging_obj: + try: + litellm_logging_obj.pre_call( + input=logging_input, + api_key="", + ) + except Exception: + verbose_logger.exception( + "Failed to run pre_call for MCP tool logging" + ) + + standard_logging_mcp_tool_call: StandardLoggingMCPToolCall = { + "name": sanitized_tool_name, + "arguments": parsed_arguments, + "namespaced_tool_name": tool_name, + } + mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name( + tool_name + ) + if mcp_server: + mcp_info = mcp_server.mcp_info or {} + standard_logging_mcp_tool_call["mcp_server_name"] = ( + mcp_info.get("server_name") + or getattr(mcp_server, "server_name", None) + or server_name + ) + logo_url = mcp_info.get("logo_url") + if logo_url: + standard_logging_mcp_tool_call["mcp_server_logo_url"] = logo_url + cost_info = mcp_info.get("mcp_server_cost_info") + if cost_info: + standard_logging_mcp_tool_call["mcp_server_cost_info"] = cost_info + + if litellm_logging_obj: + litellm_logging_obj.model_call_details[ + "mcp_tool_call_metadata" + ] = standard_logging_mcp_tool_call + litellm_logging_obj.model = f"MCP: {tool_name}" + litellm_logging_obj.call_type = CallTypes.call_mcp_tool.value + result = await global_mcp_server_manager.call_tool( server_name=server_name, name=sanitized_tool_name, @@ -526,6 +635,26 @@ class LiteLLM_Proxy_MCP_Handler: proxy_logging_obj=proxy_logging_obj, ) + if litellm_logging_obj: + try: + litellm_logging_obj.post_call(original_response=result) + end_time = datetime.now() + await litellm_logging_obj.async_post_mcp_tool_call_hook( + kwargs=litellm_logging_obj.model_call_details, + response_obj=result, + start_time=start_time, + end_time=end_time, + ) + await litellm_logging_obj.async_success_handler( + result=result, + start_time=start_time, + end_time=end_time, + ) + except Exception: + verbose_logger.exception( + "Failed to log MCP tool call success for %s", tool_name + ) + # Format result for inclusion in response result_text = LiteLLM_Proxy_MCP_Handler._parse_mcp_result(result) tool_results.append( @@ -537,6 +666,12 @@ class LiteLLM_Proxy_MCP_Handler: ) except BlockedPiiEntityError as e: + await LiteLLM_Proxy_MCP_Handler._log_mcp_tool_failure( + proxy_logging_obj=proxy_logging_obj, + user_api_key_auth=user_api_key_auth, + request_data=logging_request_data, + error=e, + ) verbose_logger.error( f"BlockedPiiEntityError in MCP tool call: {str(e)}" ) @@ -549,6 +684,12 @@ class LiteLLM_Proxy_MCP_Handler: } ) except GuardrailRaisedException as e: + await LiteLLM_Proxy_MCP_Handler._log_mcp_tool_failure( + proxy_logging_obj=proxy_logging_obj, + user_api_key_auth=user_api_key_auth, + request_data=logging_request_data, + error=e, + ) verbose_logger.error( f"GuardrailRaisedException in MCP tool call: {str(e)}" ) @@ -561,12 +702,28 @@ class LiteLLM_Proxy_MCP_Handler: } ) except HTTPException as e: + await LiteLLM_Proxy_MCP_Handler._log_mcp_tool_failure( + proxy_logging_obj=proxy_logging_obj, + user_api_key_auth=user_api_key_auth, + request_data=logging_request_data, + error=e, + ) verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}") error_message = f"Tool call failed: {str(e.detail) if hasattr(e, 'detail') else str(e)}" tool_results.append( - {"tool_call_id": tool_call_id, "result": error_message} + { + "tool_call_id": tool_call_id, + "result": error_message, + "name": tool_name, + } ) except Exception as e: + await LiteLLM_Proxy_MCP_Handler._log_mcp_tool_failure( + proxy_logging_obj=proxy_logging_obj, + user_api_key_auth=user_api_key_auth, + request_data=logging_request_data, + error=e, + ) verbose_logger.exception(f"Error executing MCP tool call: {e}") tool_results.append( { @@ -718,6 +875,33 @@ class LiteLLM_Proxy_MCP_Handler: **call_params, ) + @staticmethod + async def _log_mcp_tool_failure( + *, + proxy_logging_obj: Optional["ProxyLogging"], + user_api_key_auth: Any, + request_data: Dict[str, Any], + error: Exception, + ) -> None: + """Log MCP tool failures via proxy logging hooks.""" + + if proxy_logging_obj is None or user_api_key_auth is None: + return + + try: + traceback_str = traceback.format_exc( + limit=MAXIMUM_TRACEBACK_LINES_TO_LOG + ) + await proxy_logging_obj.post_call_failure_hook( + request_data=request_data, + original_exception=error, + user_api_key_dict=user_api_key_auth, + route="/responses/mcp/call_tool", + traceback_str=traceback_str, + ) + except Exception: + verbose_logger.exception("Failed to log MCP tool call failure") + @staticmethod def _create_mcp_streaming_response( input: Union[str, Any], @@ -758,7 +942,7 @@ class LiteLLM_Proxy_MCP_Handler: mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events tool_server_map=tool_server_map, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, - user_api_key_auth=kwargs.get("user_api_key_auth"), + user_api_key_auth=kwargs.get("user_api_key_auth") or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"), original_request_params=request_params, ) diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index ac040d3d6ec..53f39164e91 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -298,6 +298,8 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.custom_llm_provider = self.original_request_params.get( "custom_llm_provider", None ) + self.litellm_call_id = self.original_request_params.get("litellm_call_id") + self.litellm_trace_id = self.original_request_params.get("litellm_trace_id") self._extract_mcp_headers_from_params() @@ -568,6 +570,8 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): mcp_server_auth_headers=self.mcp_server_auth_headers, oauth2_headers=self.oauth2_headers, raw_headers=self.raw_headers, + litellm_call_id=self.litellm_call_id, + litellm_trace_id=self.litellm_trace_id, ) # Create completion events and output_item.done events for tool execution From 872e5b98977eac7eea95be6fd6069e751709960a Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Wed, 21 Jan 2026 14:32:08 +0900 Subject: [PATCH 07/29] feat: log mcp list_tools calls to SpendLogs --- .../proxy/_experimental/mcp_server/server.py | 226 ++++++++++++++---- .../mcp/litellm_proxy_mcp_handler.py | 12 +- litellm/types/utils.py | 2 + 3 files changed, 188 insertions(+), 52 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index e32ae85f9ad..52117d86066 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -7,6 +7,7 @@ import asyncio import contextlib from datetime import datetime import traceback +import uuid from typing import Any, AsyncIterator, Dict, List, Optional, Tuple, Union, cast from fastapi import FastAPI, HTTPException @@ -28,7 +29,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.types.mcp import MCPAuth from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall -from litellm.utils import client +from litellm.utils import Rules, client, function_setup # Check if MCP is available # "mcp" requires python 3.10 or higher, but several litellm users use python 3.8 @@ -228,6 +229,8 @@ if MCP_AVAILABLE: mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + log_list_tools_to_spendlogs=True, + list_tools_log_source="mcp_protocol", ) verbose_logger.info( f"MCP list_tools - Successfully returned {len(tools)} tools" @@ -742,6 +745,8 @@ if MCP_AVAILABLE: mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, oauth2_headers: Optional[Dict[str, str]] = None, raw_headers: Optional[Dict[str, str]] = None, + log_list_tools_to_spendlogs: bool = False, + list_tools_log_source: Optional[str] = None, ) -> List[MCPTool]: """ Helper method to fetch tools from MCP servers based on server filtering criteria. @@ -759,67 +764,186 @@ if MCP_AVAILABLE: if not MCP_AVAILABLE: return [] - allowed_mcp_servers = await _get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_servers=mcp_servers, - ) + list_tools_start_time = datetime.now() + litellm_logging_obj: Optional[LiteLLMLoggingObj] = None + list_tools_request_data: Dict[str, Any] = {} - # Decide whether to add prefix based on number of allowed servers - add_prefix = not (len(allowed_mcp_servers) == 1) + if log_list_tools_to_spendlogs: + # This is intentionally minimal: only async_success_handler / post_call_failure_hook + rules_obj = Rules() + list_tools_call_id = str(uuid.uuid4()) + spend_logs_metadata: Dict[str, Any] = { + "mcp_operation": "list_tools", + } + if isinstance(list_tools_log_source, str): + spend_logs_metadata["source"] = list_tools_log_source + if isinstance(mcp_servers, list): + spend_logs_metadata["requested_mcp_servers"] = mcp_servers - async def _fetch_and_filter_server_tools(server: MCPServer) -> List[MCPTool]: - """Fetch and filter tools from a single server with error handling.""" - if server is None: - return [] + list_tools_request_data = { + "model": "MCP: list_tools", + "call_type": CallTypes.list_mcp_tools.value, + "litellm_call_id": list_tools_call_id, + "metadata": { + "spend_logs_metadata": spend_logs_metadata, + }, + # Provide a small input payload for standard logging + "input": [ + { + "role": "system", + "content": { + "mcp_operation": "list_tools", + "requested_mcp_servers": mcp_servers, + }, + } + ], + } - server_auth_header, extra_headers = _prepare_mcp_server_headers( - server=server, - mcp_server_auth_headers=mcp_server_auth_headers, - mcp_auth_header=mcp_auth_header, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) + # Attach user identifiers when available (matches call_mcp_tool style) + if user_api_key_auth is not None: + user_api_key = getattr(user_api_key_auth, "api_key", None) + if user_api_key: + cast(dict, list_tools_request_data["metadata"])[ + "user_api_key" + ] = user_api_key + + user_identifier = getattr(user_api_key_auth, "end_user_id", None) or getattr( + user_api_key_auth, "user_id", None + ) + if user_identifier: + list_tools_request_data["user"] = user_identifier try: - tools = await global_mcp_server_manager._get_tools_from_server( + litellm_logging_obj, _ = function_setup( + original_function="list_mcp_tools", + rules_obj=rules_obj, + start_time=list_tools_start_time, + **list_tools_request_data, + ) + if litellm_logging_obj: + litellm_logging_obj.call_type = CallTypes.list_mcp_tools.value + litellm_logging_obj.model = "MCP: list_tools" + except Exception as logging_error: + verbose_logger.debug( + "Failed to initialize logging for MCP list_tools: %s", logging_error + ) + litellm_logging_obj = None + + try: + allowed_mcp_servers = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + ) + + # Decide whether to add prefix based on number of allowed servers + add_prefix = not (len(allowed_mcp_servers) == 1) + + async def _fetch_and_filter_server_tools(server: MCPServer) -> List[MCPTool]: + """Fetch and filter tools from a single server with error handling.""" + if server is None: + return [] + + server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=add_prefix, + mcp_server_auth_headers=mcp_server_auth_headers, + mcp_auth_header=mcp_auth_header, + oauth2_headers=oauth2_headers, raw_headers=raw_headers, ) - filtered_tools = filter_tools_by_allowed_tools(tools, server) - filtered_tools = await filter_tools_by_key_team_permissions( - tools=filtered_tools, - server_id=server.server_id, - user_api_key_auth=user_api_key_auth, + try: + tools = await global_mcp_server_manager._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=add_prefix, + raw_headers=raw_headers, + ) + filtered_tools = filter_tools_by_allowed_tools(tools, server) + + filtered_tools = await filter_tools_by_key_team_permissions( + tools=filtered_tools, + server_id=server.server_id, + user_api_key_auth=user_api_key_auth, + ) + + verbose_logger.debug( + f"Successfully fetched {len(tools)} tools from server {server.name}, {len(filtered_tools)} after filtering" + ) + return filtered_tools + except Exception as e: + verbose_logger.exception( + f"Error getting tools from server {server.name}: {str(e)}" + ) + return [] + + # Fetch tools from all servers in parallel + tasks = [ + _fetch_and_filter_server_tools(server) for server in allowed_mcp_servers + ] + results = await asyncio.gather(*tasks) + + # Flatten results into single list + all_tools: List[MCPTool] = [tool for tools in results for tool in tools] + + # If logging is enabled, enrich spend_logs_metadata with counts + if litellm_logging_obj: + per_server_tool_counts: Dict[str, int] = {} + for server, server_tools in zip(allowed_mcp_servers, results): + if server is None: + continue + server_key = ( + getattr(server, "server_name", None) + or getattr(server, "alias", None) + or getattr(server, "name", None) + or "unknown" + ) + per_server_tool_counts[str(server_key)] = len(server_tools) + + metadata_dict = litellm_logging_obj.model_call_details.get("metadata") + if isinstance(metadata_dict, dict): + spend_meta = metadata_dict.get("spend_logs_metadata") + if not isinstance(spend_meta, dict): + spend_meta = {} + metadata_dict["spend_logs_metadata"] = spend_meta + spend_meta["allowed_server_count"] = len(allowed_mcp_servers) + spend_meta["tool_count_total"] = len(all_tools) + spend_meta["per_server_tool_counts"] = per_server_tool_counts + + end_time = datetime.now() + await litellm_logging_obj.async_success_handler( + result=all_tools, + start_time=list_tools_start_time, + end_time=end_time, ) - verbose_logger.debug( - f"Successfully fetched {len(tools)} tools from server {server.name}, {len(filtered_tools)} after filtering" - ) - return filtered_tools - except Exception as e: - verbose_logger.exception( - f"Error getting tools from server {server.name}: {str(e)}" - ) - return [] + verbose_logger.info( + f"Successfully fetched {len(all_tools)} tools total from all MCP servers" + ) - # Fetch tools from all servers in parallel - tasks = [ - _fetch_and_filter_server_tools(server) for server in allowed_mcp_servers - ] - results = await asyncio.gather(*tasks) + return all_tools + except Exception as e: + # Only fire failure hook if logging was requested for this list-tools execution + if log_list_tools_to_spendlogs and user_api_key_auth is not None: + try: + from litellm.proxy.proxy_server import proxy_logging_obj - # Flatten results into single list - all_tools: List[MCPTool] = [tool for tools in results for tool in tools] - - verbose_logger.info( - f"Successfully fetched {len(all_tools)} tools total from all MCP servers" - ) - - return all_tools + if proxy_logging_obj: + traceback_str = traceback.format_exc( + limit=MAXIMUM_TRACEBACK_LINES_TO_LOG + ) + await proxy_logging_obj.post_call_failure_hook( + request_data=list_tools_request_data or {}, + original_exception=e, + user_api_key_dict=user_api_key_auth, + route="/mcp/list_tools", + traceback_str=traceback_str, + ) + except Exception: + verbose_logger.debug( + "Failed to log MCP list_tools failure via post_call_failure_hook" + ) + raise async def _get_prompts_from_mcp_servers( user_api_key_auth: Optional[UserAPIKeyAuth], @@ -1052,6 +1176,8 @@ if MCP_AVAILABLE: mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, oauth2_headers: Optional[Dict[str, str]] = None, raw_headers: Optional[Dict[str, str]] = None, + log_list_tools_to_spendlogs: bool = False, + list_tools_log_source: Optional[str] = None, ) -> List[MCPTool]: """ List all available MCP tools. @@ -1077,6 +1203,8 @@ if MCP_AVAILABLE: mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + log_list_tools_to_spendlogs=log_list_tools_to_spendlogs, + list_tools_log_source=list_tools_log_source, ) verbose_logger.debug( f"Successfully fetched {len(managed_tools)} tools from managed MCP servers" diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index f5757e4d52e..e22898ae0ff 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -18,7 +18,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.proxy._experimental.mcp_server.utils import split_server_prefix_from_name from litellm.responses.main import aresponses from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator -from litellm.types.llms.openai import ResponsesAPIResponse, ToolParam +from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import CallTypes, Choices, ModelResponse, StandardLoggingMCPToolCall from litellm.utils import Rules, function_setup @@ -28,6 +28,10 @@ if TYPE_CHECKING: else: MCPTool = Any +# NOTE: We intentionally keep ToolParam as a broad type here to avoid tight coupling +# to optional OpenAI SDK typing symbols in environments that may not have them available. +ToolParam = Dict[str, Any] + LITELLM_PROXY_MCP_SERVER_URL = "litellm_proxy" LITELLM_PROXY_MCP_SERVER_URL_PREFIX = f"{LITELLM_PROXY_MCP_SERVER_URL}/mcp/" @@ -123,6 +127,8 @@ class LiteLLM_Proxy_MCP_Handler: mcp_auth_header=None, mcp_servers=mcp_servers, mcp_server_auth_headers=None, + log_list_tools_to_spendlogs=True, + list_tools_log_source="responses", ) allowed_mcp_server_ids = ( await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth) @@ -495,7 +501,7 @@ class LiteLLM_Proxy_MCP_Handler: rules_obj = Rules() for tool_call in tool_calls: logging_request_data: Dict[str, Any] = {} - tool_name: str = "" + tool_name: Optional[str] = None try: ( tool_name, @@ -539,7 +545,7 @@ class LiteLLM_Proxy_MCP_Handler: } ] tool_logging_call_id = litellm_call_id or str(uuid.uuid4()) - logging_request_data: Dict[str, Any] = { + logging_request_data = { "model": f"MCP: {tool_name}", "metadata": { "tool_call_id": tool_call_id, diff --git a/litellm/types/utils.py b/litellm/types/utils.py index f5e217d8b46..324380db2b4 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -384,6 +384,7 @@ class CallTypes(str, Enum): # MCP Call Types ######################################################### call_mcp_tool = "call_mcp_tool" + list_mcp_tools = "list_mcp_tools" ######################################################### # A2A Call Types @@ -448,6 +449,7 @@ CallTypesLiteral = Literal[ "vector_store_file_delete", "avector_store_file_delete", "call_mcp_tool", + "list_mcp_tools", "asend_message", "send_message", "aresponses", From 2441b0570048d965026edcc01def2ee09bf4ae90 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 20 Jan 2026 21:38:58 -0800 Subject: [PATCH 08/29] Fix MCP Server resetting back to Overview --- .../components/mcp_tools/mcp_server_view.tsx | 31 ++- .../src/components/mcp_tools/mcp_servers.tsx | 200 +++++++++--------- .../src/components/mcp_tools/mcp_tools.tsx | 5 +- 3 files changed, 119 insertions(+), 117 deletions(-) diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx index 6a4e9c105ff..809095c048c 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx @@ -36,6 +36,8 @@ export const MCPServerView: React.FC = ({ const [editing, setEditing] = useState(isEditing); const [showFullUrl, setShowFullUrl] = useState(false); const [copiedStates, setCopiedStates] = useState>({}); + const [selectedTabIndex, setSelectedTabIndex] = useState(0); + const handleSuccess = (updated: MCPServer) => { setEditing(false); onBack(); @@ -72,11 +74,10 @@ export const MCPServerView: React.FC = ({ size="small" icon={copiedStates["mcp-server_name"] ? : } onClick={() => copyToClipboard(mcpServer.server_name, "mcp-server_name")} - className={`left-2 z-10 transition-all duration-200 ${ - copiedStates["mcp-server_name"] - ? "text-green-600 bg-green-50 border-green-200" - : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" - }`} + className={`left-2 z-10 transition-all duration-200 ${copiedStates["mcp-server_name"] + ? "text-green-600 bg-green-50 border-green-200" + : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" + }`} /> {mcpServer.alias && ( <> @@ -87,11 +88,10 @@ export const MCPServerView: React.FC = ({ size="small" icon={copiedStates["mcp-alias"] ? : } onClick={() => copyToClipboard(mcpServer.alias, "mcp-alias")} - className={`left-2 z-10 transition-all duration-200 ${ - copiedStates["mcp-alias"] - ? "text-green-600 bg-green-50 border-green-200" - : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" - }`} + className={`left-2 z-10 transition-all duration-200 ${copiedStates["mcp-alias"] + ? "text-green-600 bg-green-50 border-green-200" + : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" + }`} /> )} @@ -103,18 +103,17 @@ export const MCPServerView: React.FC = ({ size="small" icon={copiedStates["mcp-server-id"] ? : } onClick={() => copyToClipboard(mcpServer.server_id, "mcp-server-id")} - className={`left-2 z-10 transition-all duration-200 ${ - copiedStates["mcp-server-id"] - ? "text-green-600 bg-green-50 border-green-200" - : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" - }`} + className={`left-2 z-10 transition-all duration-200 ${copiedStates["mcp-server-id"] + ? "text-green-600 bg-green-50 border-green-200" + : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" + }`} /> {/* TODO: magic number for index */} - + {[ Overview, diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx index f6669fb2829..a0cd1720c88 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx @@ -212,103 +212,28 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) return
Missing required authentication parameters.
; } - const ServersTab = () => - selectedServerId ? ( - server.server_id === selectedServerId) || { - server_id: "", - server_name: "", - alias: "", - url: "", - transport: "", - auth_type: "", - created_at: "", - created_by: "", - updated_at: "", - updated_by: "", - } - } - onBack={() => { - setEditServer(false); - setSelectedServerId(null); - refetch(); - }} - isProxyAdmin={isAdminRole(userRole)} - isEditing={editServer} - accessToken={accessToken} - userID={userID} - userRole={userRole} - availableAccessGroups={uniqueMcpAccessGroups} - /> - ) : ( -
-
-
-
-
- Current Team: - - - Access Group: - - - - - -
-
-
-
-
-
} - getRowCanExpand={() => false} - isLoading={isLoadingServers} - noDataMessage="No MCP servers configured" - loadingMessage="🚅 Loading MCP servers..." - /> -
-
- ); + // Memoize the selected server to prevent unnecessary re-renders + const selectedServer = React.useMemo(() => { + return filteredServers.find((server: MCPServer) => server.server_id === selectedServerId) || { + server_id: "", + server_name: "", + alias: "", + url: "", + transport: "", + auth_type: "", + created_at: "", + created_by: "", + updated_at: "", + updated_by: "", + }; + }, [filteredServers, selectedServerId]); + + // Memoize the onBack callback to prevent unnecessary re-renders + const handleBack = React.useCallback(() => { + setEditServer(false); + setSelectedServerId(null); + refetch(); + }, [refetch]); return (
@@ -381,7 +306,86 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) - + {selectedServerId ? ( + + ) : ( +
+
+
+
+
+ Current Team: + + + Access Group: + + + + + +
+
+
+
+
+
} + getRowCanExpand={() => false} + isLoading={isLoadingServers} + noDataMessage="No MCP servers configured" + loadingMessage="🚅 Loading MCP servers..." + /> +
+
+ )}
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx index 37c61023f23..7ee1e64a228 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx @@ -127,11 +127,10 @@ const MCPToolsViewer = ({ {toolsData.map((tool: MCPTool) => (
{ setSelectedTool(tool); setToolResult(null); From 90bd80a4c2ad1afca94b43eed546c2b09335e947 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 20 Jan 2026 21:46:21 -0800 Subject: [PATCH 09/29] fixing build --- .../src/components/mcp_tools/mcp_servers.tsx | 40 +++++++++---------- 1 file changed, 20 insertions(+), 20 deletions(-) diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx index a0cd1720c88..77034309c91 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx @@ -2,7 +2,7 @@ import { isAdminRole } from "@/utils/roles"; import { QuestionCircleOutlined } from "@ant-design/icons"; import { Button, Tab, TabGroup, TabList, TabPanel, TabPanels, Text, Title } from "@tremor/react"; import { Descriptions, Modal, Select, Tooltip, Typography } from "antd"; -import React, { useEffect, useState, useMemo } from "react"; +import React, { useEffect, useState, useMemo, useCallback } from "react"; import { useMCPServers } from "../../app/(dashboard)/hooks/mcpServers/useMCPServers"; import { useMCPServerHealth } from "../../app/(dashboard)/hooks/mcpServers/useMCPServerHealth"; import NotificationsManager from "../molecules/notifications_manager"; @@ -115,20 +115,8 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) ); }, [serversWithHealth]); - // Handle team filter change - const handleTeamChange = (teamId: string) => { - setSelectedTeam(teamId); - filterServers(teamId, selectedMcpAccessGroup); - }; - - // Handle MCP access group filter change - const handleMcpAccessGroupChange = (group: string) => { - setSelectedMcpAccessGroup(group); - filterServers(selectedTeam, group); - }; - // Filtering logic for both team and access group - const filterServers = (teamId: string, group: string) => { + const filterServers = useCallback((teamId: string, group: string) => { if (!serversWithHealth) return setFilteredServers([]); let filtered = serversWithHealth; if (teamId === "personal") { @@ -144,12 +132,24 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) ); } setFilteredServers(filtered); + }, [serversWithHealth]); + + // Handle team filter change + const handleTeamChange = (teamId: string) => { + setSelectedTeam(teamId); + filterServers(teamId, selectedMcpAccessGroup); + }; + + // Handle MCP access group filter change + const handleMcpAccessGroupChange = (group: string) => { + setSelectedMcpAccessGroup(group); + filterServers(selectedTeam, group); }; // Initial and effect-based filtering (trigger on query data updates and health data updates) useEffect(() => { filterServers(selectedTeam, selectedMcpAccessGroup); - }, [serversWithHealth, selectedTeam, selectedMcpAccessGroup]); + }, [serversWithHealth, selectedTeam, selectedMcpAccessGroup, filterServers]); const columns = React.useMemo( () => @@ -207,11 +207,6 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) setModalVisible(false); }; - if (!accessToken || !userRole || !userID) { - console.log("Missing required authentication parameters", { accessToken, userRole, userID }); - return
Missing required authentication parameters.
; - } - // Memoize the selected server to prevent unnecessary re-renders const selectedServer = React.useMemo(() => { return filteredServers.find((server: MCPServer) => server.server_id === selectedServerId) || { @@ -235,6 +230,11 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) refetch(); }, [refetch]); + if (!accessToken || !userRole || !userID) { + console.log("Missing required authentication parameters", { accessToken, userRole, userID }); + return
Missing required authentication parameters.
; + } + return (
Date: Tue, 20 Jan 2026 21:50:28 -0800 Subject: [PATCH 10/29] Adding tests --- .../components/mcp_tools/mcp_servers.test.tsx | 118 +++++++++++++++++- 1 file changed, 117 insertions(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx index 4b8698b9762..b77e962f5f3 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx @@ -1,5 +1,5 @@ import React from "react"; -import { render, waitFor } from "@testing-library/react"; +import { render, waitFor, screen, fireEvent, act } from "@testing-library/react"; import { describe, it, expect, vi, beforeEach } from "vitest"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import MCPServers from "./mcp_servers"; @@ -228,4 +228,120 @@ describe("MCPServers", () => { expect(networking.fetchMCPServerHealth).toHaveBeenCalled(); }); }); + + it("should filter servers by team when a team is selected", async () => { + // Mock MCP servers with different teams + const mockServers = [ + { + server_id: "server-1", + server_name: "Team A Server", + alias: "team-a-server", + url: "https://example.com/mcp", + transport: "http", + auth_type: "none", + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + teams: [{ team_id: "team-a", team_alias: "Team A" }], + mcp_access_groups: [], + }, + { + server_id: "server-2", + server_name: "Team B Server", + alias: "team-b-server", + url: "https://example2.com/mcp", + transport: "sse", + auth_type: "api_key", + created_at: "2024-01-02T00:00:00Z", + created_by: "user-2", + updated_at: "2024-01-02T00:00:00Z", + updated_by: "user-2", + teams: [{ team_id: "team-b", team_alias: "Team B" }], + mcp_access_groups: [], + }, + { + server_id: "server-3", + server_name: "Team A Server 2", + alias: "team-a-server-2", + url: "https://example3.com/mcp", + transport: "http", + auth_type: "none", + created_at: "2024-01-03T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-03T00:00:00Z", + updated_by: "user-1", + teams: [{ team_id: "team-a", team_alias: "Team A" }], + mcp_access_groups: [], + }, + ]; + + vi.mocked(networking.fetchMCPServers).mockResolvedValue(mockServers); + vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue([]); + + const queryClient = createQueryClient(); + render( + + + , + ); + + // Wait for the component to load + await waitFor(() => { + expect(screen.getByText("MCP Servers")).toBeInTheDocument(); + }); + + // Wait for servers to be rendered + await waitFor(() => { + expect(screen.getByText("Team A Server")).toBeInTheDocument(); + }); + + // Verify all servers are initially displayed + expect(screen.getByText("Team A Server")).toBeInTheDocument(); + expect(screen.getByText("Team B Server")).toBeInTheDocument(); + expect(screen.getByText("Team A Server 2")).toBeInTheDocument(); + + // Find the team select dropdown by looking for the "Current Team:" label + const teamLabel = screen.getByText("Current Team:"); + const teamSelectContainer = teamLabel.closest("div")?.querySelector(".ant-select"); + expect(teamSelectContainer).toBeTruthy(); + + // Open the dropdown by clicking on the selector + const selectSelector = teamSelectContainer?.querySelector(".ant-select-selector"); + expect(selectSelector).toBeTruthy(); + + act(() => { + fireEvent.mouseDown(selectSelector!); + }); + + // Wait for dropdown to open + await waitFor( + () => { + const dropdownOptions = document.querySelectorAll(".ant-select-item-option"); + expect(dropdownOptions.length).toBeGreaterThan(0); + }, + { timeout: 5000 }, + ); + + // Find and click on "Team A" option + const dropdownOptions = document.querySelectorAll(".ant-select-item-option"); + const teamAOption = Array.from(dropdownOptions).find((option) => + option.textContent?.includes("Team A"), + ); + expect(teamAOption).toBeTruthy(); + + act(() => { + fireEvent.click(teamAOption!); + }); + + // Wait for filtering to complete + await waitFor(() => { + // Team A servers should still be visible + expect(screen.getByText("Team A Server")).toBeInTheDocument(); + expect(screen.getByText("Team A Server 2")).toBeInTheDocument(); + }); + + // Team B server should not be visible + expect(screen.queryByText("Team B Server")).not.toBeInTheDocument(); + }); }); From 032b1a8cd2d9c2742443a1e067a780a302024d2d Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 20 Jan 2026 21:51:13 -0800 Subject: [PATCH 11/29] Fixing tests --- .../src/components/mcp_tools/mcp_servers.test.tsx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx index b77e962f5f3..4d80b383703 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx @@ -208,7 +208,7 @@ describe("MCPServers", () => { vi.mocked(networking.fetchMCPServers).mockResolvedValue(mockServers); // Mock health check to never resolve (to test loading state) vi.mocked(networking.fetchMCPServerHealth).mockImplementation( - () => new Promise(() => {}), // Never resolves + () => new Promise(() => { }), // Never resolves ); const queryClient = createQueryClient(); From caf5f7f8aeb248b64fa95d246b7376a8428656ce Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Wed, 21 Jan 2026 14:51:56 +0900 Subject: [PATCH 12/29] test: add test --- .../mcp_server/test_mcp_server.py | 135 ++++++++++++++++ .../mcp/test_litellm_proxy_mcp_handler.py | 151 +++++++++++++++++- 2 files changed, 283 insertions(+), 3 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 7aa176aaac5..8c29cf6a590 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -1817,3 +1817,138 @@ class TestMCPServerManagerReload: mock_get_all.assert_awaited_once() mock_build.assert_awaited_once_with(db_row) assert manager.registry["server-1"] is rebuilt_server + + +@pytest.mark.asyncio +async def test_call_mcp_tool_logs_failure_via_post_call_failure_hook(): + """ + Regression test for 6267f168...: + Ensure proxy-side `call_mcp_tool` logs failures via `proxy_logging_obj.post_call_failure_hook`. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + call_mcp_tool, + global_mcp_server_manager, + ) + from litellm.proxy._types import MCPTransport, UserAPIKeyAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + except ImportError: + pytest.skip("MCP server not available") + + mock_server = MCPServer( + server_id="server-123", + name="test_server", + alias="test_server", + server_name="test_server", + url="https://test-server.com/mcp", + transport=MCPTransport.http, + mcp_info={"server_name": "test_server"}, + ) + + proxy_logging_mock = MagicMock() + proxy_logging_mock.post_call_failure_hook = AsyncMock() + + user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + + with patch.object( + global_mcp_server_manager, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[mock_server.server_id], + ), patch.object( + global_mcp_server_manager, + "get_mcp_server_by_id", + return_value=mock_server, + ), patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names", + new_callable=AsyncMock, + return_value=[mock_server], + ), patch( + "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", + new_callable=AsyncMock, + side_effect=Exception("boom"), + ), patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + proxy_logging_mock, + ): + with pytest.raises(Exception): + await call_mcp_tool( + name="test_server-any_tool", + arguments={"x": 1}, + user_api_key_auth=user_auth, + litellm_call_id="cid", + ) + + proxy_logging_mock.post_call_failure_hook.assert_awaited_once() + assert ( + proxy_logging_mock.post_call_failure_hook.await_args.kwargs.get("route") + == "/mcp/call_tool" + ) + + +@pytest.mark.asyncio +async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enabled(): + """ + Regression test for 872e5b98...: + Ensure list-tools logging path calls `async_success_handler` when enabled. + """ + try: + from litellm.proxy._experimental.mcp_server.server import _get_tools_from_mcp_servers + from litellm.proxy._types import UserAPIKeyAuth + except ImportError: + pytest.skip("MCP server not available") + + user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + + server_a = MagicMock(name="server_a_obj") + server_a.name = "server_a" + server_a.alias = "server_a" + server_a.server_name = "server_a" + server_a.server_id = "a" + server_a.auth_type = None + server_a.extra_headers = None + + tool_1 = MagicMock() + tool_1.name = "server_a-tool_1" + + dummy_logging_obj = MagicMock() + dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}} + dummy_logging_obj.async_success_handler = AsyncMock() + + with patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server_a]), + ), patch( + "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + return_value=(None, None), + ), patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + side_effect=lambda tools, _server: tools, + ), patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + new=AsyncMock(side_effect=lambda tools, **_: tools), + ), patch( + "litellm.proxy._experimental.mcp_server.server.function_setup", + return_value=(dummy_logging_obj, None), + ): + mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) + + tools = await _get_tools_from_mcp_servers( + user_api_key_auth=user_auth, + mcp_auth_header=None, + mcp_servers=["server_a"], + mcp_server_auth_headers=None, + log_list_tools_to_spendlogs=True, + list_tools_log_source="mcp_protocol", + ) + + assert tools == [tool_1] + dummy_logging_obj.async_success_handler.assert_awaited_once() + assert dummy_logging_obj.async_success_handler.await_args.kwargs["result"] == [tool_1] + + spend_meta = dummy_logging_obj.model_call_details["metadata"]["spend_logs_metadata"] + assert spend_meta["tool_count_total"] == 1 + assert spend_meta["allowed_server_count"] == 1 + assert spend_meta["per_server_tool_counts"]["server_a"] == 1 diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index b632e72f567..15fdc7bd0c4 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -1,12 +1,15 @@ import sys import types -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, MagicMock import pytest +from fastapi import HTTPException +import importlib from litellm.responses.mcp.litellm_proxy_mcp_handler import ( LiteLLM_Proxy_MCP_Handler, ) +from typing import Any, cast from litellm.types.utils import ModelResponse from litellm.types.responses.main import OutputFunctionToolCall @@ -22,7 +25,9 @@ def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module) fake_manager = types.SimpleNamespace( - call_tool=AsyncMock(return_value=_DummyMCPResult()) + call_tool=AsyncMock(return_value=_DummyMCPResult()), + # Newer logging path calls this to enrich spend logs metadata + _get_mcp_server_from_tool_name=MagicMock(return_value=None), ) monkeypatch.setattr( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -31,6 +36,15 @@ def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: return fake_manager.call_tool +def _setup_proxy_logging(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: + """Patch proxy_logging_obj so failure hook can be asserted.""" + proxy_logging_obj = MagicMock() + proxy_logging_obj.post_call_failure_hook = AsyncMock() + proxy_module = types.SimpleNamespace(proxy_logging_obj=proxy_logging_obj) + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module) + return proxy_logging_obj.post_call_failure_hook + + def test_deduplicate_mcp_tools_single_allowed_server(): tools = [{"name": "search"}, {"name": "search"}] # duplicate on purpose @@ -184,7 +198,7 @@ def test_create_follow_up_input_handles_response_function_tool_call(): ) follow_up = LiteLLM_Proxy_MCP_Handler._create_follow_up_input( - response=response, + response=cast(Any, response), tool_results=[], original_input=None, ) @@ -216,6 +230,8 @@ async def test_execute_tool_calls_strips_server_prefix(monkeypatch): user_api_key_auth=None, ) + assert call_tool_mock.await_count == 1 + assert call_tool_mock.await_args is not None assert call_tool_mock.await_args.kwargs["name"] == "read_wiki_structure" @@ -236,6 +252,8 @@ async def test_execute_tool_calls_keeps_tool_name_without_prefix(monkeypatch): user_api_key_auth=None, ) + assert call_tool_mock.await_count == 1 + assert call_tool_mock.await_args is not None assert call_tool_mock.await_args.kwargs["name"] == tool_name @@ -256,4 +274,131 @@ async def test_execute_tool_calls_keeps_tool_name_when_equal_to_server(monkeypat user_api_key_auth=None, ) + assert call_tool_mock.await_count == 1 + assert call_tool_mock.await_args is not None assert call_tool_mock.await_args.kwargs["name"] == tool_name + + +@pytest.mark.asyncio +async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkeypatch): + """ + Regression test for ae4d92ad...: + Ensure responses-side MCP tool execution logs failures via proxy_logging_obj.post_call_failure_hook. + """ + post_call_failure_hook = _setup_proxy_logging(monkeypatch) + + fake_manager = types.SimpleNamespace( + call_tool=AsyncMock( + side_effect=HTTPException(status_code=500, detail="boom") + ) + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + fake_manager, + ) + + tool_name = "deepwiki-read_wiki_structure" + tool_calls = [ + {"id": "call-err", "function": {"name": tool_name, "arguments": "{}"}} + ] + + user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user") + + results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map={tool_name: "deepwiki"}, + tool_calls=tool_calls, + user_api_key_auth=user_auth, + litellm_call_id="cid", + litellm_trace_id="tid", + ) + + assert len(results) == 1 + assert results[0]["tool_call_id"] == "call-err" + assert results[0]["name"] == tool_name + + post_call_failure_hook.assert_awaited_once() + assert post_call_failure_hook.await_args is not None + assert ( + post_call_failure_hook.await_args.kwargs.get("route") + == "/responses/mcp/call_tool" + ) + + +@pytest.mark.asyncio +async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_function_setup( + monkeypatch, +): + """ + Regression test for ae4d92ad...: + Ensure litellm_call_id / litellm_trace_id are forwarded into function_setup kwargs. + """ + _setup_proxy_logging(monkeypatch) + call_tool_mock = _setup_mcp_call_environment(monkeypatch) + + captured = {} + + def fake_function_setup(*_args, **kwargs): + captured.update(kwargs) + return None, None + + # NOTE: Don't patch via dotted string path here because `litellm.responses` + # is a function attribute on the `litellm` package (shadowing the submodule), + # which breaks monkeypatch's importpath resolution. + handler_module = importlib.import_module( + "litellm.responses.mcp.litellm_proxy_mcp_handler" + ) + monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) + + tool_name = "deepwiki-read_wiki_structure" + tool_calls = [ + {"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}} + ] + + await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map={tool_name: "deepwiki"}, + tool_calls=tool_calls, + user_api_key_auth=None, + litellm_call_id="cid", + litellm_trace_id="tid", + ) + + # Ensure the tool call was attempted (sanity) + assert call_tool_mock.await_count == 1 + + assert captured.get("litellm_call_id") == "cid" + assert captured.get("litellm_trace_id") == "tid" + + +@pytest.mark.asyncio +async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch): + """ + Regression test for 872e5b98...: + Ensure responses-side tool discovery enables list-tools SpendLogs logging flags. + """ + mock_get_tools = AsyncMock(return_value=[]) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server._get_tools_from_mcp_servers", + mock_get_tools, + ) + + # Patch manager methods used by _get_mcp_tools_from_manager to avoid needing full UserAPIKeyAuth fields. + fake_manager = types.SimpleNamespace( + get_allowed_mcp_servers=AsyncMock(return_value=[]), + get_mcp_servers_from_ids=MagicMock(return_value=[]), + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + fake_manager, + ) + + user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user") + tools, _server_names = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( + user_api_key_auth=user_auth, + mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}], + ) + + assert tools == [] + assert mock_get_tools.await_count == 1 + assert mock_get_tools.await_args is not None + assert mock_get_tools.await_args.kwargs["log_list_tools_to_spendlogs"] is True + assert mock_get_tools.await_args.kwargs["list_tools_log_source"] == "responses" From 72cbc295e597621b08bcfb7c1a89cd721d5e7685 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Wed, 21 Jan 2026 14:54:50 +0900 Subject: [PATCH 13/29] chore: format --- .../proxy/_experimental/mcp_server/server.py | 16 ++-- litellm/responses/main.py | 75 ++++++++++--------- .../mcp/litellm_proxy_mcp_handler.py | 24 +++--- .../responses/mcp/mcp_streaming_iterator.py | 67 ++++++++++------- litellm/types/utils.py | 10 +-- 5 files changed, 110 insertions(+), 82 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 52117d86066..4e5c73be806 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -807,9 +807,9 @@ if MCP_AVAILABLE: "user_api_key" ] = user_api_key - user_identifier = getattr(user_api_key_auth, "end_user_id", None) or getattr( - user_api_key_auth, "user_id", None - ) + user_identifier = getattr( + user_api_key_auth, "end_user_id", None + ) or getattr(user_api_key_auth, "user_id", None) if user_identifier: list_tools_request_data["user"] = user_identifier @@ -838,7 +838,9 @@ if MCP_AVAILABLE: # Decide whether to add prefix based on number of allowed servers add_prefix = not (len(allowed_mcp_servers) == 1) - async def _fetch_and_filter_server_tools(server: MCPServer) -> List[MCPTool]: + async def _fetch_and_filter_server_tools( + server: MCPServer, + ) -> List[MCPTool]: """Fetch and filter tools from a single server with error handling.""" if server is None: return [] @@ -1517,11 +1519,9 @@ if MCP_AVAILABLE: **kwargs, ) except Exception as e: - traceback_str = traceback.format_exc( - limit=MAXIMUM_TRACEBACK_LINES_TO_LOG - ) + traceback_str = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) from litellm.proxy.proxy_server import proxy_logging_obj - + if proxy_logging_obj and user_api_key_auth: await proxy_logging_obj.post_call_failure_hook( request_data=kwargs, diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 78d358a1e39..83c23a58500 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -169,7 +169,9 @@ async def aresponses_api_with_mcp( # Process MCP tools through the complete pipeline (fetch + filter + deduplicate + transform) # Extract user_api_key_auth from litellm_metadata (where it's added by add_user_api_key_auth_to_request_metadata) - user_api_key_auth = kwargs.get("user_api_key_auth") or kwargs.get("litellm_metadata", {}).get("user_api_key_auth") + user_api_key_auth = kwargs.get("user_api_key_auth") or kwargs.get( + "litellm_metadata", {} + ).get("user_api_key_auth") # Get original MCP tools (for events) and OpenAI tools (for LLM) by reusing existing methods ( @@ -280,7 +282,7 @@ async def aresponses_api_with_mcp( user_api_key_auth = kwargs.get("litellm_metadata", {}).get( "user_api_key_auth" ) - + # Extract MCP auth headers from the request to pass to MCP server secret_fields: Optional[Dict[str, Any]] = kwargs.get("secret_fields") ( @@ -292,7 +294,7 @@ async def aresponses_api_with_mcp( secret_fields=secret_fields, tools=tools, ) - + tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=tool_server_map, tool_calls=tool_calls, @@ -590,9 +592,12 @@ def responses( ######################################################### # Update input with provider-specific file IDs if managed files are used ######################################################### - input = cast(Union[str, ResponseInputParam], update_responses_input_with_model_file_ids(input=input)) + input = cast( + Union[str, ResponseInputParam], + update_responses_input_with_model_file_ids(input=input), + ) local_vars["input"] = input - + ######################################################### # Native MCP Responses API ######################################################### @@ -627,11 +632,11 @@ def responses( ) # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=model, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=model, + provider=litellm.LlmProviders(custom_llm_provider), ) local_vars.update(kwargs) @@ -826,11 +831,11 @@ def delete_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=None, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=None, + provider=litellm.LlmProviders(custom_llm_provider), ) if responses_api_provider_config is None: @@ -1006,11 +1011,11 @@ def get_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=None, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=None, + provider=litellm.LlmProviders(custom_llm_provider), ) if responses_api_provider_config is None: @@ -1163,11 +1168,11 @@ def list_input_items( if custom_llm_provider is None: raise ValueError("custom_llm_provider is required but passed as None") - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=None, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=None, + provider=litellm.LlmProviders(custom_llm_provider), ) if responses_api_provider_config is None: @@ -1321,11 +1326,11 @@ def cancel_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=None, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=None, + provider=litellm.LlmProviders(custom_llm_provider), ) if responses_api_provider_config is None: @@ -1503,11 +1508,11 @@ def compact_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=model, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=model, + provider=litellm.LlmProviders(custom_llm_provider), ) if responses_api_provider_config is None: diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index e22898ae0ff..7cc821f5e02 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -19,7 +19,12 @@ from litellm.proxy._experimental.mcp_server.utils import split_server_prefix_fro from litellm.responses.main import aresponses from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator from litellm.types.llms.openai import ResponsesAPIResponse -from litellm.types.utils import CallTypes, Choices, ModelResponse, StandardLoggingMCPToolCall +from litellm.types.utils import ( + CallTypes, + Choices, + ModelResponse, + StandardLoggingMCPToolCall, +) from litellm.utils import Rules, function_setup if TYPE_CHECKING: @@ -564,9 +569,9 @@ class LiteLLM_Proxy_MCP_Handler: if user_api_key: logging_request_data["metadata"]["user_api_key"] = user_api_key - user_identifier = getattr(user_api_key_auth, "end_user_id", None) or getattr( - user_api_key_auth, "user_id", None - ) + user_identifier = getattr( + user_api_key_auth, "end_user_id", None + ) or getattr(user_api_key_auth, "user_id", None) if user_identifier: logging_request_data["user"] = user_identifier @@ -620,7 +625,9 @@ class LiteLLM_Proxy_MCP_Handler: standard_logging_mcp_tool_call["mcp_server_logo_url"] = logo_url cost_info = mcp_info.get("mcp_server_cost_info") if cost_info: - standard_logging_mcp_tool_call["mcp_server_cost_info"] = cost_info + standard_logging_mcp_tool_call[ + "mcp_server_cost_info" + ] = cost_info if litellm_logging_obj: litellm_logging_obj.model_call_details[ @@ -895,9 +902,7 @@ class LiteLLM_Proxy_MCP_Handler: return try: - traceback_str = traceback.format_exc( - limit=MAXIMUM_TRACEBACK_LINES_TO_LOG - ) + traceback_str = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) await proxy_logging_obj.post_call_failure_hook( request_data=request_data, original_exception=error, @@ -948,7 +953,8 @@ class LiteLLM_Proxy_MCP_Handler: mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events tool_server_map=tool_server_map, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, - user_api_key_auth=kwargs.get("user_api_key_auth") or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"), + user_api_key_auth=kwargs.get("user_api_key_auth") + or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"), original_request_params=request_params, ) diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index 53f39164e91..731aa5c692b 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -273,9 +273,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.finished = False # Event queues and generation flags - self.mcp_discovery_events: List[ResponsesAPIStreamingResponse] = ( - mcp_events # Pre-generated MCP discovery events - ) + self.mcp_discovery_events: List[ + ResponsesAPIStreamingResponse + ] = mcp_events # Pre-generated MCP discovery events self.tool_execution_events: List[ResponsesAPIStreamingResponse] = [] self.mcp_discovery_generated = True # Events are already generated self.mcp_events = ( @@ -284,9 +284,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.tool_server_map = tool_server_map # Iterator references - self.base_iterator: Optional[Union[Any, ResponsesAPIResponse]] = ( - base_iterator # Will be created when needed - ) + self.base_iterator: Optional[ + Union[Any, ResponsesAPIResponse] + ] = base_iterator # Will be created when needed self.follow_up_iterator: Optional[Any] = None # Response collection for tool execution @@ -305,7 +305,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): # Mark as async iterator self.is_async = True - + def _extract_mcp_headers_from_params(self) -> None: """Extract MCP headers from original request params to pass to tool calls""" from typing import Dict, Optional @@ -313,25 +313,31 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) - + # Extract headers from secret_fields in original_request_params raw_headers_from_request: Optional[Dict[str, str]] = None secret_fields = self.original_request_params.get("secret_fields") if secret_fields and isinstance(secret_fields, dict): raw_headers_from_request = secret_fields.get("raw_headers") - + # Extract MCP-specific headers self.mcp_auth_header: Optional[str] = None self.mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None self.oauth2_headers: Optional[Dict[str, str]] = None self.raw_headers: Optional[Dict[str, str]] = raw_headers_from_request - + if raw_headers_from_request: headers_obj = Headers(raw_headers_from_request) - self.mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers_obj) - self.mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers_obj) - self.oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers_obj) - + self.mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers( + headers_obj + ) + self.mcp_server_auth_headers = ( + MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers_obj) + ) + self.oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers( + headers_obj + ) + # Also check if headers are provided in tools array (from request body) tools = self.original_request_params.get("tools") if tools: @@ -341,17 +347,26 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): if tool_headers and isinstance(tool_headers, dict): # Merge tool headers into mcp_server_auth_headers headers_obj_from_tool = Headers(tool_headers) - tool_mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers_obj_from_tool) - + tool_mcp_server_auth_headers = ( + MCPRequestHandler._get_mcp_server_auth_headers_from_headers( + headers_obj_from_tool + ) + ) + if tool_mcp_server_auth_headers: if self.mcp_server_auth_headers is None: self.mcp_server_auth_headers = {} # Merge the headers from tool into existing headers - for server_alias, headers_dict in tool_mcp_server_auth_headers.items(): + for ( + server_alias, + headers_dict, + ) in tool_mcp_server_auth_headers.items(): if server_alias not in self.mcp_server_auth_headers: self.mcp_server_auth_headers[server_alias] = {} - self.mcp_server_auth_headers[server_alias].update(headers_dict) - + self.mcp_server_auth_headers[server_alias].update( + headers_dict + ) + # Also merge raw headers if self.raw_headers is None: self.raw_headers = {} @@ -489,9 +504,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): # Use the pre-fetched all_tools from original_request_params (no re-processing needed) params_for_llm = {} for key, value in params.items(): - params_for_llm[key] = ( - value # Copy all params as-is since tools are already processed - ) + params_for_llm[ + key + ] = value # Copy all params as-is since tools are already processed tools_count = ( len(params_for_llm.get("tools", [])) @@ -545,9 +560,11 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): return for tool_call in tool_calls: - tool_name, tool_arguments, tool_call_id = ( - LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call) - ) + ( + tool_name, + tool_arguments, + tool_call_id, + ) = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call) if tool_name and tool_call_id: # Create MCP call events for this tool execution call_events = create_mcp_call_events( diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 324380db2b4..cac2fe85541 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -63,11 +63,11 @@ def _generate_id(): # private helper function return "chatcmpl-" + str(uuid.uuid4()) - class SafeAttributeModel: """ A base model that provides safe attribute access. """ + def __delattr__(self, name): try: super().__delattr__(name) @@ -125,13 +125,14 @@ class SearchContextCostPerQuery(TypedDict, total=False): class AgenticLoopParams(TypedDict, total=False): """ Parameters passed to agentic loop hooks (e.g., WebSearch interception). - + Stored in logging_obj.model_call_details["agentic_loop_params"] to provide agentic hooks with the original request context needed for follow-up calls. """ + model: str """The model string with provider prefix (e.g., 'bedrock/invoke/...')""" - + custom_llm_provider: str """The LLM provider name (e.g., 'bedrock', 'anthropic')""" @@ -1345,8 +1346,7 @@ class CacheCreationTokenDetails(BaseModel): class PromptTokensDetailsWrapper( - SafeAttributeModel, - PromptTokensDetails + SafeAttributeModel, PromptTokensDetails ): # extends with image generation fields (text_tokens, image_tokens) text_tokens: Optional[int] = None """Text tokens sent to the model.""" From 165081059a5405d849f1daf01b976237f43a21bf Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Wed, 21 Jan 2026 14:57:14 +0900 Subject: [PATCH 14/29] chore: ruff fix --- litellm/proxy/_experimental/mcp_server/server.py | 2 +- litellm/responses/mcp/litellm_proxy_mcp_handler.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 4e5c73be806..03652ae155e 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -738,7 +738,7 @@ if MCP_AVAILABLE: return server_auth_header, extra_headers - async def _get_tools_from_mcp_servers( + async def _get_tools_from_mcp_servers( # noqa: PLR0915 user_api_key_auth: Optional[UserAPIKeyAuth], mcp_auth_header: Optional[str], mcp_servers: Optional[List[str]], diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 7cc821f5e02..2419126fe2f 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -479,7 +479,7 @@ class LiteLLM_Proxy_MCP_Handler: return result_text or "Tool executed successfully" @staticmethod - async def _execute_tool_calls( + async def _execute_tool_calls( # noqa: PLR0915 tool_server_map: dict[str, str], tool_calls: List[Any], user_api_key_auth: Any, From a312ef31e39af511688a4d49101edacb519e14d6 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Wed, 21 Jan 2026 15:04:24 +0900 Subject: [PATCH 15/29] chore: fix mypy --- litellm/responses/mcp/litellm_proxy_mcp_handler.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 2419126fe2f..4376e076a95 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -35,7 +35,9 @@ else: # NOTE: We intentionally keep ToolParam as a broad type here to avoid tight coupling # to optional OpenAI SDK typing symbols in environments that may not have them available. -ToolParam = Dict[str, Any] +# `Any` is used to keep mypy compatible with the broader OpenAI tool union types +# passed around in Responses API while still allowing dict-style access at runtime. +ToolParam = Any LITELLM_PROXY_MCP_SERVER_URL = "litellm_proxy" LITELLM_PROXY_MCP_SERVER_URL_PREFIX = f"{LITELLM_PROXY_MCP_SERVER_URL}/mcp/" From b5a7d2ab34740619b1f5de52663e2617015bcf36 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 21 Jan 2026 10:44:18 -0800 Subject: [PATCH 16/29] Paginating model/info endpoint --- litellm/proxy/proxy_server.py | 31 ++- tests/test_litellm/proxy/test_proxy_server.py | 225 ++++++++++++++++++ 2 files changed, 254 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index eef2af89799..c4f2b45f388 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7522,6 +7522,8 @@ async def model_info_v2( False, description="Return all models across all teams user is in." ), debug: Optional[bool] = False, + page: int = Query(1, description="Page number", ge=1), + size: int = Query(50, description="Page size", ge=1), ): """ BETA ENDPOINT. Might change unexpectedly. Use `/v1/model/info` for now. @@ -7530,7 +7532,13 @@ async def model_info_v2( # Return empty data array when no models are configured (graceful handling for fresh installs) if llm_router is None or not llm_router.model_list: - return {"data": []} + return { + "data": [], + "total_count": 0, + "current_page": page, + "total_pages": 0, + "size": size, + } if prisma_client is None: raise HTTPException( @@ -7620,7 +7628,26 @@ async def model_info_v2( ) verbose_proxy_logger.debug("all_models: %s", all_models) - return {"data": all_models} + + total_count = len(all_models) + + skip = (page - 1) * size + + total_pages = -(-total_count // size) if total_count > 0 else 0 + + paginated_models = all_models[skip : skip + size] + + verbose_proxy_logger.debug( + f"Pagination: skip={skip}, take={size}, total_count={total_count}, total_pages={total_pages}" + ) + + return { + "data": paginated_models, + "total_count": total_count, + "current_page": page, + "total_pages": total_pages, + "size": size, + } @router.get( diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 751a9033871..d8970a76a9a 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3216,3 +3216,228 @@ async def test_get_hierarchical_router_settings(): prisma_client=mock_prisma_client, ) assert result is None + + +@pytest.mark.asyncio +async def test_model_info_v2_pagination_basic(monkeypatch): + """ + Test basic pagination functionality for /v2/model/info endpoint. + Tests multiple pages with different page sizes. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth + + # Create 75 mock models for testing pagination + mock_models = [ + { + "model_name": f"model-{i}", + "litellm_params": {"model": f"gpt-{i}"}, + "model_info": {"id": f"model-{i}"}, + } + for i in range(1, 76) # 75 models total + ] + + # Mock llm_router + mock_router = MagicMock() + mock_router.model_list = mock_models + + # Mock prisma_client + mock_prisma_client = MagicMock() + + # Mock proxy_config.get_config + mock_get_config = AsyncMock(return_value={}) + + # Mock user authentication + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key_dict.user_id = "test-user" + mock_user_api_key_dict.api_key = "test-key" + mock_user_api_key_dict.team_models = [] + mock_user_api_key_dict.models = [] + + # Apply monkeypatches + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) + + # Override auth dependency + original_overrides = app.dependency_overrides.copy() + app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict + + client = TestClient(app) + try: + # Test page 1 with size 25 (should return models 1-25) + response = client.get("/v2/model/info", params={"page": 1, "size": 25}) + assert response.status_code == 200 + data = response.json() + assert data["total_count"] == 75 + assert data["current_page"] == 1 + assert data["size"] == 25 + assert data["total_pages"] == 3 # ceil(75/25) = 3 + assert len(data["data"]) == 25 + assert data["data"][0]["model_name"] == "model-1" + assert data["data"][24]["model_name"] == "model-25" + + # Test page 2 with size 25 (should return models 26-50) + response = client.get("/v2/model/info", params={"page": 2, "size": 25}) + assert response.status_code == 200 + data = response.json() + assert data["total_count"] == 75 + assert data["current_page"] == 2 + assert data["size"] == 25 + assert data["total_pages"] == 3 + assert len(data["data"]) == 25 + assert data["data"][0]["model_name"] == "model-26" + assert data["data"][24]["model_name"] == "model-50" + + # Test page 3 with size 25 (should return models 51-75) + response = client.get("/v2/model/info", params={"page": 3, "size": 25}) + assert response.status_code == 200 + data = response.json() + assert data["total_count"] == 75 + assert data["current_page"] == 3 + assert data["size"] == 25 + assert data["total_pages"] == 3 + assert len(data["data"]) == 25 + assert data["data"][0]["model_name"] == "model-51" + assert data["data"][24]["model_name"] == "model-75" + + # Test different page size (size 10) + response = client.get("/v2/model/info", params={"page": 1, "size": 10}) + assert response.status_code == 200 + data = response.json() + assert data["total_count"] == 75 + assert data["current_page"] == 1 + assert data["size"] == 10 + assert data["total_pages"] == 8 # ceil(75/10) = 8 + assert len(data["data"]) == 10 + + finally: + app.dependency_overrides = original_overrides + + +@pytest.mark.asyncio +async def test_model_info_v2_pagination_edge_cases(monkeypatch): + """ + Test edge cases for pagination in /v2/model/info endpoint. + Tests empty results, last page with partial results, and boundary conditions. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth + + # Mock prisma_client + mock_prisma_client = MagicMock() + + # Mock user authentication + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key_dict.user_id = "test-user" + mock_user_api_key_dict.api_key = "test-key" + mock_user_api_key_dict.team_models = [] + mock_user_api_key_dict.models = [] + + # Mock proxy_config.get_config + mock_get_config = AsyncMock(return_value={}) + + # Apply monkeypatches + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) + + # Override auth dependency + original_overrides = app.dependency_overrides.copy() + app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict + + client = TestClient(app) + try: + # Test Case 1: Empty model list (no models configured) + mock_router_empty = MagicMock() + mock_router_empty.model_list = [] + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router_empty) + + response = client.get("/v2/model/info", params={"page": 1, "size": 25}) + assert response.status_code == 200 + data = response.json() + assert data["total_count"] == 0 + assert data["current_page"] == 1 + assert data["size"] == 25 + assert data["total_pages"] == 0 + assert len(data["data"]) == 0 + + # Test Case 2: Last page with partial results (23 models, page size 10) + mock_models_partial = [ + { + "model_name": f"model-{i}", + "litellm_params": {"model": f"gpt-{i}"}, + "model_info": {"id": f"model-{i}"}, + } + for i in range(1, 24) # 23 models total + ] + mock_router_partial = MagicMock() + mock_router_partial.model_list = mock_models_partial + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router_partial) + + # Page 1 should have 10 models + response = client.get("/v2/model/info", params={"page": 1, "size": 10}) + assert response.status_code == 200 + data = response.json() + assert data["total_count"] == 23 + assert data["current_page"] == 1 + assert data["total_pages"] == 3 # ceil(23/10) = 3 + assert len(data["data"]) == 10 + + # Page 2 should have 10 models + response = client.get("/v2/model/info", params={"page": 2, "size": 10}) + assert response.status_code == 200 + data = response.json() + assert data["total_count"] == 23 + assert data["current_page"] == 2 + assert data["total_pages"] == 3 + assert len(data["data"]) == 10 + + # Page 3 (last page) should have only 3 models + response = client.get("/v2/model/info", params={"page": 3, "size": 10}) + assert response.status_code == 200 + data = response.json() + assert data["total_count"] == 23 + assert data["current_page"] == 3 + assert data["total_pages"] == 3 + assert len(data["data"]) == 3 + assert data["data"][0]["model_name"] == "model-21" + assert data["data"][2]["model_name"] == "model-23" + + # Test Case 3: Page beyond available pages (should return empty data) + response = client.get("/v2/model/info", params={"page": 4, "size": 10}) + assert response.status_code == 200 + data = response.json() + assert data["total_count"] == 23 + assert data["current_page"] == 4 + assert data["total_pages"] == 3 + assert len(data["data"]) == 0 # No data for page beyond total_pages + + # Test Case 4: Single model with page size 1 + mock_models_single = [ + { + "model_name": "single-model", + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "single-model"}, + } + ] + mock_router_single = MagicMock() + mock_router_single.model_list = mock_models_single + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router_single) + + response = client.get("/v2/model/info", params={"page": 1, "size": 1}) + assert response.status_code == 200 + data = response.json() + assert data["total_count"] == 1 + assert data["current_page"] == 1 + assert data["total_pages"] == 1 + assert len(data["data"]) == 1 + assert data["data"][0]["model_name"] == "single-model" + + finally: + app.dependency_overrides = original_overrides From d0e35751a11db388b49a1c8f85ad9b6bb2f123f6 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 21 Jan 2026 11:02:39 -0800 Subject: [PATCH 17/29] Fixing tests and linting --- litellm/proxy/proxy_server.py | 126 +++++++----- .../proxy/test_empty_model_list.py | 58 +++++- tests/test_litellm/proxy/test_proxy_server.py | 184 ++++++++++++++++++ 3 files changed, 312 insertions(+), 56 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index c4f2b45f388..d168290e8a1 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7503,6 +7503,77 @@ async def get_all_team_and_direct_access_models( return all_models +def _enrich_model_info_with_litellm_data( + model: Dict[str, Any], debug: bool = False, llm_router: Optional[Router] = None +) -> Dict[str, Any]: + """ + Enrich a model dictionary with litellm model info (pricing, context window, etc.) + and remove sensitive information. + + Args: + model: Model dictionary to enrich + debug: Whether to include debug information like openai_client + llm_router: Optional router instance for debug info + + Returns: + Enriched model dictionary with sensitive info removed + """ + # provided model_info in config.yaml + model_info = model.get("model_info", {}) + if debug is True: + _openai_client = "None" + if llm_router is not None: + _openai_client = ( + llm_router._get_client( + deployment=model, kwargs={}, client_type="async" + ) + or "None" + ) + else: + _openai_client = "llm_router_is_None" + openai_client = str(_openai_client) + model["openai_client"] = openai_client + + # read litellm model_prices_and_context_window.json to get the following: + # input_cost_per_token, output_cost_per_token, max_tokens + litellm_model_info = get_litellm_model_info(model=model) + + # 2nd pass on the model, try seeing if we can find model in litellm model_cost map + if litellm_model_info == {}: + # use litellm_param model_name to get model_info + litellm_params = model.get("litellm_params", {}) + litellm_model = litellm_params.get("model", None) + try: + litellm_model_info = litellm.get_model_info(model=litellm_model) + except Exception: + litellm_model_info = {} + # 3rd pass on the model, try seeing if we can find model but without the "/" in model cost map + if litellm_model_info == {}: + # use litellm_param model_name to get model_info + litellm_params = model.get("litellm_params", {}) + litellm_model = litellm_params.get("model", None) + if litellm_model: + split_model = litellm_model.split("/") + if len(split_model) > 0: + litellm_model = split_model[-1] + try: + litellm_model_info = litellm.get_model_info( + model=litellm_model, custom_llm_provider=split_model[0] + ) + except Exception: + litellm_model_info = {} + for k, v in litellm_model_info.items(): + if k not in model_info: + model_info[k] = v + model["model_info"] = model_info + # don't return the api key / vertex credentials + # don't return the llm credentials + model = remove_sensitive_info_from_deployment( + model, excluded_keys={"litellm_credential_name"} + ) + return model + + @router.get( "/v2/model/info", description="v2 - returns models available to the user based on their API key permissions. Shows model info from config.yaml (except api key and api base). Filter to just user-added models with ?user_models_only=true", @@ -7573,58 +7644,9 @@ async def model_info_v2( all_models=all_models, ) # fill in model info based on config.yaml and litellm model_prices_and_context_window.json - for _model in all_models: - # provided model_info in config.yaml - model_info = _model.get("model_info", {}) - if debug is True: - _openai_client = "None" - if llm_router is not None: - _openai_client = ( - llm_router._get_client( - deployment=_model, kwargs={}, client_type="async" - ) - or "None" - ) - else: - _openai_client = "llm_router_is_None" - openai_client = str(_openai_client) - _model["openai_client"] = openai_client - - # read litellm model_prices_and_context_window.json to get the following: - # input_cost_per_token, output_cost_per_token, max_tokens - litellm_model_info = get_litellm_model_info(model=_model) - - # 2nd pass on the model, try seeing if we can find model in litellm model_cost map - if litellm_model_info == {}: - # use litellm_param model_name to get model_info - litellm_params = _model.get("litellm_params", {}) - litellm_model = litellm_params.get("model", None) - try: - litellm_model_info = litellm.get_model_info(model=litellm_model) - except Exception: - litellm_model_info = {} - # 3rd pass on the model, try seeing if we can find model but without the "/" in model cost map - if litellm_model_info == {}: - # use litellm_param model_name to get model_info - litellm_params = _model.get("litellm_params", {}) - litellm_model = litellm_params.get("model", None) - split_model = litellm_model.split("/") - if len(split_model) > 0: - litellm_model = split_model[-1] - try: - litellm_model_info = litellm.get_model_info( - model=litellm_model, custom_llm_provider=split_model[0] - ) - except Exception: - litellm_model_info = {} - for k, v in litellm_model_info.items(): - if k not in model_info: - model_info[k] = v - _model["model_info"] = model_info - # don't return the api key / vertex credentials - # don't return the llm credentials - _model = remove_sensitive_info_from_deployment( - _model, excluded_keys={"litellm_credential_name"} + for i, _model in enumerate(all_models): + all_models[i] = _enrich_model_info_with_litellm_data( + model=_model, debug=debug, llm_router=llm_router ) verbose_proxy_logger.debug("all_models: %s", all_models) diff --git a/tests/test_litellm/proxy/test_empty_model_list.py b/tests/test_litellm/proxy/test_empty_model_list.py index 6b3e59d3194..dd900d3eb53 100644 --- a/tests/test_litellm/proxy/test_empty_model_list.py +++ b/tests/test_litellm/proxy/test_empty_model_list.py @@ -32,7 +32,7 @@ class TestEmptyModelListHandling: self, client, monkeypatch ): """ - Test that /v2/model/info returns {"data": []} instead of 500 + Test that /v2/model/info returns paginated empty response instead of 500 when llm_router is None. """ monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) @@ -56,13 +56,18 @@ class TestEmptyModelListHandling: ) assert response.status_code == 200 - assert response.json() == {"data": []} + data = response.json() + assert data["data"] == [] + assert data["total_count"] == 0 + assert data["current_page"] == 1 + assert data["total_pages"] == 0 + assert data["size"] == 50 # default page size def test_v2_model_info_returns_empty_data_when_model_list_empty( self, client, monkeypatch ): """ - Test that /v2/model/info returns {"data": []} instead of 500 + Test that /v2/model/info returns paginated empty response instead of 500 when llm_router exists but model_list is empty. """ mock_router = MagicMock() @@ -89,7 +94,52 @@ class TestEmptyModelListHandling: ) assert response.status_code == 200 - assert response.json() == {"data": []} + data = response.json() + assert data["data"] == [] + assert data["total_count"] == 0 + assert data["current_page"] == 1 + assert data["total_pages"] == 0 + assert data["size"] == 50 # default page size + + def test_v2_model_info_pagination_with_empty_results( + self, client, monkeypatch + ): + """ + Test that /v2/model/info pagination parameters work correctly + when there are no models (empty results). + """ + mock_router = MagicMock() + mock_router.model_list = [] + + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_model_list", []) + monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + + with patch( + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth", + return_value=MagicMock( + user_id="test-user", + team_id=None, + team_models=[], + models=[], + user_role="proxy_admin", + ), + ): + # Test with custom pagination parameters + response = client.get( + "/v2/model/info", + params={"page": 2, "size": 25}, + headers={"Authorization": "Bearer sk-test"}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["data"] == [] + assert data["total_count"] == 0 + assert data["current_page"] == 2 # Should respect the page parameter + assert data["total_pages"] == 0 + assert data["size"] == 25 # Should respect the size parameter def test_model_group_info_returns_empty_data_when_model_list_none( self, client, monkeypatch diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index d8970a76a9a..f1854380efe 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3441,3 +3441,187 @@ async def test_model_info_v2_pagination_edge_cases(monkeypatch): finally: app.dependency_overrides = original_overrides + + +def test_enrich_model_info_with_litellm_data(): + """ + Test the _enrich_model_info_with_litellm_data helper function. + Tests model info enrichment, debug mode, and sensitive info removal. + """ + from unittest.mock import MagicMock, patch + + from litellm.proxy.proxy_server import _enrich_model_info_with_litellm_data + + # Test Case 1: Basic model enrichment without debug + model = { + "model_name": "test-model", + "litellm_params": {"model": "gpt-3.5-turbo"}, + "model_info": {"id": "test-model"}, + "api_key": "sk-secret-key", # Should be removed + } + + with patch("litellm.proxy.proxy_server.get_litellm_model_info") as mock_get_info, patch( + "litellm.proxy.proxy_server.remove_sensitive_info_from_deployment" + ) as mock_remove_sensitive: + mock_get_info.return_value = { + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + "max_tokens": 4096, + } + mock_remove_sensitive.return_value = { + "model_name": "test-model", + "litellm_params": {"model": "gpt-3.5-turbo"}, + "model_info": { + "id": "test-model", + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + "max_tokens": 4096, + }, + } + + result = _enrich_model_info_with_litellm_data(model=model, debug=False) + + # Verify get_litellm_model_info was called + mock_get_info.assert_called_once_with(model=model) + # Verify remove_sensitive_info_from_deployment was called + mock_remove_sensitive.assert_called_once() + # Verify result doesn't have api_key + assert "api_key" not in result + # Verify model_info was enriched + assert "input_cost_per_token" in result["model_info"] + + # Test Case 2: Model enrichment with debug mode + model_with_debug = { + "model_name": "test-model-debug", + "litellm_params": {"model": "gpt-4"}, + "model_info": {}, + } + + mock_router = MagicMock() + mock_client = MagicMock() + mock_router._get_client.return_value = mock_client + + with patch("litellm.proxy.proxy_server.get_litellm_model_info") as mock_get_info, patch( + "litellm.proxy.proxy_server.remove_sensitive_info_from_deployment" + ) as mock_remove_sensitive: + mock_get_info.return_value = {} + mock_remove_sensitive.return_value = { + "model_name": "test-model-debug", + "litellm_params": {"model": "gpt-4"}, + "model_info": {}, + "openai_client": str(mock_client), + } + + result = _enrich_model_info_with_litellm_data( + model=model_with_debug, debug=True, llm_router=mock_router + ) + + # Verify debug info was added + mock_remove_sensitive.assert_called_once() + call_args = mock_remove_sensitive.call_args[0][0] + assert "openai_client" in call_args + # Verify router._get_client was called for debug + mock_router._get_client.assert_called_once() + + # Test Case 3: Model with fallback to litellm.get_model_info + model_fallback = { + "model_name": "test-model-fallback", + "litellm_params": {"model": "claude-3-opus"}, + "model_info": {}, + } + + with patch("litellm.proxy.proxy_server.get_litellm_model_info") as mock_get_info, patch( + "litellm.get_model_info" + ) as mock_litellm_info, patch( + "litellm.proxy.proxy_server.remove_sensitive_info_from_deployment" + ) as mock_remove_sensitive: + # First call returns empty, triggering fallback + mock_get_info.return_value = {} + mock_litellm_info.return_value = { + "input_cost_per_token": 0.015, + "output_cost_per_token": 0.075, + "max_tokens": 200000, + } + mock_remove_sensitive.return_value = { + "model_name": "test-model-fallback", + "litellm_params": {"model": "claude-3-opus"}, + "model_info": { + "input_cost_per_token": 0.015, + "output_cost_per_token": 0.075, + "max_tokens": 200000, + }, + } + + result = _enrich_model_info_with_litellm_data(model=model_fallback, debug=False) + + # Verify fallback was attempted + mock_litellm_info.assert_called_once_with(model="claude-3-opus") + # Verify model_info was enriched with fallback data + call_args = mock_remove_sensitive.call_args[0][0] + assert call_args["model_info"]["input_cost_per_token"] == 0.015 + + # Test Case 4: Model with split model name fallback + model_split = { + "model_name": "test-model-split", + "litellm_params": {"model": "azure/gpt-4"}, + "model_info": {}, + } + + with patch("litellm.proxy.proxy_server.get_litellm_model_info") as mock_get_info, patch( + "litellm.get_model_info" + ) as mock_litellm_info, patch( + "litellm.proxy.proxy_server.remove_sensitive_info_from_deployment" + ) as mock_remove_sensitive: + # Both first and second pass return empty, triggering third pass + mock_get_info.return_value = {} + # Second pass (no split) + mock_litellm_info.side_effect = [ + {}, # First call returns empty + {"max_tokens": 8192}, # Third pass with split succeeds + ] + mock_remove_sensitive.return_value = { + "model_name": "test-model-split", + "litellm_params": {"model": "azure/gpt-4"}, + "model_info": {"max_tokens": 8192}, + } + + result = _enrich_model_info_with_litellm_data(model=model_split, debug=False) + + # Verify third pass was attempted with split model name + assert mock_litellm_info.call_count == 2 + # Check that second call used split model name + second_call = mock_litellm_info.call_args_list[1] + assert second_call[1]["model"] == "gpt-4" + assert second_call[1]["custom_llm_provider"] == "azure" + + # Test Case 5: Model with existing model_info (should preserve existing keys) + model_existing = { + "model_name": "test-model-existing", + "litellm_params": {"model": "gpt-3.5-turbo"}, + "model_info": {"id": "existing-id", "custom_key": "custom_value"}, + } + + with patch("litellm.proxy.proxy_server.get_litellm_model_info") as mock_get_info, patch( + "litellm.proxy.proxy_server.remove_sensitive_info_from_deployment" + ) as mock_remove_sensitive: + mock_get_info.return_value = { + "input_cost_per_token": 0.001, + "id": "new-id", # Should not override existing "id" + } + mock_remove_sensitive.return_value = { + "model_name": "test-model-existing", + "litellm_params": {"model": "gpt-3.5-turbo"}, + "model_info": { + "id": "existing-id", # Existing key preserved + "custom_key": "custom_value", # Existing key preserved + "input_cost_per_token": 0.001, # New key added + }, + } + + result = _enrich_model_info_with_litellm_data(model=model_existing, debug=False) + + # Verify existing keys are preserved + call_args = mock_remove_sensitive.call_args[0][0] + assert call_args["model_info"]["id"] == "existing-id" + assert call_args["model_info"]["custom_key"] == "custom_value" + assert call_args["model_info"]["input_cost_per_token"] == 0.001 From 3075b0e5a25462b76824f5236b0416c39c10de38 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 21 Jan 2026 11:22:20 -0800 Subject: [PATCH 18/29] fixing mypy linting --- litellm/proxy/proxy_server.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index d168290e8a1..8ae07121177 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7646,7 +7646,7 @@ async def model_info_v2( # fill in model info based on config.yaml and litellm model_prices_and_context_window.json for i, _model in enumerate(all_models): all_models[i] = _enrich_model_info_with_litellm_data( - model=_model, debug=debug, llm_router=llm_router + model=_model, debug=debug if debug is not None else False, llm_router=llm_router ) verbose_proxy_logger.debug("all_models: %s", all_models) From 7cf80a928335207c00b96c40adb7de99ee9fbcb0 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 21 Jan 2026 11:53:15 -0800 Subject: [PATCH 19/29] Paginate All Models Tab --- .../hooks/models/useModels.test.ts | 627 ++++++++++++++++++ .../app/(dashboard)/hooks/models/useModels.ts | 16 +- .../ModelsAndEndpointsView.tsx | 3 +- .../components/AllModelsTab.test.tsx | 348 ++++++---- .../components/AllModelsTab.tsx | 94 +-- .../src/components/networking.tsx | 6 +- 6 files changed, 937 insertions(+), 157 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts new file mode 100644 index 00000000000..a97c309ca91 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts @@ -0,0 +1,627 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { renderHook, waitFor } from "@testing-library/react"; +import React, { ReactNode } from "react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { + useModelsInfo, + useModelHub, + useAllProxyModels, + useSelectedTeamModels, + type ProxyModel, + type AllProxyModelsResponse, + type PaginatedModelInfoResponse, +} from "./useModels"; + +vi.mock("@/components/networking", () => ({ + modelInfoCall: vi.fn(), + modelHubCall: vi.fn(), + modelAvailableCall: vi.fn(), +})); + +const mockUseAuthorized = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: () => mockUseAuthorized(), +})); + +import { modelInfoCall, modelHubCall, modelAvailableCall } from "@/components/networking"; + +const mockProxyModel: ProxyModel = { + id: "model-1", + object: "model", + created: 1234567890, + owned_by: "openai", +}; + +const mockPaginatedModelInfoResponse: PaginatedModelInfoResponse = { + data: [{ id: "model-1", name: "Test Model" }], + total_count: 1, + current_page: 1, + total_pages: 1, + size: 50, +}; + +const mockAllProxyModelsResponse: AllProxyModelsResponse = { + data: [mockProxyModel], +}; + +describe("useModelsInfo", () => { + let queryClient: QueryClient; + + beforeEach(() => { + queryClient = new QueryClient({ + defaultOptions: { + queries: { + retry: false, + }, + }, + }); + + vi.clearAllMocks(); + + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId: "test-user-id", + userRole: "Admin", + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + }); + + const wrapper = ({ children }: { children: ReactNode }) => + React.createElement(QueryClientProvider, { client: queryClient }, children); + + it("should render without crashing", () => { + (modelInfoCall as any).mockResolvedValue(mockPaginatedModelInfoResponse); + + const { result } = renderHook(() => useModelsInfo(), { wrapper }); + + expect(result.current).toBeDefined(); + }); + + it("should return models data when query is successful", async () => { + (modelInfoCall as any).mockResolvedValue(mockPaginatedModelInfoResponse); + + const { result } = renderHook(() => useModelsInfo(), { wrapper }); + + expect(result.current.isLoading).toBe(true); + expect(result.current.data).toBeUndefined(); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + expect(result.current.isSuccess).toBe(true); + }); + + expect(result.current.data).toEqual(mockPaginatedModelInfoResponse); + expect(result.current.error).toBeNull(); + expect(modelInfoCall).toHaveBeenCalledWith( + "test-access-token", + "test-user-id", + "Admin", + 1, + 50 + ); + expect(modelInfoCall).toHaveBeenCalledTimes(1); + }); + + it("should use custom page and size parameters", async () => { + (modelInfoCall as any).mockResolvedValue(mockPaginatedModelInfoResponse); + + const { result } = renderHook(() => useModelsInfo(2, 25), { wrapper }); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + }); + + expect(modelInfoCall).toHaveBeenCalledWith( + "test-access-token", + "test-user-id", + "Admin", + 2, + 25 + ); + }); + + it("should handle error when modelInfoCall fails", async () => { + const errorMessage = "Failed to fetch models"; + const testError = new Error(errorMessage); + + (modelInfoCall as any).mockRejectedValue(testError); + + const { result } = renderHook(() => useModelsInfo(), { wrapper }); + + expect(result.current.isLoading).toBe(true); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + expect(result.current.isError).toBe(true); + }); + + expect(result.current.error).toEqual(testError); + expect(result.current.data).toBeUndefined(); + expect(modelInfoCall).toHaveBeenCalledTimes(1); + }); + + it("should not execute query when accessToken is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: null, + userId: "test-user-id", + userRole: "Admin", + token: null, + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useModelsInfo(), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelInfoCall).not.toHaveBeenCalled(); + }); + + it("should not execute query when userId is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId: null, + userRole: "Admin", + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useModelsInfo(), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelInfoCall).not.toHaveBeenCalled(); + }); + + it("should not execute query when userRole is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId: "test-user-id", + userRole: null, + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useModelsInfo(), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelInfoCall).not.toHaveBeenCalled(); + }); + + it("should not execute query when all required auth values are missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: null, + userId: null, + userRole: null, + token: null, + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useModelsInfo(), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelInfoCall).not.toHaveBeenCalled(); + }); +}); + +describe("useModelHub", () => { + let queryClient: QueryClient; + + beforeEach(() => { + queryClient = new QueryClient({ + defaultOptions: { + queries: { + retry: false, + }, + }, + }); + + vi.clearAllMocks(); + + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId: "test-user-id", + userRole: "Admin", + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + }); + + const wrapper = ({ children }: { children: ReactNode }) => + React.createElement(QueryClientProvider, { client: queryClient }, children); + + it("should render without crashing", () => { + (modelHubCall as any).mockResolvedValue({ data: [] }); + + const { result } = renderHook(() => useModelHub(), { wrapper }); + + expect(result.current).toBeDefined(); + }); + + it("should return model hub data when query is successful", async () => { + const mockHubData = { data: [{ id: "hub-1", name: "Test Hub" }] }; + (modelHubCall as any).mockResolvedValue(mockHubData); + + const { result } = renderHook(() => useModelHub(), { wrapper }); + + expect(result.current.isLoading).toBe(true); + expect(result.current.data).toBeUndefined(); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + expect(result.current.isSuccess).toBe(true); + }); + + expect(result.current.data).toEqual(mockHubData); + expect(result.current.error).toBeNull(); + expect(modelHubCall).toHaveBeenCalledWith("test-access-token"); + expect(modelHubCall).toHaveBeenCalledTimes(1); + }); + + it("should handle error when modelHubCall fails", async () => { + const errorMessage = "Failed to fetch model hub"; + const testError = new Error(errorMessage); + + (modelHubCall as any).mockRejectedValue(testError); + + const { result } = renderHook(() => useModelHub(), { wrapper }); + + expect(result.current.isLoading).toBe(true); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + expect(result.current.isError).toBe(true); + }); + + expect(result.current.error).toEqual(testError); + expect(result.current.data).toBeUndefined(); + expect(modelHubCall).toHaveBeenCalledTimes(1); + }); + + it("should not execute query when accessToken is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: null, + userId: "test-user-id", + userRole: "Admin", + token: null, + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useModelHub(), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelHubCall).not.toHaveBeenCalled(); + }); +}); + +describe("useAllProxyModels", () => { + let queryClient: QueryClient; + + beforeEach(() => { + queryClient = new QueryClient({ + defaultOptions: { + queries: { + retry: false, + }, + }, + }); + + vi.clearAllMocks(); + + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId: "test-user-id", + userRole: "Admin", + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + }); + + const wrapper = ({ children }: { children: ReactNode }) => + React.createElement(QueryClientProvider, { client: queryClient }, children); + + it("should render without crashing", () => { + (modelAvailableCall as any).mockResolvedValue(mockAllProxyModelsResponse); + + const { result } = renderHook(() => useAllProxyModels(), { wrapper }); + + expect(result.current).toBeDefined(); + }); + + it("should return all proxy models data when query is successful", async () => { + (modelAvailableCall as any).mockResolvedValue(mockAllProxyModelsResponse); + + const { result } = renderHook(() => useAllProxyModels(), { wrapper }); + + expect(result.current.isLoading).toBe(true); + expect(result.current.data).toBeUndefined(); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + expect(result.current.isSuccess).toBe(true); + }); + + expect(result.current.data).toEqual(mockAllProxyModelsResponse); + expect(result.current.error).toBeNull(); + expect(modelAvailableCall).toHaveBeenCalledWith( + "test-access-token", + "test-user-id", + "Admin", + true + ); + expect(modelAvailableCall).toHaveBeenCalledTimes(1); + }); + + it("should handle error when modelAvailableCall fails", async () => { + const errorMessage = "Failed to fetch proxy models"; + const testError = new Error(errorMessage); + + (modelAvailableCall as any).mockRejectedValue(testError); + + const { result } = renderHook(() => useAllProxyModels(), { wrapper }); + + expect(result.current.isLoading).toBe(true); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + expect(result.current.isError).toBe(true); + }); + + expect(result.current.error).toEqual(testError); + expect(result.current.data).toBeUndefined(); + expect(modelAvailableCall).toHaveBeenCalledTimes(1); + }); + + it("should not execute query when accessToken is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: null, + userId: "test-user-id", + userRole: "Admin", + token: null, + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useAllProxyModels(), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelAvailableCall).not.toHaveBeenCalled(); + }); + + it("should not execute query when userId is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId: null, + userRole: "Admin", + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useAllProxyModels(), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelAvailableCall).not.toHaveBeenCalled(); + }); + + it("should not execute query when userRole is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId: "test-user-id", + userRole: null, + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useAllProxyModels(), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelAvailableCall).not.toHaveBeenCalled(); + }); +}); + +describe("useSelectedTeamModels", () => { + let queryClient: QueryClient; + + beforeEach(() => { + queryClient = new QueryClient({ + defaultOptions: { + queries: { + retry: false, + }, + }, + }); + + vi.clearAllMocks(); + + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId: "test-user-id", + userRole: "Admin", + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + }); + + const wrapper = ({ children }: { children: ReactNode }) => + React.createElement(QueryClientProvider, { client: queryClient }, children); + + it("should render without crashing", () => { + (modelAvailableCall as any).mockResolvedValue(mockAllProxyModelsResponse); + + const { result } = renderHook(() => useSelectedTeamModels("team-1"), { wrapper }); + + expect(result.current).toBeDefined(); + }); + + it("should return team models data when query is successful", async () => { + (modelAvailableCall as any).mockResolvedValue(mockAllProxyModelsResponse); + + const { result } = renderHook(() => useSelectedTeamModels("team-1"), { wrapper }); + + expect(result.current.isLoading).toBe(true); + expect(result.current.data).toBeUndefined(); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + expect(result.current.isSuccess).toBe(true); + }); + + expect(result.current.data).toEqual(mockAllProxyModelsResponse); + expect(result.current.error).toBeNull(); + expect(modelAvailableCall).toHaveBeenCalledWith( + "test-access-token", + "test-user-id", + "Admin", + true, + "team-1" + ); + expect(modelAvailableCall).toHaveBeenCalledTimes(1); + }); + + it("should handle error when modelAvailableCall fails", async () => { + const errorMessage = "Failed to fetch team models"; + const testError = new Error(errorMessage); + + (modelAvailableCall as any).mockRejectedValue(testError); + + const { result } = renderHook(() => useSelectedTeamModels("team-1"), { wrapper }); + + expect(result.current.isLoading).toBe(true); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + expect(result.current.isError).toBe(true); + }); + + expect(result.current.error).toEqual(testError); + expect(result.current.data).toBeUndefined(); + expect(modelAvailableCall).toHaveBeenCalledTimes(1); + }); + + it("should not execute query when teamID is null", () => { + const { result } = renderHook(() => useSelectedTeamModels(null), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelAvailableCall).not.toHaveBeenCalled(); + }); + + it("should not execute query when accessToken is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: null, + userId: "test-user-id", + userRole: "Admin", + token: null, + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useSelectedTeamModels("team-1"), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelAvailableCall).not.toHaveBeenCalled(); + }); + + it("should not execute query when userId is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId: null, + userRole: "Admin", + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useSelectedTeamModels("team-1"), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelAvailableCall).not.toHaveBeenCalled(); + }); + + it("should not execute query when userRole is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId: "test-user-id", + userRole: null, + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useSelectedTeamModels("team-1"), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelAvailableCall).not.toHaveBeenCalled(); + }); + + it("should not execute query when teamID is missing and other auth values are present", () => { + const { result } = renderHook(() => useSelectedTeamModels(null), { wrapper }); + + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + expect(modelAvailableCall).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts index fa7ab911ecd..b9f67c1cac8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts @@ -14,21 +14,31 @@ export interface AllProxyModelsResponse { data: ProxyModel[]; } +export interface PaginatedModelInfoResponse { + data: any[]; + total_count: number; + current_page: number; + total_pages: number; + size: number; +} + const modelKeys = createQueryKeys("models"); const modelHubKeys = createQueryKeys("modelHub"); const allProxyModelsKeys = createQueryKeys("allProxyModels"); const selectedTeamModelsKeys = createQueryKeys("selectedTeamModels"); -export const useModelsInfo = () => { +export const useModelsInfo = (page: number = 1, size: number = 50) => { const { accessToken, userId, userRole } = useAuthorized(); - return useQuery({ + return useQuery({ queryKey: modelKeys.list({ filters: { ...(userId && { userId }), ...(userRole && { userRole }), + page, + size, }, }), - queryFn: async () => await modelInfoCall(accessToken!, userId!, userRole!), + queryFn: async () => await modelInfoCall(accessToken!, userId!, userRole!, page, size), enabled: Boolean(accessToken && userId && userRole), }); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx index 1cce704467a..f707ef04ffb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx @@ -94,7 +94,8 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te }, [modelDataResponse?.data]); const allModelsOnProxy = useMemo(() => { - return modelDataResponse?.data?.map((model: any) => model.model_name); + if (!modelDataResponse?.data) return []; + return modelDataResponse.data.map((model: any) => model.model_name); }, [modelDataResponse?.data]); const getProviderFromModel = (model: string) => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx index ae376701a9a..f20c3563c83 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx @@ -1,13 +1,18 @@ import * as useAuthorizedModule from "@/app/(dashboard)/hooks/useAuthorized"; import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import AllModelsTab from "./AllModelsTab"; // Mock the useModelsInfo hook -const mockUseModelsInfo = vi.fn(() => ({ data: { data: [] } })) as any; +const mockUseModelsInfo = vi.fn(() => ({ + data: { data: [], total_count: 0, current_page: 1, total_pages: 1, size: 50 }, + isLoading: false, + error: null, +})) as any; vi.mock("../../hooks/models/useModels", () => ({ - useModelsInfo: () => mockUseModelsInfo(), + useModelsInfo: (page?: number, size?: number) => mockUseModelsInfo(page, size), })); // Mock the useModelCostMap hook @@ -51,6 +56,21 @@ const createModelCostMapMock = (data: Record) => ({ error: null, }); +// Helper function to create paginated model data mock +const createPaginatedModelData = ( + models: any[], + totalCount: number = models.length, + currentPage: number = 1, + totalPages: number = 1, + size: number = 50, +) => ({ + data: models, + total_count: totalCount, + current_page: currentPage, + total_pages: totalPages, + size: size, +}); + describe("AllModelsTab", () => { const mockSetSelectedModelGroup = vi.fn(); const mockSetSelectedModelId = vi.fn(); @@ -84,7 +104,11 @@ describe("AllModelsTab", () => { }); it("should render with empty data", () => { - mockUseModelsInfo.mockReturnValueOnce({ data: { data: [] } }); + mockUseModelsInfo.mockReturnValueOnce({ + data: createPaginatedModelData([], 0, 1, 1, 50), + isLoading: false, + error: null, + }); mockUseTeams.mockReturnValueOnce({ data: [], @@ -130,28 +154,26 @@ describe("AllModelsTab", () => { }), ); - const modelData = { - data: [ - { - model_name: "gpt-4-accessible", - model_info: { - id: "model-1", - access_via_team_ids: ["team-456"], - access_groups: [], - }, + const modelData = createPaginatedModelData([ + { + model_name: "gpt-4-accessible", + model_info: { + id: "model-1", + access_via_team_ids: ["team-456"], + access_groups: [], }, - { - model_name: "gpt-3.5-turbo-blocked", - model_info: { - id: "model-2", - access_via_team_ids: ["team-789"], - access_groups: [], - }, + }, + { + model_name: "gpt-3.5-turbo-blocked", + model_info: { + id: "model-2", + access_via_team_ids: ["team-789"], + access_groups: [], }, - ], - }; + }, + ], 2, 1, 1, 50); - mockUseModelsInfo.mockReturnValue({ data: modelData }); + mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null }); render(); @@ -191,28 +213,26 @@ describe("AllModelsTab", () => { }), ); - const modelData = { - data: [ - { - model_name: "gpt-4-sales", - model_info: { - id: "model-sales-1", - access_via_team_ids: [], - access_groups: ["sales-model-group"], - }, + const modelData = createPaginatedModelData([ + { + model_name: "gpt-4-sales", + model_info: { + id: "model-sales-1", + access_via_team_ids: [], + access_groups: ["sales-model-group"], }, - { - model_name: "gpt-4-engineering", - model_info: { - id: "model-eng-1", - access_via_team_ids: [], - access_groups: ["engineering-model-group"], - }, + }, + { + model_name: "gpt-4-engineering", + model_info: { + id: "model-eng-1", + access_via_team_ids: [], + access_groups: ["engineering-model-group"], }, - ], - }; + }, + ], 2, 1, 1, 50); - mockUseModelsInfo.mockReturnValue({ data: modelData }); + mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null }); render(); @@ -236,30 +256,28 @@ describe("AllModelsTab", () => { }), ); - const modelData = { - data: [ - { - model_name: "gpt-4-personal", - model_info: { - id: "model-personal-1", - direct_access: true, - access_via_team_ids: [], - access_groups: [], - }, + const modelData = createPaginatedModelData([ + { + model_name: "gpt-4-personal", + model_info: { + id: "model-personal-1", + direct_access: true, + access_via_team_ids: [], + access_groups: [], }, - { - model_name: "gpt-4-team-only", - model_info: { - id: "model-team-1", - direct_access: false, - access_via_team_ids: ["team-123"], - access_groups: [], - }, + }, + { + model_name: "gpt-4-team-only", + model_info: { + id: "model-team-1", + direct_access: false, + access_via_team_ids: ["team-123"], + access_groups: [], }, - ], - }; + }, + ], 2, 1, 1, 50); - mockUseModelsInfo.mockReturnValue({ data: modelData }); + mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null }); render(); @@ -283,42 +301,40 @@ describe("AllModelsTab", () => { }), ); - const modelData = { - data: [ - { - model_name: "gpt-4-config", - litellm_model_name: "gpt-4-config", - provider: "openai", - model_info: { - id: "model-config-1", - db_model: false, - direct_access: true, - access_via_team_ids: [], - access_groups: [], - created_by: "user-123", - created_at: "2024-01-01", - updated_at: "2024-01-01", - }, + const modelData = createPaginatedModelData([ + { + model_name: "gpt-4-config", + litellm_model_name: "gpt-4-config", + provider: "openai", + model_info: { + id: "model-config-1", + db_model: false, + direct_access: true, + access_via_team_ids: [], + access_groups: [], + created_by: "user-123", + created_at: "2024-01-01", + updated_at: "2024-01-01", }, - { - model_name: "gpt-4-db", - litellm_model_name: "gpt-4-db", - provider: "openai", - model_info: { - id: "model-db-1", - db_model: true, - direct_access: true, - access_via_team_ids: [], - access_groups: [], - created_by: "user-123", - created_at: "2024-01-01", - updated_at: "2024-01-01", - }, + }, + { + model_name: "gpt-4-db", + litellm_model_name: "gpt-4-db", + provider: "openai", + model_info: { + id: "model-db-1", + db_model: true, + direct_access: true, + access_via_team_ids: [], + access_groups: [], + created_by: "user-123", + created_at: "2024-01-01", + updated_at: "2024-01-01", }, - ], - }; + }, + ], 2, 1, 1, 50); - mockUseModelsInfo.mockReturnValue({ data: modelData }); + mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null }); render(); @@ -342,27 +358,25 @@ describe("AllModelsTab", () => { }), ); - const modelData = { - data: [ - { - model_name: "gpt-4-config", - litellm_model_name: "gpt-4-config", - provider: "openai", - model_info: { - id: "model-config-1", - db_model: false, - direct_access: true, - access_via_team_ids: [], - access_groups: [], - created_by: "user-123", - created_at: "2024-01-01", - updated_at: "2024-01-01", - }, + const modelData = createPaginatedModelData([ + { + model_name: "gpt-4-config", + litellm_model_name: "gpt-4-config", + provider: "openai", + model_info: { + id: "model-config-1", + db_model: false, + direct_access: true, + access_via_team_ids: [], + access_groups: [], + created_by: "user-123", + created_at: "2024-01-01", + updated_at: "2024-01-01", }, - ], - }; + }, + ], 1, 1, 1, 50); - mockUseModelsInfo.mockReturnValue({ data: modelData }); + mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null }); render(); @@ -370,4 +384,110 @@ describe("AllModelsTab", () => { expect(screen.getByText("Defined in config")).toBeInTheDocument(); }); }); + + it("should handle pagination: Previous button is disabled on first page and Next button works", async () => { + mockUseTeams.mockReturnValue({ + data: [], + isLoading: false, + error: null, + refetch: vi.fn(), + }); + + mockUseModelCostMap.mockReturnValue( + createModelCostMapMock({ + "gpt-4-page1": { litellm_provider: "openai" }, + "gpt-4-page2": { litellm_provider: "openai" }, + }), + ); + + // Mock first page response (page 1 of 2) + const page1Data = createPaginatedModelData( + [ + { + model_name: "gpt-4-page1", + model_info: { + id: "model-page1-1", + direct_access: true, + access_via_team_ids: [], + access_groups: [], + }, + }, + ], + 2, // total_count + 1, // current_page + 2, // total_pages + 50, // size + ); + + // Set up mock to return page1Data for page 1 + mockUseModelsInfo.mockImplementation((page: number = 1) => { + return { data: page1Data, isLoading: false, error: null }; + }); + + render(); + + await waitFor(() => { + expect(screen.getByText("Showing 1 - 1 of 1 results")).toBeInTheDocument(); + }); + + // Check that Previous button is disabled on first page + const previousButton = screen.getByRole("button", { name: /previous/i }); + expect(previousButton).toBeDisabled(); + + // Check that Next button is enabled (since we're on page 1 of 2) + const nextButton = screen.getByRole("button", { name: /next/i }); + expect(nextButton).not.toBeDisabled(); + }); + + it("should handle pagination: Next button is disabled on last page", async () => { + mockUseTeams.mockReturnValue({ + data: [], + isLoading: false, + error: null, + refetch: vi.fn(), + }); + + mockUseModelCostMap.mockReturnValue( + createModelCostMapMock({ + "gpt-4-page2": { litellm_provider: "openai" }, + }), + ); + + // Mock single page response (page 1 of 1 - last page) + const singlePageData = createPaginatedModelData( + [ + { + model_name: "gpt-4-page2", + model_info: { + id: "model-page2-1", + direct_access: true, + access_via_team_ids: [], + access_groups: [], + }, + }, + ], + 1, // total_count + 1, // current_page + 1, // total_pages (only 1 page, so this is the last page) + 50, // size + ); + + mockUseModelsInfo.mockImplementation(() => { + return { data: singlePageData, isLoading: false, error: null }; + }); + + render(); + + await waitFor(() => { + expect(screen.getByText("Showing 1 - 1 of 1 results")).toBeInTheDocument(); + }); + + // When there's only 1 page (last page), Next should be disabled + const nextButton = screen.getByRole("button", { name: /next/i }); + expect(nextButton).toBeDisabled(); + + // Previous should also be disabled on the first (and only) page + const previousButton = screen.getByRole("button", { name: /previous/i }); + expect(previousButton).toBeDisabled(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx index a85e6516585..04300b7fd16 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx @@ -31,11 +31,26 @@ const AllModelsTab = ({ setSelectedModelId, setSelectedTeamId, }: AllModelsTabProps) => { - const { data: rawModelData, isLoading: isLoadingModelsInfo } = useModelsInfo(); const { data: modelCostMapData, isLoading: isLoadingModelCostMap } = useModelCostMap(); const { userId, userRole, premiumUser } = useAuthorized(); const { data: teams } = useTeams(); + const [modelNameSearch, setModelNameSearch] = useState(""); + const [modelViewMode, setModelViewMode] = useState("current_team"); + const [currentTeam, setCurrentTeam] = useState("personal"); + const [showFilters, setShowFilters] = useState(false); + const [selectedModelAccessGroupFilter, setSelectedModelAccessGroupFilter] = useState(null); + const [expandedRows, setExpandedRows] = useState>(new Set()); + const [currentPage, setCurrentPage] = useState(1); + const [pageSize] = useState(50); + const [pagination, setPagination] = useState({ + pageIndex: 0, + pageSize: 50, + }); + + const { data: rawModelData, isLoading: isLoadingModelsInfo } = useModelsInfo(currentPage, pageSize); + const isLoading = isLoadingModelsInfo || isLoadingModelCostMap; + const getProviderFromModel = (model: string) => { if (modelCostMapData !== null && modelCostMapData !== undefined) { if (typeof modelCostMapData == "object" && model in modelCostMapData) { @@ -50,18 +65,23 @@ const AllModelsTab = ({ return transformModelData(rawModelData, getProviderFromModel); }, [rawModelData, modelCostMapData]); - const [modelNameSearch, setModelNameSearch] = useState(""); - const [modelViewMode, setModelViewMode] = useState("current_team"); - const [currentTeam, setCurrentTeam] = useState("personal"); - const [showFilters, setShowFilters] = useState(false); - const [selectedModelAccessGroupFilter, setSelectedModelAccessGroupFilter] = useState(null); - const [expandedRows, setExpandedRows] = useState>(new Set()); - const [pagination, setPagination] = useState({ - pageIndex: 0, - pageSize: 50, - }); - - const isLoading = isLoadingModelsInfo || isLoadingModelCostMap; + // Get pagination metadata from the response + const paginationMeta = useMemo(() => { + if (!rawModelData) { + return { + total_count: 0, + current_page: 1, + total_pages: 1, + size: pageSize, + }; + } + return { + total_count: rawModelData.total_count ?? 0, + current_page: rawModelData.current_page ?? 1, + total_pages: rawModelData.total_pages ?? 1, + size: rawModelData.size ?? pageSize, + }; + }, [rawModelData, pageSize]); const filteredData = useMemo(() => { if (!modelData || !modelData.data || modelData.data.length === 0) { @@ -114,6 +134,7 @@ const AllModelsTab = ({ setSelectedModelAccessGroupFilter(null); setCurrentTeam("personal"); setModelViewMode("current_team"); + setCurrentPage(1); setPagination({ pageIndex: 0, pageSize: 50 }); }; @@ -334,10 +355,7 @@ const AllModelsTab = ({ ) : ( {filteredData.length > 0 - ? `Showing ${pagination.pageIndex * pagination.pageSize + 1} - ${Math.min( - (pagination.pageIndex + 1) * pagination.pageSize, - filteredData.length, - )} of ${filteredData.length} results` + ? `Showing 1 - ${filteredData.length} of ${filteredData.length} results` : "Showing 0 results"} )} @@ -347,15 +365,16 @@ const AllModelsTab = ({ ) : ( @@ -365,15 +384,16 @@ const AllModelsTab = ({ ) : ( @@ -391,8 +411,8 @@ const AllModelsTab = ({ setSelectedModelId, setSelectedTeamId, getDisplayModelName, - () => {}, - () => {}, + () => { }, + () => { }, expandedRows, setExpandedRows, )} diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 1fe72996cfd..46561bc9e1e 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -2007,15 +2007,17 @@ export const regenerateKeyCall = async (accessToken: string, keyToRegenerate: st let ModelListerrorShown = false; let errorTimer: NodeJS.Timeout | null = null; -export const modelInfoCall = async (accessToken: string, userID: string, userRole: string) => { +export const modelInfoCall = async (accessToken: string, userID: string, userRole: string, page: number = 1, size: number = 50) => { /** * Get all models on proxy */ try { - console.log("modelInfoCall:", accessToken, userID, userRole); + console.log("modelInfoCall:", accessToken, userID, userRole, page, size); let url = proxyBaseUrl ? `${proxyBaseUrl}/v2/model/info` : `/v2/model/info`; const params = new URLSearchParams(); params.append("include_team_models", "true"); + params.append("page", page.toString()); + params.append("size", size.toString()); if (params.toString()) { url += `?${params.toString()}`; } From 0c5f40fffeb70388f15f4e28176139ab063e40aa Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 21 Jan 2026 11:54:26 -0800 Subject: [PATCH 20/29] fixing build --- .../models-and-endpoints/components/AllModelsTab.test.tsx | 1 - 1 file changed, 1 deletion(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx index f20c3563c83..8a2298361c5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx @@ -1,6 +1,5 @@ import * as useAuthorizedModule from "@/app/(dashboard)/hooks/useAuthorized"; import { render, screen, waitFor } from "@testing-library/react"; -import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import AllModelsTab from "./AllModelsTab"; From 5cb5969a2656ba932fe35167b937d2cb341d7c40 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Wed, 21 Jan 2026 12:01:33 -0800 Subject: [PATCH 21/29] [Fix] LiteLLM VertexAI Pass through - ensuring incoming headers are forwarded down to target (#19524) * test_vertex_passthrough_forwards_anthropic_beta_header * add_incoming_headers --- .../llm_passthrough_endpoints.py | 31 +++++- .../test_vertex_passthrough_load_balancing.py | 98 ++++++++++++++++++- 2 files changed, 124 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index e48fd22bc8d..0a94fc95342 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -1369,6 +1369,27 @@ def get_vertex_base_url(vertex_location: Optional[str]) -> str: return f"https://{vertex_location}-aiplatform.googleapis.com/" +def add_incoming_headers(request: Request, auth_header: str) -> dict: + """ + Build headers from incoming request, preserving headers like anthropic-beta, + while removing headers that should not be forwarded and adding authorization. + + Args: + request: The FastAPI request object + auth_header: The authorization token to add + + Returns: + dict: Headers dictionary with authorization added + """ + headers = dict(request.headers) or {} + # Remove headers that should not be forwarded + headers.pop("content-length", None) + headers.pop("host", None) + # Add/override the Authorization header + headers["Authorization"] = f"Bearer {auth_header}" + return headers + + def get_vertex_pass_through_handler( call_type: Literal["discovery", "aiplatform"], ) -> BaseVertexAIPassThroughHandler: @@ -1512,9 +1533,13 @@ async def _prepare_vertex_auth_headers( api_base="", ) - headers = { - "Authorization": f"Bearer {auth_header}", - } + # Start with incoming request headers to preserve headers like anthropic-beta + headers = dict(request.headers) or {} + # Remove headers that should not be forwarded + headers.pop("content-length", None) + headers.pop("host", None) + # Add/override the Authorization header + headers["Authorization"] = f"Bearer {auth_header}" if base_target_url is not None: base_target_url = get_vertex_pass_through_handler.update_base_target_url_with_credential_location( diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py index ceb231eb4cb..a6701451f20 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py @@ -1,9 +1,14 @@ +from unittest.mock import AsyncMock, MagicMock, patch + import pytest -from unittest.mock import MagicMock, AsyncMock, patch -from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import _base_vertex_proxy_route + +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + _base_vertex_proxy_route, +) from litellm.types.router import DeploymentTypedDict + @pytest.mark.asyncio async def test_vertex_passthrough_load_balancing(): """ @@ -220,3 +225,92 @@ async def test_async_get_available_deployment_for_pass_through(): assert deployment is not None assert deployment["litellm_params"]["use_in_pass_through"] is True + +@pytest.mark.asyncio +async def test_vertex_passthrough_forwards_anthropic_beta_header(): + """ + Test that _prepare_vertex_auth_headers forwards the anthropic-beta header + (and other important headers) from the incoming request when credentials are available. + + This test validates the fix for the issue where the 1M context window header + (anthropic-beta: context-1m-2025-08-07) was being dropped when forwarding + requests to Vertex AI. + """ + from starlette.datastructures import Headers + + from litellm.llms.vertex_ai.vertex_llm_base import VertexBase + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + _prepare_vertex_auth_headers, + ) + + # Create a mock request with anthropic-beta header + mock_request = MagicMock() + mock_request.headers = Headers({ + "authorization": "Bearer old-token", + "anthropic-beta": "context-1m-2025-08-07", + "content-type": "application/json", + "user-agent": "test-client", + "content-length": "1234", # Should be removed + "host": "localhost:4000", # Should be removed + }) + + # Create mock vertex credentials + mock_vertex_credentials = MagicMock() + mock_vertex_credentials.vertex_project = "test-project" + mock_vertex_credentials.vertex_location = "us-central1" + mock_vertex_credentials.vertex_credentials = "test-credentials" + + # Create mock handler + mock_handler = MagicMock() + mock_handler.update_base_target_url_with_credential_location.return_value = ( + "https://us-central1-aiplatform.googleapis.com" + ) + + with patch.object( + VertexBase, + "_ensure_access_token_async", + new_callable=AsyncMock, + return_value=("test-auth-header", "test-project"), + ) as mock_ensure_token, patch.object( + VertexBase, + "_get_token_and_url", + return_value=("new-access-token", None), + ) as mock_get_token: + + # Call the function + ( + headers, + base_target_url, + headers_passed_through, + vertex_project, + vertex_location, + ) = await _prepare_vertex_auth_headers( + request=mock_request, + vertex_credentials=mock_vertex_credentials, + router_credentials=None, + vertex_project="test-project", + vertex_location="us-central1", + base_target_url="https://us-central1-aiplatform.googleapis.com", + get_vertex_pass_through_handler=mock_handler, + ) + + # Verify that the anthropic-beta header is preserved + assert "anthropic-beta" in headers + assert headers["anthropic-beta"] == "context-1m-2025-08-07" + + # Verify that other headers are preserved + assert "content-type" in headers + assert headers["content-type"] == "application/json" + assert "user-agent" in headers + + # Verify that the Authorization header was updated + assert "authorization" in headers + assert headers["authorization"] == "Bearer new-access-token" + + # Verify that content-length and host headers were removed + assert "content-length" not in headers + assert "host" not in headers + + # Verify that headers_passed_through is False (since we have credentials) + assert headers_passed_through is False + From c0007bd418b5e46c516eb8b1f7b861e8a35e7c83 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Thu, 22 Jan 2026 05:40:16 +0900 Subject: [PATCH 22/29] fix: Send litellm_trace_id to Langfuse to link LiteLLM logs with Langfuse logs --- litellm/integrations/langfuse/langfuse.py | 22 +--------------------- 1 file changed, 1 insertion(+), 21 deletions(-) diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index 7e62613a7e4..8087c17cafe 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -593,30 +593,10 @@ class LangFuseLogger: trace_id = clean_metadata.pop("trace_id", None) # Use standard_logging_object.trace_id if available (when trace_id from metadata is None) # This allows standard trace_id to be used when provided in standard_logging_object - # However, we skip standard_logging_object.trace_id if it's a UUID (from litellm_trace_id default), - # as we want to fall back to litellm_call_id instead for better traceability. - # Note: Users can still explicitly set a UUID trace_id via metadata["trace_id"] (highest priority) if trace_id is None and standard_logging_object is not None: - standard_trace_id = cast( + trace_id = cast( Optional[str], standard_logging_object.get("trace_id") ) - # Only use standard_logging_object.trace_id if it's not a UUID - # UUIDs are 36 characters with hyphens in format: xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx - # We check for this specific pattern to avoid rejecting valid trace_ids that happen to have hyphens - # This primarily filters out default litellm_trace_id UUIDs, while still allowing user-provided - # trace_ids via metadata["trace_id"] (which is checked first and not affected by this logic) - if standard_trace_id is not None: - # Check if it's a UUID: 36 chars, 4 hyphens, specific pattern - is_uuid = ( - len(standard_trace_id) == 36 - and standard_trace_id.count("-") == 4 - and standard_trace_id[8] == "-" - and standard_trace_id[13] == "-" - and standard_trace_id[18] == "-" - and standard_trace_id[23] == "-" - ) - if not is_uuid: - trace_id = standard_trace_id # Fallback to litellm_call_id if no trace_id found if trace_id is None: trace_id = litellm_call_id From 898cc3ff4ff1a911b772bc04eec2dc3a8f3d9e54 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Thu, 22 Jan 2026 06:19:43 +0900 Subject: [PATCH 23/29] test: update langfuse trace_id tests to use litellm_trace_id --- tests/litellm_utils_tests/test_utils.py | 4 ++-- tests/logging_callback_tests/test_alerting.py | 6 ++---- 2 files changed, 4 insertions(+), 6 deletions(-) diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py index e1751c8e2f8..f585aadbfcb 100644 --- a/tests/litellm_utils_tests/test_utils.py +++ b/tests/litellm_utils_tests/test_utils.py @@ -973,11 +973,11 @@ def test_logging_trace_id(langfuse_trace_id, langfuse_existing_trace_id): litellm_logging_obj._get_trace_id(service_name="langfuse") == langfuse_trace_id ) - ## if existing_trace_id exists + ## if no trace_id or existing_trace_id is provided, use litellm_trace_id else: assert ( litellm_logging_obj._get_trace_id(service_name="langfuse") - == litellm_call_id + == litellm_logging_obj.litellm_trace_id ) diff --git a/tests/logging_callback_tests/test_alerting.py b/tests/logging_callback_tests/test_alerting.py index 8a691e7618d..524cc00d5f7 100644 --- a/tests/logging_callback_tests/test_alerting.py +++ b/tests/logging_callback_tests/test_alerting.py @@ -866,11 +866,9 @@ async def test_langfuse_trace_id(): assert trace_url is not None - returned_trace_id = int(trace_url.split("/")[-1]) + returned_trace_id = trace_url.split("/")[-1] - assert returned_trace_id == int( - litellm_logging_obj._get_trace_id(service_name="langfuse") - ) + assert returned_trace_id == litellm_logging_obj._get_trace_id(service_name="langfuse") @pytest.mark.asyncio From b23e77585f85df0ad9d2eae660025099b3acd3ca Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 21 Jan 2026 14:49:33 -0800 Subject: [PATCH 24/29] Fix virtual keys table sorting --- .../src/app/(dashboard)/hooks/keys/useKeys.ts | 10 ++-- .../VirtualKeysPage/VirtualKeysTable.tsx | 51 +++++++++++++++---- .../key_team_helpers/filter_logic.tsx | 18 ++++--- 3 files changed, 59 insertions(+), 20 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts index 73daf954c12..cf477a2e556 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts @@ -101,12 +101,16 @@ const keyListCall = async ( } }; -export const useKeys = (page: number, pageSize: number): UseQueryResult => { +export const useKeys = ( + page: number, + pageSize: number, + options: KeyListCallOptions = {}, +): UseQueryResult => { const { accessToken } = useAuthorized(); return useQuery({ - queryKey: keyKeys.list({ page, limit: pageSize }), - queryFn: async () => await keyListCall(accessToken!, page, pageSize), + queryKey: keyKeys.list({ page, limit: pageSize, ...options }), + queryFn: async () => await keyListCall(accessToken!, page, pageSize, options), enabled: Boolean(accessToken), staleTime: 30000, // 30 seconds placeholderData: keepPreviousData, diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx index c9d11c778f8..6d968047920 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx @@ -71,12 +71,19 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo pageSize: 50, }); + // Extract sort parameters from sorting state + const sortBy = sorting.length > 0 ? sorting[0].id : null; + const sortOrder = sorting.length > 0 ? (sorting[0].desc ? "desc" : "asc") : null; + const { data: keys, isPending: isLoading, isFetching, refetch, - } = useKeys(tablePagination.pageIndex + 1, tablePagination.pageSize); + } = useKeys(tablePagination.pageIndex + 1, tablePagination.pageSize, { + sortBy: sortBy || undefined, + sortOrder: sortOrder || undefined, + }); const totalCount = keys?.total_count || 0; const [expandedAccordions, setExpandedAccordions] = useState>({}); @@ -110,6 +117,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo id: "expander", header: () => null, size: 40, + enableSorting: false, cell: ({ row }) => row.getCanExpand() ? (