mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge branch 'main' into litellm_add_output_format_claude
This commit is contained in:
commit
936fe3c94d
48 changed files with 4057 additions and 635 deletions
|
|
@ -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.
|
||||
:::
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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="",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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[
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
7
proxy_config.yaml
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
model_list:
|
||||
- model_name: "*"
|
||||
litellm_params:
|
||||
model: "*"
|
||||
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -1 +1,2 @@
|
|||
export const ADMIN_STORAGE_PATH = "admin.storageState.json";
|
||||
export const INTERNAL_USER_VIEWER_STORAGE_PATH = "internalViewer.storageState.json";
|
||||
|
|
@ -7,4 +7,8 @@ export const users = {
|
|||
email: "admin",
|
||||
password: isCI ? "gm" : "sk-1234",
|
||||
},
|
||||
[Role.InternalUserViewer]: {
|
||||
email: "internalViewer@test.com",
|
||||
password: "test",
|
||||
},
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
|
||||
|
|
|
|||
22
ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts
Normal file
22
ui/litellm-dashboard/e2e_tests/tests/keys/createKey.spec.ts
Normal 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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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" });
|
||||
|
|
|
|||
|
|
@ -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`, () => {
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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),
|
||||
});
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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) => {
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -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() ? (
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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 = () => {
|
||||
|
|
|
|||
|
|
@ -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>,
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 />
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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()}`;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue