Merge branch 'main' into litellm_add_output_format_claude

This commit is contained in:
Ishaan Jaffer 2026-01-21 19:20:48 -08:00
commit 936fe3c94d
48 changed files with 4057 additions and 635 deletions

View file

@ -461,3 +461,48 @@ generateContent();
</TabItem>
</Tabs>
### Using Anthropic Beta Features on Vertex AI
When using Anthropic models via Vertex AI passthrough (e.g., Claude on Vertex), you can enable Anthropic beta features like extended context windows.
The `anthropic-beta` header is automatically forwarded to Vertex AI when calling Anthropic models.
```bash
curl http://localhost:4000/vertex_ai/v1/projects/${PROJECT_ID}/locations/us-east5/publishers/anthropic/models/claude-3-5-sonnet:rawPredict \
-H "Content-Type: application/json" \
-H "Authorization: Bearer sk-1234" \
-H "anthropic-beta: context-1m-2025-08-07" \
-d '{
"anthropic_version": "vertex-2023-10-16",
"messages": [{"role": "user", "content": "Hello"}],
"max_tokens": 500
}'
```
### Forwarding Custom Headers with `x-pass-` Prefix
You can forward any custom header to the provider by prefixing it with `x-pass-`. The prefix is stripped before the header is sent to the provider.
For example:
- `x-pass-anthropic-beta: value` becomes `anthropic-beta: value`
- `x-pass-custom-header: value` becomes `custom-header: value`
This is useful when you need to send provider-specific headers that aren't in the default allowlist.
```bash
curl http://localhost:4000/vertex_ai/v1/projects/${PROJECT_ID}/locations/us-east5/publishers/anthropic/models/claude-3-5-sonnet:rawPredict \
-H "Content-Type: application/json" \
-H "Authorization: Bearer sk-1234" \
-H "x-pass-anthropic-beta: context-1m-2025-08-07" \
-H "x-pass-custom-feature: enabled" \
-d '{
"anthropic_version": "vertex-2023-10-16",
"messages": [{"role": "user", "content": "Hello"}],
"max_tokens": 500
}'
```
:::info
The `x-pass-` prefix works for all LLM pass-through endpoints, not just Vertex AI.
:::

View file

@ -1122,6 +1122,20 @@ BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES = [
"generateQuery/",
"optimize-prompt/",
]
# Headers that are safe to forward from incoming requests to Vertex AI
# Using an allowlist approach for security - only forward headers we explicitly trust
ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS = {
"anthropic-beta", # Required for Anthropic features like extended context windows
"content-type", # Required for request body parsing
}
# Prefix for headers that should be forwarded to the provider with the prefix stripped
# e.g., 'x-pass-anthropic-beta: value' becomes 'anthropic-beta: value'
# Works for all LLM pass-through endpoints (Vertex AI, Anthropic, Bedrock, etc.)
PASS_THROUGH_HEADER_PREFIX = "x-pass-"
BASE_MCP_ROUTE = "/mcp"
BATCH_STATUS_POLL_INTERVAL_SECONDS = int(

View file

@ -16,6 +16,21 @@ import openai
from litellm.types.utils import LiteLLMCommonStrings
_MINIMAL_ERROR_RESPONSE: Optional[httpx.Response] = None
def _get_minimal_error_response() -> httpx.Response:
"""Get a cached minimal httpx.Response object for error cases."""
global _MINIMAL_ERROR_RESPONSE
if _MINIMAL_ERROR_RESPONSE is None:
_MINIMAL_ERROR_RESPONSE = httpx.Response(
status_code=400,
request=httpx.Request(
method="GET", url="https://litellm.ai"
),
)
return _MINIMAL_ERROR_RESPONSE
class AuthenticationError(openai.AuthenticationError): # type: ignore
def __init__(
@ -127,16 +142,15 @@ class BadRequestError(openai.BadRequestError): # type: ignore
self.litellm_debug_info = litellm_debug_info
self.max_retries = max_retries
self.num_retries = num_retries
_response_headers = (
getattr(response, "headers", None) if response is not None else None
)
self.response = httpx.Response(
status_code=self.status_code,
headers=_response_headers,
request=httpx.Request(
method="GET", url="https://litellm.ai"
), # mock request object
)
if (
response is not None
and isinstance(response, httpx.Response)
and hasattr(response, "request")
and response.request is not None
):
self.response = response
else:
self.response = _get_minimal_error_response()
super().__init__(
self.message, response=self.response, body=body
) # Call the base class constructor with the parameters it needs

View file

@ -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

View file

@ -93,8 +93,9 @@ def get_litellm_params(
"text_completion": text_completion,
"azure_ad_token_provider": azure_ad_token_provider,
"user_continue_message": user_continue_message,
"base_model": base_model
or _get_base_model_from_litellm_call_metadata(metadata=metadata),
"base_model": base_model or (
_get_base_model_from_litellm_call_metadata(metadata=metadata) if metadata else None
),
"litellm_trace_id": litellm_trace_id,
"litellm_session_id": litellm_session_id,
"hf_model_name": hf_model_name,

View file

@ -1,7 +1,5 @@
from typing import Optional, Tuple
import httpx
import litellm
from litellm.constants import REPLICATE_MODEL_NAME_WITH_ID_LENGTH
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
@ -453,11 +451,7 @@ def get_llm_provider( # noqa: PLR0915
raise litellm.exceptions.BadRequestError( # type: ignore
message=error_str,
model=model,
response=httpx.Response(
status_code=400,
content=error_str,
request=httpx.Request(method="completion", url="https://github.com/BerriAI/litellm"), # type: ignore
),
response=None,
llm_provider="",
)
if api_base is not None and not isinstance(api_base, str):
@ -481,11 +475,7 @@ def get_llm_provider( # noqa: PLR0915
raise litellm.exceptions.BadRequestError( # type: ignore
message=f"GetLLMProvider Exception - {str(e)}\n\noriginal model: {model}",
model=model,
response=httpx.Response(
status_code=400,
content=error_str,
request=httpx.Request(method="completion", url="https://github.com/BerriAI/litellm"), # type: ignore
),
response=None,
llm_provider="",
)

View file

@ -325,12 +325,12 @@ class Logging(LiteLLMLoggingBaseClass):
messages = new_messages
self.model = model
self.messages = copy.deepcopy(messages)
self.messages = copy.deepcopy(messages) if messages is not None else None
self.stream = stream
self.start_time = start_time # log the call start time
self.call_type = call_type
self.litellm_call_id = litellm_call_id
self.litellm_trace_id: str = litellm_trace_id or str(uuid.uuid4())
self.litellm_trace_id: str = litellm_trace_id if litellm_trace_id else str(uuid.uuid4())
self.function_id = function_id
self.streaming_chunks: List[Any] = [] # for generating complete stream response
self.sync_streaming_chunks: List[

View file

@ -3,6 +3,8 @@ from urllib.parse import parse_qs
import httpx
from litellm.constants import PASS_THROUGH_HEADER_PREFIX
class BasePassthroughUtils:
@staticmethod
@ -27,7 +29,11 @@ class BasePassthroughUtils:
forward_headers: Optional[bool] = False,
):
"""
Helper to forward headers from original request
Helper to forward headers from original request.
Also handles 'x-pass-' prefixed headers which are always forwarded
with the prefix stripped, regardless of forward_headers setting.
e.g., 'x-pass-anthropic-beta: value' becomes 'anthropic-beta: value'
"""
if forward_headers is True:
# Header We Should NOT forward
@ -36,6 +42,14 @@ class BasePassthroughUtils:
# Combine request headers with custom headers
headers = {**request_headers, **headers}
# Always process x-pass- prefixed headers (strip prefix and forward)
for header_name, header_value in request_headers.items():
if header_name.lower().startswith(PASS_THROUGH_HEADER_PREFIX):
# Strip the 'x-pass-' prefix to get the actual header name
actual_header_name = header_name[len(PASS_THROUGH_HEADER_PREFIX) :]
headers[actual_header_name] = header_value
return headers
class CommonUtils:

View file

@ -6,6 +6,8 @@ LiteLLM MCP Server Routes
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
@ -13,6 +15,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,8 +28,8 @@ 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.utils import client
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
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
@ -226,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"
@ -733,13 +738,15 @@ 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]],
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.
@ -757,67 +764,188 @@ 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],
@ -1050,6 +1178,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.
@ -1075,6 +1205,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"
@ -1320,33 +1452,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 +1470,82 @@ 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,
)
)
except Exception as e:
traceback_str = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG)
from litellm.proxy.proxy_server import proxy_logging_obj
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 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
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,
)
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.",
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
)
# 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,
)
return response
async def mcp_get_prompt(
name: str,

View file

@ -1,5 +1,7 @@
from typing import Any, Dict, Optional, Union
from typing import Any, Dict, Optional, Union, TYPE_CHECKING
from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache
from litellm.proxy._types import (
KeyRequestBase,
LiteLLM_ManagementEndpoint_MetadataFields,
@ -11,6 +13,9 @@ from litellm.proxy._types import (
)
from litellm.proxy.utils import _premium_user_check
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient, ProxyLogging
def _user_has_admin_view(user_api_key_dict: UserAPIKeyAuth) -> bool:
return (
@ -31,6 +36,78 @@ def _is_user_team_admin(
return False
async def _user_has_admin_privileges(
user_api_key_dict: UserAPIKeyAuth,
prisma_client: Optional["PrismaClient"] = None,
user_api_key_cache: Optional["DualCache"] = None,
proxy_logging_obj: Optional["ProxyLogging"] = None,
) -> bool:
"""
Check if user has admin privileges (proxy admin, team admin, or org admin).
Args:
user_api_key_dict: User API key authentication object
prisma_client: Prisma client for database operations
user_api_key_cache: Cache for user API keys
proxy_logging_obj: Proxy logging object
Returns:
True if user is proxy admin, team admin for any team, or org admin for any organization
"""
# Check if user is proxy admin
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
return True
# If no database connection, can't check team/org admin status
if prisma_client is None or user_api_key_dict.user_id is None:
return False
# Get user object to check team and org admin status
from litellm.caching import DualCache as DualCacheImport
from litellm.proxy.auth.auth_checks import get_user_object
try:
user_obj = await get_user_object(
user_id=user_api_key_dict.user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache or DualCacheImport(),
user_id_upsert=False,
proxy_logging_obj=proxy_logging_obj,
)
if user_obj is None:
return False
# Check if user is org admin for any organization
if user_obj.organization_memberships is not None:
for membership in user_obj.organization_memberships:
if membership.user_role == LitellmUserRoles.ORG_ADMIN.value:
return True
# Check if user is team admin for any team
if user_obj.teams is not None and len(user_obj.teams) > 0:
# Get all teams user is in
teams = await prisma_client.db.litellm_teamtable.find_many(
where={"team_id": {"in": user_obj.teams}}
)
for team in teams:
team_obj = LiteLLM_TeamTable(**team.model_dump())
if _is_user_team_admin(
user_api_key_dict=user_api_key_dict, team_obj=team_obj
):
return True
except Exception as e:
# If there's an error checking, default to False for security
verbose_proxy_logger.debug(
f"Error checking admin privileges for user {user_api_key_dict.user_id}: {e}"
)
return False
return False
def _set_object_metadata_field(
object_data: Union[
LiteLLM_TeamTable,

View file

@ -17,7 +17,10 @@ from starlette.websockets import WebSocketState
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES
from litellm.constants import (
ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS,
BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES,
)
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import *
from litellm.proxy.auth.route_checks import RouteChecks
@ -1369,6 +1372,27 @@ def get_vertex_base_url(vertex_location: Optional[str]) -> str:
return f"https://{vertex_location}-aiplatform.googleapis.com/"
def get_vertex_ai_allowed_incoming_headers(request: Request) -> dict:
"""
Extract only the allowed headers from incoming request for Vertex AI pass-through.
Uses an allowlist approach for security - only forwards headers we explicitly trust.
This prevents accidentally forwarding sensitive headers like the LiteLLM auth token.
Args:
request: The FastAPI request object
Returns:
dict: Headers dictionary with only allowed headers
"""
incoming_headers = dict(request.headers) or {}
headers = {}
for header_name in ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS:
if header_name in incoming_headers:
headers[header_name] = incoming_headers[header_name]
return headers
def get_vertex_pass_through_handler(
call_type: Literal["discovery", "aiplatform"],
) -> BaseVertexAIPassThroughHandler:
@ -1512,9 +1536,10 @@ async def _prepare_vertex_auth_headers(
api_base="",
)
headers = {
"Authorization": f"Bearer {auth_header}",
}
# Use allowlist approach - only forward specific safe headers
headers = get_vertex_ai_allowed_incoming_headers(request)
# Add the Authorization header with vendor credentials
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(

View file

@ -5070,6 +5070,7 @@ async def model_list(
only_model_access_groups: Optional[bool] = False,
include_metadata: Optional[bool] = False,
fallback_type: Optional[str] = None,
scope: Optional[str] = None,
):
"""
Use `/model/info` - to get detailed model information, example - pricing, mode, etc.
@ -5080,14 +5081,85 @@ async def model_list(
- include_metadata: Include additional metadata in the response with fallback information
- fallback_type: Type of fallbacks to include ("general", "context_window", "content_policy")
Defaults to "general" when include_metadata=true
- scope: Optional scope parameter. Currently only accepts "expand".
When scope=expand is passed, proxy admins, team admins, and org admins
will receive all proxy models as if they are a proxy admin.
"""
global llm_model_list, general_settings, llm_router, prisma_client, user_api_key_cache, proxy_logging_obj
from litellm.proxy.management_endpoints.common_utils import (
_user_has_admin_privileges,
)
from litellm.proxy.utils import (
create_model_info_response,
get_available_models_for_user,
)
# Validate scope parameter if provided
if scope is not None and scope != "expand":
raise HTTPException(
status_code=400,
detail=f"Invalid scope parameter. Only 'expand' is currently supported. Received: {scope}",
)
# Check if scope=expand is requested and user has admin privileges
should_expand_scope = False
if scope == "expand":
should_expand_scope = await _user_has_admin_privileges(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# If scope=expand and user has admin privileges, return all proxy models
if should_expand_scope:
# Get all proxy models as if user is a proxy admin
if llm_router is None:
proxy_model_list = []
model_access_groups = {}
else:
proxy_model_list = llm_router.get_model_names()
model_access_groups = llm_router.get_model_access_groups()
# Include model access groups if requested
if include_model_access_groups:
proxy_model_list = list(set(proxy_model_list + list(model_access_groups.keys())))
# Get complete model list including wildcard routes if requested
from litellm.proxy.auth.model_checks import get_complete_model_list
all_models = get_complete_model_list(
key_models=[],
team_models=[],
proxy_model_list=proxy_model_list,
user_model=None,
infer_model_from_keys=False,
return_wildcard_routes=return_wildcard_routes or False,
llm_router=llm_router,
model_access_groups=model_access_groups,
include_model_access_groups=include_model_access_groups or False,
only_model_access_groups=only_model_access_groups or False,
)
# Build response data with all proxy models
model_data = []
for model in all_models:
model_info = create_model_info_response(
model_id=model,
provider="openai",
include_metadata=include_metadata or False,
fallback_type=fallback_type,
llm_router=llm_router,
)
model_data.append(model_info)
return dict(
data=model_data,
object="list",
)
# Otherwise, use the normal behavior (current implementation)
# Get available models for the user
all_models = await get_available_models_for_user(
user_api_key_dict=user_api_key_dict,
@ -7503,6 +7575,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",
@ -7522,6 +7665,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 +7675,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(
@ -7565,62 +7716,32 @@ 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 if debug is not None else False, llm_router=llm_router
)
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(

View file

@ -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,
@ -301,6 +303,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 +353,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)
@ -587,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
#########################################################
@ -624,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)
@ -823,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:
@ -1003,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:
@ -1160,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:
@ -1318,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:
@ -1500,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:

View file

@ -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:

View file

@ -1,3 +1,5 @@
import traceback
from datetime import datetime
from typing import (
TYPE_CHECKING,
Any,
@ -11,17 +13,32 @@ 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.llms.openai import ResponsesAPIResponse
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
# 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.
# `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/"
@ -117,6 +134,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)
@ -462,7 +481,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,
@ -470,6 +489,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 +499,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: Optional[str] = None
try:
(
tool_name,
@ -514,6 +541,103 @@ 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 = {
"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 +650,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 +681,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 +699,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 +717,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 +890,31 @@ 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 +955,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"),
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,
)

View file

@ -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
@ -298,12 +298,14 @@ 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()
# 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
@ -311,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:
@ -339,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 = {}
@ -487,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", []))
@ -543,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(
@ -568,6 +587,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

View file

@ -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')"""
@ -384,6 +385,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 +450,7 @@ CallTypesLiteral = Literal[
"vector_store_file_delete",
"avector_store_file_delete",
"call_mcp_tool",
"list_mcp_tools",
"asend_message",
"send_message",
"aresponses",
@ -1343,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."""

View file

@ -771,7 +771,8 @@ def function_setup( # noqa: PLR0915
function_id: Optional[str] = kwargs["id"] if "id" in kwargs else None
## LAZY LOAD COROUTINE CHECKER ##
get_coroutine_checker = getattr(sys.modules[__name__], "get_coroutine_checker")
get_coroutine_checker_fn = getattr(sys.modules[__name__], "get_coroutine_checker")
coroutine_checker = get_coroutine_checker_fn()
## DYNAMIC CALLBACKS ##
dynamic_callbacks: Optional[
@ -825,7 +826,7 @@ def function_setup( # noqa: PLR0915
if len(litellm.input_callback) > 0:
removed_async_items = []
for index, callback in enumerate(litellm.input_callback): # type: ignore
if get_coroutine_checker().is_async_callable(callback):
if coroutine_checker.is_async_callable(callback):
litellm._async_input_callback.append(callback)
removed_async_items.append(index)
@ -835,7 +836,7 @@ def function_setup( # noqa: PLR0915
if len(litellm.success_callback) > 0:
removed_async_items = []
for index, callback in enumerate(litellm.success_callback): # type: ignore
if get_coroutine_checker().is_async_callable(callback):
if coroutine_checker.is_async_callable(callback):
litellm.logging_callback_manager.add_litellm_async_success_callback(
callback
)
@ -860,7 +861,7 @@ def function_setup( # noqa: PLR0915
if len(litellm.failure_callback) > 0:
removed_async_items = []
for index, callback in enumerate(litellm.failure_callback): # type: ignore
if get_coroutine_checker().is_async_callable(callback):
if coroutine_checker.is_async_callable(callback):
litellm.logging_callback_manager.add_litellm_async_failure_callback(
callback
)
@ -893,7 +894,7 @@ def function_setup( # noqa: PLR0915
removed_async_items = []
for index, callback in enumerate(kwargs["success_callback"]):
if (
get_coroutine_checker().is_async_callable(callback)
coroutine_checker.is_async_callable(callback)
or callback == "dynamodb"
or callback == "s3"
):

7
proxy_config.yaml Normal file
View file

@ -0,0 +1,7 @@
model_list:
- model_name: "*"
litellm_params:
model: "*"
general_settings:
master_key: sk-1234

View file

@ -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
)

View file

@ -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

View file

@ -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

View file

@ -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):
"""

View file

@ -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,225 @@ 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 allowlisted headers are preserved
assert "anthropic-beta" in headers
assert headers["anthropic-beta"] == "context-1m-2025-08-07"
assert "content-type" in headers
assert headers["content-type"] == "application/json"
# Verify that the Authorization header is set with vendor credentials
assert "Authorization" in headers
assert headers["Authorization"] == "Bearer new-access-token"
# Verify that non-allowlisted headers are NOT forwarded (security)
# Only anthropic-beta, content-type, and Authorization should be present
assert "authorization" not in headers # lowercase auth token not forwarded
assert "user-agent" not in headers # not in allowlist
assert "content-length" not in headers # not in allowlist
assert "host" not in headers # not in allowlist
# Verify that headers_passed_through is False (since we have credentials)
assert headers_passed_through is False
@pytest.mark.asyncio
async def test_vertex_passthrough_does_not_forward_litellm_auth_token():
"""
Test that the LiteLLM authorization header is NOT forwarded to Vertex AI.
This test validates the fix for the issue where both the LiteLLM auth token
(lowercase 'authorization') and the Vertex AI token (uppercase 'Authorization')
were being sent, causing 401 errors on the vendor side.
The incoming request has:
- authorization: Bearer <litellm_token> (should NOT be forwarded)
The outgoing request should only have:
- Authorization: Bearer <vertex_token> (vendor credentials)
"""
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 ONLY the litellm auth token (no other headers)
mock_request = MagicMock()
mock_request.headers = Headers({
"authorization": "Bearer sk-litellm-secret-key", # LiteLLM token - should NOT be forwarded
"Authorization": "Bearer sk-litellm-secret-key-uppercase", # Also try uppercase
})
# 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"),
), patch.object(
VertexBase,
"_get_token_and_url",
return_value=("vertex-access-token", None),
):
(
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,
)
# The ONLY Authorization header should be the Vertex token
assert headers["Authorization"] == "Bearer vertex-access-token"
# The LiteLLM token should NOT be present (neither lowercase nor as a duplicate)
assert "authorization" not in headers
assert headers.get("Authorization") != "Bearer sk-litellm-secret-key"
assert headers.get("Authorization") != "Bearer sk-litellm-secret-key-uppercase"
# Verify we only have the expected headers (Authorization + any allowlisted ones present)
# Since the request only had auth headers, only Authorization should be in output
assert set(headers.keys()) == {"Authorization"}
def test_forward_headers_from_request_x_pass_prefix():
"""
Test that headers with 'x-pass-' prefix are forwarded with the prefix stripped.
This allows users to force-forward arbitrary headers to the vendor API:
- 'x-pass-anthropic-beta: value' becomes 'anthropic-beta: value'
- 'x-pass-custom-header: value' becomes 'custom-header: value'
This is tested on BasePassthroughUtils.forward_headers_from_request which is used
by all pass-through endpoints (not just Vertex AI).
"""
from litellm.passthrough.utils import BasePassthroughUtils
# Simulate incoming request headers
request_headers = {
"x-pass-anthropic-beta": "context-1m-2025-08-07",
"x-pass-custom-header": "custom-value",
"x-pass-another-header": "another-value",
"authorization": "Bearer sk-litellm-key",
"x-litellm-api-key": "sk-1234",
"content-type": "application/json",
}
# Start with empty headers dict (simulating custom headers from endpoint config)
headers = {}
# Call the method with forward_headers=False (default behavior)
# x-pass- headers should still be forwarded
result = BasePassthroughUtils.forward_headers_from_request(
request_headers=request_headers,
headers=headers,
forward_headers=False,
)
# Verify x-pass- prefixed headers are forwarded with prefix stripped
assert "anthropic-beta" in result
assert result["anthropic-beta"] == "context-1m-2025-08-07"
assert "custom-header" in result
assert result["custom-header"] == "custom-value"
assert "another-header" in result
assert result["another-header"] == "another-value"
# Verify other headers are NOT forwarded (since forward_headers=False)
assert "authorization" not in result
assert "x-litellm-api-key" not in result
assert "content-type" not in result
# Verify original x-pass- prefixed headers are NOT in output (only stripped versions)
assert "x-pass-anthropic-beta" not in result
assert "x-pass-custom-header" not in result

View file

@ -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

View file

@ -3216,3 +3216,851 @@ 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
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
@pytest.mark.asyncio
async def test_model_list_scope_parameter_validation(monkeypatch):
"""Test that invalid scope parameter raises HTTPException"""
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth, LitellmUserRoles
from litellm.proxy.proxy_server import model_list
mock_user_api_key_dict = UserAPIKeyAuth(
user_id="test-user",
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="test-key",
)
# Test invalid scope parameter
with pytest.raises(HTTPException) as exc_info:
await model_list(
user_api_key_dict=mock_user_api_key_dict,
scope="invalid_scope",
)
assert exc_info.value.status_code == 400
assert "Invalid scope parameter" in exc_info.value.detail
assert "Only 'expand' is currently supported" in exc_info.value.detail
@pytest.mark.asyncio
async def test_model_list_scope_expand_proxy_admin(monkeypatch):
"""Test that proxy admin with scope=expand returns all proxy models"""
from litellm.proxy._types import UserAPIKeyAuth, LitellmUserRoles, LiteLLM_UserTable
from litellm.proxy.proxy_server import model_list
# Mock user API key dict for proxy admin
mock_user_api_key_dict = UserAPIKeyAuth(
user_id="proxy-admin-user",
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="test-key",
)
# Mock llm_router with proxy models
mock_router = MagicMock()
mock_router.get_model_names.return_value = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"]
mock_router.get_model_access_groups.return_value = {}
# Mock prisma_client
mock_prisma_client = MagicMock()
# Mock user_api_key_cache
mock_user_api_key_cache = MagicMock()
# Mock proxy_logging_obj
mock_proxy_logging_obj = MagicMock()
# Mock get_complete_model_list
mock_all_models = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"]
# Mock create_model_info_response
def mock_create_model_info_response(model_id, provider, include_metadata=False, fallback_type=None, llm_router=None):
return {"id": model_id, "object": "model"}
# 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_api_key_cache", mock_user_api_key_cache)
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
monkeypatch.setattr(
"litellm.proxy.auth.model_checks.get_complete_model_list",
lambda **kwargs: mock_all_models,
)
monkeypatch.setattr(
"litellm.proxy.utils.create_model_info_response",
mock_create_model_info_response,
)
# Call model_list with scope=expand
result = await model_list(
user_api_key_dict=mock_user_api_key_dict,
scope="expand",
)
# Verify result contains all proxy models
assert result["object"] == "list"
assert len(result["data"]) == 3
assert all(model["id"] in mock_all_models for model in result["data"])
# Verify router methods were called
mock_router.get_model_names.assert_called_once()
mock_router.get_model_access_groups.assert_called_once()
@pytest.mark.asyncio
async def test_model_list_scope_expand_org_admin(monkeypatch):
"""Test that org admin with scope=expand returns all proxy models"""
from litellm.proxy._types import (
UserAPIKeyAuth,
LitellmUserRoles,
LiteLLM_UserTable,
)
from litellm.proxy.proxy_server import model_list
# Mock user API key dict for org admin
mock_user_api_key_dict = UserAPIKeyAuth(
user_id="org-admin-user",
user_role=LitellmUserRoles.INTERNAL_USER, # Not proxy admin, but org admin
api_key="test-key",
)
# Mock user object with org admin membership
from litellm.proxy._types import LiteLLM_OrganizationMembershipTable
from datetime import datetime
mock_user_obj = LiteLLM_UserTable(
user_id="org-admin-user",
user_email="org-admin@example.com",
organization_memberships=[
LiteLLM_OrganizationMembershipTable(
user_id="org-admin-user",
organization_id="org-123",
user_role=LitellmUserRoles.ORG_ADMIN.value,
spend=0.0,
created_at=datetime.now(),
updated_at=datetime.now(),
)
],
teams=[],
)
# Mock llm_router with proxy models
mock_router = MagicMock()
mock_router.get_model_names.return_value = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"]
mock_router.get_model_access_groups.return_value = {}
# Mock prisma_client
mock_prisma_client = MagicMock()
# Mock user_api_key_cache
mock_user_api_key_cache = MagicMock()
# Mock proxy_logging_obj
mock_proxy_logging_obj = MagicMock()
# Mock get_user_object to return user with org admin role
async def mock_get_user_object(*args, **kwargs):
return mock_user_obj
# Mock get_complete_model_list
mock_all_models = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"]
# Mock create_model_info_response
def mock_create_model_info_response(model_id, provider, include_metadata=False, fallback_type=None, llm_router=None):
return {"id": model_id, "object": "model"}
# 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_api_key_cache", mock_user_api_key_cache)
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
monkeypatch.setattr(
"litellm.proxy.auth.auth_checks.get_user_object",
mock_get_user_object,
)
monkeypatch.setattr(
"litellm.proxy.auth.model_checks.get_complete_model_list",
lambda **kwargs: mock_all_models,
)
monkeypatch.setattr(
"litellm.proxy.utils.create_model_info_response",
mock_create_model_info_response,
)
# Call model_list with scope=expand
result = await model_list(
user_api_key_dict=mock_user_api_key_dict,
scope="expand",
)
# Verify result contains all proxy models
assert result["object"] == "list"
assert len(result["data"]) == 3
assert all(model["id"] in mock_all_models for model in result["data"])
# Verify router methods were called
mock_router.get_model_names.assert_called_once()
mock_router.get_model_access_groups.assert_called_once()
@pytest.mark.asyncio
async def test_model_list_scope_expand_team_admin(monkeypatch):
"""Test that team admin with scope=expand returns all proxy models"""
from litellm.proxy._types import (
UserAPIKeyAuth,
LitellmUserRoles,
LiteLLM_UserTable,
LiteLLM_TeamTable,
)
from litellm.proxy.proxy_server import model_list
# Mock user API key dict for team admin
mock_user_api_key_dict = UserAPIKeyAuth(
user_id="team-admin-user",
user_role=LitellmUserRoles.INTERNAL_USER, # Not proxy admin, but team admin
api_key="test-key",
)
# Mock team with user as admin - use dict structure that matches Prisma return
mock_team = MagicMock()
mock_team.model_dump.return_value = {
"team_id": "team-123",
"members_with_roles": [
{"user_id": "team-admin-user", "role": "admin"}
],
}
# Create team object from the dict (validator will convert members_with_roles to Member objects)
mock_team_obj = LiteLLM_TeamTable(**mock_team.model_dump())
# Mock user object with team membership
mock_user_obj = LiteLLM_UserTable(
user_id="team-admin-user",
user_email="team-admin@example.com",
organization_memberships=[],
teams=["team-123"],
)
# Mock llm_router with proxy models
mock_router = MagicMock()
mock_router.get_model_names.return_value = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"]
mock_router.get_model_access_groups.return_value = {}
# Mock prisma_client
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(
return_value=[mock_team]
)
# Mock user_api_key_cache
mock_user_api_key_cache = MagicMock()
# Mock proxy_logging_obj
mock_proxy_logging_obj = MagicMock()
# Mock get_user_object to return user with team membership
async def mock_get_user_object(*args, **kwargs):
return mock_user_obj
# Mock get_complete_model_list
mock_all_models = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"]
# Mock create_model_info_response
def mock_create_model_info_response(model_id, provider, include_metadata=False, fallback_type=None, llm_router=None):
return {"id": model_id, "object": "model"}
# 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_api_key_cache", mock_user_api_key_cache)
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
monkeypatch.setattr(
"litellm.proxy.auth.auth_checks.get_user_object",
mock_get_user_object,
)
monkeypatch.setattr(
"litellm.proxy.auth.model_checks.get_complete_model_list",
lambda **kwargs: mock_all_models,
)
monkeypatch.setattr(
"litellm.proxy.utils.create_model_info_response",
mock_create_model_info_response,
)
# Call model_list with scope=expand
result = await model_list(
user_api_key_dict=mock_user_api_key_dict,
scope="expand",
)
# Verify result contains all proxy models
assert result["object"] == "list"
assert len(result["data"]) == 3
assert all(model["id"] in mock_all_models for model in result["data"])
# Verify router methods were called
mock_router.get_model_names.assert_called_once()
mock_router.get_model_access_groups.assert_called_once()
@pytest.mark.asyncio
async def test_model_list_scope_expand_normal_user(monkeypatch):
"""Test that normal internal user with scope=expand returns only their models (not expanded)"""
from litellm.proxy._types import UserAPIKeyAuth, LitellmUserRoles, LiteLLM_UserTable
from litellm.proxy.proxy_server import model_list
# Mock user API key dict for normal internal user
mock_user_api_key_dict = UserAPIKeyAuth(
user_id="normal-user",
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="test-key",
models=["gpt-3.5-turbo"], # User only has access to this model
)
# Mock user object without admin privileges
mock_user_obj = LiteLLM_UserTable(
user_id="normal-user",
user_email="normal@example.com",
organization_memberships=[], # No org admin
teams=[], # No teams
)
# Mock llm_router
mock_router = MagicMock()
mock_router.get_model_names.return_value = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"]
# Mock prisma_client
mock_prisma_client = MagicMock()
# Mock user_api_key_cache
mock_user_api_key_cache = MagicMock()
# Mock proxy_logging_obj
mock_proxy_logging_obj = MagicMock()
# Mock get_user_object to return user without admin privileges
async def mock_get_user_object(*args, **kwargs):
return mock_user_obj
# Mock get_available_models_for_user to return only user's models
async def mock_get_available_models_for_user(*args, **kwargs):
return ["gpt-3.5-turbo"] # Only user's accessible models
# Mock create_model_info_response
def mock_create_model_info_response(model_id, provider, include_metadata=False, fallback_type=None, llm_router=None):
return {"id": model_id, "object": "model"}
# 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_api_key_cache", mock_user_api_key_cache)
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
monkeypatch.setattr(
"litellm.proxy.auth.auth_checks.get_user_object",
mock_get_user_object,
)
monkeypatch.setattr(
"litellm.proxy.utils.get_available_models_for_user",
mock_get_available_models_for_user,
)
monkeypatch.setattr(
"litellm.proxy.utils.create_model_info_response",
mock_create_model_info_response,
)
# Call model_list with scope=expand
result = await model_list(
user_api_key_dict=mock_user_api_key_dict,
scope="expand",
)
# Verify result contains only user's models (not all proxy models)
assert result["object"] == "list"
assert len(result["data"]) == 1
assert result["data"][0]["id"] == "gpt-3.5-turbo"
# Verify router methods were NOT called (normal path, not expanded)
mock_router.get_model_names.assert_not_called()
mock_router.get_model_access_groups.assert_not_called()
@pytest.mark.asyncio
async def test_model_list_no_scope_parameter(monkeypatch):
"""Test that model_list without scope parameter uses normal behavior"""
from litellm.proxy._types import UserAPIKeyAuth, LitellmUserRoles
from litellm.proxy.proxy_server import model_list
# Mock user API key dict
mock_user_api_key_dict = UserAPIKeyAuth(
user_id="test-user",
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="test-key",
models=["gpt-3.5-turbo"],
)
# Mock llm_router
mock_router = MagicMock()
# Mock prisma_client
mock_prisma_client = MagicMock()
# Mock user_api_key_cache
mock_user_api_key_cache = MagicMock()
# Mock proxy_logging_obj
mock_proxy_logging_obj = MagicMock()
# Mock get_available_models_for_user
async def mock_get_available_models_for_user(*args, **kwargs):
return ["gpt-3.5-turbo"]
# Mock create_model_info_response
def mock_create_model_info_response(model_id, provider, include_metadata=False, fallback_type=None, llm_router=None):
return {"id": model_id, "object": "model"}
# 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_api_key_cache", mock_user_api_key_cache)
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
monkeypatch.setattr(
"litellm.proxy.utils.get_available_models_for_user",
mock_get_available_models_for_user,
)
monkeypatch.setattr(
"litellm.proxy.utils.create_model_info_response",
mock_create_model_info_response,
)
# Call model_list without scope parameter
result = await model_list(
user_api_key_dict=mock_user_api_key_dict,
scope=None,
)
# Verify result uses normal behavior
assert result["object"] == "list"
assert len(result["data"]) == 1
assert result["data"][0]["id"] == "gpt-3.5-turbo"
# Verify router methods were NOT called (normal path)
mock_router.get_model_names.assert_not_called()
mock_router.get_model_access_groups.assert_not_called()

View file

@ -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"

View file

@ -1 +1,2 @@
export const ADMIN_STORAGE_PATH = "admin.storageState.json";
export const INTERNAL_USER_VIEWER_STORAGE_PATH = "internalViewer.storageState.json";

View file

@ -7,4 +7,8 @@ export const users = {
email: "admin",
password: isCI ? "gm" : "sk-1234",
},
[Role.InternalUserViewer]: {
email: "internalViewer@test.com",
password: "test",
},
};

View file

@ -1,17 +1,30 @@
import { chromium } from "@playwright/test";
import { users } from "./fixtures/users";
import { Role } from "./fixtures/roles";
import { ADMIN_STORAGE_PATH, INTERNAL_USER_VIEWER_STORAGE_PATH } from "./constants";
async function loginAndSaveState(
browser,
user,
storagePath: string
) {
const context = await browser.newContext();
const page = await context.newPage();
await page.goto("http://localhost:4000/ui/login");
await page.getByPlaceholder("Enter your username").fill(user.email);
await page.getByPlaceholder("Enter your password").fill(user.password);
await page.getByRole("button", { name: "Login" }).click();
await page.getByText('AI GATEWAY').waitFor();
await context.storageState({ path: storagePath });
await context.close();
}
async function globalSetup() {
const browser = await chromium.launch();
const page = await browser.newPage();
await page.goto("http://localhost:4000/ui/login");
await page.getByPlaceholder("Enter your username").fill(users[Role.ProxyAdmin].email);
await page.getByPlaceholder("Enter your password").fill(users[Role.ProxyAdmin].password);
const loginButton = page.getByRole("button", { name: "Login" });
await loginButton.click();
await page.waitForSelector("text=AI Gateway");
await page.context().storageState({ path: "admin.storageState.json" });
await loginAndSaveState(browser, users[Role.ProxyAdmin], ADMIN_STORAGE_PATH);
await loginAndSaveState(browser, users[Role.InternalUserViewer], INTERNAL_USER_VIEWER_STORAGE_PATH);
await browser.close();
}

View file

@ -0,0 +1,22 @@
import { test, expect } from "@playwright/test";
import { ADMIN_STORAGE_PATH } from "../../constants";
import { Page } from "../../fixtures/pages";
import { navigateToPage } from "../../helpers/navigation";
test.describe("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 page.keyboard.press("Escape");
await expect(page.getByText("e2eUITestingCreateKeyAllTeamModels")).toBeVisible();
});
});

View file

@ -3,7 +3,7 @@ import { users } from "../../fixtures/users";
import { Role } from "../../fixtures/roles";
test("user can log in", async ({ page }) => {
await page.goto("http://localhost:4000/ui/login");
await page.goto("/ui/login");
await page.getByPlaceholder("Enter your username").fill(users[Role.ProxyAdmin].email);
await page.getByPlaceholder("Enter your password").fill(users[Role.ProxyAdmin].password);
const loginButton = page.getByRole("button", { name: "Login" });

View file

@ -1,6 +1,6 @@
import test, { expect } from "@playwright/test";
import { Role } from "../../fixtures/roles";
import { ADMIN_STORAGE_PATH } from "../../constants";
import { ADMIN_STORAGE_PATH, INTERNAL_USER_VIEWER_STORAGE_PATH } from "../../constants";
import { Page } from "../../fixtures/pages";
import { menuLabelToPage } from "../../fixtures/menuMappings";
import { navigateToPage } from "../../helpers/navigation";
@ -8,17 +8,28 @@ import { navigateToPage } from "../../helpers/navigation";
const sidebarButtons = {
[Role.ProxyAdmin]: [
"Virtual Keys",
"MCP Servers",
"Playground",
"Models",
"Usage",
"Logs",
"Teams",
"Internal Users",
"API Reference",
"AI Hub",
],
[Role.InternalUserViewer]: [
"Virtual Keys",
"MCP Servers",
"Usage",
"Teams",
"Logs",
"API Reference",
"AI Hub",
],
};
const roles = [{ role: Role.ProxyAdmin, storage: ADMIN_STORAGE_PATH }];
const roles = [{ role: Role.ProxyAdmin, storage: ADMIN_STORAGE_PATH }, { role: Role.InternalUserViewer, storage: INTERNAL_USER_VIEWER_STORAGE_PATH }];
for (const { role, storage } of roles) {
test.describe(`${role} sidebar`, () => {

View file

@ -2,7 +2,7 @@ import { describe, it, expect, vi, beforeEach } from "vitest";
import { renderHook, waitFor } from "@testing-library/react";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import React, { ReactNode } from "react";
import { useKeys } from "./useKeys";
import { useKeys, useDeletedKeys } from "./useKeys";
import type { KeyResponse } from "@/components/key_team_helpers/key_list";
// Mock the networking utilities
@ -397,3 +397,293 @@ describe("useKeys", () => {
);
});
});
describe("useDeletedKeys", () => {
let queryClient: QueryClient;
beforeEach(() => {
queryClient = new QueryClient({
defaultOptions: {
queries: {
retry: false,
},
},
});
// Reset all mocks
vi.clearAllMocks();
// Set default mock for useAuthorized (enabled state)
mockUseAuthorized.mockReturnValue({
accessToken: "test-access-token",
userRole: "Admin",
userId: "test-user-id",
token: "test-token",
userEmail: "test@example.com",
premiumUser: false,
disabledPersonalKeyCreation: null,
showSSOBanner: false,
});
// Reset fetch mock
mockFetch.mockClear();
});
const wrapper = ({ children }: { children: ReactNode }) =>
React.createElement(QueryClientProvider, { client: queryClient }, children);
it("should return deleted keys data when query is successful", async () => {
// Mock successful API call
mockFetch.mockResolvedValueOnce({
ok: true,
json: async () => mockKeysResponse,
});
const { result } = renderHook(() => useDeletedKeys(1, 10), { wrapper });
// Initially loading
expect(result.current.isLoading).toBe(true);
expect(result.current.data).toBeUndefined();
// Wait for success
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
expect(result.current.isSuccess).toBe(true);
});
expect(result.current.data).toEqual(mockKeysResponse);
expect(result.current.error).toBeNull();
expect(mockFetch).toHaveBeenCalledTimes(1);
expect(mockFetch).toHaveBeenCalledWith(
"/key/list?page=1&size=10&status=deleted&return_full_object=true&include_team_keys=true&include_created_by_keys=true",
{
method: "GET",
headers: {
Authorization: "Bearer test-access-token",
"Content-Type": "application/json",
},
},
);
});
it("should pass status=deleted parameter to the API", async () => {
// Mock successful API call
mockFetch.mockResolvedValueOnce({
ok: true,
json: async () => mockKeysResponse,
});
const { result } = renderHook(() => useDeletedKeys(1, 10), { wrapper });
// Wait for success
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
});
// Verify that status=deleted is included in the URL
const callUrl = mockFetch.mock.calls[0][0];
expect(callUrl).toContain("status=deleted");
expect(result.current.data).toEqual(mockKeysResponse);
});
it("should handle error when deleted keys API call fails", async () => {
const errorMessage = "Failed to fetch deleted keys";
const errorResponse = { error: errorMessage };
// Mock failed API call
mockFetch.mockResolvedValueOnce({
ok: false,
json: async () => errorResponse,
});
const { result } = renderHook(() => useDeletedKeys(1, 10), { wrapper });
// Initially loading
expect(result.current.isLoading).toBe(true);
// Wait for error
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
expect(result.current.isError).toBe(true);
});
expect(result.current.error).toBeDefined();
expect(result.current.error?.message).toBe(errorMessage);
expect(result.current.data).toBeUndefined();
expect(mockFetch).toHaveBeenCalledTimes(1);
expect(mockFetch).toHaveBeenCalledWith(
"/key/list?page=1&size=10&status=deleted&return_full_object=true&include_team_keys=true&include_created_by_keys=true",
{
method: "GET",
headers: {
Authorization: "Bearer test-access-token",
"Content-Type": "application/json",
},
},
);
});
it("should not execute query when accessToken is missing", async () => {
// Mock missing accessToken
mockUseAuthorized.mockReturnValue({
accessToken: null,
userRole: "Admin",
userId: "test-user-id",
token: null,
userEmail: "test@example.com",
premiumUser: false,
disabledPersonalKeyCreation: null,
showSSOBanner: false,
});
const { result } = renderHook(() => useDeletedKeys(1, 10), { wrapper });
// Query should not execute
expect(result.current.isLoading).toBe(false);
expect(result.current.data).toBeUndefined();
expect(result.current.isFetched).toBe(false);
// API should not be called
expect(mockFetch).not.toHaveBeenCalled();
});
it("should pass correct page and pageSize parameters to the API", async () => {
// Mock successful API call
mockFetch.mockResolvedValueOnce({
ok: true,
json: async () => mockKeysResponse,
});
const page = 2;
const pageSize = 20;
const { result } = renderHook(() => useDeletedKeys(page, pageSize), { wrapper });
// Wait for success
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
});
expect(mockFetch).toHaveBeenCalledWith(
`/key/list?page=${page}&size=${pageSize}&status=deleted&return_full_object=true&include_team_keys=true&include_created_by_keys=true`,
{
method: "GET",
headers: {
Authorization: "Bearer test-access-token",
"Content-Type": "application/json",
},
},
);
});
it("should return empty deleted keys array when API returns empty data", async () => {
// Mock API returning empty keys array
const emptyResponse = {
keys: [],
total_count: 0,
current_page: 1,
total_pages: 0,
};
mockFetch.mockResolvedValueOnce({
ok: true,
json: async () => emptyResponse,
});
const { result } = renderHook(() => useDeletedKeys(1, 10), { wrapper });
// Wait for success
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
expect(result.current.isSuccess).toBe(true);
});
expect(result.current.data).toEqual(emptyResponse);
expect(mockFetch).toHaveBeenCalledWith(
"/key/list?page=1&size=10&status=deleted&return_full_object=true&include_team_keys=true&include_created_by_keys=true",
{
method: "GET",
headers: {
Authorization: "Bearer test-access-token",
"Content-Type": "application/json",
},
},
);
});
it("should handle network timeout error", async () => {
const timeoutError = new Error("Network timeout");
// Mock network timeout
mockFetch.mockRejectedValueOnce(timeoutError);
const { result } = renderHook(() => useDeletedKeys(1, 10), { wrapper });
// Wait for error
await waitFor(() => {
expect(result.current.isError).toBe(true);
});
expect(result.current.error).toEqual(timeoutError);
expect(result.current.data).toBeUndefined();
});
it("should handle pagination correctly", async () => {
const paginatedResponse = {
keys: [mockKeys[0]], // Only first key
total_count: 15,
current_page: 2,
total_pages: 2,
};
mockFetch.mockResolvedValueOnce({
ok: true,
json: async () => paginatedResponse,
});
const { result } = renderHook(() => useDeletedKeys(2, 10), { wrapper });
// Wait for success
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
});
expect(result.current.data).toEqual(paginatedResponse);
expect(mockFetch).toHaveBeenCalledWith(
"/key/list?page=2&size=10&status=deleted&return_full_object=true&include_team_keys=true&include_created_by_keys=true",
{
method: "GET",
headers: {
Authorization: "Bearer test-access-token",
"Content-Type": "application/json",
},
},
);
});
it("should pass additional options along with status=deleted", async () => {
// Mock successful API call
mockFetch.mockResolvedValueOnce({
ok: true,
json: async () => mockKeysResponse,
});
const options = {
organizationID: "org-1",
teamID: "team-1",
selectedKeyAlias: "test-alias",
};
const { result } = renderHook(() => useDeletedKeys(1, 10, options), { wrapper });
// Wait for success
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
});
const callUrl = mockFetch.mock.calls[0][0];
expect(callUrl).toContain("status=deleted");
expect(callUrl).toContain("organization_id=org-1");
expect(callUrl).toContain("team_id=team-1");
expect(callUrl).toContain("key_alias=test-alias");
expect(result.current.data).toEqual(mockKeysResponse);
});
});

View file

@ -101,12 +101,16 @@ const keyListCall = async (
}
};
export const useKeys = (page: number, pageSize: number): UseQueryResult<KeysResponse> => {
export const useKeys = (
page: number,
pageSize: number,
options: KeyListCallOptions = {},
): UseQueryResult<KeysResponse> => {
const { accessToken } = useAuthorized();
return useQuery<KeysResponse>({
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,

View file

@ -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();
});
});

View file

@ -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<PaginatedModelInfoResponse>({
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),
});
};

View file

@ -94,7 +94,8 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({ premiumUser, te
}, [modelDataResponse?.data]);
const allModelsOnProxy = useMemo<string[]>(() => {
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) => {

View file

@ -4,10 +4,14 @@ 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 +55,21 @@ const createModelCostMapMock = (data: Record<string, any>) => ({
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 +103,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 +153,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(<AllModelsTab {...defaultProps} />);
@ -191,28 +212,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(<AllModelsTab {...defaultProps} />);
@ -236,30 +255,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(<AllModelsTab {...defaultProps} />);
@ -283,42 +300,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(<AllModelsTab {...defaultProps} />);
@ -342,27 +357,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(<AllModelsTab {...defaultProps} />);
@ -370,4 +383,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(<AllModelsTab {...defaultProps} />);
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(<AllModelsTab {...defaultProps} />);
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();
});
});

View file

@ -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<string>("");
const [modelViewMode, setModelViewMode] = useState<ModelViewMode>("current_team");
const [currentTeam, setCurrentTeam] = useState<Team | "personal">("personal");
const [showFilters, setShowFilters] = useState<boolean>(false);
const [selectedModelAccessGroupFilter, setSelectedModelAccessGroupFilter] = useState<string | null>(null);
const [expandedRows, setExpandedRows] = useState<Set<string>>(new Set());
const [currentPage, setCurrentPage] = useState<number>(1);
const [pageSize] = useState<number>(50);
const [pagination, setPagination] = useState<PaginationState>({
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<string>("");
const [modelViewMode, setModelViewMode] = useState<ModelViewMode>("current_team");
const [currentTeam, setCurrentTeam] = useState<Team | "personal">("personal");
const [showFilters, setShowFilters] = useState<boolean>(false);
const [selectedModelAccessGroupFilter, setSelectedModelAccessGroupFilter] = useState<string | null>(null);
const [expandedRows, setExpandedRows] = useState<Set<string>>(new Set());
const [pagination, setPagination] = useState<PaginationState>({
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 = ({
) : (
<span className="text-sm text-gray-700">
{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"}
</span>
)}
@ -347,15 +365,16 @@ const AllModelsTab = ({
<Skeleton.Button active style={{ width: 84, height: 30 }} />
) : (
<button
onClick={() =>
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: prev.pageIndex - 1 }))
}
disabled={pagination.pageIndex === 0}
className={`px-3 py-1 text-sm border rounded-md ${
pagination.pageIndex === 0
? "bg-gray-100 text-gray-400 cursor-not-allowed"
: "hover:bg-gray-50"
}`}
onClick={() => {
const newPage = currentPage - 1;
setCurrentPage(newPage);
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
}}
disabled={currentPage === 1}
className={`px-3 py-1 text-sm border rounded-md ${currentPage === 1
? "bg-gray-100 text-gray-400 cursor-not-allowed"
: "hover:bg-gray-50"
}`}
>
Previous
</button>
@ -365,15 +384,16 @@ const AllModelsTab = ({
<Skeleton.Button active style={{ width: 56, height: 30 }} />
) : (
<button
onClick={() =>
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: prev.pageIndex + 1 }))
}
disabled={pagination.pageIndex >= Math.ceil(filteredData.length / pagination.pageSize) - 1}
className={`px-3 py-1 text-sm border rounded-md ${
pagination.pageIndex >= Math.ceil(filteredData.length / pagination.pageSize) - 1
? "bg-gray-100 text-gray-400 cursor-not-allowed"
: "hover:bg-gray-50"
}`}
onClick={() => {
const newPage = currentPage + 1;
setCurrentPage(newPage);
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
}}
disabled={currentPage >= paginationMeta.total_pages}
className={`px-3 py-1 text-sm border rounded-md ${currentPage >= paginationMeta.total_pages
? "bg-gray-100 text-gray-400 cursor-not-allowed"
: "hover:bg-gray-50"
}`}
>
Next
</button>
@ -391,8 +411,8 @@ const AllModelsTab = ({
setSelectedModelId,
setSelectedTeamId,
getDisplayModelName,
() => {},
() => {},
() => { },
() => { },
expandedRows,
setExpandedRows,
)}

View file

@ -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<Record<string, boolean>>({});
@ -110,6 +117,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
id: "expander",
header: () => null,
size: 40,
enableSorting: false,
cell: ({ row }) =>
row.getCanExpand() ? (
<button onClick={row.getToggleExpandedHandler()} style={{ cursor: "pointer" }}>
@ -122,6 +130,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "token",
header: "Key ID",
size: 150,
enableSorting: true,
cell: (info) => (
<div className="overflow-hidden">
<Tooltip title={info.getValue() as string}>
@ -142,6 +151,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "key_alias",
header: "Key Alias",
size: 150,
enableSorting: true,
cell: (info) => {
const value = info.getValue() as string;
const width = info.cell.column.getSize();
@ -159,6 +169,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "key_name",
header: "Secret Key",
size: 120,
enableSorting: false,
cell: (info) => <span className="font-mono text-xs">{info.getValue() as string}</span>,
},
{
@ -166,6 +177,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "team_id",
header: "Team Alias",
size: 120,
enableSorting: false,
cell: ({ row, getValue }) => {
const teamId = getValue() as string;
const team = teams?.find((t) => t.team_id === teamId);
@ -177,6 +189,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "team_id",
header: "Team ID",
size: 120,
enableSorting: false,
cell: (info) => (
<Tooltip title={info.getValue() as string}>
{info.getValue() ? `${(info.getValue() as string).slice(0, 7)}...` : "-"}
@ -188,6 +201,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "organization_id",
header: "Organization ID",
size: 140,
enableSorting: false,
cell: (info) => (info.getValue() ? info.renderValue() : "-"),
},
{
@ -195,6 +209,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "user",
header: "User Email",
size: 160,
enableSorting: false,
cell: (info) => {
const user = info.getValue() as any;
const value = user?.user_email;
@ -213,6 +228,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "user_id",
header: "User ID",
size: 120,
enableSorting: false,
cell: (info) => {
const userId = info.getValue() as string | null;
if (userId && userId.length > 15) {
@ -230,6 +246,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "created_at",
header: "Created At",
size: 120,
enableSorting: true,
cell: (info) => {
const value = info.getValue();
return value ? new Date(value as string).toLocaleDateString() : "-";
@ -240,6 +257,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "created_by",
header: "Created By",
size: 120,
enableSorting: false,
cell: (info) => {
const value = info.getValue() as string | null;
if (value && value.length > 15) {
@ -257,6 +275,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "updated_at",
header: "Updated At",
size: 120,
enableSorting: true,
cell: (info) => {
const value = info.getValue();
return value ? new Date(value as string).toLocaleDateString() : "Never";
@ -267,6 +286,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "expires",
header: "Expires",
size: 120,
enableSorting: false,
cell: (info) => {
const value = info.getValue();
return value ? new Date(value as string).toLocaleDateString() : "Never";
@ -277,6 +297,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "spend",
header: "Spend (USD)",
size: 100,
enableSorting: true,
cell: (info) => formatNumberWithCommas(info.getValue() as number, 4),
},
{
@ -284,6 +305,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "max_budget",
header: "Budget (USD)",
size: 110,
enableSorting: true,
cell: (info) => {
const maxBudget = info.getValue() as number | null;
if (maxBudget === null) {
@ -297,6 +319,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "budget_reset_at",
header: "Budget Reset",
size: 130,
enableSorting: false,
cell: (info) => {
const value = info.getValue();
return value ? new Date(value as string).toLocaleString() : "Never";
@ -307,6 +330,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
accessorKey: "models",
header: "Models",
size: 200,
enableSorting: false,
cell: (info) => {
const models = info.getValue() as string[];
return (
@ -391,6 +415,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
id: "rate_limits",
header: "Rate Limits",
size: 140,
enableSorting: false,
cell: ({ row }) => {
const key = row.original;
return (
@ -491,11 +516,16 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
const sortBy = sortState.id;
const sortOrder = sortState.desc ? "desc" : "asc";
console.log(`sortBy: ${sortBy}, sortOrder: ${sortOrder}`);
handleFilterChange({
...filters,
"Sort By": sortBy,
"Sort Order": sortOrder,
});
// Update filters state without triggering debouncedSearch
// The useKeys hook will automatically refetch with the new sort parameters
handleFilterChange(
{
...filters,
"Sort By": sortBy,
"Sort Order": sortOrder,
},
true, // skipDebounce - let useKeys handle the API call with correct page size
);
onSortChange?.(sortBy, sortOrder);
}
},
@ -601,12 +631,13 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
key={header.id}
data-header-id={header.id}
className={`py-1 h-8 relative hover:bg-gray-50 ${header.id === "actions"
? "sticky right-0 bg-white shadow-[-4px_0_8px_-6px_rgba(0,0,0,0.1)]"
: ""
? "sticky right-0 bg-white shadow-[-4px_0_8px_-6px_rgba(0,0,0,0.1)]"
: ""
}`}
style={{
width: header.getSize(),
position: "relative",
cursor: header.column.getCanSort() ? "pointer" : "default",
}}
onMouseEnter={() => {
const resizer = document.querySelector(`[data-header-id="${header.id}"] .resizer`);
@ -620,7 +651,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
(resizer as HTMLElement).style.opacity = "0";
}
}}
onClick={header.column.getToggleSortingHandler()}
onClick={header.column.getCanSort() ? header.column.getToggleSortingHandler() : undefined}
>
<div className="flex items-center justify-between gap-2">
<div className="flex items-center">
@ -628,7 +659,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
? null
: flexRender(header.column.columnDef.header, header.getContext())}
</div>
{header.id !== "actions" && (
{header.id !== "actions" && header.column.getCanSort() && (
<div className="w-4">
{header.column.getIsSorted() ? (
{

View file

@ -151,7 +151,7 @@ export function useFilterLogic({
}
}, [organizations]);
const handleFilterChange = (newFilters: Record<string, string>) => {
const handleFilterChange = (newFilters: Record<string, string>, skipDebounce: boolean = false) => {
// Update filters state
setFilters({
"Team ID": newFilters["Team ID"] || "",
@ -162,12 +162,16 @@ export function useFilterLogic({
"Sort Order": newFilters["Sort Order"] || "desc",
});
// Fetch keys based on new filters
const updatedFilters = {
...filters,
...newFilters,
};
debouncedSearch(updatedFilters);
// Only trigger debouncedSearch if skipDebounce is false
// This allows sorting to be handled by the parent component's useKeys hook
if (!skipDebounce) {
// Fetch keys based on new filters
const updatedFilters = {
...filters,
...newFilters,
};
debouncedSearch(updatedFilters);
}
};
const handleFilterReset = () => {

View file

@ -36,6 +36,8 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
const [editing, setEditing] = useState(isEditing);
const [showFullUrl, setShowFullUrl] = useState(false);
const [copiedStates, setCopiedStates] = useState<Record<string, boolean>>({});
const [selectedTabIndex, setSelectedTabIndex] = useState(0);
const handleSuccess = (updated: MCPServer) => {
setEditing(false);
onBack();
@ -72,11 +74,10 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
size="small"
icon={copiedStates["mcp-server_name"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
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<MCPServerViewProps> = ({
size="small"
icon={copiedStates["mcp-alias"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
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<MCPServerViewProps> = ({
size="small"
icon={copiedStates["mcp-server-id"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
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"
}`}
/>
</div>
</div>
</div>
{/* TODO: magic number for index */}
<TabGroup defaultIndex={editing ? 2 : 0}>
<TabGroup index={selectedTabIndex} onIndexChange={setSelectedTabIndex}>
<TabList className="mb-4">
{[
<Tab key="overview">Overview</Tab>,

View file

@ -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";
@ -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();
@ -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(
<QueryClientProvider client={queryClient}>
<MCPServers {...defaultProps} />
</QueryClientProvider>,
);
// 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();
});
});

View file

@ -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<MCPServerProps> = ({ 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<MCPServerProps> = ({ 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,109 +207,34 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
setModalVisible(false);
};
// 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]);
if (!accessToken || !userRole || !userID) {
console.log("Missing required authentication parameters", { accessToken, userRole, userID });
return <div className="p-6 text-center text-gray-500">Missing required authentication parameters.</div>;
}
const ServersTab = () =>
selectedServerId ? (
<MCPServerView
mcpServer={
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: "",
}
}
onBack={() => {
setEditServer(false);
setSelectedServerId(null);
refetch();
}}
isProxyAdmin={isAdminRole(userRole)}
isEditing={editServer}
accessToken={accessToken}
userID={userID}
userRole={userRole}
availableAccessGroups={uniqueMcpAccessGroups}
/>
) : (
<div className="w-full h-full">
<div className="w-full px-6">
<div className="flex flex-col space-y-4">
<div className="flex items-center justify-between bg-gray-50 rounded-lg p-4 border-2 border-gray-200">
<div className="flex items-center gap-4">
<Text className="text-lg font-semibold text-gray-900">Current Team:</Text>
<Select value={selectedTeam} onChange={handleTeamChange} style={{ width: 300 }}>
<Option value="all">
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-blue-500 rounded-full"></div>
<span className="font-medium">{isInternalUser ? "All Available Servers" : "All Servers"}</span>
</div>
</Option>
<Option value="personal">
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-green-500 rounded-full"></div>
<span className="font-medium">Personal</span>
</div>
</Option>
{uniqueTeams.map((team) => (
<Option key={team.team_id} value={team.team_id}>
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-green-500 rounded-full"></div>
<span className="font-medium">{team.team_alias || team.team_id}</span>
</div>
</Option>
))}
</Select>
<Text className="text-lg font-semibold text-gray-900 ml-6">
Access Group:
<Tooltip title="An MCP Access Group is a set of users or teams that have permission to access specific MCP servers. Use access groups to control and organize who can connect to which servers.">
<QuestionCircleOutlined style={{ marginLeft: 4, color: "#888" }} />
</Tooltip>
</Text>
<Select value={selectedMcpAccessGroup} onChange={handleMcpAccessGroupChange} style={{ width: 300 }}>
<Option value="all">
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-blue-500 rounded-full"></div>
<span className="font-medium">All Access Groups</span>
</div>
</Option>
{uniqueMcpAccessGroups.map((group) => (
<Option key={group} value={group}>
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-green-500 rounded-full"></div>
<span className="font-medium">{group}</span>
</div>
</Option>
))}
</Select>
</div>
</div>
</div>
</div>
<div className="w-full px-6 mt-6">
<DataTable
data={filteredServers}
columns={columns}
renderSubComponent={() => <div></div>}
getRowCanExpand={() => false}
isLoading={isLoadingServers}
noDataMessage="No MCP servers configured"
loadingMessage="🚅 Loading MCP servers..."
/>
</div>
</div>
);
return (
<div className="w-full h-full p-6">
<Modal
@ -381,7 +306,86 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
</TabList>
<TabPanels>
<TabPanel>
<ServersTab />
{selectedServerId ? (
<MCPServerView
key={selectedServerId}
mcpServer={selectedServer}
onBack={handleBack}
isProxyAdmin={isAdminRole(userRole)}
isEditing={editServer}
accessToken={accessToken}
userID={userID}
userRole={userRole}
availableAccessGroups={uniqueMcpAccessGroups}
/>
) : (
<div className="w-full h-full">
<div className="w-full px-6">
<div className="flex flex-col space-y-4">
<div className="flex items-center justify-between bg-gray-50 rounded-lg p-4 border-2 border-gray-200">
<div className="flex items-center gap-4">
<Text className="text-lg font-semibold text-gray-900">Current Team:</Text>
<Select value={selectedTeam} onChange={handleTeamChange} style={{ width: 300 }}>
<Option value="all">
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-blue-500 rounded-full"></div>
<span className="font-medium">{isInternalUser ? "All Available Servers" : "All Servers"}</span>
</div>
</Option>
<Option value="personal">
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-green-500 rounded-full"></div>
<span className="font-medium">Personal</span>
</div>
</Option>
{uniqueTeams.map((team) => (
<Option key={team.team_id} value={team.team_id}>
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-green-500 rounded-full"></div>
<span className="font-medium">{team.team_alias || team.team_id}</span>
</div>
</Option>
))}
</Select>
<Text className="text-lg font-semibold text-gray-900 ml-6">
Access Group:
<Tooltip title="An MCP Access Group is a set of users or teams that have permission to access specific MCP servers. Use access groups to control and organize who can connect to which servers.">
<QuestionCircleOutlined style={{ marginLeft: 4, color: "#888" }} />
</Tooltip>
</Text>
<Select value={selectedMcpAccessGroup} onChange={handleMcpAccessGroupChange} style={{ width: 300 }}>
<Option value="all">
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-blue-500 rounded-full"></div>
<span className="font-medium">All Access Groups</span>
</div>
</Option>
{uniqueMcpAccessGroups.map((group) => (
<Option key={group} value={group}>
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-green-500 rounded-full"></div>
<span className="font-medium">{group}</span>
</div>
</Option>
))}
</Select>
</div>
</div>
</div>
</div>
<div className="w-full px-6 mt-6">
<DataTable
data={filteredServers}
columns={columns}
renderSubComponent={() => <div></div>}
getRowCanExpand={() => false}
isLoading={isLoadingServers}
noDataMessage="No MCP servers configured"
loadingMessage="🚅 Loading MCP servers..."
/>
</div>
</div>
)}
</TabPanel>
<TabPanel>
<MCPConnect />

View file

@ -127,11 +127,10 @@ const MCPToolsViewer = ({
{toolsData.map((tool: MCPTool) => (
<div
key={tool.name}
className={`border rounded-lg p-3 cursor-pointer transition-all hover:shadow-sm ${
selectedTool?.name === tool.name
className={`border rounded-lg p-3 cursor-pointer transition-all hover:shadow-sm ${selectedTool?.name === tool.name
? "border-blue-500 bg-blue-50 ring-1 ring-blue-200"
: "border-gray-200 bg-white hover:border-gray-300"
}`}
}`}
onClick={() => {
setSelectedTool(tool);
setToolResult(null);

View file

@ -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()}`;
}

View file

@ -7,6 +7,7 @@ export default defineConfig({
setupFiles: ["tests/setupTests.ts"],
globals: true,
css: true, // lets you import CSS/modules without extra mocks
testTimeout: 10000,
coverage: {
provider: "v8",
reporter: ["text", "lcov"],