Merge branch 'BerriAI:main' into LangfuseUsageDetails

This commit is contained in:
Fabrício Ceschin 2025-09-19 13:54:56 -04:00 • committed by GitHub
commit f0490ab60d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
104 changed files with 8386 additions and 3387 deletions

View file

@ -1458,6 +1458,7 @@ jobs:
# - run: python ./tests/documentation_tests/test_general_setting_keys.py
- run: python ./tests/code_coverage_tests/check_licenses.py
- run: python ./tests/code_coverage_tests/router_code_coverage.py
- run: python ./tests/code_coverage_tests/test_chat_completion_imports.py
- run: python ./tests/code_coverage_tests/info_log_check.py
- run: python ./tests/code_coverage_tests/test_ban_set_verbose.py
- run: python ./tests/code_coverage_tests/code_qa_check_tests.py
@ -2825,8 +2826,8 @@ jobs:
source "$NVM_DIR/bash_completion"
# Install and use Node version
nvm install v18.17.0
nvm use v18.17.0
nvm install v20
nvm use v20
cd ui/litellm-dashboard
@ -2879,7 +2880,26 @@ jobs:
name: Install Playwright Browsers
command: |
npx playwright install
- run:
name: Run UI unit tests (Vitest)
command: |
# Use Node 20 (several deps require >=20)
export NVM_DIR="/opt/circleci/.nvm"
source "$NVM_DIR/nvm.sh"
nvm install 20
nvm use 20
cd ui/litellm-dashboard
npm ci || npm install
# CI run, with both LCOV (Codecov) and HTML (artifact you can click)
CI=true npm run test -- --run --coverage \
--coverage.provider=v8 \
--coverage.reporter=lcov \
--coverage.reporter=html \
--coverage.reportsDirectory=coverage/html
- run:
name: Build Docker image
command: docker build -t my-app:latest -f ./docker/Dockerfile.database .

View file

@ -114,7 +114,6 @@ mcp_servers:
description: "My custom MCP server"
auth_type: "api_key"
auth_value: "abc123"
spec_version: "2025-03-26"
```
**Configuration Options:**
@ -716,7 +715,6 @@ mcp_servers:
url: https://mcp.deepwiki.com/mcp
transport: "http"
auth_type: "none"
spec_version: "2025-03-26"
access_groups: ["dev_group"]
```

View file

@ -237,7 +237,10 @@ litellm.metadata = {
}
```
### Session Tracking and Tracing
</TabItem>
</Tabs>
## Session Tracking and Tracing
Track multi-step and agentic LLM interactions using session IDs and paths:

View file

@ -1821,6 +1821,7 @@ Here's an example of using a bedrock model with LiteLLM. For a complete list, re
| Mistral 7B Instruct | `completion(model='bedrock/mistral.mistral-7b-instruct-v0:2', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` |
| Mixtral 8x7B Instruct | `completion(model='bedrock/mistral.mixtral-8x7b-instruct-v0:1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` |
## Bedrock Embedding
### API keys
@ -1842,11 +1843,29 @@ response = embedding(
print(response)
```
#### Titan V2 - encoding_format support
```python
from litellm import embedding
# Float format (default)
response = embedding(
model="bedrock/amazon.titan-embed-text-v2:0",
input=["good morning from litellm"],
encoding_format="float" # Returns float array
)
# Binary format
response = embedding(
model="bedrock/amazon.titan-embed-text-v2:0",
input=["good morning from litellm"],
encoding_format="base64" # Returns base64 encoded binary
)
```
## Supported AWS Bedrock Embedding Models
| Model Name | Usage | Supported Additional OpenAI params |
|----------------------|---------------------------------------------|-----|
| Titan Embeddings V2 | `embedding(model="bedrock/amazon.titan-embed-text-v2:0", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py#L59) |
| Titan Embeddings V2 | `embedding(model="bedrock/amazon.titan-embed-text-v2:0", input=input)` | `dimensions`, `encoding_format` |
| Titan Embeddings - V1 | `embedding(model="bedrock/amazon.titan-embed-text-v1", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py#L53)
| Titan Multimodal Embeddings | `embedding(model="bedrock/amazon.titan-embed-image-v1", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/amazon_titan_multimodal_transformation.py#L28) |
| Cohere Embeddings - English | `embedding(model="bedrock/cohere.embed-english-v3", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/cohere_transformation.py#L18)

View file

@ -0,0 +1,95 @@
# Bedrock Embedding
## Supported Embedding Models
| Provider | LiteLLM Route | AWS Documentation |
|----------|---------------|-------------------|
| Amazon Titan | `bedrock/amazon.*` | [Amazon Titan Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/titan-embedding-models.html) |
| Cohere | `bedrock/cohere.*` | [Cohere Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-cohere-embed.html) |
| TwelveLabs | `bedrock/us.twelvelabs.*` | [TwelveLabs](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-twelvelabs.html) |
### API keys
This can be set as env variables or passed as **params to litellm.embedding()**
```python
import os
os.environ["AWS_ACCESS_KEY_ID"] = "" # Access key
os.environ["AWS_SECRET_ACCESS_KEY"] = "" # Secret access key
os.environ["AWS_REGION_NAME"] = "" # us-east-1, us-east-2, us-west-1, us-west-2
```
## Usage
### LiteLLM Python SDK
```python
from litellm import embedding
response = embedding(
model="bedrock/amazon.titan-embed-text-v1",
input=["good morning from litellm"],
)
print(response)
```
### LiteLLM Proxy Server
#### 1. Setup config.yaml
```yaml
model_list:
- model_name: titan-embed-v1
litellm_params:
model: bedrock/amazon.titan-embed-text-v1
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: us-east-1
- model_name: titan-embed-v2
litellm_params:
model: bedrock/amazon.titan-embed-text-v2:0
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: us-east-1
```
#### 2. Start Proxy
```bash
litellm --config /path/to/config.yaml
```
#### 3. Use with OpenAI Python SDK
```python
import openai
client = openai.OpenAI(
api_key="anything",
base_url="http://0.0.0.0:4000"
)
response = client.embeddings.create(
input=["good morning from litellm"],
model="titan-embed-v1"
)
print(response)
```
#### 4. Use with LiteLLM Python SDK
```python
import litellm
response = litellm.embedding(
model="titan-embed-v1", # model alias from config.yaml
input=["good morning from litellm"],
api_base="http://0.0.0.0:4000",
api_key="anything"
)
print(response)
```
## Supported AWS Bedrock Embedding Models
| Model Name | Usage | Supported Additional OpenAI params |
|----------------------|---------------------------------------------|-----|
| Titan Embeddings V2 | `embedding(model="bedrock/amazon.titan-embed-text-v2:0", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py#L59) |
| Titan Embeddings - V1 | `embedding(model="bedrock/amazon.titan-embed-text-v1", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py#L53)
| Titan Multimodal Embeddings | `embedding(model="bedrock/amazon.titan-embed-image-v1", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/amazon_titan_multimodal_transformation.py#L28) |
| TwelveLabs Marengo Embed 2.7 | `embedding(model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0", input=input)` | Supports multimodal input (text, video, audio, image) |
| Cohere Embeddings - English | `embedding(model="bedrock/cohere.embed-english-v3", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/cohere_transformation.py#L18)
| Cohere Embeddings - Multilingual | `embedding(model="bedrock/cohere.embed-multilingual-v3", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/cohere_transformation.py#L18)
### Advanced - [Drop Unsupported Params](https://docs.litellm.ai/docs/completion/drop_params#openai-proxy-usage)
### Advanced - [Pass model/provider-specific Params](https://docs.litellm.ai/docs/completion/provider_specific_params#proxy-usage)

View file

@ -29,5 +29,6 @@ Common timezone values:
- `US/Pacific` - Pacific Time
- `Europe/London` - UK Time
- `Asia/Kolkata` - Indian Standard Time (IST)
- `Asia/Bangkok` - Indochina Time (ICT)
- `Asia/Tokyo` - Japan Standard Time
- `Australia/Sydney` - Australian Eastern Time

View file

@ -31,7 +31,12 @@ This release is not yet live.
<TabItem value="docker" label="Docker">
``` showLineNumbers title="docker run litellm"
docker run \
-e STORE_MODEL_IN_DB=True \
-p 4000:4000 \
ghcr.io/berriai/litellm:main-v1.77.2.rc.2
```
</TabItem>
<TabItem value="pip" label="Pip">

View file

@ -411,6 +411,7 @@ const sidebars = {
label: "Bedrock",
items: [
"providers/bedrock",
"providers/bedrock_embedding",
"providers/bedrock_agents",
"providers/bedrock_batches",
"providers/bedrock_vector_store",

View file

@ -102,7 +102,9 @@ class PrometheusLogger(CustomLogger):
# "team",
# "team_alias",
# ],
labelnames=self.get_labels_for_metric("litellm_llm_api_time_to_first_token_metric"),
labelnames=self.get_labels_for_metric(
"litellm_llm_api_time_to_first_token_metric"
),
buckets=LATENCY_BUCKETS,
)
@ -240,14 +242,14 @@ class PrometheusLogger(CustomLogger):
self.litellm_deployment_state = self._gauge_factory(
"litellm_deployment_state",
"LLM Deployment Analytics - The state of the deployment: 0 = healthy, 1 = partial outage, 2 = complete outage",
labelnames=self.get_labels_for_metric("litellm_deployment_state")
labelnames=self.get_labels_for_metric("litellm_deployment_state"),
)
self.litellm_deployment_cooled_down = self._counter_factory(
"litellm_deployment_cooled_down",
"LLM Deployment Analytics - Number of times a deployment has been cooled down by LiteLLM load balancing logic. exception_status is the status of the exception that caused the deployment to be cooled down",
# labelnames=_logged_llm_labels + [EXCEPTION_STATUS],
labelnames=self.get_labels_for_metric("litellm_deployment_cooled_down")
labelnames=self.get_labels_for_metric("litellm_deployment_cooled_down"),
)
self.litellm_deployment_success_responses = self._counter_factory(
@ -1039,20 +1041,12 @@ class PrometheusLogger(CustomLogger):
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_proxy_total_requests_metric"
metric_name="litellm_spend_metric"
),
enum_values=enum_values,
)
self.litellm_spend_metric.labels(
end_user_id,
user_api_key,
user_api_key_alias,
model,
user_api_team,
user_api_team_alias,
user_id,
).inc(response_cost)
self.litellm_spend_metric.labels(**_labels).inc(response_cost)
def _set_virtual_key_rate_limit_metrics(
self,
@ -2280,7 +2274,9 @@ def get_custom_labels_from_metadata(metadata: dict) -> Dict[str, str]:
return result
def _tag_matches_wildcard_configured_pattern(tags: List[str], configured_tag: str) -> bool:
def _tag_matches_wildcard_configured_pattern(
tags: List[str], configured_tag: str
) -> bool:
"""
Check if any of the request tags matches a wildcard configured pattern
@ -2305,6 +2301,7 @@ def _tag_matches_wildcard_configured_pattern(tags: List[str], configured_tag: st
import re
from litellm.router_utils.pattern_match_deployments import PatternMatchRouter
pattern_router = PatternMatchRouter()
regex_pattern = pattern_router._pattern_to_regex(configured_tag)
return any(re.match(pattern=regex_pattern, string=tag) for tag in tags)
@ -2313,11 +2310,11 @@ def _tag_matches_wildcard_configured_pattern(tags: List[str], configured_tag: st
def get_custom_labels_from_tags(tags: List[str]) -> Dict[str, str]:
"""
Get custom labels from tags based on admin configuration.
Supports both exact matches and wildcard patterns:
- Exact match: "prod" matches "prod" exactly
- Wildcard pattern: "User-Agent: curl/*" matches "User-Agent: curl/7.68.0"
- Wildcard pattern: "User-Agent: curl/*" matches "User-Agent: curl/7.68.0"
Reuses PatternMatchRouter for wildcard pattern matching.
Returns dict of label_name: "true" if the tag matches the configured tag, "false" otherwise
@ -2345,17 +2342,19 @@ def get_custom_labels_from_tags(tags: List[str]) -> Dict[str, str]:
for configured_tag in configured_tags:
label_name = _sanitize_prometheus_label_name(f"tag_{configured_tag}")
# Check for exact match first (backwards compatibility)
if configured_tag in tags:
result[label_name] = "true"
continue
# Use PatternMatchRouter for wildcard pattern matching
if "*" in configured_tag and _tag_matches_wildcard_configured_pattern(tags=tags, configured_tag=configured_tag):
if "*" in configured_tag and _tag_matches_wildcard_configured_pattern(
tags=tags, configured_tag=configured_tag
):
result[label_name] = "true"
continue
# No match found
result[label_name] = "false"

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm-enterprise"
version = "0.1.19"
version = "0.1.20"
description = "Package for LiteLLM Enterprise features"
authors = ["BerriAI"]
readme = "README.md"
@ -22,7 +22,7 @@ requires = ["poetry-core"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "0.1.19"
version = "0.1.20"
version_files = [
"pyproject.toml:version",
"../requirements.txt:litellm-enterprise==",

View file

@ -0,0 +1,8 @@
/*
Warnings:
- You are about to drop the column `spec_version` on the `LiteLLM_MCPServerTable` table. All the data in the column will be lost.
*/
-- AlterTable
ALTER TABLE "public"."LiteLLM_MCPServerTable" DROP COLUMN "spec_version";

View file

@ -171,7 +171,6 @@ model LiteLLM_MCPServerTable {
description String?
url String?
transport String @default("sse")
spec_version String @default("2025-03-26")
auth_type String?
created_at DateTime? @default(now()) @map("created_at")
created_by String?

View file

@ -67,6 +67,7 @@ from litellm.constants import (
bedrock_embedding_models,
known_tokenizer_config,
BEDROCK_INVOKE_PROVIDERS_LITERAL,
BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
BEDROCK_CONVERSE_MODELS,
DEFAULT_MAX_TOKENS,
DEFAULT_SOFT_BUDGET,

View file

@ -769,6 +769,12 @@ BEDROCK_INVOKE_PROVIDERS_LITERAL = Literal[
"deepseek_r1",
]
BEDROCK_EMBEDDING_PROVIDERS_LITERAL = Literal[
"cohere",
"amazon",
"twelvelabs",
]
BEDROCK_CONVERSE_MODELS = [
"openai.gpt-oss-20b-1:0",
"openai.gpt-oss-120b-1:0",

View file

@ -19,8 +19,6 @@ from litellm._logging import verbose_logger
from litellm.types.mcp import (
MCPAuth,
MCPAuthType,
MCPSpecVersion,
MCPSpecVersionType,
MCPStdioConfig,
MCPTransport,
MCPTransportType,
@ -48,7 +46,6 @@ class MCPClient:
auth_value: Optional[str] = None,
timeout: float = 60.0,
stdio_config: Optional[MCPStdioConfig] = None,
protocol_version: MCPSpecVersionType = MCPSpecVersion.jun_2025,
):
self.server_url: str = server_url
self.transport_type: MCPTransport = transport_type
@ -62,7 +59,6 @@ class MCPClient:
self._session_ctx = None
self._task: Optional[asyncio.Task] = None
self.stdio_config: Optional[MCPStdioConfig] = stdio_config
self.protocol_version: MCPSpecVersionType = protocol_version
# handle the basic auth value if provided
if auth_value:
@ -84,22 +80,24 @@ class MCPClient:
"""Initialize the transport and session."""
if self._session:
return # Already connected
try:
if self.transport_type == MCPTransport.stdio:
# For stdio transport, use stdio_client with command-line parameters
if not self.stdio_config:
raise ValueError("stdio_config is required for stdio transport")
server_params = StdioServerParameters(
command=self.stdio_config.get("command", ""),
args=self.stdio_config.get("args", []),
env=self.stdio_config.get("env", {})
env=self.stdio_config.get("env", {}),
)
self._transport_ctx = stdio_client(server_params)
self._transport = await self._transport_ctx.__aenter__()
self._session_ctx = ClientSession(self._transport[0], self._transport[1])
self._session_ctx = ClientSession(
self._transport[0], self._transport[1]
)
self._session = await self._session_ctx.__aenter__()
await self._session.initialize()
elif self.transport_type == MCPTransport.sse:
@ -110,7 +108,9 @@ class MCPClient:
headers=headers,
)
self._transport = await self._transport_ctx.__aenter__()
self._session_ctx = ClientSession(self._transport[0], self._transport[1])
self._session_ctx = ClientSession(
self._transport[0], self._transport[1]
)
self._session = await self._session_ctx.__aenter__()
await self._session.initialize()
else: # http
@ -121,7 +121,9 @@ class MCPClient:
headers=headers,
)
self._transport = await self._transport_ctx.__aenter__()
self._session_ctx = ClientSession(self._transport[0], self._transport[1])
self._session_ctx = ClientSession(
self._transport[0], self._transport[1]
)
self._session = await self._session_ctx.__aenter__()
await self._session.initialize()
except ValueError as e:
@ -185,7 +187,7 @@ class MCPClient:
def _get_auth_headers(self) -> dict:
"""Generate authentication headers based on auth type."""
headers = {}
if self._mcp_auth_value:
if self.auth_type == MCPAuth.bearer_token:
headers["Authorization"] = f"Bearer {self._mcp_auth_value}"
@ -196,18 +198,8 @@ class MCPClient:
elif self.auth_type == MCPAuth.authorization:
headers["Authorization"] = self._mcp_auth_value
# Handle protocol version - it might be a string or enum
if hasattr(self.protocol_version, 'value'):
# It's an enum
protocol_version_str = self.protocol_version.value
else:
# It's a string
protocol_version_str = str(self.protocol_version)
headers["MCP-Protocol-Version"] = protocol_version_str
return headers
async def list_tools(self) -> List[MCPTool]:
"""List available tools from the server."""
if not self._session:
@ -216,7 +208,7 @@ class MCPClient:
except Exception as e:
verbose_logger.warning(f"MCP client connection failed: {str(e)}")
return []
if self._session is None:
verbose_logger.warning("MCP client session is not initialized")
return []
@ -245,17 +237,20 @@ class MCPClient:
except Exception as e:
verbose_logger.warning(f"MCP client connection failed: {str(e)}")
return MCPCallToolResult(
content=[TextContent(type="text", text=f"{str(e)}")],
isError=True
content=[TextContent(type="text", text=f"{str(e)}")], isError=True
)
if self._session is None:
verbose_logger.warning("MCP client session is not initialized")
return MCPCallToolResult(
content=[TextContent(type="text", text="MCP client session is not initialized")],
content=[
TextContent(
type="text", text="MCP client session is not initialized"
)
],
isError=True,
)
try:
tool_result = await self._session.call_tool(
name=call_tool_request_params.name,
@ -270,8 +265,8 @@ class MCPClient:
await self.disconnect()
# Return a default error result instead of raising
return MCPCallToolResult(
content=[TextContent(type="text", text=f"{str(e)}")], # Empty content for error case
content=[
TextContent(type="text", text=f"{str(e)}")
], # Empty content for error case
isError=True,
)

View file

@ -357,6 +357,7 @@ class CustomGuardrail(CustomLogger):
end_time: Optional[float] = None,
duration: Optional[float] = None,
masked_entity_count: Optional[Dict[str, int]] = None,
guardrail_provider: Optional[str] = None,
) -> None:
"""
Builds `StandardLoggingGuardrailInformation` and adds it to the request metadata so it can be used for logging to DataDog, Langfuse, etc.
@ -367,6 +368,7 @@ class CustomGuardrail(CustomLogger):
slg = StandardLoggingGuardrailInformation(
guardrail_name=self.guardrail_name,
guardrail_provider=guardrail_provider,
guardrail_mode=(
GuardrailMode(**self.event_hook.model_dump()) # type: ignore
if isinstance(self.event_hook, Mode)

View file

@ -498,6 +498,7 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
"guardrail_information": standard_logging_payload.get(
"guardrail_information", None
),
"is_streamed_request": self._get_stream_value_from_payload(standard_logging_payload),
}
#########################################################
@ -561,6 +562,31 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
return latency_metrics
def _get_stream_value_from_payload(self, standard_logging_payload: StandardLoggingPayload) -> bool:
"""
Extract the stream value from standard logging payload.
The stream field in StandardLoggingPayload is only set to True for completed streaming responses.
For non-streaming requests, it's None. The original stream parameter is in model_parameters.
Returns:
bool: True if this was a streaming request, False otherwise
"""
# Check top-level stream field first (only True for completed streaming)
stream_value = standard_logging_payload.get("stream")
if stream_value is True:
return True
# Fallback to model_parameters.stream for original request parameters
model_params = standard_logging_payload.get("model_parameters", {})
if isinstance(model_params, dict):
stream_value = model_params.get("stream")
if stream_value is True:
return True
# Default to False for non-streaming requests
return False
def _get_spend_metrics(
self, standard_logging_payload: StandardLoggingPayload
) -> DDLLMObsSpendMetrics:

View file

@ -39,6 +39,7 @@ class LangsmithLogger(CustomBatchLogger):
langsmith_api_key: Optional[str] = None,
langsmith_project: Optional[str] = None,
langsmith_base_url: Optional[str] = None,
langsmith_sampling_rate: Optional[float] = None,
**kwargs,
):
self.flush_lock = asyncio.Lock()
@ -49,7 +50,8 @@ class LangsmithLogger(CustomBatchLogger):
langsmith_base_url=langsmith_base_url,
)
self.sampling_rate: float = (
float(os.getenv("LANGSMITH_SAMPLING_RATE")) # type: ignore
langsmith_sampling_rate
or float(os.getenv("LANGSMITH_SAMPLING_RATE")) # type: ignore
if os.getenv("LANGSMITH_SAMPLING_RATE") is not None
and os.getenv("LANGSMITH_SAMPLING_RATE").strip().isdigit() # type: ignore
else 1.0
@ -76,26 +78,14 @@ class LangsmithLogger(CustomBatchLogger):
langsmith_base_url: Optional[str] = None,
) -> LangsmithCredentialsObject:
_credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY")
if _credentials_api_key is None:
raise Exception(
"Invalid Langsmith API Key given. _credentials_api_key=None."
)
_credentials_project = (
langsmith_project or os.getenv("LANGSMITH_PROJECT") or "litellm-completion"
)
if _credentials_project is None:
raise Exception(
"Invalid Langsmith API Key given. _credentials_project=None."
)
_credentials_base_url = (
langsmith_base_url
or os.getenv("LANGSMITH_BASE_URL")
or "https://api.smith.langchain.com"
)
if _credentials_base_url is None:
raise Exception(
"Invalid Langsmith API Key given. _credentials_base_url=None."
)
return LangsmithCredentialsObject(
LANGSMITH_API_KEY=_credentials_api_key,
@ -200,12 +190,7 @@ class LangsmithLogger(CustomBatchLogger):
def log_success_event(self, kwargs, response_obj, start_time, end_time):
try:
sampling_rate = (
float(os.getenv("LANGSMITH_SAMPLING_RATE")) # type: ignore
if os.getenv("LANGSMITH_SAMPLING_RATE") is not None
and os.getenv("LANGSMITH_SAMPLING_RATE").strip().isdigit() # type: ignore
else 1.0
)
sampling_rate = self._get_sampling_rate_to_use_for_request(kwargs=kwargs)
random_sample = random.random()
if random_sample > sampling_rate:
verbose_logger.info(
@ -219,6 +204,7 @@ class LangsmithLogger(CustomBatchLogger):
kwargs,
response_obj,
)
credentials = self._get_credentials_to_use_for_request(kwargs=kwargs)
data = self._prepare_log_data(
kwargs=kwargs,
@ -245,7 +231,7 @@ class LangsmithLogger(CustomBatchLogger):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
try:
sampling_rate = self.sampling_rate
sampling_rate = self._get_sampling_rate_to_use_for_request(kwargs=kwargs)
random_sample = random.random()
if random_sample > sampling_rate:
verbose_logger.info(
@ -286,7 +272,7 @@ class LangsmithLogger(CustomBatchLogger):
)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
sampling_rate = self.sampling_rate
sampling_rate = self._get_sampling_rate_to_use_for_request(kwargs=kwargs)
random_sample = random.random()
if random_sample > sampling_rate:
verbose_logger.info(
@ -417,6 +403,17 @@ class LangsmithLogger(CustomBatchLogger):
for queue_object in self.log_queue:
credentials = queue_object["credentials"]
# if credential missing, skip - log warning
if (
credentials["LANGSMITH_API_KEY"] is None
or credentials["LANGSMITH_PROJECT"] is None
):
verbose_logger.warning(
"Langsmith Logging - credentials missing - api_key: %s, project: %s",
credentials["LANGSMITH_API_KEY"],
credentials["LANGSMITH_PROJECT"],
)
continue
key = CredentialsKey(
api_key=credentials["LANGSMITH_API_KEY"],
project=credentials["LANGSMITH_PROJECT"],
@ -432,6 +429,19 @@ class LangsmithLogger(CustomBatchLogger):
return log_queue_by_credentials
def _get_sampling_rate_to_use_for_request(self, kwargs: Dict[str, Any]) -> float:
standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = (
kwargs.get("standard_callback_dynamic_params", None)
)
sampling_rate: float = self.sampling_rate
if standard_callback_dynamic_params is not None:
_sampling_rate = standard_callback_dynamic_params.get(
"langsmith_sampling_rate"
)
if _sampling_rate is not None:
sampling_rate = float(_sampling_rate)
return sampling_rate
def _get_credentials_to_use_for_request(
self, kwargs: Dict[str, Any]
) -> LangsmithCredentialsObject:
@ -442,9 +452,9 @@ class LangsmithLogger(CustomBatchLogger):
Otherwise, use the default credentials.
"""
standard_callback_dynamic_params: Optional[
StandardCallbackDynamicParams
] = kwargs.get("standard_callback_dynamic_params", None)
standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = (
kwargs.get("standard_callback_dynamic_params", None)
)
if standard_callback_dynamic_params is not None:
credentials = self.get_credentials_from_env(
langsmith_api_key=standard_callback_dynamic_params.get(

View file

@ -3,6 +3,7 @@ Opik Logger that logs LLM events to an Opik server
"""
import asyncio
from datetime import timezone
import json
import traceback
from typing import Dict, List
@ -291,8 +292,8 @@ class OpikLogger(CustomBatchLogger):
"project_name": project_name,
"id": trace_id,
"name": trace_name,
"start_time": start_time.isoformat() + "Z",
"end_time": end_time.isoformat() + "Z",
"start_time": start_time.astimezone(timezone.utc).isoformat().replace("+00:00", "Z"),
"end_time": end_time.astimezone(timezone.utc).isoformat().replace("+00:00", "Z"),
"input": input_data,
"output": output_data,
"metadata": metadata,
@ -312,8 +313,8 @@ class OpikLogger(CustomBatchLogger):
"parent_span_id": parent_span_id,
"name": span_name,
"type": "llm",
"start_time": start_time.isoformat() + "Z",
"end_time": end_time.isoformat() + "Z",
"start_time": start_time.astimezone(timezone.utc).isoformat().replace("+00:00", "Z"),
"end_time": end_time.astimezone(timezone.utc).isoformat().replace("+00:00", "Z"),
"input": input_data,
"output": output_data,
"metadata": metadata,

View file

@ -0,0 +1,56 @@
"""
Cached imports module for LiteLLM.
This module provides cached import functionality to avoid repeated imports
inside functions that are critical to performance.
"""
from typing import TYPE_CHECKING, Callable, Optional, Type
# Type annotations for cached imports
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.coroutine_checker import CoroutineChecker
# Global cache variables
_LiteLLMLogging: Optional[Type["Logging"]] = None
_coroutine_checker: Optional["CoroutineChecker"] = None
_set_callbacks: Optional[Callable] = None
def get_litellm_logging_class() -> Type["Logging"]:
"""Get the cached LiteLLM Logging class, initializing if needed."""
global _LiteLLMLogging
if _LiteLLMLogging is not None:
return _LiteLLMLogging
from litellm.litellm_core_utils.litellm_logging import Logging
_LiteLLMLogging = Logging
return _LiteLLMLogging
def get_coroutine_checker() -> "CoroutineChecker":
"""Get the cached coroutine checker instance, initializing if needed."""
global _coroutine_checker
if _coroutine_checker is not None:
return _coroutine_checker
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
_coroutine_checker = coroutine_checker
return _coroutine_checker
def get_set_callbacks() -> Callable:
"""Get the cached set_callbacks function, initializing if needed."""
global _set_callbacks
if _set_callbacks is not None:
return _set_callbacks
from litellm.litellm_core_utils.litellm_logging import set_callbacks
_set_callbacks = set_callbacks
return _set_callbacks
def clear_cached_imports() -> None:
"""Clear all cached imports. Useful for testing or memory management."""
global _LiteLLMLogging, _coroutine_checker, _set_callbacks
_LiteLLMLogging = None
_coroutine_checker = None
_set_callbacks = None

View file

@ -158,6 +158,7 @@ def _setup_timezone(
"US/Eastern": timezone(timedelta(hours=-4)), # EDT
"US/Pacific": timezone(timedelta(hours=-7)), # PDT
"Asia/Kolkata": timezone(timedelta(hours=5, minutes=30)), # IST
"Asia/Bangkok": timezone(timedelta(hours=7)), # ICT (Indochina Time)
"Europe/London": timezone(timedelta(hours=1)), # BST
"UTC": timezone.utc,
}

View file

@ -556,7 +556,7 @@ def exception_type( # type: ignore # noqa: PLR0915
model=model,
llm_provider="anthropic",
)
elif "overloaded_error" in error_str:
elif "overloaded_error" in error_str or "Overloaded" in error_str:
exception_mapping_worked = True
raise InternalServerError(
message="AnthropicError - {}".format(error_str),

View file

@ -1,11 +1,12 @@
# What is this?
## Helper utilities for cost_per_token()
from typing import Any, Literal, Optional, Tuple, cast
from typing import Any, Literal, Optional, Tuple, TypedDict, cast
import litellm
from litellm._logging import verbose_logger
from litellm.types.utils import (
CacheCreationTokenDetails,
CallTypes,
ImageResponse,
ModelInfo,
@ -115,7 +116,7 @@ def _generic_cost_per_character(
def _get_token_base_cost(
model_info: ModelInfo, usage: Usage
) -> Tuple[float, float, float, float]:
) -> Tuple[float, float, float, float, float]:
"""
Return prompt cost, completion cost, and cache costs for a given model and usage.
@ -134,6 +135,10 @@ def _get_token_base_cost(
cache_creation_cost = cast(
float, _get_cost_per_unit(model_info, "cache_creation_input_token_cost")
)
cache_creation_cost_above_1hr = cast(
float,
_get_cost_per_unit(model_info, "cache_creation_input_token_cost_above_1hr"),
)
cache_read_cost = cast(
float, _get_cost_per_unit(model_info, "cache_read_input_token_cost")
)
@ -194,7 +199,13 @@ def _get_token_base_cost(
except Exception:
continue
return prompt_base_cost, completion_base_cost, cache_creation_cost, cache_read_cost
return (
prompt_base_cost,
completion_base_cost,
cache_creation_cost,
cache_creation_cost_above_1hr,
cache_read_cost,
)
def calculate_cost_component(
@ -241,6 +252,196 @@ def _get_cost_per_unit(
return default_value
def calculate_cache_writing_cost(
cache_creation_tokens: int,
cache_creation_token_details: Optional[CacheCreationTokenDetails],
cache_creation_cost_above_1hr: float,
cache_creation_cost: float,
) -> float:
"""
Adjust cost of cache creation tokens based on the cache creation token details.
"""
total_cost: float = 0.0
if cache_creation_token_details is not None:
# get the number of 5m and 1h cache creation tokens
cache_creation_tokens_5m = (
cache_creation_token_details.ephemeral_5m_input_tokens
)
cache_creation_tokens_1h = (
cache_creation_token_details.ephemeral_1h_input_tokens
)
# add the number of 5m and 1h cache creation tokens to the cache creation tokens
total_cost += (
cache_creation_tokens_5m * cache_creation_cost
if cache_creation_tokens_5m is not None
else 0.0
)
total_cost += (
cache_creation_tokens_1h * cache_creation_cost_above_1hr
if cache_creation_tokens_1h is not None
else 0.0
)
else:
total_cost += cache_creation_tokens * cache_creation_cost
return total_cost
class PromptTokensDetailsResult(TypedDict):
cache_hit_tokens: int
cache_creation_tokens: int
cache_creation_token_details: Optional[CacheCreationTokenDetails]
text_tokens: int
audio_tokens: int
character_count: int
image_count: int
video_length_seconds: int
def _parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
cache_hit_tokens = (
cast(Optional[int], getattr(usage.prompt_tokens_details, "cached_tokens", 0))
or 0
)
cache_creation_tokens = (
cast(
Optional[int],
getattr(usage.prompt_tokens_details, "cache_creation_tokens", 0),
)
or 0
)
cache_creation_token_details = (
cast(
Optional[CacheCreationTokenDetails],
getattr(usage.prompt_tokens_details, "cache_creation_token_details", None),
)
or None
)
text_tokens = (
cast(Optional[int], getattr(usage.prompt_tokens_details, "text_tokens", None))
or 0 # default to prompt tokens, if this field is not set
)
audio_tokens = (
cast(Optional[int], getattr(usage.prompt_tokens_details, "audio_tokens", 0))
or 0
)
character_count = (
cast(
Optional[int],
getattr(usage.prompt_tokens_details, "character_count", 0),
)
or 0
)
image_count = (
cast(Optional[int], getattr(usage.prompt_tokens_details, "image_count", 0)) or 0
)
video_length_seconds = (
cast(
Optional[int],
getattr(usage.prompt_tokens_details, "video_length_seconds", 0),
)
or 0
)
return PromptTokensDetailsResult(
cache_hit_tokens=cache_hit_tokens,
cache_creation_tokens=cache_creation_tokens,
cache_creation_token_details=cache_creation_token_details,
text_tokens=text_tokens,
audio_tokens=audio_tokens,
character_count=character_count,
image_count=image_count,
video_length_seconds=video_length_seconds,
)
class CompletionTokensDetailsResult(TypedDict):
audio_tokens: int
text_tokens: int
reasoning_tokens: int
def _parse_completion_tokens_details(usage: Usage) -> CompletionTokensDetailsResult:
audio_tokens = (
cast(
Optional[int],
getattr(usage.completion_tokens_details, "audio_tokens", 0),
)
or 0
)
text_tokens = (
cast(
Optional[int],
getattr(usage.completion_tokens_details, "text_tokens", None),
)
or 0 # default to completion tokens, if this field is not set
)
reasoning_tokens = (
cast(
Optional[int],
getattr(usage.completion_tokens_details, "reasoning_tokens", 0),
)
or 0
)
return CompletionTokensDetailsResult(
audio_tokens=audio_tokens,
text_tokens=text_tokens,
reasoning_tokens=reasoning_tokens,
)
def _calculate_input_cost(
prompt_tokens_details: PromptTokensDetailsResult,
model_info: ModelInfo,
prompt_base_cost: float,
cache_read_cost: float,
cache_creation_cost: float,
cache_creation_cost_above_1hr: float,
) -> float:
"""
Calculates the input cost for a given model, prompt tokens, and completion tokens.
"""
prompt_cost = float(prompt_tokens_details["text_tokens"]) * prompt_base_cost
### CACHE READ COST - Now uses tiered pricing
prompt_cost += float(prompt_tokens_details["cache_hit_tokens"]) * cache_read_cost
### AUDIO COST
prompt_cost += calculate_cost_component(
model_info, "input_cost_per_audio_token", prompt_tokens_details["audio_tokens"]
)
### CACHE WRITING COST - Now uses tiered pricing
prompt_cost += calculate_cache_writing_cost(
cache_creation_tokens=prompt_tokens_details["cache_creation_tokens"],
cache_creation_token_details=prompt_tokens_details[
"cache_creation_token_details"
],
cache_creation_cost_above_1hr=cache_creation_cost_above_1hr,
cache_creation_cost=cache_creation_cost,
)
### CHARACTER COST
prompt_cost += calculate_cost_component(
model_info, "input_cost_per_character", prompt_tokens_details["character_count"]
)
### IMAGE COUNT COST
prompt_cost += calculate_cost_component(
model_info, "input_cost_per_image", prompt_tokens_details["image_count"]
)
### VIDEO LENGTH COST
prompt_cost += calculate_cost_component(
model_info,
"input_cost_per_video_per_second",
prompt_tokens_details["video_length_seconds"],
)
return prompt_cost
def generic_cost_per_token(
model: str, usage: Usage, custom_llm_provider: str
) -> Tuple[float, float]:
@ -264,97 +465,45 @@ def generic_cost_per_token(
### Cost of processing (non-cache hit + cache hit) + Cost of cache-writing (cache writing)
prompt_cost = 0.0
### PROCESSING COST
text_tokens = usage.prompt_tokens
cache_hit_tokens = 0
cache_creation_tokens = 0
audio_tokens = 0
character_count = 0
image_count = 0
video_length_seconds = 0
prompt_tokens_details = PromptTokensDetailsResult(
cache_hit_tokens=0,
cache_creation_tokens=0,
cache_creation_token_details=None,
text_tokens=usage.prompt_tokens,
audio_tokens=0,
character_count=0,
image_count=0,
video_length_seconds=0,
)
if usage.prompt_tokens_details:
cache_hit_tokens = (
cast(
Optional[int], getattr(usage.prompt_tokens_details, "cached_tokens", 0)
)
or 0
)
cache_creation_tokens = (
cast(
Optional[int],
getattr(usage.prompt_tokens_details, "cache_creation_tokens", 0),
)
or 0
)
text_tokens = (
cast(
Optional[int], getattr(usage.prompt_tokens_details, "text_tokens", None)
)
or 0 # default to prompt tokens, if this field is not set
)
audio_tokens = (
cast(Optional[int], getattr(usage.prompt_tokens_details, "audio_tokens", 0))
or 0
)
character_count = (
cast(
Optional[int],
getattr(usage.prompt_tokens_details, "character_count", 0),
)
or 0
)
image_count = (
cast(Optional[int], getattr(usage.prompt_tokens_details, "image_count", 0))
or 0
)
video_length_seconds = (
cast(
Optional[int],
getattr(usage.prompt_tokens_details, "video_length_seconds", 0),
)
or 0
)
prompt_tokens_details = _parse_prompt_tokens_details(usage)
## EDGE CASE - text tokens not set inside PromptTokensDetails
if text_tokens == 0:
if prompt_tokens_details["text_tokens"] == 0:
text_tokens = (
usage.prompt_tokens
- cache_hit_tokens
- audio_tokens
- cache_creation_tokens
- prompt_tokens_details["cache_hit_tokens"]
- prompt_tokens_details["audio_tokens"]
- prompt_tokens_details["cache_creation_tokens"]
)
prompt_tokens_details["text_tokens"] = text_tokens
prompt_base_cost, completion_base_cost, cache_creation_cost, cache_read_cost = (
_get_token_base_cost(model_info=model_info, usage=usage)
)
(
prompt_base_cost,
completion_base_cost,
cache_creation_cost,
cache_creation_cost_above_1hr,
cache_read_cost,
) = _get_token_base_cost(model_info=model_info, usage=usage)
prompt_cost = float(text_tokens) * prompt_base_cost
### CACHE READ COST - Now uses tiered pricing
prompt_cost += float(cache_hit_tokens) * cache_read_cost
### AUDIO COST
prompt_cost += calculate_cost_component(
model_info, "input_cost_per_audio_token", audio_tokens
)
### CACHE WRITING COST - Now uses tiered pricing
prompt_cost += float(cache_creation_tokens) * cache_creation_cost
### CHARACTER COST
prompt_cost += calculate_cost_component(
model_info, "input_cost_per_character", character_count
)
### IMAGE COUNT COST
prompt_cost += calculate_cost_component(
model_info, "input_cost_per_image", image_count
)
### VIDEO LENGTH COST
prompt_cost += calculate_cost_component(
model_info, "input_cost_per_video_per_second", video_length_seconds
prompt_cost = _calculate_input_cost(
prompt_tokens_details=prompt_tokens_details,
model_info=model_info,
prompt_base_cost=prompt_base_cost,
cache_read_cost=cache_read_cost,
cache_creation_cost=cache_creation_cost,
cache_creation_cost_above_1hr=cache_creation_cost_above_1hr,
)
## CALCULATE OUTPUT COST
@ -363,27 +512,10 @@ def generic_cost_per_token(
reasoning_tokens = 0
is_text_tokens_total = False
if usage.completion_tokens_details is not None:
audio_tokens = (
cast(
Optional[int],
getattr(usage.completion_tokens_details, "audio_tokens", 0),
)
or 0
)
text_tokens = (
cast(
Optional[int],
getattr(usage.completion_tokens_details, "text_tokens", None),
)
or 0 # default to completion tokens, if this field is not set
)
reasoning_tokens = (
cast(
Optional[int],
getattr(usage.completion_tokens_details, "reasoning_tokens", 0),
)
or 0
)
completion_tokens_details = _parse_completion_tokens_details(usage)
audio_tokens = completion_tokens_details["audio_tokens"]
text_tokens = completion_tokens_details["text_tokens"]
reasoning_tokens = completion_tokens_details["reasoning_tokens"]
if text_tokens == 0:
text_tokens = usage.completion_tokens

View file

@ -0,0 +1,137 @@
"""
Generic object pooling utilities for LiteLLM.
This module provides a flexible object pooling system that can be used
to pool any type of object, reducing memory allocation overhead and
improving performance for frequently created/destroyed objects.
Memory Management Strategy:
- Balanced eviction-based memory control to optimize reuse ratio
- Moderate eviction frequency (300s) to maintain high object reuse
- Conservative eviction weight (0.3) to avoid destroying useful objects
- Lower pre-warm count (5) to reduce initial memory footprint
- Always keeps at least one object available for high availability
- Unlimited pools when maxsize is not specified (eviction controls actual usage)
"""
from typing import Any, Callable, Optional, Type, TypeVar
from pond import Pond, PooledObject, PooledObjectFactory
T = TypeVar('T')
class GenericPooledObjectFactory(PooledObjectFactory):
"""Generic factory class for creating pooled objects of any type."""
def __init__(
self,
object_class: Type[T],
pooled_maxsize: Optional[int] = None, # None = unlimited pool with eviction-based memory control
least_one: bool = True, # Always keep at least one for high concurrency
initializer: Optional[Callable[[T], None]] = None
):
# Only pass maxsize to Pond if user specified it - otherwise let Pond handle unlimited pools
if pooled_maxsize is not None:
super().__init__(pooled_maxsize=pooled_maxsize, least_one=least_one)
else:
super().__init__(least_one=least_one)
self.object_class = object_class
self.initializer = initializer
self._user_maxsize = pooled_maxsize # Store original user preference
def createInstance(self) -> PooledObject:
"""Create a new instance wrapped in a PooledObject."""
# Create a properly initialized instance
obj = self.object_class()
return PooledObject(obj)
def destroy(self, pooled_object: PooledObject):
"""Destroy the pooled object."""
if hasattr(pooled_object.keeped_object, '__dict__'):
pooled_object.keeped_object.__dict__.clear()
del pooled_object
def reset(self, pooled_object: PooledObject, **kwargs: Any) -> PooledObject:
"""Reset the pooled object to a clean state."""
obj = pooled_object.keeped_object
# Reset the object by calling its reset method if it exists
if hasattr(obj, 'reset') and callable(getattr(obj, 'reset')):
obj.reset()
else:
# Fallback: clear all attributes to reset the object
if hasattr(obj, '__dict__'):
obj.__dict__.clear()
return pooled_object
def validate(self, pooled_object: PooledObject) -> bool:
"""Validate if the pooled object is still usable."""
return pooled_object.keeped_object is not None
# Global pond instances
_pools: dict[str, Pond] = {}
def get_object_pool(
pool_name: str,
object_class: Type[T],
pooled_maxsize: Optional[int] = None, # None = unlimited pool with eviction-based memory control
least_one: bool = True, # Always keep at least one
borrowed_timeout: int = 10, # Longer timeout for high concurrency
time_between_eviction_runs: int = 300, # Less frequent eviction to maintain high reuse ratio
eviction_weight: float = 0.3, # Less aggressive eviction for better reuse
prewarm_count: int = 5 # Lower pre-warm count to reduce initial memory usage
) -> Pond:
"""Get or create a global object pool instance with balanced eviction-based memory control.
Memory is controlled through moderate eviction to balance reuse ratio and memory usage:
- Moderate eviction frequency (300s) to maintain high object reuse ratio
- Conservative eviction weight (0.3) to avoid destroying useful objects
- Lower pre-warm count (5) to reduce initial memory footprint
Args:
pool_name: Unique name for the pool
object_class: The class type to pool
pooled_maxsize: Maximum number of objects in the pool (None = truly unlimited)
least_one: Whether to keep at least one object in the pool (default: True)
borrowed_timeout: Timeout for borrowing objects (seconds, default: 10)
time_between_eviction_runs: Time between eviction runs (seconds, default: 300)
eviction_weight: Weight for eviction algorithm (default: 0.3, conservative)
prewarm_count: Number of objects to pre-warm the pool with (default: 5)
Returns:
Pond instance for the specified object type
"""
if pool_name in _pools:
return _pools[pool_name]
# Create new pond
pond = Pond(
borrowed_timeout=borrowed_timeout,
time_between_eviction_runs=time_between_eviction_runs,
thread_daemon=True,
eviction_weight=eviction_weight
)
# Register the factory with user's maxsize preference
factory = GenericPooledObjectFactory(
object_class=object_class,
pooled_maxsize=pooled_maxsize,
least_one=least_one
)
pond.register(factory, name=f"{pool_name}Factory")
# Pre-warm the pool
_prewarm_pool(pond, pool_name, prewarm_count)
_pools[pool_name] = pond
return pond
def _prewarm_pool(pond: Pond, pool_name: str, prewarm_count: int = 20) -> None:
"""Pre-warm the pool with initial objects for high concurrency."""
for _ in range(prewarm_count):
try:
pooled_obj = pond.borrow(name=f"{pool_name}Factory")
pond.recycle(pooled_obj, name=f"{pool_name}Factory")
except Exception:
# If pre-warming fails, just continue
break

View file

@ -45,7 +45,10 @@ from litellm.types.llms.openai import (
OpenAIMcpServerTool,
OpenAIWebSearchOptions,
)
from litellm.types.utils import CompletionTokensDetailsWrapper
from litellm.types.utils import (
CacheCreationTokenDetails,
CompletionTokensDetailsWrapper,
)
from litellm.types.utils import Message as LitellmMessage
from litellm.types.utils import PromptTokensDetailsWrapper, ServerToolUse
from litellm.utils import (
@ -820,6 +823,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
_usage = usage_object
cache_creation_input_tokens: int = 0
cache_read_input_tokens: int = 0
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
web_search_requests: Optional[int] = None
if (
"cache_creation_input_tokens" in _usage
@ -842,8 +846,20 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
int, _usage["server_tool_use"]["web_search_requests"]
)
if "cache_creation" in _usage and _usage["cache_creation"] is not None:
cache_creation_token_details = CacheCreationTokenDetails(
ephemeral_5m_input_tokens=_usage["cache_creation"].get(
"ephemeral_5m_input_tokens"
),
ephemeral_1h_input_tokens=_usage["cache_creation"].get(
"ephemeral_1h_input_tokens"
),
)
prompt_tokens_details = PromptTokensDetailsWrapper(
cached_tokens=cache_read_input_tokens,
cache_creation_tokens=cache_read_input_tokens,
cache_creation_token_details=cache_creation_token_details,
)
completion_token_details = (
CompletionTokensDetailsWrapper(

View file

@ -20,7 +20,11 @@ from pydantic import BaseModel
from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.constants import BEDROCK_INVOKE_PROVIDERS_LITERAL, BEDROCK_MAX_POLICY_SIZE
from litellm.constants import (
BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
BEDROCK_INVOKE_PROVIDERS_LITERAL,
BEDROCK_MAX_POLICY_SIZE,
)
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.secret_managers.main import get_secret, get_secret_str
@ -327,6 +331,40 @@ class BaseAWSLLM:
return provider
return None
@staticmethod
def get_bedrock_embedding_provider(
model: str,
) -> Optional[BEDROCK_EMBEDDING_PROVIDERS_LITERAL]:
"""
Helper function to get the bedrock embedding provider from the model
Handles scenarios like:
1. model=cohere.embed-english-v3:0 -> Returns `cohere`
2. model=amazon.titan-embed-text-v1 -> Returns `amazon`
3. model=us.twelvelabs.marengo-embed-2-7-v1:0 -> Returns `twelvelabs`
4. model=twelvelabs.marengo-embed-2-7-v1:0 -> Returns `twelvelabs`
"""
# Handle regional models like us.twelvelabs.marengo-embed-2-7-v1:0
if "." in model:
parts = model.split(".")
# Check if the second part (after potential region) is a known provider
if len(parts) >= 2:
potential_provider = parts[1] # e.g., "twelvelabs" from "us.twelvelabs.marengo-embed-2-7-v1:0"
if potential_provider in get_args(BEDROCK_EMBEDDING_PROVIDERS_LITERAL):
return cast(BEDROCK_EMBEDDING_PROVIDERS_LITERAL, potential_provider)
# Check if the first part is a known provider (standard format)
potential_provider = parts[0] # e.g., "cohere" from "cohere.embed-english-v3:0"
if potential_provider in get_args(BEDROCK_EMBEDDING_PROVIDERS_LITERAL):
return cast(BEDROCK_EMBEDDING_PROVIDERS_LITERAL, potential_provider)
# Fallback: check if any provider name appears in the model string
for provider in get_args(BEDROCK_EMBEDDING_PROVIDERS_LITERAL):
if provider in model:
return cast(BEDROCK_EMBEDDING_PROVIDERS_LITERAL, provider)
return None
def _get_aws_region_name(
self,
optional_params: dict,

View file

@ -10,7 +10,7 @@ Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-tit
"""
import types
from typing import List, Optional
from typing import List, Optional, Union
from litellm.types.llms.bedrock import (
AmazonTitanV2EmbeddingRequest,
@ -30,9 +30,7 @@ class AmazonTitanV2Config:
normalize: Optional[bool] = None
dimensions: Optional[int] = None
def __init__(
self, normalize: Optional[bool] = None, dimensions: Optional[int] = None
) -> None:
def __init__(self, normalize: Optional[bool] = None, dimensions: Optional[int] = None) -> None:
locals_ = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
@ -57,32 +55,56 @@ class AmazonTitanV2Config:
}
def get_supported_openai_params(self) -> List[str]:
return ["dimensions"]
return ["dimensions", "encoding_format"]
def map_openai_params(
self, non_default_params: dict, optional_params: dict
) -> dict:
def map_openai_params(self, non_default_params: dict, optional_params: dict) -> dict:
for k, v in non_default_params.items():
if k == "dimensions":
optional_params["dimensions"] = v
elif k == "encoding_format":
# Map OpenAI encoding_format to AWS embeddingTypes
if v == "float":
optional_params["embeddingTypes"] = ["float"]
elif v == "base64":
# base64 maps to binary format in AWS
optional_params["embeddingTypes"] = ["binary"]
else:
# For any other encoding format, default to float
optional_params["embeddingTypes"] = ["float"]
return optional_params
def _transform_request(
self, input: str, inference_params: dict
) -> AmazonTitanV2EmbeddingRequest:
def _transform_request(self, input: str, inference_params: dict) -> AmazonTitanV2EmbeddingRequest:
return AmazonTitanV2EmbeddingRequest(inputText=input, **inference_params) # type: ignore
def _transform_response(
self, response_list: List[dict], model: str
) -> EmbeddingResponse:
def _transform_response(self, response_list: List[dict], model: str) -> EmbeddingResponse:
total_prompt_tokens = 0
transformed_responses: List[Embedding] = []
for index, response in enumerate(response_list):
_parsed_response = AmazonTitanV2EmbeddingResponse(**response) # type: ignore
# According to AWS docs, embeddingsByType is always present
# If binary was requested (encoding_format="base64"), use binary data
# Otherwise, use float data from embeddingsByType or fallback to embedding field
embedding_data: Union[List[float], List[int]]
if ("embeddingsByType" in _parsed_response and
"binary" in _parsed_response["embeddingsByType"]):
# Use binary data if available (for encoding_format="base64")
embedding_data = _parsed_response["embeddingsByType"]["binary"]
elif ("embeddingsByType" in _parsed_response and
"float" in _parsed_response["embeddingsByType"]):
# Use float data from embeddingsByType
embedding_data = _parsed_response["embeddingsByType"]["float"]
elif "embedding" in _parsed_response:
# Fallback to legacy embedding field
embedding_data = _parsed_response["embedding"]
else:
raise ValueError(f"No embedding data found in response: {response}")
transformed_responses.append(
Embedding(
embedding=_parsed_response["embedding"],
embedding=embedding_data,
index=index,
object="embedding",
)

View file

@ -5,11 +5,12 @@ Handles embedding calls to Bedrock's `/invoke` endpoint
import copy
import json
import urllib.parse
from typing import Any, Callable, List, Optional, Tuple, Union
from typing import Any, Callable, List, Optional, Tuple, Union, get_args
import httpx
import litellm
from litellm.constants import BEDROCK_EMBEDDING_PROVIDERS_LITERAL
from litellm.llms.cohere.embed.handler import embedding as cohere_embedding
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
@ -150,6 +151,44 @@ class BedrockEmbedding(BaseAWSLLM):
raise BedrockError(status_code=408, message="Timeout error occurred.")
return response.json()
def _transform_response(
self, response_list: List[dict], model: str, provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL
) -> Optional[EmbeddingResponse]:
"""
Transforms the response from the Bedrock embedding provider to the OpenAI format.
"""
returned_response: Optional[EmbeddingResponse] = None
if model == "amazon.titan-embed-image-v1":
returned_response = (
AmazonTitanMultimodalEmbeddingG1Config()._transform_response(
response_list=response_list, model=model
)
)
elif model == "amazon.titan-embed-text-v1":
returned_response = AmazonTitanG1Config()._transform_response(
response_list=response_list, model=model
)
elif model == "amazon.titan-embed-text-v2:0":
returned_response = AmazonTitanV2Config()._transform_response(
response_list=response_list, model=model
)
elif provider == "twelvelabs":
returned_response = TwelveLabsMarengoEmbeddingConfig()._transform_response(
response_list=response_list, model=model
)
##########################################################
# Validate returned response
##########################################################
if returned_response is None:
raise Exception(
"Unable to map model response to known provider format. model={}".format(
model
)
)
return returned_response
def _single_func_embeddings(
self,
@ -162,6 +201,7 @@ class BedrockEmbedding(BaseAWSLLM):
aws_region_name: str,
model: str,
logging_obj: Any,
provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
api_key: Optional[str] = None,
):
responses: List[dict] = []
@ -208,32 +248,9 @@ class BedrockEmbedding(BaseAWSLLM):
responses.append(response)
returned_response: Optional[EmbeddingResponse] = None
## TRANSFORM RESPONSE ##
if model == "amazon.titan-embed-image-v1":
returned_response = (
AmazonTitanMultimodalEmbeddingG1Config()._transform_response(
response_list=responses, model=model
)
)
elif model == "amazon.titan-embed-text-v1":
returned_response = AmazonTitanG1Config()._transform_response(
response_list=responses, model=model
)
elif model == "amazon.titan-embed-text-v2:0":
returned_response = AmazonTitanV2Config()._transform_response(
response_list=responses, model=model
)
if returned_response is None:
raise Exception(
"Unable to map model response to known provider format. model={}".format(
model
)
)
return returned_response
return self._transform_response(
response_list=responses, model=model, provider=provider
)
async def _async_single_func_embeddings(
self,
@ -246,6 +263,7 @@ class BedrockEmbedding(BaseAWSLLM):
aws_region_name: str,
model: str,
logging_obj: Any,
provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
api_key: Optional[str] = None,
):
responses: List[dict] = []
@ -291,33 +309,10 @@ class BedrockEmbedding(BaseAWSLLM):
)
responses.append(response)
returned_response: Optional[EmbeddingResponse] = None
## TRANSFORM RESPONSE ##
if model == "amazon.titan-embed-image-v1":
returned_response = (
AmazonTitanMultimodalEmbeddingG1Config()._transform_response(
response_list=responses, model=model
)
)
elif model == "amazon.titan-embed-text-v1":
returned_response = AmazonTitanG1Config()._transform_response(
response_list=responses, model=model
)
elif model == "amazon.titan-embed-text-v2:0":
returned_response = AmazonTitanV2Config()._transform_response(
response_list=responses, model=model
)
if returned_response is None:
raise Exception(
"Unable to map model response to known provider format. model={}".format(
model
)
)
return returned_response
return self._transform_response(
response_list=responses, model=model, provider=provider
)
def embeddings(
self,
@ -349,7 +344,12 @@ class BedrockEmbedding(BaseAWSLLM):
model_id=unencoded_model_id,
)
provider = model.split(".")[0]
provider = self.get_bedrock_embedding_provider(model)
if provider is None:
raise Exception(
f"Unable to determine bedrock embedding provider for model: {model}. "
f"Supported providers: {list(get_args(BEDROCK_EMBEDDING_PROVIDERS_LITERAL))}"
)
inference_params = copy.deepcopy(optional_params)
inference_params = {
k: v
@ -399,9 +399,7 @@ class BedrockEmbedding(BaseAWSLLM):
)
)
batch_data.append(transformed_request)
elif provider == "twelvelabs" and model in [
"twelvelabs.marengo-embed-2-7-v1:0",
]:
elif provider == "twelvelabs":
batch_data = []
for i in input:
twelvelabs_request: (
@ -438,8 +436,9 @@ class BedrockEmbedding(BaseAWSLLM):
model=model,
logging_obj=logging_obj,
api_key=api_key,
provider=provider,
)
return self._single_func_embeddings(
returned_response = self._single_func_embeddings(
client=(
client
if client is not None and isinstance(client, HTTPHandler)
@ -454,7 +453,11 @@ class BedrockEmbedding(BaseAWSLLM):
model=model,
logging_obj=logging_obj,
api_key=api_key,
provider=provider,
)
if returned_response is None:
raise Exception("Unable to map Bedrock request to provider")
return returned_response
elif data is None:
raise Exception("Unable to map Bedrock request to provider")

View file

@ -69,15 +69,6 @@ class TwelveLabsMarengoEmbeddingConfig:
if "textTruncate" not in inference_params:
transformed_request["textTruncate"] = "end"
# Set default embedding options for Phase 1 (text and image)
if "embeddingOption" not in inference_params:
if is_encoded:
# For images, return both visual-text and visual-image embeddings
transformed_request["embeddingOption"] = ["visual-text", "visual-image"]
else:
# For text, return visual-text embedding
transformed_request["embeddingOption"] = ["visual-text"]
# Apply any additional inference parameters
for k, v in inference_params.items():
if k not in [
@ -94,14 +85,32 @@ class TwelveLabsMarengoEmbeddingConfig:
) -> EmbeddingResponse:
"""
Transform TwelveLabs response to OpenAI format.
Handles multiple embedding types in the response.
Handles the actual TwelveLabs response format: {"data": [{"embedding": [...]}]}
"""
embeddings: List[Embedding] = []
total_tokens = 0
for response in response_list:
if "embedding" in response:
# Single embedding response
# TwelveLabs response format has a "data" field containing the embeddings
if "data" in response and isinstance(response["data"], list):
for item in response["data"]:
if "embedding" in item:
# Single embedding response
embedding = Embedding(
embedding=item["embedding"],
index=len(embeddings),
object="embedding",
)
embeddings.append(embedding)
# Estimate token count (rough approximation)
if "inputTextTokenCount" in item:
total_tokens += item["inputTextTokenCount"]
else:
# Rough estimate: 1 token per 4 characters for text, or use embedding size
total_tokens += len(item["embedding"]) // 4
elif "embedding" in response:
# Direct embedding response (fallback for other formats)
embedding = Embedding(
embedding=response["embedding"],
index=len(embeddings),

View file

@ -85,17 +85,25 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
) -> str:
"""
Get the complete url for the request
Google AI API format: https://generativelanguage.googleapis.com/v1beta/models/{model}:predict
Gemini 2.5 Flash Image Preview: :generateContent
Other Imagen models: :predict
"""
complete_url: str = (
api_base
or get_secret_str("GEMINI_API_BASE")
api_base
or get_secret_str("GEMINI_API_BASE")
or self.DEFAULT_BASE_URL
)
complete_url = complete_url.rstrip("/")
complete_url = f"{complete_url}/models/{model}:predict"
# Gemini 2.5 Flash Image Preview uses generateContent endpoint
if "2.5-flash-image-preview" in model:
complete_url = f"{complete_url}/models/{model}:generateContent"
else:
# All other Imagen models use predict endpoint
complete_url = f"{complete_url}/models/{model}:predict"
return complete_url
def validate_environment(
@ -128,35 +136,52 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
headers: dict,
) -> dict:
"""
Transform the image generation request to Google AI Imagen format
Google AI API format:
Transform the image generation request to Gemini format
For Gemini 2.5 Flash Image Preview, use the standard Gemini format with response_modalities:
{
"instances": [
"contents": [
{
"prompt": "Robot holding a red skateboard"
"parts": [
{"text": "Generate an image of..."}
]
}
],
"parameters": {
"sampleCount": 4,
"aspectRatio": "1:1",
"personGeneration": "allow_adult"
"generationConfig": {
"response_modalities": ["IMAGE", "TEXT"]
}
}
"""
from litellm.types.llms.gemini import (
GeminiImageGenerationInstance,
GeminiImageGenerationParameters,
)
request_body: GeminiImageGenerationRequest = GeminiImageGenerationRequest(
instances=[
GeminiImageGenerationInstance(
prompt=prompt
)
],
parameters=GeminiImageGenerationParameters(**optional_params)
)
return request_body.model_dump(exclude_none=True)
# For Gemini 2.5 Flash Image Preview, use standard Gemini format
if "2.5-flash-image-preview" in model:
request_body: dict = {
"contents": [
{
"parts": [
{"text": prompt}
]
}
],
"generationConfig": {
"response_modalities": ["IMAGE", "TEXT"]
}
}
return request_body
else:
# For other Imagen models, use the original Imagen format
from litellm.types.llms.gemini import (
GeminiImageGenerationInstance,
GeminiImageGenerationParameters,
)
request_body_obj: GeminiImageGenerationRequest = GeminiImageGenerationRequest(
instances=[
GeminiImageGenerationInstance(
prompt=prompt
)
],
parameters=GeminiImageGenerationParameters(**optional_params)
)
return request_body_obj.model_dump(exclude_none=True)
def transform_image_generation_response(
self,
@ -185,14 +210,30 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
if not model_response.data:
model_response.data = []
# Google AI returns predictions with generated images
predictions = response_data.get("predictions", [])
for prediction in predictions:
# Google AI returns base64 encoded images in the prediction
model_response.data.append(ImageObject(
b64_json=prediction.get("bytesBase64Encoded", None),
url=None, # Google AI returns base64, not URLs
))
# Handle different response formats based on model
if "2.5-flash-image-preview" in model:
# Gemini 2.5 Flash Image Preview returns in candidates format
candidates = response_data.get("candidates", [])
for candidate in candidates:
content = candidate.get("content", {})
parts = content.get("parts", [])
for part in parts:
# Look for inlineData with image
if "inlineData" in part:
inline_data = part["inlineData"]
if "data" in inline_data:
model_response.data.append(ImageObject(
b64_json=inline_data["data"],
url=None,
))
else:
# Original Imagen format - predictions with generated images
predictions = response_data.get("predictions", [])
for prediction in predictions:
# Google AI returns base64 encoded images in the prediction
model_response.data.append(ImageObject(
b64_json=prediction.get("bytesBase64Encoded", None),
url=None, # Google AI returns base64, not URLs
))
return model_response

View file

@ -305,8 +305,56 @@
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true,
"supports_multimodal_embedding": true
"supports_image_input": true
},
"us.twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"max_tokens": 77,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"eu.twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"max_tokens": 77,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"us.twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"eu.twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"amazon.titan-text-express-v1": {
"input_cost_per_token": 1.3e-06,
@ -9078,7 +9126,7 @@
"max_tokens": 65535,
"max_video_length": 1,
"max_videos_per_prompt": 10,
"mode": "chat",
"mode": "image_generation",
"output_cost_per_image": 0.039,
"output_cost_per_reasoning_token": 3e-05,
"output_cost_per_token": 3e-05,
@ -10441,7 +10489,7 @@
"max_tokens": 65535,
"max_video_length": 1,
"max_videos_per_prompt": 10,
"mode": "chat",
"mode": "image_generation",
"output_cost_per_image": 0.039,
"output_cost_per_reasoning_token": 3e-05,
"output_cost_per_token": 3e-05,

View file

@ -28,12 +28,16 @@ class MCPRequestHandler:
LITELLM_MCP_SERVERS_HEADER_NAME = SpecialHeaders.mcp_servers.value
LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value
# MCP Protocol Version header
MCP_PROTOCOL_VERSION_HEADER_NAME = "MCP-Protocol-Version"
@staticmethod
async def process_mcp_request(scope: Scope) -> Tuple[UserAPIKeyAuth, Optional[str], Optional[List[str]], Optional[Dict[str, str]], Optional[str]]:
async def process_mcp_request(
scope: Scope,
) -> Tuple[
UserAPIKeyAuth, Optional[str], Optional[List[str]], Optional[Dict[str, str]]
]:
"""
Process and validate MCP request headers from the ASGI scope.
This includes:
@ -49,7 +53,6 @@ class MCPRequestHandler:
mcp_auth_header: Optional[str] MCP auth header to be passed to the MCP server (deprecated)
mcp_servers: Optional[List[str]] List of MCP servers and access groups to use
mcp_server_auth_headers: Optional[Dict[str, str]] Server-specific auth headers in format {server_alias: auth_value}
mcp_protocol_version: Optional[str] MCP protocol version from request header
Raises:
HTTPException: If headers are invalid or missing required headers
@ -58,39 +61,50 @@ class MCPRequestHandler:
litellm_api_key = (
MCPRequestHandler.get_litellm_api_key_from_headers(headers) or ""
)
# Get the old mcp_auth_header for backward compatibility
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers)
# Get the new server-specific auth headers
mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
# Get MCP protocol version from header
mcp_protocol_version = headers.get(MCPRequestHandler.MCP_PROTOCOL_VERSION_HEADER_NAME)
# Get the new server-specific auth headers
mcp_server_auth_headers = (
MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
)
# Parse MCP servers from header
mcp_servers_header = headers.get(MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME)
mcp_servers_header = headers.get(
MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME
)
verbose_logger.debug(f"Raw MCP servers header: {mcp_servers_header}")
mcp_servers = None
if mcp_servers_header is not None:
try:
mcp_servers = [s.strip() for s in mcp_servers_header.split(",") if s.strip()]
mcp_servers = [
s.strip() for s in mcp_servers_header.split(",") if s.strip()
]
verbose_logger.debug(f"Parsed MCP servers: {mcp_servers}")
except Exception as e:
verbose_logger.debug(f"Error parsing mcp_servers header: {e}")
mcp_servers = None
if mcp_servers_header == "" or (mcp_servers is not None and len(mcp_servers) == 0):
if mcp_servers_header == "" or (
mcp_servers is not None and len(mcp_servers) == 0
):
mcp_servers = []
# Create a proper Request object with mock body method to avoid ASGI receive channel issues
request = Request(scope=scope)
async def mock_body():
return b"{}"
request.body = mock_body # type: ignore
validated_user_api_key_auth = await user_api_key_auth(
api_key=litellm_api_key, request=request
)
return validated_user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, mcp_protocol_version
return (
validated_user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
)
@staticmethod
def _get_mcp_auth_header_from_headers(headers: Headers) -> Optional[str]:
@ -104,10 +118,12 @@ class MCPRequestHandler:
Support this auth: https://docs.litellm.ai/docs/mcp#using-your-mcp-with-client-side-credentials
If you want to use a different header name, you can set the `LITELLM_MCP_CLIENT_SIDE_AUTH_HEADER_NAME` in the secret manager or `mcp_client_side_auth_header_name` in the general settings.
DEPRECATED: This method is deprecated in favor of server-specific auth headers using the format x-mcp-{{server_alias}}-{{header_name}} instead.
"""
mcp_client_side_auth_header_name: str = MCPRequestHandler._get_mcp_client_side_auth_header_name()
mcp_client_side_auth_header_name: str = (
MCPRequestHandler._get_mcp_client_side_auth_header_name()
)
auth_header = headers.get(mcp_client_side_auth_header_name)
if auth_header:
verbose_logger.warning(
@ -115,42 +131,49 @@ class MCPRequestHandler:
f"Please use server-specific auth headers in the format 'x-mcp-{{server_alias}}-{{header_name}}' instead."
)
return auth_header
@staticmethod
def _get_mcp_server_auth_headers_from_headers(headers: Headers) -> Dict[str, str]:
"""
Parse server-specific MCP auth headers from the request headers.
Looks for headers in the format: x-mcp-{server_alias}-{header_name}
Examples:
- x-mcp-github-authorization: Bearer token123
- x-mcp-zapier-x-api-key: api_key_456
- x-mcp-deepwiki-authorization: Basic base64_encoded_creds
Returns:
Dict[str, str]: Mapping of server alias to auth value
"""
server_auth_headers = {}
prefix = "x-mcp-"
for header_name, header_value in headers.items():
if header_name.lower().startswith(prefix):
# Skip the access groups header as it's not a server auth header
if header_name.lower() == MCPRequestHandler.LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME.lower() or header_name.lower() == MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME.lower():
if (
header_name.lower()
== MCPRequestHandler.LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME.lower()
or header_name.lower()
== MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME.lower()
):
continue
# Extract server_alias and header_name from x-mcp-{server_alias}-{header_name}
remaining = header_name[len(prefix):].lower()
if '-' in remaining:
remaining = header_name[len(prefix) :].lower()
if "-" in remaining:
# Split on the last dash to separate server_alias from header_name
parts = remaining.rsplit('-', 1)
parts = remaining.rsplit("-", 1)
if len(parts) == 2:
server_alias, auth_header_name = parts
server_auth_headers[server_alias] = header_value
verbose_logger.debug(f"Found server auth header: {server_alias} -> {auth_header_name}: {header_value[:10]}...")
verbose_logger.debug(
f"Found server auth header: {server_alias} -> {auth_header_name}: {header_value[:10]}..."
)
return server_auth_headers
@staticmethod
def _get_mcp_client_side_auth_header_name() -> str:
"""
@ -162,13 +185,21 @@ class MCPRequestHandler:
"""
from litellm.proxy.proxy_server import general_settings
from litellm.secret_managers.main import get_secret_str
MCP_CLIENT_SIDE_AUTH_HEADER_NAME: str = MCPRequestHandler.LITELLM_MCP_AUTH_HEADER_NAME
if get_secret_str("LITELLM_MCP_CLIENT_SIDE_AUTH_HEADER_NAME") is not None:
MCP_CLIENT_SIDE_AUTH_HEADER_NAME = get_secret_str("LITELLM_MCP_CLIENT_SIDE_AUTH_HEADER_NAME") or MCP_CLIENT_SIDE_AUTH_HEADER_NAME
elif general_settings.get("mcp_client_side_auth_header_name") is not None:
MCP_CLIENT_SIDE_AUTH_HEADER_NAME = general_settings.get("mcp_client_side_auth_header_name") or MCP_CLIENT_SIDE_AUTH_HEADER_NAME
return MCP_CLIENT_SIDE_AUTH_HEADER_NAME
MCP_CLIENT_SIDE_AUTH_HEADER_NAME: str = (
MCPRequestHandler.LITELLM_MCP_AUTH_HEADER_NAME
)
if get_secret_str("LITELLM_MCP_CLIENT_SIDE_AUTH_HEADER_NAME") is not None:
MCP_CLIENT_SIDE_AUTH_HEADER_NAME = (
get_secret_str("LITELLM_MCP_CLIENT_SIDE_AUTH_HEADER_NAME")
or MCP_CLIENT_SIDE_AUTH_HEADER_NAME
)
elif general_settings.get("mcp_client_side_auth_header_name") is not None:
MCP_CLIENT_SIDE_AUTH_HEADER_NAME = (
general_settings.get("mcp_client_side_auth_header_name")
or MCP_CLIENT_SIDE_AUTH_HEADER_NAME
)
return MCP_CLIENT_SIDE_AUTH_HEADER_NAME
@staticmethod
def get_litellm_api_key_from_headers(headers: Headers) -> Optional[str]:
@ -229,10 +260,14 @@ class MCPRequestHandler:
try:
allowed_mcp_servers: List[str] = []
allowed_mcp_servers_for_key = (
await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
await MCPRequestHandler._get_allowed_mcp_servers_for_key(
user_api_key_auth
)
)
allowed_mcp_servers_for_team = (
await MCPRequestHandler._get_allowed_mcp_servers_for_team(user_api_key_auth)
await MCPRequestHandler._get_allowed_mcp_servers_for_team(
user_api_key_auth
)
)
#########################################################
@ -274,7 +309,9 @@ class MCPRequestHandler:
try:
key_object_permission = (
await prisma_client.db.litellm_objectpermissiontable.find_unique(
where={"object_permission_id": user_api_key_auth.object_permission_id},
where={
"object_permission_id": user_api_key_auth.object_permission_id
},
)
)
if key_object_permission is None:
@ -282,17 +319,21 @@ class MCPRequestHandler:
# Get direct MCP servers
direct_mcp_servers = key_object_permission.mcp_servers or []
# Get MCP servers from access groups
access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups(
key_object_permission.mcp_access_groups or []
access_group_servers = (
await MCPRequestHandler._get_mcp_servers_from_access_groups(
key_object_permission.mcp_access_groups or []
)
)
# Combine both lists
all_servers = direct_mcp_servers + access_group_servers
return list(set(all_servers))
except Exception as e:
verbose_logger.warning(f"Failed to get allowed MCP servers for key: {str(e)}")
verbose_logger.warning(
f"Failed to get allowed MCP servers for key: {str(e)}"
)
return []
@staticmethod
@ -318,10 +359,10 @@ class MCPRequestHandler:
return []
try:
team_obj: Optional[LiteLLM_TeamTable] = (
await prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": user_api_key_auth.team_id},
)
team_obj: Optional[
LiteLLM_TeamTable
] = await prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": user_api_key_auth.team_id},
)
if team_obj is None:
verbose_logger.debug("team_obj is None")
@ -333,21 +374,27 @@ class MCPRequestHandler:
# Get direct MCP servers
direct_mcp_servers = object_permissions.mcp_servers or []
# Get MCP servers from access groups
access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
access_group_servers = (
await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
)
)
# Combine both lists
all_servers = direct_mcp_servers + access_group_servers
return list(set(all_servers))
except Exception as e:
verbose_logger.warning(f"Failed to get allowed MCP servers for team: {str(e)}")
verbose_logger.warning(
f"Failed to get allowed MCP servers for team: {str(e)}"
)
return []
@staticmethod
def _get_config_server_ids_for_access_groups(config_mcp_servers, access_groups: List[str]) -> Set[str]:
def _get_config_server_ids_for_access_groups(
config_mcp_servers, access_groups: List[str]
) -> Set[str]:
"""
Helper to get server_ids from config-loaded servers that match any of the given access groups.
"""
@ -359,7 +406,9 @@ class MCPRequestHandler:
return server_ids
@staticmethod
async def _get_db_server_ids_for_access_groups(prisma_client, access_groups: List[str]) -> Set[str]:
async def _get_db_server_ids_for_access_groups(
prisma_client, access_groups: List[str]
) -> Set[str]:
"""
Helper to get server_ids from DB servers that match any of the given access groups.
"""
@ -367,21 +416,19 @@ class MCPRequestHandler:
if access_groups and prisma_client is not None:
try:
mcp_servers = await prisma_client.db.litellm_mcpservertable.find_many(
where={
"mcp_access_groups": {
"hasSome": access_groups
}
}
where={"mcp_access_groups": {"hasSome": access_groups}}
)
for server in mcp_servers:
server_ids.add(server.server_id)
except Exception as e:
verbose_logger.debug(f"Error getting MCP servers from access groups: {e}")
verbose_logger.debug(
f"Error getting MCP servers from access groups: {e}"
)
return server_ids
@staticmethod
async def _get_mcp_servers_from_access_groups(
access_groups: List[str]
access_groups: List[str],
) -> List[str]:
"""
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers
@ -390,22 +437,28 @@ class MCPRequestHandler:
try:
# Import here to avoid circular import
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
# Use the new helper for config-loaded servers
server_ids = MCPRequestHandler._get_config_server_ids_for_access_groups(
global_mcp_server_manager.config_mcp_servers, access_groups
)
# Use the new helper for DB servers
db_server_ids = await MCPRequestHandler._get_db_server_ids_for_access_groups(
prisma_client, access_groups
db_server_ids = (
await MCPRequestHandler._get_db_server_ids_for_access_groups(
prisma_client, access_groups
)
)
server_ids.update(db_server_ids)
return list(server_ids)
except Exception as e:
verbose_logger.warning(f"Failed to get MCP servers from access groups: {str(e)}")
verbose_logger.warning(
f"Failed to get MCP servers from access groups: {str(e)}"
)
return []
@staticmethod
@ -418,8 +471,8 @@ class MCPRequestHandler:
from typing import List
access_groups: List[str] = []
access_groups_for_key = (
await MCPRequestHandler._get_mcp_access_groups_for_key(user_api_key_auth)
access_groups_for_key = await MCPRequestHandler._get_mcp_access_groups_for_key(
user_api_key_auth
)
access_groups_for_team = (
await MCPRequestHandler._get_mcp_access_groups_for_team(user_api_key_auth)
@ -482,10 +535,10 @@ class MCPRequestHandler:
verbose_logger.debug("prisma_client is None")
return []
team_obj: Optional[LiteLLM_TeamTable] = (
await prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": user_api_key_auth.team_id},
)
team_obj: Optional[
LiteLLM_TeamTable
] = await prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": user_api_key_auth.team_id},
)
if team_obj is None:
verbose_logger.debug("team_obj is None")
@ -502,10 +555,14 @@ class MCPRequestHandler:
"""
Extract and parse the x-mcp-access-groups header as a list of strings.
"""
mcp_access_groups_header = headers.get(MCPRequestHandler.LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME)
mcp_access_groups_header = headers.get(
MCPRequestHandler.LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME
)
if mcp_access_groups_header is not None:
try:
return [s.strip() for s in mcp_access_groups_header.split(",") if s.strip()]
return [
s.strip() for s in mcp_access_groups_header.split(",") if s.strip()
]
except Exception:
return None
return None
@ -516,4 +573,4 @@ class MCPRequestHandler:
Extract and parse the x-mcp-access-groups header from an ASGI scope.
"""
headers = MCPRequestHandler._safe_get_headers_from_scope(scope)
return MCPRequestHandler.get_mcp_access_groups_from_headers(headers)
return MCPRequestHandler.get_mcp_access_groups_from_headers(headers)

View file

@ -34,8 +34,6 @@ from litellm.proxy._experimental.mcp_server.utils import (
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
MCPAuthType,
MCPSpecVersion,
MCPSpecVersionType,
MCPTransport,
MCPTransportType,
UserAPIKeyAuth,
@ -70,38 +68,6 @@ def _deserialize_env_dict(env_data: Any) -> Optional[Dict[str, str]]:
return env_data
def _convert_protocol_version_to_enum(
protocol_version: Optional[str | MCPSpecVersionType],
) -> MCPSpecVersionType:
"""
Convert string protocol version to MCPSpecVersion enum.
Args:
protocol_version: String protocol version, enum, or None
Returns:
MCPSpecVersionType: The enum value
"""
if not protocol_version:
return cast(MCPSpecVersionType, MCPSpecVersion.jun_2025)
# If it's already an MCPSpecVersion enum, return it
if isinstance(protocol_version, MCPSpecVersion):
return cast(MCPSpecVersionType, protocol_version)
# If it's a string, try to match it to enum values
if isinstance(protocol_version, str):
for version in MCPSpecVersion:
if version.value == protocol_version:
return cast(MCPSpecVersionType, version)
# If no match found, return default
verbose_logger.warning(
f"Unknown protocol version '{protocol_version}', using default"
)
return cast(MCPSpecVersionType, MCPSpecVersion.jun_2025)
class MCPServerManager:
def __init__(self):
self.registry: Dict[str, MCPServer] = {}
@ -113,8 +79,7 @@ class MCPServerManager:
"name": "zapier_mcp_server",
"url": "https://actions.zapier.com/mcp/sk-ak-2ew3bofIeQIkNoeKIdXrF1Hhhp/sse"
"transport": "sse",
"auth_type": "api_key",
"spec_version": "2025-03-26"
"auth_type": "api_key"
},
"uuid-2": {
"name": "google_drive_mcp_server",
@ -223,7 +188,6 @@ class MCPServerManager:
server_name=server_name,
url=server_config.get("url", None) or "",
transport=server_config.get("transport", MCPTransport.http),
spec_version=server_config.get("spec_version", MCPSpecVersion.jun_2025),
auth_type=server_config.get("auth_type", None),
alias=alias,
)
@ -239,7 +203,6 @@ class MCPServerManager:
env=server_config.get("env", None) or {},
# TODO: utility fn the default values
transport=server_config.get("transport", MCPTransport.http),
spec_version=server_config.get("spec_version", MCPSpecVersion.jun_2025),
auth_type=server_config.get("auth_type", None),
authentication_token=server_config.get(
"authentication_token", server_config.get("auth_value", None)
@ -287,7 +250,6 @@ class MCPServerManager:
server_name=getattr(mcp_server, "server_name", None),
url=mcp_server.url,
transport=cast(MCPTransportType, mcp_server.transport),
spec_version=_convert_protocol_version_to_enum(mcp_server.spec_version),
auth_type=cast(MCPAuthType, mcp_server.auth_type),
mcp_info=MCPInfo(
server_name=mcp_server.server_name or mcp_server.server_id,
@ -350,7 +312,6 @@ class MCPServerManager:
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
mcp_protocol_version: Optional[str] = None,
) -> List[MCPTool]:
"""
List all tools available across all MCP Servers.
@ -390,7 +351,6 @@ class MCPServerManager:
tools = await self._get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
mcp_protocol_version=mcp_protocol_version,
)
list_tools_result.extend(tools)
verbose_logger.info(
@ -414,7 +374,6 @@ class MCPServerManager:
self,
server: MCPServer,
mcp_auth_header: Optional[str] = None,
protocol_version: Optional[str] = None,
) -> MCPClient:
"""
Create an MCPClient instance for the given server.
@ -422,18 +381,12 @@ class MCPServerManager:
Args:
server (MCPServer): The server configuration
mcp_auth_header: MCP auth header to be passed to the MCP server. This is optional and will be used if provided.
protocol_version: Optional MCP protocol version to use. If not provided, uses server's default.
Returns:
MCPClient: Configured MCP client instance
"""
transport = server.transport or MCPTransport.sse
# Convert protocol version string to enum
protocol_version_enum = _convert_protocol_version_to_enum(
protocol_version or server.spec_version
)
# Handle stdio transport
if transport == MCPTransport.stdio:
# For stdio, we need to get the stdio config from the server
@ -450,7 +403,6 @@ class MCPServerManager:
auth_value=mcp_auth_header or server.authentication_token,
timeout=60.0,
stdio_config=stdio_config,
protocol_version=protocol_version_enum,
)
else:
# For HTTP/SSE transports
@ -461,14 +413,12 @@ class MCPServerManager:
auth_type=server.auth_type,
auth_value=mcp_auth_header or server.authentication_token,
timeout=60.0,
protocol_version=protocol_version_enum,
)
async def _get_tools_from_server(
self,
server: MCPServer,
mcp_auth_header: Optional[str] = None,
mcp_protocol_version: Optional[str] = None,
) -> List[MCPTool]:
"""
Helper method to get tools from a single MCP server with prefixed names.
@ -483,22 +433,18 @@ class MCPServerManager:
verbose_logger.debug(f"Connecting to url: {server.url}")
verbose_logger.info(f"_get_tools_from_server for {server.name}...")
protocol_version = (
mcp_protocol_version if mcp_protocol_version else server.spec_version
)
client = None
try:
client = self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
protocol_version=protocol_version,
)
tools = await self._fetch_tools_with_timeout(client, server.name)
prefixed_tools = self._create_prefixed_tools(tools, server)
return prefixed_tools
except Exception as e:
@ -530,7 +476,7 @@ class MCPServerManager:
async def _list_tools_task():
try:
await client.connect()
tools = await client.list_tools()
verbose_logger.debug(f"Tools from {server_name}: {tools}")
return tools
@ -609,7 +555,6 @@ class MCPServerManager:
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
mcp_protocol_version: Optional[str] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> CallToolResult:
"""
@ -660,32 +605,54 @@ class MCPServerManager:
"arguments": arguments,
"server_name": server_name_from_prefix,
"user_api_key_auth": user_api_key_auth,
"user_api_key_user_id": getattr(user_api_key_auth, 'user_id', None) if user_api_key_auth else None,
"user_api_key_team_id": getattr(user_api_key_auth, 'team_id', None) if user_api_key_auth else None,
"user_api_key_end_user_id": getattr(user_api_key_auth, 'end_user_id', None) if user_api_key_auth else None,
"user_api_key_hash": getattr(user_api_key_auth, 'api_key_hash', None) if user_api_key_auth else None,
"user_api_key_user_id": getattr(user_api_key_auth, "user_id", None)
if user_api_key_auth
else None,
"user_api_key_team_id": getattr(user_api_key_auth, "team_id", None)
if user_api_key_auth
else None,
"user_api_key_end_user_id": getattr(
user_api_key_auth, "end_user_id", None
)
if user_api_key_auth
else None,
"user_api_key_hash": getattr(user_api_key_auth, "api_key_hash", None)
if user_api_key_auth
else None,
}
# Create MCP request object for processing
mcp_request_obj = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs)
mcp_request_obj = proxy_logging_obj._create_mcp_request_object_from_kwargs(
pre_hook_kwargs
)
# Convert to LLM format for existing guardrail compatibility
synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs)
synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(
mcp_request_obj, pre_hook_kwargs
)
try:
# Use standard pre_call_hook with call_type="mcp_call"
modified_data = await proxy_logging_obj.pre_call_hook(
user_api_key_dict=user_api_key_auth, #type: ignore
user_api_key_dict=user_api_key_auth, # type: ignore
data=synthetic_llm_data,
call_type="mcp_call" #type: ignore
call_type="mcp_call", # type: ignore
)
if modified_data:
# Convert response back to MCP format and apply modifications
modified_kwargs = proxy_logging_obj._convert_mcp_hook_response_to_kwargs(modified_data, pre_hook_kwargs)
modified_kwargs = (
proxy_logging_obj._convert_mcp_hook_response_to_kwargs(
modified_data, pre_hook_kwargs
)
)
if modified_kwargs.get("arguments") != arguments:
arguments = modified_kwargs["arguments"]
except (BlockedPiiEntityError, GuardrailRaisedException, HTTPException) as e:
except (
BlockedPiiEntityError,
GuardrailRaisedException,
HTTPException,
) as e:
# Re-raise guardrail exceptions to properly fail the MCP call
verbose_logger.error(
f"Guardrail blocked MCP tool call pre call: {str(e)}"
@ -706,11 +673,9 @@ class MCPServerManager:
client = self._create_mcp_client(
server=mcp_server,
mcp_auth_header=server_auth_header,
protocol_version=mcp_protocol_version,
)
async with client:
# Use the original tool name (without prefix) for the actual call
call_tool_params = MCPCallToolRequestParams(
name=original_tool_name,
@ -721,7 +686,7 @@ class MCPServerManager:
# Create synthetic LLM data for during hook processing
from litellm.types.llms.base import HiddenParams
from litellm.types.mcp import MCPDuringCallRequestObject
request_obj = MCPDuringCallRequestObject(
tool_name=name,
arguments=arguments,
@ -729,28 +694,29 @@ class MCPServerManager:
start_time=start_time.timestamp() if start_time else None,
hidden_params=HiddenParams(),
)
during_hook_kwargs = {
"name": name,
"arguments": arguments,
"server_name": server_name_from_prefix,
"user_api_key_auth": user_api_key_auth,
}
synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(request_obj, during_hook_kwargs)
synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(
request_obj, during_hook_kwargs
)
during_hook_task = asyncio.create_task(
proxy_logging_obj.during_call_hook(
user_api_key_dict=user_api_key_auth,
data=synthetic_llm_data,
call_type="mcp_call" #type: ignore
call_type="mcp_call", # type: ignore
)
)
tasks.append(during_hook_task)
tasks.append(asyncio.create_task(client.call_tool(call_tool_params)))
try:
mcp_responses = await asyncio.gather(*tasks)
# If proxy_logging_obj is None, the tool call result is at index 0
@ -839,19 +805,21 @@ class MCPServerManager:
)
verbose_logger.info("Loading MCP servers from database into registry...")
# perform authz check to filter the mcp servers user has access to
prisma_client = get_prisma_client_or_throw(
"Database not connected. Connect a database to your proxy"
)
db_mcp_servers = await get_all_mcp_servers(prisma_client)
verbose_logger.info(f"Found {len(db_mcp_servers)} MCP servers in database")
# ensure the global_mcp_server_manager is up to date with the db
for server in db_mcp_servers:
verbose_logger.debug(f"Adding server to registry: {server.server_id} ({server.server_name})")
verbose_logger.debug(
f"Adding server to registry: {server.server_id} ({server.server_name})"
)
self.add_update_server(server)
verbose_logger.info(f"Registry now contains {len(self.get_registry())} servers")
def get_mcp_server_by_id(self, server_id: str) -> Optional[MCPServer]:
@ -869,7 +837,6 @@ class MCPServerManager:
server_name: str,
url: str,
transport: str,
spec_version: str,
auth_type: Optional[str] = None,
alias: Optional[str] = None,
) -> str:
@ -885,7 +852,6 @@ class MCPServerManager:
server_name: Name of the server
url: Server URL
transport: Transport type (sse, http, etc.)
spec_version: MCP spec version
auth_type: Authentication type (optional)
alias: Server alias (optional)
@ -893,7 +859,9 @@ class MCPServerManager:
A deterministic server ID string
"""
# Create a string from all the identifying parameters
params_string = f"{server_name}|{url}|{transport}|{spec_version}|{auth_type or ''}|{alias or ''}"
params_string = (
f"{server_name}|{url}|{transport}|{auth_type or ''}|{alias or ''}"
)
# Generate SHA-256 hash
hash_object = hashlib.sha256(params_string.encode("utf-8"))
@ -1050,11 +1018,12 @@ class MCPServerManager:
alias=_server_config.alias,
url=_server_config.url,
transport=_server_config.transport,
spec_version=_server_config.spec_version,
auth_type=_server_config.auth_type,
created_at=datetime.datetime.now(),
updated_at=datetime.datetime.now(),
description=_server_config.mcp_info.get("description") if _server_config.mcp_info else None,
description=_server_config.mcp_info.get("description")
if _server_config.mcp_info
else None,
mcp_info=_server_config.mcp_info,
mcp_access_groups=_server_config.access_groups or [],
# Stdio-specific fields
@ -1111,7 +1080,6 @@ class MCPServerManager:
description=server.description,
url=server.url,
transport=server.transport,
spec_version=server.spec_version,
auth_type=server.auth_type,
created_at=server.created_at,
created_by=server.created_by,

View file

@ -23,7 +23,6 @@ router = APIRouter(
if MCP_AVAILABLE:
from litellm.experimental_mcp_client.client import MCPTool
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_convert_protocol_version_to_enum,
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.server import (
@ -34,18 +33,24 @@ if MCP_AVAILABLE:
########################################################
############ MCP Server REST API Routes #################
def _get_server_auth_header(
server, mcp_server_auth_headers: Optional[Dict[str, str]], mcp_auth_header: Optional[str]
server,
mcp_server_auth_headers: Optional[Dict[str, str]],
mcp_auth_header: Optional[str],
) -> Optional[str]:
"""Helper function to get server-specific auth header with case-insensitive matching."""
if mcp_server_auth_headers and server.alias:
normalized_server_alias = server.alias.lower()
normalized_headers = {k.lower(): v for k, v in mcp_server_auth_headers.items()}
normalized_headers = {
k.lower(): v for k, v in mcp_server_auth_headers.items()
}
server_auth = normalized_headers.get(normalized_server_alias)
if server_auth is not None:
return server_auth
elif mcp_server_auth_headers and server.server_name:
normalized_server_name = server.server_name.lower()
normalized_headers = {k.lower(): v for k, v in mcp_server_auth_headers.items()}
normalized_headers = {
k.lower(): v for k, v in mcp_server_auth_headers.items()
}
server_auth = normalized_headers.get(normalized_server_name)
if server_auth is not None:
return server_auth
@ -63,12 +68,11 @@ if MCP_AVAILABLE:
for tool in tools
]
async def _get_tools_for_single_server(server, server_auth_header, mcp_protocol_version):
async def _get_tools_for_single_server(server, server_auth_header):
"""Helper function to get tools for a single server."""
tools = await global_mcp_server_manager._get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
mcp_protocol_version=mcp_protocol_version,
)
return _create_tool_response_objects(tools, server.mcp_info)
@ -104,17 +108,20 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
try:
# Extract auth headers from request
headers = request.headers
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers)
mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
mcp_protocol_version = headers.get(MCPRequestHandler.MCP_PROTOCOL_VERSION_HEADER_NAME)
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(
headers
)
mcp_server_auth_headers = (
MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
)
list_tools_result = []
error_message = None
# If server_id is specified, only query that specific server
if server_id:
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
@ -122,49 +129,67 @@ if MCP_AVAILABLE:
return {
"tools": [],
"error": "server_not_found",
"message": f"Server with id {server_id} not found"
"message": f"Server with id {server_id} not found",
}
server_auth_header = _get_server_auth_header(server, mcp_server_auth_headers, mcp_auth_header)
server_auth_header = _get_server_auth_header(
server, mcp_server_auth_headers, mcp_auth_header
)
try:
list_tools_result = await _get_tools_for_single_server(server, server_auth_header, mcp_protocol_version)
list_tools_result = await _get_tools_for_single_server(
server, server_auth_header
)
except Exception as e:
verbose_logger.exception(f"Error getting tools from {server.name}: {e}")
verbose_logger.exception(
f"Error getting tools from {server.name}: {e}"
)
return {
"tools": [],
"error": "server_error",
"message": f"Failed to get tools from server {server.name}: {str(e)}"
"message": f"Failed to get tools from server {server.name}: {str(e)}",
}
else:
# Query all servers
errors = []
for server in global_mcp_server_manager.get_registry().values():
server_auth_header = _get_server_auth_header(server, mcp_server_auth_headers, mcp_auth_header)
server_auth_header = _get_server_auth_header(
server, mcp_server_auth_headers, mcp_auth_header
)
try:
tools_result = await _get_tools_for_single_server(server, server_auth_header, mcp_protocol_version)
tools_result = await _get_tools_for_single_server(
server, server_auth_header
)
list_tools_result.extend(tools_result)
except Exception as e:
verbose_logger.exception(f"Error getting tools from {server.name}: {e}")
verbose_logger.exception(
f"Error getting tools from {server.name}: {e}"
)
errors.append(f"{server.name}: {str(e)}")
continue
if errors and not list_tools_result:
error_message = "Failed to get tools from servers: " + "; ".join(errors)
error_message = "Failed to get tools from servers: " + "; ".join(
errors
)
return {
"tools": list_tools_result,
"error": "partial_failure" if error_message else None,
"message": error_message if error_message else "Successfully retrieved tools"
"message": error_message
if error_message
else "Successfully retrieved tools",
}
except Exception as e:
verbose_logger.exception("Unexpected error in list_tool_rest_api: %s", str(e))
verbose_logger.exception(
"Unexpected error in list_tool_rest_api: %s", str(e)
)
return {
"tools": [],
"error": "unexpected_error",
"message": f"An unexpected error occurred: {str(e)}"
"message": f"An unexpected error occurred: {str(e)}",
}
@router.post("/tools/call", dependencies=[Depends(user_api_key_auth)])
@ -196,9 +221,9 @@ if MCP_AVAILABLE:
detail={
"error": "blocked_pii_entity",
"message": str(e),
"entity_type": getattr(e, 'entity_type', None),
"guardrail_name": getattr(e, 'guardrail_name', None)
}
"entity_type": getattr(e, "entity_type", None),
"guardrail_name": getattr(e, "guardrail_name", None),
},
)
except GuardrailRaisedException as e:
verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}")
@ -207,8 +232,8 @@ if MCP_AVAILABLE:
detail={
"error": "guardrail_violation",
"message": str(e),
"guardrail_name": getattr(e, 'guardrail_name', None)
}
"guardrail_name": getattr(e, "guardrail_name", None),
},
)
except HTTPException as e:
# Re-raise HTTPException as-is to preserve status code and detail
@ -220,10 +245,10 @@ if MCP_AVAILABLE:
status_code=500,
detail={
"error": "internal_server_error",
"message": f"An unexpected error occurred: {str(e)}"
}
"message": f"An unexpected error occurred: {str(e)}",
},
)
########################################################
# MCP Connection testing routes
# /health -> Test if we can connect to the MCP server
@ -234,15 +259,15 @@ if MCP_AVAILABLE:
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
NewMCPServerRequest,
)
async def _execute_with_mcp_client(request: NewMCPServerRequest, operation):
"""
Common helper to create MCP client, execute operation, and ensure proper cleanup.
Args:
request: MCP server configuration
operation: Async function that takes a client and returns the operation result
Returns:
Operation result or error response
"""
@ -254,15 +279,14 @@ if MCP_AVAILABLE:
name=request.alias or request.server_name or "",
url=request.url,
transport=request.transport,
spec_version=_convert_protocol_version_to_enum(request.spec_version),
auth_type=request.auth_type,
mcp_info=request.mcp_info,
),
mcp_auth_header=None,
)
return await operation(client)
except Exception as e:
verbose_logger.error(f"Error in MCP operation: {e}", exc_info=True)
return {"status": "error", "message": "An internal error has occurred."}
@ -273,6 +297,7 @@ if MCP_AVAILABLE:
await client.disconnect()
except Exception as e:
verbose_logger.warning(f"Error disconnecting MCP client: {e}")
@router.post("/test/connection")
async def test_connection(
request: NewMCPServerRequest,
@ -280,13 +305,13 @@ if MCP_AVAILABLE:
"""
Test if we can connect to the provided MCP server before adding it
"""
async def _test_connection_operation(client):
await client.connect()
return {"status": "ok"}
return await _execute_with_mcp_client(request, _test_connection_operation)
@router.post("/test/tools/list")
async def test_tools_list(
request: NewMCPServerRequest,
@ -295,13 +320,16 @@ if MCP_AVAILABLE:
"""
Preview tools available from MCP server before adding it
"""
async def _list_tools_operation(client):
list_tools_result: List[MCPTool] = await client.list_tools()
model_dumped_tools: List[dict] = [tool.model_dump() for tool in list_tools_result]
model_dumped_tools: List[dict] = [
tool.model_dump() for tool in list_tools_result
]
return {
"tools": model_dumped_tools,
"error": None,
"message": "Successfully retrieved tools"
"message": "Successfully retrieved tools",
}
return await _execute_with_mcp_client(request, _list_tools_operation)

View file

@ -130,7 +130,9 @@ if MCP_AVAILABLE:
await _sse_session_manager_cm.__aenter__()
_SESSION_MANAGERS_INITIALIZED = True
verbose_logger.info("MCP Server started with StreamableHTTP and SSE session managers!")
verbose_logger.info(
"MCP Server started with StreamableHTTP and SSE session managers!"
)
async def shutdown_session_managers():
"""Shutdown the session managers."""
@ -171,11 +173,18 @@ if MCP_AVAILABLE:
"""
try:
# Get user authentication from context variable
user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, mcp_protocol_version = (
get_auth_context()
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
) = get_auth_context()
verbose_logger.debug(
f"MCP list_tools - User API Key Auth from context: {user_api_key_auth}"
)
verbose_logger.debug(
f"MCP list_tools - MCP servers from context: {mcp_servers}"
)
verbose_logger.debug(f"MCP list_tools - User API Key Auth from context: {user_api_key_auth}")
verbose_logger.debug(f"MCP list_tools - MCP servers from context: {mcp_servers}")
verbose_logger.debug(
f"MCP list_tools - MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}"
)
@ -186,9 +195,10 @@ if MCP_AVAILABLE:
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_protocol_version=mcp_protocol_version,
)
verbose_logger.info(f"MCP list_tools - Successfully returned {len(tools)} tools")
verbose_logger.info(
f"MCP list_tools - Successfully returned {len(tools)} tools"
)
return tools
except Exception as e:
verbose_logger.exception(f"Error in list_tools endpoint: {str(e)}")
@ -220,9 +230,16 @@ if MCP_AVAILABLE:
from litellm.proxy.proxy_server import proxy_config
# Validate arguments
user_api_key_auth, mcp_auth_header, _, mcp_server_auth_headers, mcp_protocol_version = get_auth_context()
(
user_api_key_auth,
mcp_auth_header,
_,
mcp_server_auth_headers,
) = get_auth_context()
verbose_logger.debug(f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}")
verbose_logger.debug(
f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}"
)
try:
# Create a body date for logging
body_data = {"name": name, "arguments": arguments}
@ -249,17 +266,22 @@ if MCP_AVAILABLE:
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_protocol_version=mcp_protocol_version,
**data, # for logging
)
except BlockedPiiEntityError as e:
verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}")
# Return error as text content for MCP protocol
return [TextContent(text=f"Error: Blocked PII entity detected - {str(e)}", type="text")]
return [
TextContent(
text=f"Error: Blocked PII entity detected - {str(e)}", type="text"
)
]
except GuardrailRaisedException as e:
verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}")
# Return error as text content for MCP protocol
return [TextContent(text=f"Error: Guardrail violation - {str(e)}", type="text")]
return [
TextContent(text=f"Error: Guardrail violation - {str(e)}", type="text")
]
except HTTPException as e:
verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}")
# Return error as text content for MCP protocol
@ -287,6 +309,7 @@ if MCP_AVAILABLE:
Get the filtered MCP servers from the MCP server names
"""
from typing import Set
filtered_server_ids: Set[str] = set()
# Filter servers based on mcp_servers parameter if provided
if mcp_servers is not None:
@ -297,7 +320,11 @@ if MCP_AVAILABLE:
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
if server:
match_list = [s.lower() for s in [server.alias, server.server_name, server_id] if s is not None]
match_list = [
s.lower()
for s in [server.alias, server.server_name, server_id]
if s is not None
]
if server_or_group.lower() in match_list:
filtered_server_ids.add(server_id)
@ -306,19 +333,23 @@ if MCP_AVAILABLE:
if not server_name_matched:
try:
access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups(
[server_or_group]
access_group_server_ids = (
await MCPRequestHandler._get_mcp_servers_from_access_groups(
[server_or_group]
)
)
# Only include servers that the user has access to
for server_id in access_group_server_ids:
if server_id in allowed_mcp_servers:
filtered_server_ids.add(server_id)
except Exception as e:
verbose_logger.debug(f"Could not resolve '{server_or_group}' as access group: {e}")
verbose_logger.debug(
f"Could not resolve '{server_or_group}' as access group: {e}"
)
if filtered_server_ids:
allowed_mcp_servers = list(filtered_server_ids)
return allowed_mcp_servers
async def _get_tools_from_mcp_servers(
@ -326,7 +357,6 @@ if MCP_AVAILABLE:
mcp_auth_header: Optional[str],
mcp_servers: Optional[List[str]],
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
mcp_protocol_version: Optional[str] = None,
) -> List[MCPTool]:
"""
Helper method to fetch tools from MCP servers based on server filtering criteria.
@ -344,7 +374,9 @@ if MCP_AVAILABLE:
return []
# Get allowed MCP servers based on user permissions
allowed_mcp_servers = await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth)
allowed_mcp_servers = await global_mcp_server_manager.get_allowed_mcp_servers(
user_api_key_auth
)
if mcp_servers is not None:
allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names(
@ -352,7 +384,6 @@ if MCP_AVAILABLE:
allowed_mcp_servers=allowed_mcp_servers,
)
# Get tools from each allowed server
all_tools = []
for server_id in allowed_mcp_servers:
@ -375,15 +406,20 @@ if MCP_AVAILABLE:
tools = await global_mcp_server_manager._get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
mcp_protocol_version=mcp_protocol_version,
)
all_tools.extend(tools)
verbose_logger.debug(f"Successfully fetched {len(tools)} tools from server {server.name}")
verbose_logger.debug(
f"Successfully fetched {len(tools)} tools from server {server.name}"
)
except Exception as e:
verbose_logger.exception(f"Error getting tools from server {server.name}: {str(e)}")
verbose_logger.exception(
f"Error getting tools from server {server.name}: {str(e)}"
)
# Continue with other servers instead of failing completely
verbose_logger.info(f"Successfully fetched {len(all_tools)} tools total from all MCP servers")
verbose_logger.info(
f"Successfully fetched {len(all_tools)} tools total from all MCP servers"
)
return all_tools
async def _list_mcp_tools(
@ -391,7 +427,6 @@ if MCP_AVAILABLE:
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
mcp_protocol_version: Optional[str] = None,
) -> List[MCPTool]:
"""
List all available MCP tools.
@ -415,11 +450,14 @@ if MCP_AVAILABLE:
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_protocol_version=mcp_protocol_version,
)
verbose_logger.debug(f"Successfully fetched {len(managed_tools)} tools from managed MCP servers")
verbose_logger.debug(
f"Successfully fetched {len(managed_tools)} tools from managed MCP servers"
)
except Exception as e:
verbose_logger.exception(f"Error getting tools from managed MCP servers: {str(e)}")
verbose_logger.exception(
f"Error getting tools from managed MCP servers: {str(e)}"
)
# Continue with empty managed tools list instead of failing completely
# Get tools from local registry
@ -430,10 +468,16 @@ if MCP_AVAILABLE:
# Convert local tools to MCPTool format
for tool in local_tools_raw:
# Convert from litellm.types.mcp_server.tool_registry.MCPTool to mcp.types.Tool
mcp_tool = MCPTool(name=tool.name, description=tool.description, inputSchema=tool.input_schema)
mcp_tool = MCPTool(
name=tool.name,
description=tool.description,
inputSchema=tool.input_schema,
)
local_tools.append(mcp_tool)
except Exception as e:
verbose_logger.exception(f"Error getting tools from local registry: {str(e)}")
verbose_logger.exception(
f"Error getting tools from local registry: {str(e)}"
)
# Continue with empty local tools list instead of failing completely
# Combine all tools
@ -448,7 +492,6 @@ if MCP_AVAILABLE:
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
mcp_protocol_version: Optional[str] = None,
**kwargs: Any,
) -> List[Union[TextContent, ImageContent, EmbeddedResource]]:
"""
@ -456,35 +499,46 @@ if MCP_AVAILABLE:
"""
start_time = datetime.now()
if arguments is None:
raise HTTPException(status_code=400, detail="Request arguments are required")
raise HTTPException(
status_code=400, detail="Request arguments are required"
)
# Remove prefix from tool name for logging and processing
original_tool_name, server_name_from_prefix = get_server_name_prefix_tool_mcp(name)
standard_logging_mcp_tool_call: StandardLoggingMCPToolCall = _get_standard_logging_mcp_tool_call(
name=original_tool_name, # Use original name for logging
arguments=arguments,
server_name=server_name_from_prefix,
original_tool_name, server_name_from_prefix = get_server_name_prefix_tool_mcp(
name
)
standard_logging_mcp_tool_call: StandardLoggingMCPToolCall = (
_get_standard_logging_mcp_tool_call(
name=original_tool_name, # Use original name for logging
arguments=arguments,
server_name=server_name_from_prefix,
)
)
litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get(
"litellm_logging_obj", None
)
litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj", None)
if litellm_logging_obj:
litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = standard_logging_mcp_tool_call
litellm_logging_obj.model_call_details[
"mcp_tool_call_metadata"
] = standard_logging_mcp_tool_call
litellm_logging_obj.model = f"MCP: {name}"
# Try managed server tool first (pass the full prefixed name)
# Primary and recommended way to use MCP servers
#########################################################
mcp_server: Optional[MCPServer] = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
mcp_server: Optional[
MCPServer
] = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
if mcp_server:
standard_logging_mcp_tool_call["mcp_server_cost_info"] = (mcp_server.mcp_info or {}).get(
"mcp_server_cost_info"
)
standard_logging_mcp_tool_call["mcp_server_cost_info"] = (
mcp_server.mcp_info or {}
).get("mcp_server_cost_info")
response = await _handle_managed_mcp_tool(
name=name, # Pass the full name (potentially prefixed)
arguments=arguments,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_protocol_version=mcp_protocol_version,
litellm_logging_obj=litellm_logging_obj,
)
@ -537,7 +591,6 @@ if MCP_AVAILABLE:
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
mcp_protocol_version: Optional[str] = None,
litellm_logging_obj: Optional[Any] = None,
) -> List[Union[TextContent, ImageContent, EmbeddedResource]]:
"""Handle tool execution for managed server tools"""
@ -577,6 +630,7 @@ if MCP_AVAILABLE:
Get the MCP servers from the path
"""
import re
mcp_servers_from_path: Optional[List[str]] = None
# Match /mcp/<servers>/<optional_path>
# Where <servers> can be comma-separated list of server names
@ -627,7 +681,6 @@ if MCP_AVAILABLE:
mcp_auth_header,
_,
mcp_server_auth_headers,
mcp_protocol_version,
) = await MCPRequestHandler.process_mcp_request(scope)
mcp_servers = mcp_servers_from_path
else:
@ -636,11 +689,12 @@ if MCP_AVAILABLE:
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
mcp_protocol_version,
) = await MCPRequestHandler.process_mcp_request(scope)
return user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, mcp_protocol_version
return user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers
async def handle_streamable_http_mcp(scope: Scope, receive: Receive, send: Send) -> None:
async def handle_streamable_http_mcp(
scope: Scope, receive: Receive, send: Send
) -> None:
"""Handle MCP requests through StreamableHTTP."""
try:
path = scope.get("path", "")
@ -649,20 +703,19 @@ if MCP_AVAILABLE:
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
mcp_protocol_version,
) = await extract_mcp_auth_context(scope, path)
verbose_logger.debug(f"MCP request mcp_servers (header/path): {mcp_servers}")
verbose_logger.debug(
f"MCP request mcp_servers (header/path): {mcp_servers}"
)
verbose_logger.debug(
f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}"
)
verbose_logger.debug(f"MCP protocol version: {mcp_protocol_version}")
# Set the auth context variable for easy access in MCP functions
set_auth_context(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_protocol_version=mcp_protocol_version,
)
# Ensure session managers are initialized
@ -686,7 +739,9 @@ if MCP_AVAILABLE:
)
await error_response(scope, receive, send)
except Exception as response_error:
verbose_logger.exception(f"Failed to send error response: {response_error}")
verbose_logger.exception(
f"Failed to send error response: {response_error}"
)
# If we can't send a proper response, re-raise the original error
raise e
@ -699,19 +754,18 @@ if MCP_AVAILABLE:
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
mcp_protocol_version,
) = await extract_mcp_auth_context(scope, path)
verbose_logger.debug(f"MCP request mcp_servers (header/path): {mcp_servers}")
verbose_logger.debug(
f"MCP request mcp_servers (header/path): {mcp_servers}"
)
verbose_logger.debug(
f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}"
)
verbose_logger.debug(f"MCP protocol version: {mcp_protocol_version}")
set_auth_context(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_protocol_version=mcp_protocol_version,
)
if not _SESSION_MANAGERS_INITIALIZED:
@ -733,7 +787,9 @@ if MCP_AVAILABLE:
)
await error_response(scope, receive, send)
except Exception as response_error:
verbose_logger.exception(f"Failed to send error response: {response_error}")
verbose_logger.exception(
f"Failed to send error response: {response_error}"
)
# If we can't send a proper response, re-raise the original error
raise e
@ -769,7 +825,6 @@ if MCP_AVAILABLE:
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
mcp_protocol_version: Optional[str] = None,
) -> None:
"""
Set the UserAPIKeyAuth in the auth context variable.
@ -785,13 +840,17 @@ if MCP_AVAILABLE:
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_protocol_version=mcp_protocol_version,
)
auth_context_var.set(auth_user)
def get_auth_context() -> Tuple[
Optional[UserAPIKeyAuth], Optional[str], Optional[List[str]], Optional[Dict[str, str]], Optional[str]
]:
def get_auth_context() -> (
Tuple[
Optional[UserAPIKeyAuth],
Optional[str],
Optional[List[str]],
Optional[Dict[str, str]],
]
):
"""
Get the UserAPIKeyAuth from the auth context variable.
@ -806,9 +865,8 @@ if MCP_AVAILABLE:
auth_user.mcp_auth_header,
auth_user.mcp_servers,
auth_user.mcp_server_auth_headers,
auth_user.mcp_protocol_version,
)
return None, None, None, None, None
return None, None, None, None
########################################################
############ End of Auth Context Functions #############

View file

@ -15,18 +15,3 @@ model_list:
model: hosted_vllm/whisper-v3
api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5"
api_key: dummy
guardrails:
- guardrail_name: "intel-bedrock-guard-cfg"
litellm_params:
guardrail: bedrock
mode: [pre_call, post_call]
guardrailIdentifier: "1234"
guardrailVersion: "1"
aws_access_key_id: "os.environ/AWS_ACCESS_KEY_ID"
aws_secret_access_key: "os.environ/AWS_SECRET_ACCESS_KEY"
aws_bedrock_runtime_endpoint: "os.environ/AWS_BEDROCK_RUNTIME_ENDPOINT"
default_on: true

View file

@ -28,8 +28,6 @@ from litellm.types.integrations.slack_alerting import AlertType
from litellm.types.llms.openai import AllMessageValues, OpenAIFileObject
from litellm.types.mcp import (
MCPAuthType,
MCPSpecVersion,
MCPSpecVersionType,
MCPTransport,
MCPTransportType,
)
@ -748,9 +746,9 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
allowed_cache_controls: Optional[list] = []
config: Optional[dict] = {}
permissions: Optional[dict] = {}
model_max_budget: Optional[dict] = (
{}
) # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
model_max_budget: Optional[
dict
] = {} # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
model_config = ConfigDict(protected_namespaces=())
model_rpm_limit: Optional[dict] = None
@ -789,6 +787,7 @@ class GenerateKeyRequest(KeyRequestBase):
description="Type of key that determines default allowed routes.",
)
class GenerateKeyResponse(KeyRequestBase):
key: str # type: ignore
key_name: Optional[str] = None
@ -916,7 +915,6 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
alias: Optional[str] = None
description: Optional[str] = None
transport: MCPTransportType = MCPTransport.sse
spec_version: MCPSpecVersionType = MCPSpecVersion.jun_2025
auth_type: Optional[MCPAuthType] = None
url: Optional[str] = None
mcp_info: Optional[MCPInfo] = None
@ -948,7 +946,6 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
alias: Optional[str] = None
description: Optional[str] = None
transport: MCPTransportType = MCPTransport.sse
spec_version: MCPSpecVersionType = MCPSpecVersion.jun_2025
auth_type: Optional[MCPAuthType] = None
url: Optional[str] = None
mcp_info: Optional[MCPInfo] = None
@ -983,7 +980,6 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
description: Optional[str] = None
url: Optional[str] = None
transport: MCPTransportType
spec_version: MCPSpecVersionType
auth_type: Optional[MCPAuthType] = None
created_at: Optional[datetime] = None
created_by: Optional[str] = None
@ -1150,12 +1146,12 @@ class NewCustomerRequest(BudgetNewRequest):
blocked: bool = False # allow/disallow requests for this end-user
budget_id: Optional[str] = None # give either a budget_id or max_budget
spend: Optional[float] = None
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
@model_validator(mode="before")
@classmethod
@ -1177,12 +1173,12 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase):
blocked: bool = False # allow/disallow requests for this end-user
max_budget: Optional[float] = None
budget_id: Optional[str] = None # give either a budget_id or max_budget
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
class DeleteCustomerRequest(LiteLLMPydanticObjectBase):
@ -1256,15 +1252,15 @@ class NewTeamRequest(TeamBase):
guardrails: Optional[List[str]] = None
prompts: Optional[List[str]] = None
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
team_member_budget: Optional[float] = (
None # allow user to set a budget for all team members
)
team_member_rpm_limit: Optional[int] = (
None # allow user to set RPM limit for all team members
)
team_member_tpm_limit: Optional[int] = (
None # allow user to set TPM limit for all team members
)
team_member_budget: Optional[
float
] = None # allow user to set a budget for all team members
team_member_rpm_limit: Optional[
int
] = None # allow user to set RPM limit for all team members
team_member_tpm_limit: Optional[
int
] = None # allow user to set TPM limit for all team members
team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m"
model_config = ConfigDict(protected_namespaces=())
@ -1343,9 +1339,9 @@ class BlockKeyRequest(LiteLLMPydanticObjectBase):
class AddTeamCallback(LiteLLMPydanticObjectBase):
callback_name: str
callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = (
"success_and_failure"
)
callback_type: Optional[
Literal["success", "failure", "success_and_failure"]
] = "success_and_failure"
callback_vars: Dict[str, str]
@model_validator(mode="before")
@ -1614,15 +1610,16 @@ class ConfigList(LiteLLMPydanticObjectBase):
stored_in_db: Optional[bool]
field_default_value: Any
premium_field: bool = False
nested_fields: Optional[List[FieldDetail]] = (
None # For nested dictionary or Pydantic fields
)
nested_fields: Optional[
List[FieldDetail]
] = None # For nested dictionary or Pydantic fields
class UserHeaderMapping(LiteLLMPydanticObjectBase):
"""
Map an incoming HTTP header to a LiteLLM user role.
"""
header_name: str
litellm_user_role: Literal[
LitellmUserRoles.INTERNAL_USER,
@ -1633,6 +1630,7 @@ class UserHeaderMapping(LiteLLMPydanticObjectBase):
"extra": "forbid",
}
class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"""
Documents all the fields supported by `general_settings` in config.yaml
@ -1943,9 +1941,9 @@ class LiteLLM_OrganizationMembershipTable(LiteLLMPydanticObjectBase):
budget_id: Optional[str] = None
created_at: datetime
updated_at: datetime
user: Optional[Any] = (
None # You might want to replace 'Any' with a more specific type if available
)
user: Optional[
Any
] = None # You might want to replace 'Any' with a more specific type if available
litellm_budget_table: Optional[LiteLLM_BudgetTable] = None
model_config = ConfigDict(protected_namespaces=())
@ -2840,9 +2838,9 @@ class TeamModelDeleteRequest(BaseModel):
# Organization Member Requests
class OrganizationMemberAddRequest(OrgMemberAddRequest):
organization_id: str
max_budget_in_organization: Optional[float] = (
None # Users max budget within the organization
)
max_budget_in_organization: Optional[
float
] = None # Users max budget within the organization
class OrganizationMemberDeleteRequest(MemberDeleteRequest):
@ -2941,10 +2939,12 @@ class LitellmDataForBackendLLMCall(TypedDict, total=False):
user: Optional[str]
num_retries: Optional[int]
class LitellmMetadataFromRequestHeaders(TypedDict, total=False):
"""
Headers a user can pass that will get added to litellm metadata for the request
"""
spend_logs_metadata: Optional[dict]
@ -3050,9 +3050,9 @@ class ProviderBudgetResponse(LiteLLMPydanticObjectBase):
Maps provider names to their budget configs.
"""
providers: Dict[str, ProviderBudgetResponseObject] = (
{}
) # Dictionary mapping provider names to their budget configurations
providers: Dict[
str, ProviderBudgetResponseObject
] = {} # Dictionary mapping provider names to their budget configurations
class ProxyStateVariables(TypedDict):
@ -3186,9 +3186,9 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
enforce_rbac: bool = False
roles_jwt_field: Optional[str] = None # v2 on role mappings
role_mappings: Optional[List[RoleMapping]] = None
object_id_jwt_field: Optional[str] = (
None # can be either user / team, inferred from the role mapping
)
object_id_jwt_field: Optional[
str
] = None # can be either user / team, inferred from the role mapping
scope_mappings: Optional[List[ScopeMapping]] = None
enforce_scope_based_access: bool = False
enforce_team_based_model_access: bool = False

View file

@ -1,4 +1,5 @@
import json
import re
from typing import Any, Dict, List, Optional
import orjson
@ -51,8 +52,6 @@ async def _read_request_body(request: Optional[Request]) -> Dict:
body_str = body.decode("utf-8") if isinstance(body, bytes) else body
# Replace invalid surrogate pairs
import re
# This regex finds incomplete surrogate pairs
body_str = re.sub(
r"[\uD800-\uDBFF](?![\uDC00-\uDFFF])", "", body_str

View file

@ -24,7 +24,6 @@ from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.custom_httpx.http_handler import (
@ -111,6 +110,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
self.guardrailIdentifier = guardrailIdentifier
self.guardrailVersion = guardrailVersion
self.guardrail_provider = "bedrock"
# store kwargs as optional_params
self.optional_params = kwargs
@ -372,6 +372,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# Add guardrail information to request trace
#########################################################
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response=response.json(),
request_data=request_data or {},
guardrail_status=self._get_bedrock_guardrail_response_status(
@ -403,6 +404,39 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return bedrock_guardrail_response
def _check_bedrock_response_for_exception(self, response) -> bool:
"""
Return True if the Bedrock ApplyGuardrail response indicates an exception.
Works with real httpx.Response objects and MagicMock responses used in tests.
"""
payload = None
try:
json_method = getattr(response, "json", None)
if callable(json_method):
payload = json_method()
except Exception:
payload = None
if payload is None:
try:
raw = getattr(response, "content", None)
if isinstance(raw, (bytes, bytearray)):
payload = json.loads(raw.decode("utf-8"))
else:
text = getattr(response, "text", None)
if isinstance(text, str):
payload = json.loads(text)
except Exception:
# Can't parse -> assume no explicit Exception marker
return False
if not isinstance(payload, dict):
return False
return "Exception" in payload.get("Output", {}).get("__type", "")
def _get_bedrock_guardrail_response_status(
self, response: httpx.Response
) -> Literal["success", "failure"]:
@ -410,6 +444,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
Get the status of the bedrock guardrail response.
"""
if response.status_code == 200:
if self._check_bedrock_response_for_exception(response):
return "failure"
return "success"
return "failure"
@ -516,7 +552,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# This means all actions were ANONYMIZED or NONE, so don't raise exception
return False
@log_guardrail_information
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -562,11 +597,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
#########################################################
########## 2. Update the messages with the guardrail response ##########
#########################################################
data["messages"] = (
self._update_messages_with_updated_bedrock_guardrail_response(
messages=new_messages,
bedrock_guardrail_response=bedrock_guardrail_response,
)
data[
"messages"
] = self._update_messages_with_updated_bedrock_guardrail_response(
messages=new_messages,
bedrock_guardrail_response=bedrock_guardrail_response,
)
#########################################################
@ -617,11 +652,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
#########################################################
########## 2. Update the messages with the guardrail response ##########
#########################################################
data["messages"] = (
self._update_messages_with_updated_bedrock_guardrail_response(
messages=new_messages,
bedrock_guardrail_response=bedrock_guardrail_response,
)
data[
"messages"
] = self._update_messages_with_updated_bedrock_guardrail_response(
messages=new_messages,
bedrock_guardrail_response=bedrock_guardrail_response,
)
#########################################################

View file

@ -76,6 +76,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
self.logging_only = True
kwargs["event_hook"] = GuardrailEventHooks.logging_only
super().__init__(**kwargs)
self.guardrail_provider = "presidio"
self.pii_tokens: dict = (
{}
) # mapping of PII token to original text - only used with Presidio `replace` operation
@ -369,6 +370,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
else:
guardrail_json_response = exception_str
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response=guardrail_json_response,
request_data=request_data,
guardrail_status=status,

View file

@ -21,7 +21,9 @@ from litellm.proxy._types import (
)
# Cache special headers as a frozenset for O(1) lookup performance
_SPECIAL_HEADERS_CACHE = frozenset(v.value.lower() for v in SpecialHeaders._member_map_.values())
_SPECIAL_HEADERS_CACHE = frozenset(
v.value.lower() for v in SpecialHeaders._member_map_.values()
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.router import Router
from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS
@ -64,6 +66,7 @@ LITELLM_METADATA_ROUTES = (
"files",
)
def _get_metadata_variable_name(request: Request) -> str:
"""
Helper to return what the "metadata" field should be called in the request data
@ -157,6 +160,7 @@ class KeyAndTeamLoggingSettings:
@staticmethod
def get_team_dynamic_logging_settings(user_api_key_dict: UserAPIKeyAuth):
if (
user_api_key_dict.team_metadata is not None
and "logging" in user_api_key_dict.team_metadata
@ -169,12 +173,12 @@ def _get_dynamic_logging_metadata(
user_api_key_dict: UserAPIKeyAuth, proxy_config: ProxyConfig
) -> Optional[TeamCallbackMetadata]:
callback_settings_obj: Optional[TeamCallbackMetadata] = None
key_dynamic_logging_settings: Optional[
dict
] = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict)
team_dynamic_logging_settings: Optional[
dict
] = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict)
key_dynamic_logging_settings: Optional[dict] = (
KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict)
)
team_dynamic_logging_settings: Optional[dict] = (
KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict)
)
#########################################################################################
# Key-based callbacks
#########################################################################################
@ -234,12 +238,16 @@ def clean_headers(
Removes litellm api key from headers
"""
clean_headers = {}
litellm_key_lower = litellm_key_header_name.lower() if litellm_key_header_name is not None else None
litellm_key_lower = (
litellm_key_header_name.lower() if litellm_key_header_name is not None else None
)
for header, value in headers.items():
header_lower = header.lower()
# Check if header should be excluded: either in special headers cache or matches custom litellm key
if (header_lower not in _SPECIAL_HEADERS_CACHE and (litellm_key_lower is None or header_lower != litellm_key_lower)):
if header_lower not in _SPECIAL_HEADERS_CACHE and (
litellm_key_lower is None or header_lower != litellm_key_lower
):
clean_headers[header] = value
return clean_headers
@ -614,11 +622,11 @@ class LiteLLMProxyRequestSetup:
## KEY-LEVEL SPEND LOGS / TAGS
if "tags" in key_metadata and key_metadata["tags"] is not None:
data[_metadata_variable_name][
"tags"
] = LiteLLMProxyRequestSetup._merge_tags(
request_tags=data[_metadata_variable_name].get("tags"),
tags_to_add=key_metadata["tags"],
data[_metadata_variable_name]["tags"] = (
LiteLLMProxyRequestSetup._merge_tags(
request_tags=data[_metadata_variable_name].get("tags"),
tags_to_add=key_metadata["tags"],
)
)
if "spend_logs_metadata" in key_metadata and isinstance(
key_metadata["spend_logs_metadata"], dict
@ -847,9 +855,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915
data[_metadata_variable_name]["litellm_api_version"] = version
if general_settings is not None:
data[_metadata_variable_name][
"global_max_parallel_requests"
] = general_settings.get("global_max_parallel_requests", None)
data[_metadata_variable_name]["global_max_parallel_requests"] = (
general_settings.get("global_max_parallel_requests", None)
)
### KEY-LEVEL Controls
key_metadata = user_api_key_dict.metadata

View file

@ -1,5 +1,3 @@
"""
1. Allow proxy admin to perform create, update, and delete operations on MCP servers in the db.
2. Allows users to view the mcp servers they have access to.
@ -82,7 +80,6 @@ if MCP_AVAILABLE:
return True
return False
# Router to fetch all MCP tools available for the current key
@router.get(
@ -91,18 +88,18 @@ if MCP_AVAILABLE:
dependencies=[Depends(user_api_key_auth)],
)
async def get_mcp_tools(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Get all MCP tools available for the current key, including those from access groups
"""
from litellm.proxy._experimental.mcp_server.server import _list_mcp_tools
tools = await _list_mcp_tools(
user_api_key_auth=user_api_key_dict,
mcp_auth_header=None,
mcp_servers=None,
mcp_server_auth_headers=None,
mcp_protocol_version=None,
)
dumped_tools = [dict(tool) for tool in tools]
@ -114,7 +111,7 @@ if MCP_AVAILABLE:
dependencies=[Depends(user_api_key_auth)],
)
async def get_mcp_access_groups(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Get all available MCP access groups from the database AND config
@ -136,7 +133,10 @@ if MCP_AVAILABLE:
try:
mcp_servers = await prisma_client.db.litellm_mcpservertable.find_many()
for server in mcp_servers:
if hasattr(server, 'mcp_access_groups') and server.mcp_access_groups:
if (
hasattr(server, "mcp_access_groups")
and server.mcp_access_groups
):
access_groups.update(server.mcp_access_groups)
except Exception as e:
verbose_proxy_logger.debug(f"Error getting MCP access groups: {e}")
@ -194,10 +194,14 @@ if MCP_AVAILABLE:
# Perform health check using server manager
try:
health_result = await global_mcp_server_manager.health_check_server(server_id)
health_result = await global_mcp_server_manager.health_check_server(
server_id
)
return health_result
except Exception as e:
verbose_proxy_logger.exception(f"Error performing health check on MCP server {server_id}: {str(e)}")
verbose_proxy_logger.exception(
f"Error performing health check on MCP server {server_id}: {str(e)}"
)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Error performing health check: {str(e)}"},
@ -220,19 +224,33 @@ if MCP_AVAILABLE:
"""
# Use server manager to get health checks for allowed servers
try:
all_health_results = await global_mcp_server_manager.health_check_allowed_servers(
user_api_key_auth=user_api_key_dict
all_health_results = (
await global_mcp_server_manager.health_check_allowed_servers(
user_api_key_auth=user_api_key_dict
)
)
return {
"total_servers": len(all_health_results),
"healthy_count": len([r for r in all_health_results.values() if r["status"] == "healthy"]),
"unhealthy_count": len([r for r in all_health_results.values() if r["status"] == "unhealthy"]),
"unknown_count": len([r for r in all_health_results.values() if r["status"] == "unknown"]),
"servers": all_health_results
"healthy_count": len(
[r for r in all_health_results.values() if r["status"] == "healthy"]
),
"unhealthy_count": len(
[
r
for r in all_health_results.values()
if r["status"] == "unhealthy"
]
),
"unknown_count": len(
[r for r in all_health_results.values() if r["status"] == "unknown"]
),
"servers": all_health_results,
}
except Exception as e:
verbose_proxy_logger.exception(f"Error performing health checks on MCP servers: {str(e)}")
verbose_proxy_logger.exception(
f"Error performing health checks on MCP servers: {str(e)}"
)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Error performing health checks: {str(e)}"},
@ -256,8 +274,10 @@ if MCP_AVAILABLE:
```
"""
# Use server manager to get all servers with health and team data
return await global_mcp_server_manager.get_all_mcp_servers_with_health_and_teams(
user_api_key_auth=user_api_key_dict
return (
await global_mcp_server_manager.get_all_mcp_servers_with_health_and_teams(
user_api_key_auth=user_api_key_dict
)
)
@router.get(
@ -293,13 +313,23 @@ if MCP_AVAILABLE:
# Perform health check on the server using server manager
try:
health_result = await global_mcp_server_manager.health_check_server(server_id)
health_result = await global_mcp_server_manager.health_check_server(
server_id
)
# Update the server object with health check results
mcp_server.status = health_result.get("status", "unknown")
mcp_server.last_health_check = datetime.fromisoformat(health_result.get("last_health_check", datetime.now().isoformat())) if health_result.get("last_health_check") else None
mcp_server.last_health_check = (
datetime.fromisoformat(
health_result.get("last_health_check", datetime.now().isoformat())
)
if health_result.get("last_health_check")
else None
)
mcp_server.health_check_error = health_result.get("error")
except Exception as e:
verbose_proxy_logger.debug(f"Error performing health check on server {server_id}: {e}")
verbose_proxy_logger.debug(
f"Error performing health check on server {server_id}: {e}"
)
mcp_server.status = "unknown"
mcp_server.last_health_check = datetime.now()
mcp_server.health_check_error = str(e)
@ -390,7 +420,7 @@ if MCP_AVAILABLE:
touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
)
global_mcp_server_manager.add_update_server(new_mcp_server)
# Ensure registry is up to date by reloading from database
await global_mcp_server_manager.reload_servers_from_database()
except Exception as e:
@ -451,7 +481,7 @@ if MCP_AVAILABLE:
detail={"error": f"MCP Server not found, passed server_id={server_id}"},
)
global_mcp_server_manager.remove_server(mcp_server_record_deleted)
# Ensure registry is up to date by reloading from database
await global_mcp_server_manager.reload_servers_from_database()
@ -526,7 +556,7 @@ if MCP_AVAILABLE:
},
)
global_mcp_server_manager.add_update_server(mcp_server_record_updated)
# Ensure registry is up to date by reloading from database
await global_mcp_server_manager.reload_servers_from_database()
@ -535,4 +565,3 @@ if MCP_AVAILABLE:
pass
return mcp_server_record_updated

View file

@ -6,7 +6,7 @@ Use this when each team should control its own callbacks
import json
import traceback
from typing import Optional
from typing import List, Optional
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
@ -79,10 +79,14 @@ async def add_team_callbacks(
"""
try:
from litellm.proxy._types import CommonProxyErrors
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail={"error": "No db connected"})
raise HTTPException(
status_code=500,
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
# Check if team_id exists already
_existing_team = await prisma_client.get_data(
@ -98,65 +102,30 @@ async def add_team_callbacks(
# store team callback settings in metadata
team_metadata = _existing_team.metadata
team_callback_settings = team_metadata.get("callback_settings", {})
# expect callback settings to be
team_callback_settings_obj = TeamCallbackMetadata(**team_callback_settings)
if data.callback_type == "success":
if team_callback_settings_obj.success_callback is None:
team_callback_settings_obj.success_callback = []
team_callback_settings: List[dict] = team_metadata.get(
"logging"
) # will be dict of type AddTeamCallback
if team_callback_settings is None or not isinstance(
team_callback_settings, list
):
team_callback_settings = []
if data.callback_name in team_callback_settings_obj.success_callback:
## check if it already exists, for the same callback event
for callback in team_callback_settings:
if (
callback.get("callback_name") == data.callback_name
and callback.get("callback_type") == data.callback_type
):
raise ProxyException(
message=f"callback_name = {data.callback_name} already exists in failure_callback, for team_id = {team_id}. \n Existing failure_callback = {team_callback_settings_obj.success_callback}",
message=f"callback_name = {data.callback_name} already exists in team_callback_settings, for team_id = {team_id} and event = {data.callback_type}",
code=status.HTTP_400_BAD_REQUEST,
type=ProxyErrorTypes.bad_request_error,
param="callback_name",
)
team_callback_settings_obj.success_callback.append(data.callback_name)
elif data.callback_type == "failure":
if team_callback_settings_obj.failure_callback is None:
team_callback_settings_obj.failure_callback = []
team_callback_settings.append(data.model_dump())
if data.callback_name in team_callback_settings_obj.failure_callback:
raise ProxyException(
message=f"callback_name = {data.callback_name} already exists in failure_callback, for team_id = {team_id}. \n Existing failure_callback = {team_callback_settings_obj.failure_callback}",
code=status.HTTP_400_BAD_REQUEST,
type=ProxyErrorTypes.bad_request_error,
param="callback_name",
)
team_callback_settings_obj.failure_callback.append(data.callback_name)
elif data.callback_type == "success_and_failure":
if team_callback_settings_obj.success_callback is None:
team_callback_settings_obj.success_callback = []
if team_callback_settings_obj.failure_callback is None:
team_callback_settings_obj.failure_callback = []
if data.callback_name in team_callback_settings_obj.success_callback:
raise ProxyException(
message=f"callback_name = {data.callback_name} already exists in success_callback, for team_id = {team_id}. \n Existing success_callback = {team_callback_settings_obj.success_callback}",
code=status.HTTP_400_BAD_REQUEST,
type=ProxyErrorTypes.bad_request_error,
param="callback_name",
)
if data.callback_name in team_callback_settings_obj.failure_callback:
raise ProxyException(
message=f"callback_name = {data.callback_name} already exists in failure_callback, for team_id = {team_id}. \n Existing failure_callback = {team_callback_settings_obj.failure_callback}",
code=status.HTTP_400_BAD_REQUEST,
type=ProxyErrorTypes.bad_request_error,
param="callback_name",
)
team_callback_settings_obj.success_callback.append(data.callback_name)
team_callback_settings_obj.failure_callback.append(data.callback_name)
for var, value in data.callback_vars.items():
if team_callback_settings_obj.callback_vars is None:
team_callback_settings_obj.callback_vars = {}
team_callback_settings_obj.callback_vars[var] = value
team_callback_settings_obj_dict = team_callback_settings_obj.model_dump()
team_metadata["callback_settings"] = team_callback_settings_obj_dict
team_metadata["logging"] = team_callback_settings
team_metadata_json = json.dumps(team_metadata) # update team_metadata
new_team_row = await prisma_client.db.litellm_teamtable.update(
@ -168,22 +137,16 @@ async def add_team_callbacks(
"data": new_team_row,
}
except HTTPException as e:
raise e
except ProxyException as e:
raise e
except Exception as e:
verbose_proxy_logger.error(
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.add_team_callbacks(): Exception occured - {}".format(
str(e)
)
)
verbose_proxy_logger.debug(traceback.format_exc())
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "detail", f"Internal Server Error({str(e)})"),
type=ProxyErrorTypes.internal_server_error.value,
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR),
)
elif isinstance(e, ProxyException):
raise e
raise ProxyException(
message="Internal Server Error, " + str(e),
type=ProxyErrorTypes.internal_server_error.value,

View file

@ -171,7 +171,6 @@ model LiteLLM_MCPServerTable {
description String?
url String?
transport String @default("sse")
spec_version String @default("2025-03-26")
auth_type String?
created_at DateTime? @default(now()) @map("created_at")
created_by String?

View file

@ -13,6 +13,7 @@ else:
LITELLM_PROXY_MCP_SERVER_URL = "litellm_proxy"
LITELLM_PROXY_MCP_SERVER_URL_PREFIX = f"{LITELLM_PROXY_MCP_SERVER_URL}/mcp/"
class LiteLLM_Proxy_MCP_Handler:
"""
Helper class with static methods for MCP integration with Responses API.
@ -29,34 +30,40 @@ class LiteLLM_Proxy_MCP_Handler:
for tool in tools:
if isinstance(tool, dict) and tool.get("type") == "mcp":
server_url = tool.get("server_url", "")
if isinstance(server_url, str) and server_url.startswith(LITELLM_PROXY_MCP_SERVER_URL):
if isinstance(server_url, str) and server_url.startswith(
LITELLM_PROXY_MCP_SERVER_URL
):
return True
return False
@staticmethod
def _parse_mcp_tools(tools: Optional[Iterable[ToolParam]]) -> Tuple[List[ToolParam], List[Any]]:
def _parse_mcp_tools(
tools: Optional[Iterable[ToolParam]],
) -> Tuple[List[ToolParam], List[Any]]:
"""
Parse tools and separate MCP tools with litellm_proxy from other tools.
Returns:
Tuple of (mcp_tools_with_litellm_proxy, other_tools)
"""
mcp_tools_with_litellm_proxy: List[ToolParam] = []
other_tools: List[Any] = []
if tools:
for tool in tools:
if isinstance(tool, dict) and tool.get("type") == "mcp":
server_url = tool.get("server_url", "")
if isinstance(server_url, str) and server_url.startswith(LITELLM_PROXY_MCP_SERVER_URL):
if isinstance(server_url, str) and server_url.startswith(
LITELLM_PROXY_MCP_SERVER_URL
):
mcp_tools_with_litellm_proxy.append(tool)
else:
other_tools.append(tool)
else:
other_tools.append(tool)
return mcp_tools_with_litellm_proxy, other_tools
@staticmethod
async def _get_mcp_tools_from_manager(
user_api_key_auth: Any,
@ -64,7 +71,7 @@ class LiteLLM_Proxy_MCP_Handler:
) -> List[MCPTool]:
"""
Get available tools from the MCP server manager.
Args:
user_api_key_auth: User authentication info for access control
mcp_tools_with_litellm_proxy: ToolParam objects with server_url starting with "litellm_proxy"
@ -72,48 +79,57 @@ class LiteLLM_Proxy_MCP_Handler:
from litellm.proxy._experimental.mcp_server.server import (
_get_tools_from_mcp_servers,
)
mcp_servers: List[str] = []
if mcp_tools_with_litellm_proxy:
for _tool in mcp_tools_with_litellm_proxy:
# if user specifies servers as server_url: litellm_proxy/mcp/zapier,github then return zapier,github
server_url = _tool.get("server_url", "") if isinstance(_tool, dict) else ""
if isinstance(server_url, str) and server_url.startswith(LITELLM_PROXY_MCP_SERVER_URL_PREFIX):
server_url = (
_tool.get("server_url", "") if isinstance(_tool, dict) else ""
)
if isinstance(server_url, str) and server_url.startswith(
LITELLM_PROXY_MCP_SERVER_URL_PREFIX
):
mcp_servers.append(server_url.split("/")[-1])
return await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=None,
mcp_servers=mcp_servers,
mcp_server_auth_headers=None,
mcp_protocol_version=None,
)
@staticmethod
def _deduplicate_mcp_tools(mcp_tools: List[Any]) -> List[Any]:
"""
Deduplicate MCP tools by name, keeping the first occurrence of each tool.
Args:
mcp_tools: List of MCP tools that may contain duplicates
Returns:
List of deduplicated MCP tools
"""
seen_names = set()
deduplicated_tools = []
for tool in mcp_tools:
tool_name = getattr(tool, 'name', None) if hasattr(tool, 'name') else tool.get('name') if isinstance(tool, dict) else None
tool_name = (
getattr(tool, "name", None)
if hasattr(tool, "name")
else tool.get("name")
if isinstance(tool, dict)
else None
)
if tool_name and tool_name not in seen_names:
seen_names.add(tool_name)
deduplicated_tools.append(tool)
return deduplicated_tools
@staticmethod
def _filter_mcp_tools_by_allowed_tools(
mcp_tools: List[Any],
mcp_tools_with_litellm_proxy: List[ToolParam]
mcp_tools: List[Any], mcp_tools_with_litellm_proxy: List[ToolParam]
) -> List[Any]:
"""Filter MCP tools based on allowed_tools parameter from the original tool configs."""
# Collect all allowed tool names from all MCP tool configs
@ -123,127 +139,135 @@ class LiteLLM_Proxy_MCP_Handler:
allowed_tools = tool_config.get("allowed_tools", [])
if isinstance(allowed_tools, list):
allowed_tool_names.update(allowed_tools)
# If no allowed_tools specified, return all tools
if not allowed_tool_names:
return mcp_tools
# Filter tools based on allowed names
filtered_tools = []
for mcp_tool in mcp_tools:
tool_name = getattr(mcp_tool, 'name', None) if hasattr(mcp_tool, 'name') else mcp_tool.get('name') if isinstance(mcp_tool, dict) else None
tool_name = (
getattr(mcp_tool, "name", None)
if hasattr(mcp_tool, "name")
else mcp_tool.get("name")
if isinstance(mcp_tool, dict)
else None
)
if tool_name and tool_name in allowed_tool_names:
filtered_tools.append(mcp_tool)
return filtered_tools
@staticmethod
async def _process_mcp_tools_to_openai_format(
user_api_key_auth: Any,
mcp_tools_with_litellm_proxy: List[ToolParam]
user_api_key_auth: Any, mcp_tools_with_litellm_proxy: List[ToolParam]
) -> List[Any]:
"""
Centralized method to process MCP tools through the complete pipeline:
1. Fetch tools from MCP manager
2. Filter based on allowed_tools parameter
2. Filter based on allowed_tools parameter
3. Deduplicate tools by name
4. Transform to OpenAI format
Args:
user_api_key_auth: User authentication info for access control
mcp_tools_with_litellm_proxy: ToolParam objects with server_url starting with "litellm_proxy"
Returns:
List of tools in OpenAI format ready to be sent to the LLM
"""
if not mcp_tools_with_litellm_proxy:
return []
# Step 1: Fetch MCP tools from manager
mcp_tools_fetched = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager(
user_api_key_auth=user_api_key_auth,
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
)
# Step 2: Filter tools based on allowed_tools parameter
filtered_mcp_tools = LiteLLM_Proxy_MCP_Handler._filter_mcp_tools_by_allowed_tools(
mcp_tools=mcp_tools_fetched,
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
filtered_mcp_tools = (
LiteLLM_Proxy_MCP_Handler._filter_mcp_tools_by_allowed_tools(
mcp_tools=mcp_tools_fetched,
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
)
)
# Step 3: Deduplicate tools after filtering
deduplicated_mcp_tools = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(
filtered_mcp_tools
)
# Step 4: Transform to OpenAI format
openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(
deduplicated_mcp_tools
)
return openai_tools
@staticmethod
async def _process_mcp_tools_without_openai_transform(
user_api_key_auth: Any,
mcp_tools_with_litellm_proxy: List[ToolParam]
user_api_key_auth: Any, mcp_tools_with_litellm_proxy: List[ToolParam]
) -> List[Any]:
"""
Process MCP tools through filtering and deduplication pipeline without OpenAI transformation.
This is useful for cases where we need the original MCP tool objects (e.g., for events).
Args:
user_api_key_auth: User authentication info for access control
mcp_tools_with_litellm_proxy: ToolParam objects with server_url starting with "litellm_proxy"
Returns:
List of filtered and deduplicated MCP tools in their original format
"""
if not mcp_tools_with_litellm_proxy:
return []
# Step 1: Fetch MCP tools from manager
mcp_tools_fetched = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager(
user_api_key_auth=user_api_key_auth,
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
)
# Step 2: Filter tools based on allowed_tools parameter
filtered_mcp_tools = LiteLLM_Proxy_MCP_Handler._filter_mcp_tools_by_allowed_tools(
mcp_tools=mcp_tools_fetched,
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
filtered_mcp_tools = (
LiteLLM_Proxy_MCP_Handler._filter_mcp_tools_by_allowed_tools(
mcp_tools=mcp_tools_fetched,
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
)
)
# Step 3: Deduplicate tools after filtering
deduplicated_mcp_tools = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(
filtered_mcp_tools
)
return deduplicated_mcp_tools
@staticmethod
def _transform_mcp_tools_to_openai(mcp_tools: List[Any]) -> List[Any]:
"""Transform MCP tools to OpenAI-compatible format."""
from litellm.experimental_mcp_client.tools import (
transform_mcp_tool_to_openai_responses_api_tool,
)
openai_tools = []
for mcp_tool in mcp_tools:
openai_tool = transform_mcp_tool_to_openai_responses_api_tool(mcp_tool)
openai_tools.append(openai_tool)
return openai_tools
@staticmethod
def _should_auto_execute_tools(
mcp_tools_with_litellm_proxy: Union[List[Dict[str, Any]], List[ToolParam]],
mcp_tools_with_litellm_proxy: Union[List[Dict[str, Any]], List[ToolParam]],
) -> bool:
"""Check if we should auto-execute tool calls.
Only auto-execute tools if user passed a MCP tool with require_approval set to "never".
"""
for tool in mcp_tools_with_litellm_proxy:
if isinstance(tool, dict):
@ -252,24 +276,31 @@ class LiteLLM_Proxy_MCP_Handler:
elif getattr(tool, "require_approval", None) == "never":
return True
return False
@staticmethod
def _extract_tool_calls_from_response(response: ResponsesAPIResponse) -> List[Any]:
"""Extract tool calls from the response output."""
tool_calls: List[Any] = []
for output_item in response.output:
# Check if this is a function call output item
if (isinstance(output_item, dict) and
output_item.get("type") == "function_call"):
if (
isinstance(output_item, dict)
and output_item.get("type") == "function_call"
):
tool_calls.append(output_item)
elif hasattr(output_item, 'type') and getattr(output_item, 'type') == "function_call":
elif (
hasattr(output_item, "type")
and getattr(output_item, "type") == "function_call"
):
# Handle pydantic model case
tool_calls.append(output_item)
return tool_calls
@staticmethod
def _extract_tool_call_details(tool_call) -> Tuple[Optional[str], Optional[str], Optional[str]]:
def _extract_tool_call_details(
tool_call,
) -> Tuple[Optional[str], Optional[str], Optional[str]]:
"""Extract tool name, arguments, and call_id from a tool call."""
if isinstance(tool_call, dict):
tool_name = tool_call.get("name")
@ -278,15 +309,17 @@ class LiteLLM_Proxy_MCP_Handler:
else:
tool_name = getattr(tool_call, "name", None)
tool_arguments = getattr(tool_call, "arguments", None)
tool_call_id = getattr(tool_call, "call_id", None) or getattr(tool_call, "id", None)
tool_call_id = getattr(tool_call, "call_id", None) or getattr(
tool_call, "id", None
)
return tool_name, tool_arguments, tool_call_id
@staticmethod
def _parse_tool_arguments(tool_arguments: Any) -> Dict[str, Any]:
"""Parse tool arguments, handling both string and dict formats."""
import json
if isinstance(tool_arguments, str):
try:
return json.loads(tool_arguments)
@ -294,23 +327,23 @@ class LiteLLM_Proxy_MCP_Handler:
return {}
else:
return tool_arguments or {}
@staticmethod
def _parse_mcp_result(result: Any) -> str:
"""Parse MCP tool call result and extract meaningful content."""
if not result or not hasattr(result, 'content') or not result.content:
if not result or not hasattr(result, "content") or not result.content:
return "Tool executed successfully"
# Import MCP content types for isinstance checks
try:
from mcp.types import EmbeddedResource, ImageContent, TextContent
except ImportError:
# Fallback to generic handling if MCP types not available
return "Tool executed successfully"
text_parts = []
other_content_types = []
for content_item in result.content:
if isinstance(content_item, TextContent):
# Text content - extract the text
@ -325,21 +358,20 @@ class LiteLLM_Proxy_MCP_Handler:
# Other unknown content types
content_type = type(content_item).__name__
other_content_types.append(content_type)
# Combine text parts if any
result_text = " ".join(text_parts) if text_parts else ""
# Add info about other content types
if other_content_types:
other_info = f"[Generated {', '.join(other_content_types)}]"
result_text = f"{result_text} {other_info}".strip()
return result_text or "Tool executed successfully"
@staticmethod
async def _execute_tool_calls(
tool_calls: List[Any],
user_api_key_auth: Any
tool_calls: List[Any], user_api_key_auth: Any
) -> List[Dict[str, Any]]:
"""Execute tool calls and return results."""
from fastapi import HTTPException
@ -348,107 +380,115 @@ class LiteLLM_Proxy_MCP_Handler:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
tool_results = []
tool_call_id: Optional[str] = None
for tool_call in tool_calls:
try:
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 not tool_name:
verbose_logger.warning(f"Tool call missing name: {tool_call}")
continue
parsed_arguments = LiteLLM_Proxy_MCP_Handler._parse_tool_arguments(tool_arguments)
parsed_arguments = LiteLLM_Proxy_MCP_Handler._parse_tool_arguments(
tool_arguments
)
# Import here to avoid circular import
from litellm.proxy.proxy_server import proxy_logging_obj
result = await global_mcp_server_manager.call_tool(
name=tool_name,
arguments=parsed_arguments,
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
)
# Format result for inclusion in response
result_text = LiteLLM_Proxy_MCP_Handler._parse_mcp_result(result)
tool_results.append({
"tool_call_id": tool_call_id,
"result": result_text
})
tool_results.append(
{"tool_call_id": tool_call_id, "result": result_text}
)
except BlockedPiiEntityError as e:
verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}")
verbose_logger.error(
f"BlockedPiiEntityError in MCP tool call: {str(e)}"
)
error_message = f"Tool call blocked: PII entity '{getattr(e, 'entity_type', 'unknown')}' detected by guardrail '{getattr(e, 'guardrail_name', 'unknown')}'. {str(e)}"
tool_results.append({
"tool_call_id": tool_call_id,
"result": error_message
})
tool_results.append(
{"tool_call_id": tool_call_id, "result": error_message}
)
except GuardrailRaisedException as e:
verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}")
verbose_logger.error(
f"GuardrailRaisedException in MCP tool call: {str(e)}"
)
error_message = f"Tool call blocked: Guardrail '{getattr(e, 'guardrail_name', 'unknown')}' violation. {str(e)}"
tool_results.append({
"tool_call_id": tool_call_id,
"result": error_message
})
tool_results.append(
{"tool_call_id": tool_call_id, "result": error_message}
)
except HTTPException as 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_results.append(
{"tool_call_id": tool_call_id, "result": error_message}
)
except Exception as e:
verbose_logger.exception(f"Error executing MCP tool call: {e}")
tool_results.append({
"tool_call_id": tool_call_id,
"result": f"Error executing tool: {str(e)}"
})
tool_results.append(
{
"tool_call_id": tool_call_id,
"result": f"Error executing tool: {str(e)}",
}
)
return tool_results
@staticmethod
def _create_follow_up_input(
response: ResponsesAPIResponse,
tool_results: List[Dict[str, Any]],
original_input: Any = None
response: ResponsesAPIResponse,
tool_results: List[Dict[str, Any]],
original_input: Any = None,
) -> List[Any]:
"""Create follow-up input with tool results in proper format."""
follow_up_input: List[Any] = []
# Add original user input if available to maintain conversation context
if original_input:
if isinstance(original_input, str):
follow_up_input.append({
"type": "message",
"role": "user",
"content": original_input
})
follow_up_input.append(
{"type": "message", "role": "user", "content": original_input}
)
elif isinstance(original_input, list):
follow_up_input.extend(original_input)
else:
follow_up_input.append(original_input)
# Add the assistant message with function calls
assistant_message_content: List[Any] = []
function_calls: List[Dict[str, Any]] = []
for output_item in response.output:
if isinstance(output_item, dict):
if output_item.get("type") == "function_call":
call_id = output_item.get("call_id") or output_item.get("id")
name = output_item.get("name")
arguments = output_item.get("arguments")
# Only add if we have required fields
if call_id and name:
function_calls.append({
"type": "function_call",
"call_id": call_id,
"name": name,
"arguments": arguments
})
function_calls.append(
{
"type": "function_call",
"call_id": call_id,
"name": name,
"arguments": arguments,
}
)
elif output_item.get("type") == "message":
# Extract content from message
content = output_item.get("content", [])
@ -456,36 +496,40 @@ class LiteLLM_Proxy_MCP_Handler:
assistant_message_content.extend(content)
else:
assistant_message_content.append(content)
# Add assistant message with content and function calls
if assistant_message_content or function_calls:
follow_up_input.append({
"type": "message",
"role": "assistant",
"content": assistant_message_content
})
follow_up_input.append(
{
"type": "message",
"role": "assistant",
"content": assistant_message_content,
}
)
# Add function calls after assistant message
for function_call in function_calls:
follow_up_input.append(function_call)
# Add tool results (function call outputs)
for tool_result in tool_results:
follow_up_input.append({
"type": "function_call_output",
"call_id": tool_result["tool_call_id"],
"output": tool_result["result"]
})
follow_up_input.append(
{
"type": "function_call_output",
"call_id": tool_result["tool_call_id"],
"output": tool_result["result"],
}
)
return follow_up_input
@staticmethod
async def _make_follow_up_call(
follow_up_input: List[Any],
model: str,
all_tools: Optional[List[Any]],
response_id: str,
**call_params: Any
**call_params: Any,
) -> Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]:
"""Make follow-up response API call with tool results."""
return await aresponses(
@ -493,7 +537,7 @@ class LiteLLM_Proxy_MCP_Handler:
model=model,
tools=all_tools, # Keep tools for potential future calls
previous_response_id=response_id, # Link to previous response
**call_params
**call_params,
)
@staticmethod
@ -505,11 +549,11 @@ class LiteLLM_Proxy_MCP_Handler:
mcp_discovery_events: List[Any],
call_params: Dict[str, Any],
previous_response_id: Optional[str],
**kwargs
**kwargs,
) -> Any:
"""
Create MCP enhanced streaming response that handles the full MCP workflow.
This creates a streaming iterator that:
1. Immediately emits MCP discovery events
2. Makes the LLM call and streams the response
@ -526,16 +570,16 @@ class LiteLLM_Proxy_MCP_Handler:
all_tools=all_tools,
call_params=call_params,
previous_response_id=previous_response_id,
**kwargs
**kwargs,
)
# Create the enhanced streaming iterator that will handle everything
return MCPEnhancedStreamingIterator(
base_iterator=None, # Will be created internally
mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
user_api_key_auth=kwargs.get("user_api_key_auth"),
original_request_params=request_params
original_request_params=request_params,
)
@staticmethod
@ -545,126 +589,127 @@ class LiteLLM_Proxy_MCP_Handler:
all_tools: Optional[List[Any]],
call_params: Dict[str, Any],
previous_response_id: Optional[str],
**kwargs
**kwargs,
) -> Dict[str, Any]:
"""
Build a clean request parameters dictionary for MCP streaming.
Combines input, model, tools with call_params and additional kwargs
in a clean, maintainable way.
"""
# Start with the core required parameters
request_params = {
'input': input,
'model': model,
'tools': all_tools,
"input": input,
"model": model,
"tools": all_tools,
}
# Add previous_response_id if provided
if previous_response_id is not None:
request_params['previous_response_id'] = previous_response_id
request_params["previous_response_id"] = previous_response_id
# Merge in all call_params (which contains most of the API parameters)
request_params.update(call_params)
# Merge in any additional kwargs
request_params.update(kwargs)
return request_params
@staticmethod
def _create_tool_execution_events(
tool_calls: List[Any],
tool_results: List[Dict[str, Any]]
tool_calls: List[Any], tool_results: List[Dict[str, Any]]
) -> List[Any]:
"""
Create MCP tool execution events for streaming.
Args:
tool_calls: List of tool calls from the LLM response
tool_results: List of tool execution results
Returns:
List of MCP tool execution events for streaming
"""
import uuid
from litellm.responses.mcp.mcp_streaming_iterator import create_mcp_call_events
tool_execution_events: List[Any] = []
# Create events for each tool execution
for tool_result in tool_results:
tool_call_id = tool_result.get("tool_call_id", "unknown")
result_text = tool_result.get("result", "")
# Extract tool name and arguments from tool calls
tool_name = "unknown"
tool_arguments = "{}"
for tool_call in tool_calls:
name, args, call_id = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call)
(
name,
args,
call_id,
) = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call)
if call_id == tool_call_id:
tool_name = name or "unknown"
tool_arguments = args or "{}"
break
execution_events = create_mcp_call_events(
tool_name=tool_name,
tool_call_id=tool_call_id,
arguments=tool_arguments, # Use actual arguments
result=result_text,
base_item_id=f"mcp_{uuid.uuid4().hex[:8]}", # Unique ID for each tool call
sequence_start=len(tool_execution_events) + 1
sequence_start=len(tool_execution_events) + 1,
)
tool_execution_events.extend(execution_events)
return tool_execution_events
@staticmethod
def _prepare_initial_call_params(
call_params: Dict[str, Any],
should_auto_execute: bool
call_params: Dict[str, Any], should_auto_execute: bool
) -> Dict[str, Any]:
"""
Prepare call parameters for the initial LLM call.
For auto-execute scenarios, we need to disable streaming for the initial call
so we can process the tool calls before streaming the final response.
"""
initial_params = call_params.copy()
if should_auto_execute:
# Disable streaming for initial call when auto-executing tools
initial_params["stream"] = False
return initial_params
@staticmethod
def _prepare_follow_up_call_params(
call_params: Dict[str, Any],
original_stream_setting: bool
call_params: Dict[str, Any], original_stream_setting: bool
) -> Dict[str, Any]:
"""
Prepare call parameters for the follow-up LLM call after tool execution.
Restores the original streaming setting and removes tool_choice since
we're now providing tool results, not requesting tool calls.
"""
follow_up_params = call_params.copy()
# Restore original streaming setting for follow-up call
follow_up_params["stream"] = original_stream_setting
# Remove tool_choice since we're providing results, not requesting tool calls
follow_up_params.pop("tool_choice", None)
return follow_up_params
@staticmethod
def _add_mcp_output_elements_to_response(
response: ResponsesAPIResponse,
mcp_tools_fetched: List[Any],
tool_results: List[Dict[str, Any]]
tool_results: List[Dict[str, Any]],
) -> ResponsesAPIResponse:
"""Add custom output elements to the final response for MCP tool execution."""
# Import the required classes for creating output items
@ -683,11 +728,11 @@ class LiteLLM_Proxy_MCP_Handler:
OutputText(
type="output_text",
text=json.dumps(mcp_tools_fetched, indent=2, default=str),
annotations=[]
annotations=[],
)
]
],
)
# Create output element for tool execution results
tool_results_output = GenericResponseOutputItem(
type="tool_execution_results",
@ -698,13 +743,13 @@ class LiteLLM_Proxy_MCP_Handler:
OutputText(
type="output_text",
text=json.dumps(tool_results, indent=2, default=str),
annotations=[]
annotations=[],
)
]
],
)
# Add the new output elements to the response
response.output.append(mcp_tools_output.model_dump()) # type: ignore
response.output.append(tool_results_output.model_dump()) # type: ignore
return response

View file

@ -23,7 +23,7 @@ from litellm.types.llms.openai import (
ResponseText,
)
from litellm.types.responses.main import DecodedResponseId
from litellm.types.utils import SpecialEnums, Usage
from litellm.types.utils import PromptTokensDetails, SpecialEnums, Usage
class ResponsesAPIRequestUtils:
@ -375,8 +375,15 @@ class ResponseAPILoggingUtils:
)
prompt_tokens: int = response_api_usage.input_tokens or 0
completion_tokens: int = response_api_usage.output_tokens or 0
prompt_tokens_details: Optional[PromptTokensDetails] = None
if response_api_usage.input_tokens_details:
prompt_tokens_details = PromptTokensDetails(
cached_tokens=response_api_usage.input_tokens_details.cached_tokens,
audio_tokens=response_api_usage.input_tokens_details.audio_tokens,
)
return Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
)
prompt_tokens_details=prompt_tokens_details,
)

View file

@ -360,9 +360,9 @@ class Router:
) # names of models under litellm_params. ex. azure/chatgpt-v-2
self.deployment_latency_map = {}
### CACHING ###
cache_type: Literal[
"local", "redis", "redis-semantic", "s3", "disk"
] = "local" # default to an in-memory cache
cache_type: Literal["local", "redis", "redis-semantic", "s3", "disk"] = (
"local" # default to an in-memory cache
)
redis_cache = None
cache_config: Dict[str, Any] = {}
@ -404,9 +404,9 @@ class Router:
self.default_max_parallel_requests = default_max_parallel_requests
self.provider_default_deployment_ids: List[str] = []
self.pattern_router = PatternMatchRouter()
self.team_pattern_routers: Dict[
str, PatternMatchRouter
] = {} # {"TEAM_ID": PatternMatchRouter}
self.team_pattern_routers: Dict[str, PatternMatchRouter] = (
{}
) # {"TEAM_ID": PatternMatchRouter}
self.auto_routers: Dict[str, "AutoRouter"] = {}
if model_list is not None:
@ -588,9 +588,9 @@ class Router:
)
)
self.model_group_retry_policy: Optional[
Dict[str, RetryPolicy]
] = model_group_retry_policy
self.model_group_retry_policy: Optional[Dict[str, RetryPolicy]] = (
model_group_retry_policy
)
self.allowed_fails_policy: Optional[AllowedFailsPolicy] = None
if allowed_fails_policy is not None:
@ -1212,7 +1212,10 @@ class Router:
async def _acompletion(
self, model: str, messages: List[Dict[str, str]], **kwargs
) -> Union[ModelResponse, CustomStreamWrapper,]:
) -> Union[
ModelResponse,
CustomStreamWrapper,
]:
"""
- Get an available deployment
- call it with a semaphore over the call
@ -3156,9 +3159,9 @@ class Router:
healthy_deployments=healthy_deployments, responses=responses
)
returned_response = cast(OpenAIFileObject, responses[0])
returned_response._hidden_params[
"model_file_id_mapping"
] = model_file_id_mapping
returned_response._hidden_params["model_file_id_mapping"] = (
model_file_id_mapping
)
return returned_response
except Exception as e:
verbose_router_logger.exception(
@ -3721,11 +3724,11 @@ class Router:
if isinstance(e, litellm.ContextWindowExceededError):
if context_window_fallbacks is not None:
context_window_fallback_model_group: Optional[
List[str]
] = self._get_fallback_model_group_from_fallbacks(
fallbacks=context_window_fallbacks,
model_group=model_group,
context_window_fallback_model_group: Optional[List[str]] = (
self._get_fallback_model_group_from_fallbacks(
fallbacks=context_window_fallbacks,
model_group=model_group,
)
)
if context_window_fallback_model_group is None:
raise original_exception
@ -3757,11 +3760,11 @@ class Router:
e.message += "\n{}".format(error_message)
elif isinstance(e, litellm.ContentPolicyViolationError):
if content_policy_fallbacks is not None:
content_policy_fallback_model_group: Optional[
List[str]
] = self._get_fallback_model_group_from_fallbacks(
fallbacks=content_policy_fallbacks,
model_group=model_group,
content_policy_fallback_model_group: Optional[List[str]] = (
self._get_fallback_model_group_from_fallbacks(
fallbacks=content_policy_fallbacks,
model_group=model_group,
)
)
if content_policy_fallback_model_group is None:
raise original_exception
@ -4415,7 +4418,7 @@ class Router:
return tpm_key
except Exception as e:
verbose_router_logger.exception(
verbose_router_logger.debug(
"litellm.router.Router::deployment_callback_on_success(): Exception occured - {}".format(
str(e)
)
@ -4993,26 +4996,26 @@ class Router:
"""
from litellm.router_strategy.auto_router.auto_router import AutoRouter
auto_router_config_path: Optional[
str
] = deployment.litellm_params.auto_router_config_path
auto_router_config_path: Optional[str] = (
deployment.litellm_params.auto_router_config_path
)
auto_router_config: Optional[str] = deployment.litellm_params.auto_router_config
if auto_router_config_path is None and auto_router_config is None:
raise ValueError(
"auto_router_config_path or auto_router_config is required for auto-router deployments. Please set it in the litellm_params"
)
default_model: Optional[
str
] = deployment.litellm_params.auto_router_default_model
default_model: Optional[str] = (
deployment.litellm_params.auto_router_default_model
)
if default_model is None:
raise ValueError(
"auto_router_default_model is required for auto-router deployments. Please set it in the litellm_params"
)
embedding_model: Optional[
str
] = deployment.litellm_params.auto_router_embedding_model
embedding_model: Optional[str] = (
deployment.litellm_params.auto_router_embedding_model
)
if embedding_model is None:
raise ValueError(
"auto_router_embedding_model is required for auto-router deployments. Please set it in the litellm_params"

View file

@ -28,8 +28,8 @@ class LangsmithInputs(BaseModel):
class LangsmithCredentialsObject(TypedDict):
LANGSMITH_API_KEY: str
LANGSMITH_PROJECT: str
LANGSMITH_API_KEY: Optional[str]
LANGSMITH_PROJECT: Optional[str]
LANGSMITH_BASE_URL: str

View file

@ -328,15 +328,22 @@ class CohereEmbeddingResponse(TypedDict):
texts: List[str]
class AmazonTitanV2EmbeddingRequest(TypedDict):
inputText: str
class AmazonTitanV2EmbeddingRequest(TypedDict, total=False):
inputText: Required[str]
dimensions: int
normalize: bool
embeddingTypes: List[Literal["float", "binary"]]
class AmazonTitanV2EmbeddingResponse(TypedDict):
embedding: List[float]
inputTextTokenCount: int
class AmazonTitanV2EmbeddingsByType(TypedDict, total=False):
binary: List[int] # Array of integers for binary format
float: List[float] # Array of floats for float format
class AmazonTitanV2EmbeddingResponse(TypedDict, total=False):
embedding: List[float] # Legacy field - array of floats (backward compatibility)
embeddingsByType: AmazonTitanV2EmbeddingsByType # New format per AWS schema
inputTextTokenCount: Required[int] # Always present in AWS response
class AmazonTitanG1EmbeddingRequest(TypedDict):

View file

@ -1,9 +1,9 @@
from typing import TYPE_CHECKING, Dict, List, Optional
from typing import Dict, List, Optional
from pydantic import BaseModel, ConfigDict
from typing_extensions import TypedDict
from litellm.proxy._types import MCPAuthType, MCPSpecVersionType, MCPTransportType
from litellm.proxy._types import MCPAuthType, MCPTransportType
from litellm.types.mcp import MCPServerCostInfo
@ -21,7 +21,6 @@ class MCPServer(BaseModel):
server_name: Optional[str] = None
url: Optional[str] = None
transport: MCPTransportType
spec_version: MCPSpecVersionType
auth_type: Optional[MCPAuthType] = None
authentication_token: Optional[str] = None
mcp_info: Optional[MCPInfo] = None

View file

@ -862,6 +862,11 @@ class CompletionTokensDetailsWrapper(
"""Text tokens generated by the model."""
class CacheCreationTokenDetails(BaseModel):
ephemeral_5m_input_tokens: Optional[int] = None
ephemeral_1h_input_tokens: Optional[int] = None
class PromptTokensDetailsWrapper(
PromptTokensDetails
): # wrapper for older openai versions
@ -886,6 +891,9 @@ class PromptTokensDetailsWrapper(
cache_creation_tokens: Optional[int] = None
"""Number of cache creation tokens sent to the model. Used for Anthropic prompt caching."""
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
"""Details of cache creation tokens sent to the model. Used for tracking 5m/1h cache creation tokens for Anthropic prompt caching."""
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
if self.character_count is None:
@ -898,6 +906,8 @@ class PromptTokensDetailsWrapper(
del self.web_search_requests
if self.cache_creation_tokens is None:
del self.cache_creation_tokens
if self.cache_creation_token_details is None:
del self.cache_creation_token_details
class ServerToolUse(BaseModel):
@ -2015,6 +2025,7 @@ class GuardrailMode(TypedDict, total=False):
class StandardLoggingGuardrailInformation(TypedDict, total=False):
guardrail_name: Optional[str]
guardrail_provider: Optional[str]
guardrail_mode: Optional[
Union[GuardrailEventHooks, List[GuardrailEventHooks], GuardrailMode]
]
@ -2130,6 +2141,7 @@ class StandardCallbackDynamicParams(TypedDict, total=False):
langsmith_api_key: Optional[str]
langsmith_project: Optional[str]
langsmith_base_url: Optional[str]
langsmith_sampling_rate: Optional[float]
# Humanloop dynamic params
humanloop_api_key: Optional[str]

View file

@ -59,6 +59,12 @@ import litellm.litellm_core_utils.audio_utils.utils
import litellm.litellm_core_utils.json_validation_rule
import litellm.llms
import litellm.llms.gemini
# Import cached imports utilities
from litellm.litellm_core_utils.cached_imports import (
get_coroutine_checker,
get_litellm_logging_class,
get_set_callbacks,
)
from litellm.caching._internal_lru_cache import lru_cache_wrapper
from litellm.caching.caching import DualCache
from litellm.caching.caching_handler import CachingHandlerResponse, LLMCachingHandler
@ -222,6 +228,7 @@ from typing import (
get_args,
)
from openai import OpenAIError as OriginalError
from litellm.litellm_core_utils.thread_pool_executor import executor
@ -521,16 +528,12 @@ def get_dynamic_callbacks(
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
def function_setup( # noqa: PLR0915
original_function: str, rules_obj, start_time, *args, **kwargs
): # just run once to check if user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
### NOTICES ###
from litellm import Logging as LiteLLMLogging
from litellm.litellm_core_utils.litellm_logging import set_callbacks
if litellm.set_verbose is True:
verbose_logger.warning(
"`litellm.set_verbose` is deprecated. Please set `os.environ['LITELLM_LOG'] = 'DEBUG'` for debug logs."
@ -593,12 +596,12 @@ def function_setup( # noqa: PLR0915
+ litellm.failure_callback
)
)
set_callbacks(callback_list=callback_list, function_id=function_id)
get_set_callbacks()(callback_list=callback_list, function_id=function_id)
## ASYNC CALLBACKS
if len(litellm.input_callback) > 0:
removed_async_items = []
for index, callback in enumerate(litellm.input_callback): # type: ignore
if coroutine_checker.is_async_callable(callback):
if get_coroutine_checker().is_async_callable(callback):
litellm._async_input_callback.append(callback)
removed_async_items.append(index)
@ -608,7 +611,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 coroutine_checker.is_async_callable(callback):
if get_coroutine_checker().is_async_callable(callback):
litellm.logging_callback_manager.add_litellm_async_success_callback(
callback
)
@ -633,7 +636,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 coroutine_checker.is_async_callable(callback):
if get_coroutine_checker().is_async_callable(callback):
litellm.logging_callback_manager.add_litellm_async_failure_callback(
callback
)
@ -666,7 +669,7 @@ def function_setup( # noqa: PLR0915
removed_async_items = []
for index, callback in enumerate(kwargs["success_callback"]):
if (
coroutine_checker.is_async_callable(callback)
get_coroutine_checker().is_async_callable(callback)
or callback == "dynamodb"
or callback == "s3"
):
@ -790,7 +793,7 @@ def function_setup( # noqa: PLR0915
call_type=call_type,
):
stream = True
logging_obj = LiteLLMLogging(
logging_obj = get_litellm_logging_class()( # Victim for object pool
model=model, # type: ignore
messages=messages,
stream=stream,
@ -903,7 +906,7 @@ def client(original_function): # noqa: PLR0915
rules_obj = Rules()
def check_coroutine(value) -> bool:
return coroutine_checker.is_async_callable(value)
return get_coroutine_checker().is_async_callable(value)
async def async_pre_call_deployment_hook(kwargs: Dict[str, Any], call_type: str):
"""
@ -1597,7 +1600,7 @@ def client(original_function): # noqa: PLR0915
setattr(e, "timeout", timeout)
raise e
is_coroutine = coroutine_checker.is_async_callable(original_function)
is_coroutine = get_coroutine_checker().is_async_callable(original_function)
# Return the appropriate wrapper based on the original function type
if is_coroutine:
@ -4877,6 +4880,9 @@ def _get_model_info_helper( # noqa: PLR0915
cache_read_input_token_cost=_model_info.get(
"cache_read_input_token_cost", None
),
cache_creation_input_token_cost_above_1hr=_model_info.get(
"cache_creation_input_token_cost_above_1hr", None
),
input_cost_per_character=_model_info.get(
"input_cost_per_character", None
),

View file

@ -305,8 +305,56 @@
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true,
"supports_multimodal_embedding": true
"supports_image_input": true
},
"us.twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"max_tokens": 77,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"eu.twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"max_tokens": 77,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"us.twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"eu.twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"amazon.titan-text-express-v1": {
"input_cost_per_token": 1.3e-06,
@ -9078,7 +9126,7 @@
"max_tokens": 65535,
"max_video_length": 1,
"max_videos_per_prompt": 10,
"mode": "chat",
"mode": "image_generation",
"output_cost_per_image": 0.039,
"output_cost_per_reasoning_token": 3e-05,
"output_cost_per_token": 3e-05,
@ -10441,7 +10489,7 @@
"max_tokens": 65535,
"max_video_length": 1,
"max_videos_per_prompt": 10,
"mode": "chat",
"mode": "image_generation",
"output_cost_per_image": 0.039,
"output_cost_per_reasoning_token": 3e-05,
"output_cost_per_token": 3e-05,

1683
poetry.lock generated

File diff suppressed because it is too large Load diff

View file

@ -32,6 +32,7 @@ jinja2 = "^3.1.2"
aiohttp = ">=3.10"
pydantic = "^2.5.0"
jsonschema = "^4.22.0"
pondpond = "^1.4.1"
numpydoc = {version = "*", optional = true} # used in utils.py
uvicorn = {version = "^0.29.0", optional = true}
@ -61,7 +62,7 @@ redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.
mcp = {version = "^1.10.0", optional = true, python = ">=3.10"}
litellm-proxy-extras = {version = "0.2.18", optional = true}
rich = {version = "13.7.1", optional = true}
litellm-enterprise = {version = "0.1.19", optional = true}
litellm-enterprise = {version = "0.1.20", optional = true}
diskcache = {version = "^5.6.1", optional = true}
polars = {version = "^1.31.0", optional = true, python = ">=3.10"}
semantic-router = {version = "*", optional = true, python = ">=3.9"}

View file

@ -58,8 +58,9 @@ tenacity==8.2.3 # for retrying requests, when litellm.num_retries set
pydantic==2.10.2 # proxy + openai req.
jsonschema==4.22.0 # validating json schema
websockets==13.1.0 # for realtime API
pondpond==1.4.1 # for object pooling
########################
# LITELLM ENTERPRISE DEPENDENCIES
########################
litellm-enterprise==0.1.19
litellm-enterprise==0.1.20

View file

@ -171,7 +171,6 @@ model LiteLLM_MCPServerTable {
description String?
url String?
transport String @default("sse")
spec_version String @default("2025-03-26")
auth_type String?
created_at DateTime? @default(now()) @map("created_at")
created_by String?

View file

@ -134,4 +134,5 @@ jsonschema: >=4.22.0 # Unknown license
websockets: >=13.1.0 # Unknown license
polars: >=1.31.0 # Unknown license, the license.md allows free of charge use
semantic_router: >=0.1.10 # Unknown license
pondpond: >=1.4.1 # Apache 2.0 License

View file

@ -0,0 +1,43 @@
## Tests that chat_completion endpoint has no imports inside function bodies
## This is critical for performance optimization in the hot path
import ast
from pathlib import Path
def test_chat_completion_no_imports():
"""Test that chat_completion endpoint has no imports in function bodies."""
# Path to the proxy server file
proxy_server_path = Path(__file__).parent.parent.parent / "litellm" / "proxy" / "proxy_server.py"
with open(proxy_server_path, 'r') as f:
content = f.read()
# Parse the AST
tree = ast.parse(content)
# Find the chat_completion function
chat_completion_func = None
for node in ast.walk(tree):
if (isinstance(node, ast.AsyncFunctionDef) and node.name == "chat_completion"):
chat_completion_func = node
break
assert chat_completion_func is not None, "chat_completion function not found"
# Check for imports inside the function body
import_violations = []
for node in ast.walk(chat_completion_func):
if isinstance(node, (ast.Import, ast.ImportFrom)):
# Get line number
line_num = node.lineno
import_violations.append(line_num)
# Assert no import violations found
if import_violations:
print(f"Found {len(import_violations)} import violations in chat_completion endpoint:")
for line_num in import_violations:
print(f" - Line {line_num}: Import statement found")
print("\nchat_completion endpoint should not contain imports for optimal performance.")
raise Exception("Import violations found in chat_completion endpoint")

View file

@ -0,0 +1,106 @@
"""
Simplified tests for object pooling utilities in litellm.
"""
import pytest
from litellm.litellm_core_utils.object_pooling import (
get_object_pool,
_pools
)
class SimpleObject:
"""Simple test object with internal reset tracking."""
def __init__(self):
self.data = {}
self.reset_count = 0
self.creation_id = id(self)
def reset(self):
"""Reset method that tracks how many times it's called."""
self.data.clear()
self.reset_count += 1
def set_data(self, key, value):
"""Set data to verify reset works."""
self.data[key] = value
class SimpleObjectNoReset:
"""Test object without reset method."""
def __init__(self):
self.data = {}
self.creation_id = id(self)
class TestObjectPooling:
"""Simplified test suite for object pooling."""
def setup_method(self):
"""Clear pools before each test."""
_pools.clear()
def test_reset_method_works(self):
"""Test that reset method is called when recycling objects."""
pool_name = "reset_test"
pool = get_object_pool(pool_name, SimpleObject, pooled_maxsize=1, prewarm_count=0)
# Get an object and modify it
obj = pool.borrow(name=f"{pool_name}Factory")
obj.keeped_object.set_data("test", "value")
initial_reset_count = obj.keeped_object.reset_count
# Return to pool (should trigger reset)
pool.recycle(obj, name=f"{pool_name}Factory")
# Get the same object back
obj2 = pool.borrow(name=f"{pool_name}Factory")
# Verify reset was called
assert obj2.keeped_object.reset_count == initial_reset_count + 1
assert obj2.keeped_object.data == {} # Data should be cleared
assert obj.keeped_object.creation_id == obj2.keeped_object.creation_id # Same object
def test_fallback_reset_works(self):
"""Test fallback reset when no reset method exists."""
pool_name = "fallback_test"
pool = get_object_pool(pool_name, SimpleObjectNoReset, pooled_maxsize=1, prewarm_count=0)
# Get an object and modify it
obj = pool.borrow(name=f"{pool_name}Factory")
obj.keeped_object.data["test"] = "value"
# Return to pool (should trigger fallback reset)
pool.recycle(obj, name=f"{pool_name}Factory")
# Get the same object back
obj2 = pool.borrow(name=f"{pool_name}Factory")
# Verify fallback reset worked - all attributes should be cleared by __dict__.clear()
assert obj2.keeped_object.__dict__ == {}, "All attributes should be cleared by fallback reset"
assert obj is obj2, "Should be the same pooled object instance"
def test_pool_reuses_objects(self):
"""Test that pool actually reuses objects instead of creating new ones."""
pool_name = "reuse_test"
pool = get_object_pool(pool_name, SimpleObject, pooled_maxsize=1, prewarm_count=0)
# Get first object
obj1 = pool.borrow(name=f"{pool_name}Factory")
creation_id1 = obj1.keeped_object.creation_id
# Return it
pool.recycle(obj1, name=f"{pool_name}Factory")
# Get second object
obj2 = pool.borrow(name=f"{pool_name}Factory")
creation_id2 = obj2.keeped_object.creation_id
# Should be the same object (reused)
assert creation_id1 == creation_id2, "Pool should reuse objects"
assert obj1.keeped_object is obj2.keeped_object, "Should be same object instance"
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -76,3 +76,90 @@ def test_bedrock_embedding_models(model, input_type, embed_response):
except Exception as e:
pytest.fail(f"Error occurred: {e}")
def test_e2e_bedrock_embedding():
"""
Test text embedding with TwelveLabs Marengo.
Validates that the transformation properly extracts embedding data from TwelveLabs response format.
"""
print("Testing text embedding...")
litellm._turn_on_debug()
response = litellm.embedding(
model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0",
input=["Hello world from LiteLLM with TwelveLabs Marengo!"],
aws_region_name="us-east-1"
)
# Validate response structure
assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type"
assert hasattr(response, 'data'), "Response should have 'data' attribute"
assert len(response.data) > 0, "Response data should not be empty"
# Validate first embedding
embedding_obj = response.data[0]
assert hasattr(embedding_obj, 'embedding'), "Embedding object should have 'embedding' attribute"
assert isinstance(embedding_obj.embedding, list), "Embedding should be a list of floats"
assert len(embedding_obj.embedding) > 0, "Embedding vector should not be empty"
assert all(isinstance(x, (int, float)) for x in embedding_obj.embedding), "All embedding values should be numeric"
# Validate embedding properties
assert embedding_obj.index == 0, "First embedding should have index 0"
assert embedding_obj.object == "embedding", "Embedding object type should be 'embedding'"
# Validate usage information
assert hasattr(response, 'usage'), "Response should have usage information"
assert response.usage is not None, "Usage should not be None"
assert response.usage.total_tokens >= 0, "Total tokens should be non-negative"
print(f"Text embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}")
def test_e2e_bedrock_embedding_image_twelvelabs_marengo():
"""
Test image embedding with TwelveLabs Marengo.
Validates that the transformation properly extracts embedding data from TwelveLabs response format for images.
"""
print("Testing image embedding...")
litellm._turn_on_debug()
# Load duck.png and convert to base64
duck_img_path = os.path.join(os.path.dirname(__file__), "duck.png")
with open(duck_img_path, "rb") as img_file:
duck_img_data = base64.b64encode(img_file.read()).decode('utf-8')
duck_img_base64 = f"data:image/png;base64,{duck_img_data}"
response = litellm.embedding(
model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0",
input=[duck_img_base64],
aws_region_name="us-east-1"
)
# Validate response structure
assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type"
assert hasattr(response, 'data'), "Response should have 'data' attribute"
assert len(response.data) > 0, "Response data should not be empty"
# Validate first embedding
embedding_obj = response.data[0]
assert hasattr(embedding_obj, 'embedding'), "Embedding object should have 'embedding' attribute"
assert isinstance(embedding_obj.embedding, list), "Embedding should be a list of floats"
assert len(embedding_obj.embedding) > 0, "Embedding vector should not be empty"
assert all(isinstance(x, (int, float)) for x in embedding_obj.embedding), "All embedding values should be numeric"
# Validate embedding properties
assert embedding_obj.index == 0, "First embedding should have index 0"
assert embedding_obj.object == "embedding", "Embedding object type should be 'embedding'"
# Validate usage information
assert hasattr(response, 'usage'), "Response should have usage information"
assert response.usage is not None, "Usage should not be None"
assert response.usage.total_tokens >= 0, "Total tokens should be non-negative"
# TwelveLabs Marengo should return 1024-dimensional embeddings
expected_dimension = 1024
assert len(embedding_obj.embedding) == expected_dimension, f"TwelveLabs Marengo should return {expected_dimension}-dimensional embeddings, got {len(embedding_obj.embedding)}"
print(f"Image embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}")

View file

@ -272,6 +272,119 @@ def test_gemini_image_generation():
assert response.choices[0].message.images[0]["image_url"]["url"].startswith("data:image/png;base64,")
def test_gemini_2_5_flash_image_preview():
"""
Test for GitHub issue #14120 - gemini-2.5-flash-image-preview model routing fix
Validates that the model correctly routes to image generation instead of chat completion
"""
from unittest.mock import patch, MagicMock
from litellm.types.utils import ImageResponse, ImageObject
# Mock successful response to avoid API limits
mock_response = ImageResponse()
mock_response.data = [ImageObject(b64_json="test_base64_data", url=None)]
with patch("litellm.llms.custom_httpx.llm_http_handler.HTTPHandler.post") as mock_post:
# Mock successful HTTP response
mock_http_response = MagicMock()
mock_http_response.json.return_value = {
"candidates": [
{
"content": {
"parts": [
{
"inlineData": {
"data": "test_base64_image_data"
}
}
]
}
}
]
}
mock_http_response.status_code = 200
mock_post.return_value = mock_http_response
# Test that the function works without throwing the original 400 error
response = litellm.image_generation(
model="gemini/gemini-2.5-flash-image-preview",
prompt="Generate a simple test image",
api_key="test_api_key"
)
# Validate response structure
assert response is not None
assert hasattr(response, 'data')
assert response.data is not None
assert len(response.data) > 0
# Validate the correct endpoint was called
mock_post.assert_called_once()
call_args = mock_post.call_args
called_url = call_args[0][0] if call_args[0] else call_args.kwargs.get('url', '')
# Verify it uses generateContent endpoint for gemini-2.5-flash-image-preview (not predict)
assert ":generateContent" in called_url
assert "gemini-2.5-flash-image-preview" in called_url
# Verify request format is Gemini format (not Imagen)
request_data = call_args.kwargs.get('json', {})
assert "contents" in request_data
assert "parts" in request_data["contents"][0]
# Verify response_modalities is set correctly for image generation
assert "generationConfig" in request_data
assert "response_modalities" in request_data["generationConfig"]
assert request_data["generationConfig"]["response_modalities"] == ["IMAGE", "TEXT"]
def test_gemini_imagen_models_use_predict_endpoint():
"""
Test that Imagen models still use :predict endpoint (not broken by gemini-2.5-flash-image-preview fix)
"""
from unittest.mock import patch, MagicMock
from litellm.types.utils import ImageResponse, ImageObject
with patch("litellm.llms.custom_httpx.llm_http_handler.HTTPHandler.post") as mock_post:
# Mock successful HTTP response for Imagen
mock_http_response = MagicMock()
mock_http_response.json.return_value = {
"predictions": [
{
"bytesBase64Encoded": "test_base64_image_data"
}
]
}
mock_http_response.status_code = 200
mock_post.return_value = mock_http_response
# Test an Imagen model
response = litellm.image_generation(
model="gemini/imagen-3.0-generate-001",
prompt="Generate a simple test image",
api_key="test_api_key"
)
# Validate response structure
assert response is not None
assert hasattr(response, 'data')
# Validate the correct endpoint was called for Imagen models
mock_post.assert_called_once()
call_args = mock_post.call_args
called_url = call_args[0][0] if call_args[0] else call_args.kwargs.get('url', '')
# Verify Imagen models use predict endpoint (not generateContent)
assert ":predict" in called_url
assert "imagen-3.0-generate-001" in called_url
assert ":generateContent" not in called_url
# Verify request format is Imagen format (not Gemini)
request_data = call_args.kwargs.get('json', {})
assert "instances" in request_data
assert "parameters" in request_data
def test_gemini_thinking():
litellm._turn_on_debug()
from litellm.types.utils import Message, CallTypes

View file

@ -6,6 +6,7 @@ import pytest
# sys.path.insert(
# 0, os.path.abspath("../..")
# ) # noqa
# ) # Adds the parent directory to the system path
from base_llm_unit_tests import BaseLLMChatTest

View file

@ -320,7 +320,7 @@ async def test_anthropic_api_prompt_caching_basic_with_cache_creation():
random_id
)
* 400,
"cache_control": {"type": "ephemeral"},
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
],
},
@ -331,7 +331,7 @@ async def test_anthropic_api_prompt_caching_basic_with_cache_creation():
{
"type": "text",
"text": "What are the key terms and conditions in this agreement?",
"cache_control": {"type": "ephemeral"},
"cache_control": {"type": "ephemeral", "ttl": "5m"},
}
],
},
@ -580,7 +580,6 @@ async def test_anthropic_api_prompt_caching_streaming():
if hasattr(chunk, "usage") and hasattr(
chunk.usage, "cache_creation_input_tokens"
):
print("chunk.usage", chunk.usage)
is_cache_creation_input_tokens_in_usage = True
idx += 1

View file

@ -68,4 +68,4 @@ def test_bedrock_embed_v2_with_drop_params():
custom_llm_provider=custom_llm_provider,
)
print(f"received optional_params: {optional_params}")
assert optional_params == {"dimensions": 512}
assert optional_params == {"dimensions": 512, "embeddingTypes": ["binary"]}

View file

@ -23,22 +23,20 @@ def test_mcp_server_works_without_config_auth_value():
name="Test MCP Server No Config Auth",
server_name="test_server_no_config",
alias="test_no_config",
url="https://api.example.com/mcp",
url="https://api.example.com/mcp",
transport=MCPTransport.http,
spec_version=MCPSpecVersion.jun_2025,
auth_type=MCPAuth.authorization,
authentication_token=None # No config auth
authentication_token=None, # No config auth
)
manager = MCPServerManager()
# Test that it works with only header auth
client = manager._create_mcp_client(
server=server_without_config_auth,
mcp_auth_header="Bearer token_from_header_only",
protocol_version="2025-06-18"
)
# Verify header token is used
assert client._mcp_auth_value == "Bearer token_from_header_only"
assert client.auth_type == MCPAuth.authorization

View file

@ -17,7 +17,7 @@ from mcp.types import Tool as MCPTool, CallToolResult as MCPCallToolResult
class TestMCPClientUnitTests:
"""Unit tests for MCPClient functionality."""
def test_init_with_auth(self):
"""Test initialization with authentication."""
client = MCPClient(
@ -25,43 +25,48 @@ class TestMCPClientUnitTests:
transport_type=MCPTransport.sse,
auth_type=MCPAuth.bearer_token,
auth_value="test_token",
timeout=30.0
timeout=30.0,
)
assert client.server_url == "http://example.com"
assert client.transport_type == MCPTransport.sse
assert client.auth_type == MCPAuth.bearer_token
assert client.timeout == 30.0
assert client._mcp_auth_value == "test_token"
def test_get_auth_headers(self):
"""Test authentication header generation for different auth types."""
# Bearer token
client = MCPClient(
"http://example.com",
"http://example.com",
auth_type=MCPAuth.bearer_token,
auth_value="test_token"
auth_value="test_token",
)
headers = client._get_auth_headers()
assert headers == {"Authorization": "Bearer test_token", "MCP-Protocol-Version": "2025-06-18"}
assert headers == {
"Authorization": "Bearer test_token",
"MCP-Protocol-Version": "2025-06-18",
}
# Basic auth
client = MCPClient(
"http://example.com",
auth_type=MCPAuth.basic,
auth_value="user:pass"
"http://example.com", auth_type=MCPAuth.basic, auth_value="user:pass"
)
expected_encoded = base64.b64encode("user:pass".encode("utf-8")).decode()
headers = client._get_auth_headers()
assert headers == {"Authorization": f"Basic {expected_encoded}", "MCP-Protocol-Version": "2025-06-18"}
assert headers == {
"Authorization": f"Basic {expected_encoded}",
"MCP-Protocol-Version": "2025-06-18",
}
# API key
client = MCPClient(
"http://example.com",
auth_type=MCPAuth.api_key,
auth_value="api_key_123"
"http://example.com", auth_type=MCPAuth.api_key, auth_value="api_key_123"
)
headers = client._get_auth_headers()
assert headers == {"X-API-Key": "api_key_123", "MCP-Protocol-Version": "2025-06-18"}
assert headers == {
"X-API-Key": "api_key_123",
"MCP-Protocol-Version": "2025-06-18",
}
# Custom authorization header
client = MCPClient(
@ -74,24 +79,15 @@ class TestMCPClientUnitTests:
"Authorization": "Token custom_token",
"MCP-Protocol-Version": "2025-06-18",
}
# No auth
client = MCPClient("http://example.com")
headers = client._get_auth_headers()
assert headers == {"MCP-Protocol-Version": "2025-06-18"}
# Custom protocol version
from litellm.types.mcp import MCPSpecVersion
client = MCPClient(
"http://example.com",
protocol_version=MCPSpecVersion.mar_2025
)
headers = client._get_auth_headers()
assert headers == {"MCP-Protocol-Version": "2025-03-26"}
@pytest.mark.asyncio
@patch('litellm.experimental_mcp_client.client.streamablehttp_client')
@patch('litellm.experimental_mcp_client.client.ClientSession')
@patch("litellm.experimental_mcp_client.client.streamablehttp_client")
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_connect(self, mock_session_class, mock_transport):
"""Test connecting to MCP server with authentication."""
# Setup mocks
@ -99,30 +95,33 @@ class TestMCPClientUnitTests:
mock_transport.return_value = mock_transport_ctx
mock_transport_instance = MagicMock()
mock_transport_ctx.__aenter__ = AsyncMock(return_value=mock_transport_instance)
mock_session_ctx = AsyncMock()
mock_session_class.return_value = mock_session_ctx
mock_session_instance = AsyncMock()
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
client = MCPClient(
"http://example.com",
auth_type=MCPAuth.bearer_token,
auth_value="test_token"
auth_value="test_token",
)
await client.connect()
# Verify transport was created with auth headers
call_args = mock_transport.call_args
assert call_args[1]['headers'] == {"Authorization": "Bearer test_token", "MCP-Protocol-Version": "2025-06-18"}
assert call_args[1]["headers"] == {
"Authorization": "Bearer test_token",
"MCP-Protocol-Version": "2025-06-18",
}
# Verify session was initialized
mock_session_instance.initialize.assert_called_once()
assert client._session == mock_session_instance
@pytest.mark.asyncio
@patch('litellm.experimental_mcp_client.client.streamablehttp_client')
@patch('litellm.experimental_mcp_client.client.ClientSession')
@patch("litellm.experimental_mcp_client.client.streamablehttp_client")
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_list_tools(self, mock_session_class, mock_transport):
"""Test listing tools from the server."""
# Setup mocks
@ -130,70 +129,71 @@ class TestMCPClientUnitTests:
mock_transport.return_value = mock_transport_ctx
mock_transport_instance = MagicMock()
mock_transport_ctx.__aenter__ = AsyncMock(return_value=mock_transport_instance)
mock_session_ctx = AsyncMock()
mock_session_class.return_value = mock_session_ctx
mock_session_instance = AsyncMock()
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
mock_tools = [
MCPTool(
name="test_tool",
name="test_tool",
description="Test tool",
inputSchema={
"type": "object",
"properties": {"arg1": {"type": "string"}},
"required": ["arg1"]
}
"required": ["arg1"],
},
)
]
mock_result = MagicMock()
mock_result.tools = mock_tools
mock_session_instance.list_tools.return_value = mock_result
client = MCPClient("http://example.com")
result = await client.list_tools()
assert result == mock_tools
mock_session_instance.initialize.assert_called_once()
mock_session_instance.list_tools.assert_called_once()
@pytest.mark.asyncio
@patch('litellm.experimental_mcp_client.client.streamablehttp_client')
@patch('litellm.experimental_mcp_client.client.ClientSession')
@patch("litellm.experimental_mcp_client.client.streamablehttp_client")
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_call_tool(self, mock_session_class, mock_transport):
"""Test calling a tool."""
from mcp.types import CallToolRequestParams
# Setup mocks
mock_transport_ctx = AsyncMock()
mock_transport.return_value = mock_transport_ctx
mock_transport_instance = MagicMock()
mock_transport_ctx.__aenter__ = AsyncMock(return_value=mock_transport_instance)
mock_session_ctx = AsyncMock()
mock_session_class.return_value = mock_session_ctx
mock_session_instance = AsyncMock()
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
mock_result = MCPCallToolResult(content=[])
mock_session_instance.call_tool.return_value = mock_result
client = MCPClient("http://example.com")
params = CallToolRequestParams(name="test_tool", arguments={"arg1": "value1"})
result = await client.call_tool(params)
assert result == mock_result
mock_session_instance.initialize.assert_called_once()
mock_session_instance.call_tool.assert_called_once_with(
name="test_tool",
arguments={"arg1": "value1"}
name="test_tool", arguments={"arg1": "value1"}
)
def test_protocol_version_header_extraction(self):
"""Test that MCP protocol version header is correctly extracted from requests."""
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
# Mock scope with headers
mock_scope = {
"type": "http",
@ -203,26 +203,37 @@ class TestMCPClientUnitTests:
(b"authorization", b"Bearer test_token"),
(b"mcp-protocol-version", b"2025-06-18"),
(b"content-type", b"application/json"),
]
],
}
# Mock the user_api_key_auth function
with patch('litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth') as mock_auth:
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth"
) as mock_auth:
mock_auth.return_value = MagicMock()
# Call process_mcp_request
import asyncio
result = asyncio.run(MCPRequestHandler.process_mcp_request(mock_scope))
# Verify the protocol version is extracted
user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, mcp_protocol_version = result
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
mcp_protocol_version,
) = result
assert mcp_protocol_version == "2025-06-18"
def test_protocol_version_header_missing(self):
"""Test that MCP protocol version header is None when not provided."""
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
# Mock scope without protocol version header
mock_scope = {
"type": "http",
@ -231,22 +242,31 @@ class TestMCPClientUnitTests:
"headers": [
(b"authorization", b"Bearer test_token"),
(b"content-type", b"application/json"),
]
],
}
# Mock the user_api_key_auth function
with patch('litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth') as mock_auth:
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth"
) as mock_auth:
mock_auth.return_value = MagicMock()
# Call process_mcp_request
import asyncio
result = asyncio.run(MCPRequestHandler.process_mcp_request(mock_scope))
# Verify the protocol version is None
user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, mcp_protocol_version = result
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
mcp_protocol_version,
) = result
assert mcp_protocol_version is None
if __name__ == "__main__":
pytest.main([__file__, "-v"])
pytest.main([__file__, "-v"])

File diff suppressed because it is too large Load diff

View file

@ -109,12 +109,16 @@ async def test_proxy_failure_metrics():
expected_metric_pattern = 'litellm_proxy_failed_requests_metric_total{api_key_alias="None",end_user="None",exception_class="Openai.RateLimitError",exception_status="429",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",requested_model="fake-azure-endpoint",route="/chat/completions",team="None",team_alias="None",user="default_user_id",user_email="None"}'
# Check if the pattern is in metrics (this metric doesn't include user_email field)
assert any(expected_metric_pattern in line for line in metrics.split('\n')), f"Expected failure metric pattern not found in /metrics. Pattern: {expected_metric_pattern}"
assert any(
expected_metric_pattern in line for line in metrics.split("\n")
), f"Expected failure metric pattern not found in /metrics. Pattern: {expected_metric_pattern}"
# Check total requests metric which includes user_email
total_requests_pattern = 'litellm_proxy_total_requests_metric_total{api_key_alias="None",end_user="None",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",requested_model="fake-azure-endpoint",route="/chat/completions",status_code="429",team="None",team_alias="None",user="default_user_id",user_email="None"}'
assert any(total_requests_pattern in line for line in metrics.split('\n')), f"Expected total requests metric pattern not found in /metrics. Pattern: {total_requests_pattern}"
assert any(
total_requests_pattern in line for line in metrics.split("\n")
), f"Expected total requests metric pattern not found in /metrics. Pattern: {total_requests_pattern}"
@pytest.mark.asyncio
@ -252,8 +256,8 @@ async def create_test_team(
async def create_test_user(
session: aiohttp.ClientSession, user_data: Dict[str, Any]
) -> str:
"""Create a new user and return the user_id"""
) -> Dict[str, Any]:
"""Create a new user and return the user info"""
url = "http://0.0.0.0:4000/user/new"
headers = {
"Authorization": "Bearer sk-1234",
@ -378,7 +382,9 @@ async def test_team_budget_metrics():
assert first_budget["total"] == 10.0, "Total budget metric is incorrect"
print("first_budget['remaining_hours']", first_budget["remaining_hours"])
# Budget should have positive remaining hours, up to 7 days
assert 0 < first_budget["remaining_hours"] <= 168, "Budget should have positive remaining hours, up to 7 days"
assert (
0 < first_budget["remaining_hours"] <= 168
), "Budget should have positive remaining hours, up to 7 days"
# Get team info and verify spend matches prometheus metrics
team_info = await get_team_info(session, team_id)
@ -510,7 +516,9 @@ async def test_key_budget_metrics():
print("first_budget['remaining_hours']", first_budget["remaining_hours"])
# The budget reset time is now standardized - for "7d" it resets on Monday at midnight
# So we'll check if it's within a reasonable range (0-7 days depending on current day of week)
assert 0 <= first_budget["remaining_hours"] <= 168, "Budget remaining hours should be within a reasonable range (0-7 days depending on day of week)"
assert (
0 <= first_budget["remaining_hours"] <= 168
), "Budget remaining hours should be within a reasonable range (0-7 days depending on day of week)"
# Get key info and verify spend matches prometheus metrics
key_info = await get_key_info(session, key)
@ -570,6 +578,7 @@ async def test_user_email_metrics():
user_email in metrics_after_first
), "user_email should be tracked correctly"
@pytest.mark.asyncio
async def test_user_email_in_all_required_metrics():
"""
@ -607,20 +616,25 @@ async def test_user_email_in_all_required_metrics():
# Check that user_email appears in all the required metrics
required_metrics_with_user_email = [
"litellm_proxy_total_requests_metric_total",
"litellm_input_tokens_metric_total",
"litellm_output_tokens_metric_total",
"litellm_requests_metric_total",
"litellm_spend_metric_total"
# "litellm_proxy_total_requests_metric_total",
# "litellm_input_tokens_metric_total",
# "litellm_output_tokens_metric_total",
# "litellm_requests_metric_total",
"litellm_spend_metric_total",
]
import re
for metric_name in required_metrics_with_user_email:
# Check that the metric exists and contains user_email label
import re
# Look for the metric with user_email in its labels
pattern = rf'{metric_name}{{[^}}]*user_email="{re.escape(user_email)}"[^}}]*}}'
pattern = (
rf'{metric_name}{{[^}}]*user_email="{re.escape(user_email)}"[^}}]*}}'
)
matches = re.findall(pattern, metrics_text)
assert len(matches) > 0, f"Metric {metric_name} should contain user_email={user_email} but was not found in metrics"
assert (
len(matches) > 0
), f"Metric {metric_name} should contain user_email={user_email} but was not found in metrics"
# Also test failure metric by making a bad request
try:
@ -639,4 +653,6 @@ async def test_user_email_in_all_required_metrics():
# Check that failure metric also contains user_email
failure_pattern = rf'litellm_proxy_failed_requests_metric_total{{[^}}]*user_email="{re.escape(user_email)}"[^}}]*}}'
failure_matches = re.findall(failure_pattern, metrics_text)
assert len(failure_matches) > 0, f"litellm_proxy_failed_requests_metric_total should contain user_email={user_email}"
assert (
len(failure_matches) > 0
), f"litellm_proxy_failed_requests_metric_total should contain user_email={user_email}"

View file

@ -3428,6 +3428,16 @@ async def test_list_keys(prisma_client):
),
page=1,
size=10,
user_id=None,
team_id=None,
organization_id=None,
key_hash=None,
key_alias=None,
return_full_object=False,
include_team_keys=False,
include_created_by_keys=False,
sort_by=None,
sort_order="desc",
)
print("response=", response)
assert "keys" in response
@ -3442,6 +3452,16 @@ async def test_list_keys(prisma_client):
UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value),
page=1,
size=2,
user_id=None,
team_id=None,
organization_id=None,
key_hash=None,
key_alias=None,
return_full_object=False,
include_team_keys=False,
include_created_by_keys=False,
sort_by=None,
sort_order="desc",
)
print("pagination response=", response)
assert len(response["keys"]) == 2
@ -3470,9 +3490,18 @@ async def test_list_keys(prisma_client):
response = await list_keys(
request,
UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value),
user_id=user_id,
page=1,
size=10,
user_id=user_id,
team_id=None,
organization_id=None,
key_hash=None,
key_alias=None,
return_full_object=False,
include_team_keys=False,
include_created_by_keys=False,
sort_by=None,
sort_order="desc",
)
print("filtered user_id response=", response)
assert len(response["keys"]) == 1
@ -3482,9 +3511,18 @@ async def test_list_keys(prisma_client):
response = await list_keys(
request,
UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value),
key_alias=key_alias,
page=1,
size=10,
user_id=None,
team_id=None,
organization_id=None,
key_hash=None,
key_alias=key_alias,
return_full_object=False,
include_team_keys=False,
include_created_by_keys=False,
sort_by=None,
sort_order="desc",
)
assert len(response["keys"]) == 1
assert _key in response["keys"]

View file

@ -11,15 +11,27 @@ from fastapi import FastAPI
from starlette import status
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
from litellm.proxy._types import MCPSpecVersion, MCPSpecVersionType, MCPTransportType, MCPTransport, NewMCPServerRequest, LiteLLM_MCPServerTable, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import (
MCPSpecVersion,
MCPSpecVersionType,
MCPTransportType,
MCPTransport,
NewMCPServerRequest,
LiteLLM_MCPServerTable,
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.types.mcp import MCPAuth
from litellm.proxy.management_endpoints.mcp_management_endpoints import does_mcp_server_exist
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
does_mcp_server_exist,
)
TEST_MASTER_KEY = os.getenv("LITELLM_MASTER_KEY", "sk-1234")
def generate_mcpserver_record(url: Optional[str] = None,
transport: Optional[MCPTransportType] = None,
spec_version: Optional[MCPSpecVersionType] = None) -> LiteLLM_MCPServerTable:
def generate_mcpserver_record(
url: Optional[str] = None, transport: Optional[MCPTransportType] = None
) -> LiteLLM_MCPServerTable:
"""
Generate a mock record for testing.
"""
@ -30,11 +42,11 @@ def generate_mcpserver_record(url: Optional[str] = None,
alias="Test Server",
url=url or "http://localhost.com:8080/mcp",
transport=transport or MCPTransport.sse,
spec_version=spec_version or MCPSpecVersion.mar_2025,
created_at=now,
updated_at=now,
)
# Cheers SO
def is_valid_uuid(val):
try:
@ -43,11 +55,12 @@ def is_valid_uuid(val):
except ValueError:
return False
def generate_mcpserver_create_request(
server_id: Optional[str] = None,
url: Optional[str] = None,
transport: Optional[MCPTransportType] = None,
spec_version: Optional[MCPSpecVersionType] = None) -> NewMCPServerRequest:
server_id: Optional[str] = None,
url: Optional[str] = None,
transport: Optional[MCPTransportType] = None,
) -> NewMCPServerRequest:
"""
Generate a mock create request for testing.
"""
@ -56,10 +69,12 @@ def generate_mcpserver_create_request(
alias="Test Server",
url=url or "http://localhost.com:8080/mcp",
transport=transport or MCPTransport.sse,
spec_version=spec_version or MCPSpecVersion.mar_2025,
)
def assert_mcp_server_record_same(mcp_server: NewMCPServerRequest, resp: LiteLLM_MCPServerTable):
def assert_mcp_server_record_same(
mcp_server: NewMCPServerRequest, resp: LiteLLM_MCPServerTable
):
"""
Assert that the mcp server record is created correctly.
"""
@ -71,7 +86,6 @@ def assert_mcp_server_record_same(mcp_server: NewMCPServerRequest, resp: LiteLLM
assert resp.url == mcp_server.url
assert resp.description == mcp_server.description
assert resp.transport == mcp_server.transport
assert resp.spec_version == mcp_server.spec_version
assert resp.auth_type == mcp_server.auth_type
assert resp.created_at is not None
assert resp.updated_at is not None
@ -83,224 +97,263 @@ def test_does_mcp_server_exist():
"""
Unit Test if the MCP server exists in the list.
"""
mcp_server_records: List[LiteLLM_MCPServerTable] = [generate_mcpserver_record(), generate_mcpserver_record()]
mcp_server_records: List[LiteLLM_MCPServerTable] = [
generate_mcpserver_record(),
generate_mcpserver_record(),
]
# test all records are found
for record in mcp_server_records:
assert does_mcp_server_exist(mcp_server_records, record.server_id)
# test record not found
not_found_record = str(uuid.uuid4())
assert False == does_mcp_server_exist(mcp_server_records, not_found_record)
@pytest.mark.asyncio
async def test_create_mcp_server_direct():
"""
Direct test of the MCP server creation logic without HTTP calls.
"""
# Mock the database functions directly
with mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE", True), \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw") as mock_get_prisma, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server", new_callable=mock.AsyncMock) as mock_create, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server", new_callable=mock.AsyncMock) as mock_get_server, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager") as mock_manager:
with mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE",
True,
), mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
) as mock_get_prisma, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server",
new_callable=mock.AsyncMock,
) as mock_create, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
new_callable=mock.AsyncMock,
) as mock_get_server, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager"
) as mock_manager:
# Import after mocking
from litellm.proxy.management_endpoints.mcp_management_endpoints import add_mcp_server
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
add_mcp_server,
)
# Mock database client
mock_prisma = mock.Mock()
mock_get_prisma.return_value = mock_prisma
# Mock server manager
mock_manager.add_update_server = mock.Mock()
mock_manager.reload_servers_from_database = mock.AsyncMock()
# Set up test data
server_id = str(uuid.uuid4())
mcp_server_request = generate_mcpserver_create_request(server_id=server_id)
# The function will normalize the alias by replacing spaces with underscores
expected_alias = mcp_server_request.alias.replace(' ', '_') if mcp_server_request.alias else None
expected_alias = (
mcp_server_request.alias.replace(" ", "_")
if mcp_server_request.alias
else None
)
expected_response = LiteLLM_MCPServerTable(
server_id=server_id,
alias=expected_alias, # Use the normalized alias
description=mcp_server_request.description,
url=mcp_server_request.url,
transport=mcp_server_request.transport,
spec_version=mcp_server_request.spec_version,
auth_type=mcp_server_request.auth_type,
created_at=datetime.now(),
updated_at=datetime.now(),
created_by=LITELLM_PROXY_ADMIN_NAME,
updated_by=LITELLM_PROXY_ADMIN_NAME,
teams=[]
teams=[],
)
# Mock the database calls
mock_get_server.return_value = None # Server doesn't exist yet
# Set up async mock for create_mcp_server using AsyncMock
mock_create.return_value = expected_response
# Create mock user auth
user_auth = UserAPIKeyAuth(
api_key=TEST_MASTER_KEY,
user_id="test-user",
user_role=LitellmUserRoles.PROXY_ADMIN
user_role=LitellmUserRoles.PROXY_ADMIN,
)
# Call the function directly
result = await add_mcp_server(
payload=mcp_server_request,
user_api_key_dict=user_auth
payload=mcp_server_request, user_api_key_dict=user_auth
)
# Verify the result
assert result.server_id == server_id
assert result.alias == expected_alias # Check against normalized alias
assert result.url == mcp_server_request.url
assert result.transport == mcp_server_request.transport
assert result.spec_version == mcp_server_request.spec_version
# Verify mocks were called
mock_get_server.assert_called_once_with(mock_prisma, server_id)
mock_create.assert_called_once()
mock_manager.add_update_server.assert_called_once_with(expected_response)
@pytest.mark.asyncio
@pytest.mark.asyncio
async def test_create_duplicate_mcp_server():
"""
Test that creating a duplicate MCP server fails appropriately.
"""
# Mock the database functions directly
with mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE", True), \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw") as mock_get_prisma, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server", new_callable=mock.AsyncMock) as mock_get_server:
with mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE",
True,
), mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
) as mock_get_prisma, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
new_callable=mock.AsyncMock,
) as mock_get_server:
# Import after mocking
from litellm.proxy.management_endpoints.mcp_management_endpoints import add_mcp_server
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
add_mcp_server,
)
from fastapi import HTTPException
# Mock database client
mock_prisma = mock.Mock()
mock_get_prisma.return_value = mock_prisma
# Set up test data
server_id = str(uuid.uuid4())
mcp_server_request = generate_mcpserver_create_request(server_id=server_id)
existing_server = LiteLLM_MCPServerTable(
server_id=server_id,
alias="Existing Server",
url="http://existing.com",
transport=MCPTransport.sse,
spec_version=MCPSpecVersion.mar_2025,
created_at=datetime.now(),
updated_at=datetime.now(),
teams=[]
teams=[],
)
# Mock that server already exists
mock_get_server.return_value = existing_server
# Create mock user auth
user_auth = UserAPIKeyAuth(
api_key=TEST_MASTER_KEY,
user_id="test-user",
user_role=LitellmUserRoles.PROXY_ADMIN
user_role=LitellmUserRoles.PROXY_ADMIN,
)
# Expect HTTPException to be raised
with pytest.raises(HTTPException) as exc_info:
await add_mcp_server(
payload=mcp_server_request,
user_api_key_dict=user_auth
payload=mcp_server_request, user_api_key_dict=user_auth
)
# Verify the exception details
assert exc_info.value.status_code == 400
assert "already exists" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_create_mcp_server_auth_failure():
"""
Test that non-admin users cannot create MCP servers.
"""
# Mock the database functions directly
with mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE", True), \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw") as mock_get_prisma:
with mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE",
True,
), mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
) as mock_get_prisma:
# Import after mocking
from litellm.proxy.management_endpoints.mcp_management_endpoints import add_mcp_server
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
add_mcp_server,
)
from fastapi import HTTPException
# Mock database client
# Mock database client
mock_prisma = mock.Mock()
mock_get_prisma.return_value = mock_prisma
# Set up test data
server_id = str(uuid.uuid4())
mcp_server_request = generate_mcpserver_create_request(server_id=server_id)
# Create mock user auth without admin role
user_auth = UserAPIKeyAuth(
api_key=TEST_MASTER_KEY,
user_id="test-user",
user_role=LitellmUserRoles.INTERNAL_USER # Not an admin
user_role=LitellmUserRoles.INTERNAL_USER, # Not an admin
)
# Expect HTTPException to be raised
with pytest.raises(HTTPException) as exc_info:
await add_mcp_server(
payload=mcp_server_request,
user_api_key_dict=user_auth
payload=mcp_server_request, user_api_key_dict=user_auth
)
# Verify the exception details
assert exc_info.value.status_code == 403
assert "permission" in str(exc_info.value.detail)
@pytest.mark.asyncio
@pytest.mark.asyncio
async def test_create_mcp_server_invalid_alias():
"""
Test that creating an MCP server with a '-' in the alias fails with the correct error.
"""
with mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE", True), \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw") as mock_get_prisma, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server") as mock_get_server, \
mock.patch("litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server") as mock_create:
from litellm.proxy.management_endpoints.mcp_management_endpoints import add_mcp_server
with mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE",
True,
), mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
) as mock_get_prisma, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server"
) as mock_get_server, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server"
) as mock_create:
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
add_mcp_server,
)
from fastapi import HTTPException
mock_prisma = mock.Mock()
mock_get_prisma.return_value = mock_prisma
# Set up test data with invalid alias
server_id = str(uuid.uuid4())
mcp_server_request = generate_mcpserver_create_request(server_id=server_id)
mcp_server_request.alias = "invalid-alias" # This should trigger the validation error
mcp_server_request.alias = (
"invalid-alias" # This should trigger the validation error
)
# Mock that server does not exist
mock_get_server.return_value = None
# Mock create_mcp_server to prevent 500 error (this should not be called due to validation)
mock_create.return_value = None
user_auth = UserAPIKeyAuth(
api_key=TEST_MASTER_KEY,
user_id="test-user",
user_role=LitellmUserRoles.PROXY_ADMIN
user_role=LitellmUserRoles.PROXY_ADMIN,
)
with pytest.raises(HTTPException) as exc_info:
await add_mcp_server(
payload=mcp_server_request,
user_api_key_dict=user_auth
payload=mcp_server_request, user_api_key_dict=user_auth
)
assert exc_info.value.status_code == 400
assert "Server name cannot contain '-'. Use an alternative character instead Found: invalid-alias" in str(exc_info.value.detail)
assert (
"Server name cannot contain '-'. Use an alternative character instead Found: invalid-alias"
in str(exc_info.value.detail)
)
def test_validate_mcp_server_name_direct():
"""
@ -308,16 +361,16 @@ def test_validate_mcp_server_name_direct():
"""
from litellm.proxy._experimental.mcp_server.utils import validate_mcp_server_name
from fastapi import HTTPException
# Test that valid names pass
validate_mcp_server_name("valid_name")
validate_mcp_server_name("valid name")
# Test that invalid names with hyphens raise exceptions
with pytest.raises(Exception) as exc_info:
validate_mcp_server_name("invalid-name")
assert "cannot contain" in str(exc_info.value)
# Test that invalid names with hyphens raise HTTPException when requested
with pytest.raises(HTTPException) as exc_info:
validate_mcp_server_name("invalid-name", raise_http_exception=True)

View file

@ -203,6 +203,9 @@ class TestDataDogLLMObsLogger:
assert metadata["cache_hit"] is True
assert metadata["cache_key"] == "test-cache-key-789"
# Test 4: Verify is_streamed_request is in metadata
assert metadata["is_streamed_request"] is True
def test_cache_metadata_fields(self, mock_env_vars, mock_response_obj):
"""Test that cache-related metadata fields are correctly tracked"""
with patch(

View file

@ -0,0 +1,134 @@
import os
import sys
from unittest.mock import MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../.."))
from litellm.integrations.langsmith import LangsmithLogger
class TestLangsmithLoggerInit:
"""Test cases for LangSmith logger initialization, particularly sampling rate handling.
These tests verify that the sampling_rate attribute is set during initialization.
Note: The current implementation has some edge cases in the sampling rate logic.
"""
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "1"}, clear=False)
def test_langsmith_sampling_rate_parameter_respected_with_valid_env(
self, mock_create_task
):
"""Test that langsmith_sampling_rate parameter is properly set when env var condition is met."""
# When there's a valid integer in env var, the parameter should be used due to 'or' logic
sampling_rate = 0.5
logger = LangsmithLogger(
langsmith_api_key="test-key",
langsmith_project="test-project",
langsmith_sampling_rate=sampling_rate,
)
# With the current 'or' logic and valid env var, the parameter should be used
assert (
logger.sampling_rate == sampling_rate
), f"Expected sampling_rate to be {sampling_rate}, got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "1"}, clear=False)
def test_langsmith_sampling_rate_zero_parameter_falls_back_to_env(
self, mock_create_task
):
"""Test that 0.0 parameter falls back to env var due to falsy value."""
# This demonstrates the current behavior where 0.0 is falsy and falls back to env
logger = LangsmithLogger(
langsmith_api_key="test-key",
langsmith_project="test-project",
langsmith_sampling_rate=0.0, # This is falsy!
)
# Due to current 'or' logic, 0.0 falls back to env var
assert (
logger.sampling_rate == 1.0
), f"Expected sampling_rate to fall back to 1.0 from env, got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "1"}, clear=False)
def test_langsmith_sampling_rate_from_integer_env_var(self, mock_create_task):
"""Test that sampling rate uses environment variable when parameter not provided and env var is integer."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
)
# Should use env var since it's a valid integer
assert (
logger.sampling_rate == 1.0
), f"Expected sampling_rate to be 1.0 from env var, got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "0.8"}, clear=False)
def test_langsmith_sampling_rate_decimal_env_var_ignored(self, mock_create_task):
"""Test that decimal environment variables are ignored due to isdigit() check."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
)
# Decimal env vars are ignored due to isdigit() check, falls back to 1.0
assert (
logger.sampling_rate == 1.0
), f"Expected sampling_rate to default to 1.0 (decimal env ignored), got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {}, clear=True)
def test_langsmith_sampling_rate_default_value(self, mock_create_task):
"""Test that sampling rate defaults to 1.0 when no parameter or env var provided."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
)
assert (
logger.sampling_rate == 1.0
), f"Expected default sampling_rate to be 1.0, got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "invalid"}, clear=False)
def test_langsmith_sampling_rate_invalid_env_var_defaults(self, mock_create_task):
"""Test that invalid environment variable falls back to default value."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
)
assert (
logger.sampling_rate == 1.0
), f"Expected sampling_rate to default to 1.0 with invalid env var, got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": ""}, clear=False)
def test_langsmith_sampling_rate_empty_env_var_defaults(self, mock_create_task):
"""Test that empty environment variable falls back to default value."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
)
assert (
logger.sampling_rate == 1.0
), f"Expected sampling_rate to default to 1.0 with empty env var, got {logger.sampling_rate}"
@patch("asyncio.create_task")
def test_langsmith_sampling_rate_attribute_exists(self, mock_create_task):
"""Test that the sampling_rate attribute is always set on the logger instance."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
)
# Verify the attribute exists and is a float
assert hasattr(
logger, "sampling_rate"
), "LangsmithLogger should have sampling_rate attribute"
assert isinstance(
logger.sampling_rate, float
), f"sampling_rate should be a float, got {type(logger.sampling_rate)}"
assert (
logger.sampling_rate >= 0.0
), f"sampling_rate should be non-negative, got {logger.sampling_rate}"

View file

@ -22,8 +22,11 @@ sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
from litellm.types.utils import Usage
from litellm.litellm_core_utils.llm_cost_calc.utils import (
calculate_cache_writing_cost,
generic_cost_per_token,
)
from litellm.types.utils import CacheCreationTokenDetails, Usage
def test_reasoning_tokens_no_price_set():
@ -385,3 +388,79 @@ def test_string_cost_values_with_threshold():
assert round(prompt_cost, 12) == round(expected_prompt_cost, 12)
assert round(completion_cost, 12) == round(expected_completion_cost, 12)
def test_calculate_cache_writing_cost():
"""Test the calculate_cache_writing_cost function with detailed cache creation token breakdown."""
# Test case 1: With cache creation token details (matching the provided input)
cache_creation_tokens = 14055
cache_creation_token_details = CacheCreationTokenDetails(
ephemeral_5m_input_tokens=56, ephemeral_1h_input_tokens=13999
)
cache_creation_cost_above_1hr = 6e-06
cache_creation_cost = 3.75e-06
result = calculate_cache_writing_cost(
cache_creation_tokens=cache_creation_tokens,
cache_creation_token_details=cache_creation_token_details,
cache_creation_cost_above_1hr=cache_creation_cost_above_1hr,
cache_creation_cost=cache_creation_cost,
)
# Expected calculation:
# 5m tokens: 56 * 3.75e-06 = 0.00021
# 1h tokens: 13999 * 6e-06 = 0.083994
# Total: 0.00021 + 0.083994 = 0.084204
expected_cost = (56 * 3.75e-06) + (13999 * 6e-06)
assert round(result, 6) == round(expected_cost, 6)
assert round(result, 6) == 0.084204
# Test case 2: Without cache creation token details (fallback behavior)
cache_creation_tokens_no_details = 1000
cache_creation_token_details_none = None
cache_creation_cost_fallback = 5e-06
result_no_details = calculate_cache_writing_cost(
cache_creation_tokens=cache_creation_tokens_no_details,
cache_creation_token_details=cache_creation_token_details_none,
cache_creation_cost_above_1hr=cache_creation_cost_above_1hr,
cache_creation_cost=cache_creation_cost_fallback,
)
# Expected calculation when no details: 1000 * 5e-06 = 0.005
expected_cost_no_details = 1000 * 5e-06
assert round(result_no_details, 6) == round(expected_cost_no_details, 6)
assert result_no_details == 0.005
# Test case 3: With cache creation token details but None values
cache_creation_token_details_partial = CacheCreationTokenDetails(
ephemeral_5m_input_tokens=None, ephemeral_1h_input_tokens=100
)
result_partial = calculate_cache_writing_cost(
cache_creation_tokens=500,
cache_creation_token_details=cache_creation_token_details_partial,
cache_creation_cost_above_1hr=6e-06,
cache_creation_cost=3e-06,
)
# Expected calculation: 0 (for None 5m tokens) + (100 * 6e-06) = 0.0006
expected_cost_partial = (0.0) + (100 * 6e-06)
assert round(result_partial, 6) == round(expected_cost_partial, 6)
assert round(result_partial, 6) == 0.0006
# Test case 4: Zero costs
result_zero = calculate_cache_writing_cost(
cache_creation_tokens=1000,
cache_creation_token_details=CacheCreationTokenDetails(
ephemeral_5m_input_tokens=50, ephemeral_1h_input_tokens=950
),
cache_creation_cost_above_1hr=0.0,
cache_creation_cost=0.0,
)
assert result_zero == 0.0

View file

@ -88,6 +88,12 @@ class TestStandardizedResetTime(unittest.TestCase):
)
self.assertEqual(london_result, london_expected)
# Test Bangkok timezone (UTC+7): 5:30 AM next day, so next reset is midnight the day after
bangkok = ZoneInfo("Asia/Bangkok")
bangkok_expected = datetime(2023, 5, 17, 0, 0, 0, tzinfo=bangkok)
bangkok_result = get_next_standardized_reset_time("1d", base_time, "Asia/Bangkok")
self.assertEqual(bangkok_result, bangkok_expected)
def test_edge_cases(self):
"""Test edge cases and boundary conditions"""
# Exactly on hour boundary

View file

@ -2,11 +2,12 @@ import json
import os
import sys
from unittest.mock import Mock, patch
import pytest
sys.path.insert(0, os.path.abspath("../../../../..")) # Adds the parent directory to the system path
import litellm
from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
# Mock responses for different embedding models
titan_embedding_response = {
@ -146,7 +147,7 @@ def test_bedrock_embedding_with_sigv4():
"""Test embedding falls back to SigV4 auth when no bearer token is provided"""
litellm.set_verbose = True
model = "bedrock/amazon.titan-embed-text-v1"
with patch("litellm.llms.bedrock.embed.embedding.BedrockEmbedding.embeddings") as mock_bedrock_embed:
mock_embedding_response = litellm.EmbeddingResponse()
mock_embedding_response.data = [{"embedding": [0.1, 0.2, 0.3]}]
@ -159,4 +160,85 @@ def test_bedrock_embedding_with_sigv4():
)
assert isinstance(response, litellm.EmbeddingResponse)
mock_bedrock_embed.assert_called_once()
mock_bedrock_embed.assert_called_once()
def test_bedrock_titan_v2_encoding_format_float():
"""Test amazon.titan-embed-text-v2:0 with encoding_format=float parameter"""
litellm.set_verbose = True
client = HTTPHandler()
test_api_key = "test-bearer-token-12345"
model = "bedrock/amazon.titan-embed-text-v2:0"
# Mock response with embeddingsByType for binary format (addressing issue #14680)
titan_v2_response = {
"embedding": [0.1, 0.2, 0.3],
"inputTextTokenCount": 10
}
with patch.object(client, "post") as mock_post:
mock_response = Mock()
mock_response.status_code = 200
mock_response.text = json.dumps(titan_v2_response)
mock_response.json = lambda: json.loads(mock_response.text)
mock_post.return_value = mock_response
response = litellm.embedding(
model=model,
input=test_input,
encoding_format="float", # This should work but currently throws UnsupportedParamsError
client=client,
aws_region_name="us-east-1",
aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com",
api_key=test_api_key
)
assert isinstance(response, litellm.EmbeddingResponse)
assert isinstance(response.data[0]['embedding'], list)
assert len(response.data[0]['embedding']) == 3
# Verify that the request contains embeddingTypes: ["float"] instead of encoding_format
request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}"))
assert "embeddingTypes" in request_body
assert request_body["embeddingTypes"] == ["float"]
assert "encoding_format" not in request_body
def test_bedrock_titan_v2_encoding_format_base64():
"""Test amazon.titan-embed-text-v2:0 with encoding_format=base64 parameter (maps to binary)"""
litellm.set_verbose = True
client = HTTPHandler()
test_api_key = "test-bearer-token-12345"
model = "bedrock/amazon.titan-embed-text-v2:0"
# Mock response with embeddingsByType for binary format
titan_v2_binary_response = {
"embeddingsByType": {
"binary": "YmluYXJ5X2VtYmVkZGluZ19kYXRh" # base64 encoded binary data
},
"inputTextTokenCount": 10
}
with patch.object(client, "post") as mock_post:
mock_response = Mock()
mock_response.status_code = 200
mock_response.text = json.dumps(titan_v2_binary_response)
mock_response.json = lambda: json.loads(mock_response.text)
mock_post.return_value = mock_response
response = litellm.embedding(
model=model,
input=test_input,
encoding_format="base64", # This should map to embeddingTypes: ["binary"]
client=client,
aws_region_name="us-east-1",
aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com",
api_key=test_api_key
)
assert isinstance(response, litellm.EmbeddingResponse)
# Verify that the request contains embeddingTypes: ["binary"] for base64 encoding
request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}"))
assert "embeddingTypes" in request_body
assert request_body["embeddingTypes"] == ["binary"]

View file

@ -23,7 +23,6 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@pytest.mark.asyncio
class TestMCPRequestHandler:
@pytest.mark.parametrize(
"user_api_key_auth,object_permission_id,prisma_client_available,db_result,expected_result",
[
@ -172,7 +171,6 @@ class TestMCPRequestHandler:
with patch.object(
MCPRequestHandler, "_get_allowed_mcp_servers_for_team"
) as mock_team_servers:
# Configure mocks to return the test data
mock_key_servers.return_value = key_servers
mock_team_servers.return_value = team_servers
@ -366,7 +364,6 @@ class TestMCPRequestHandler:
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
mcp_protocol_version,
) = await MCPRequestHandler.process_mcp_request(scope)
# Assert the results
@ -377,7 +374,6 @@ class TestMCPRequestHandler:
assert mcp_server_auth_headers == expected_server_auth_headers
# For these tests, mcp_servers should be None
assert mcp_servers is None
assert mcp_protocol_version is None
@pytest.mark.parametrize(
"headers,expected_result",
@ -547,15 +543,12 @@ class TestMCPRequestHandler:
mcp_auth_header,
mcp_servers_result,
mcp_server_auth_headers,
mcp_protocol_version,
) = await MCPRequestHandler.process_mcp_request(scope)
assert auth_result == mock_auth_result
assert mcp_auth_header == expected_result["mcp_auth"]
assert mcp_servers_result == expected_result["mcp_servers"]
# For these tests, mcp_server_auth_headers should be empty
assert mcp_server_auth_headers == {}
# For these tests, mcp_protocol_version should be None
assert mcp_protocol_version is None
class TestMCPCustomHeaderName:
@ -588,7 +581,6 @@ class TestMCPCustomHeaderName:
with patch(
"litellm.proxy.proxy_server.general_settings"
) as mock_general_settings:
# Configure mocks
mock_get_secret.return_value = env_var
mock_general_settings.get.return_value = general_setting
@ -692,7 +684,6 @@ class TestMCPCustomHeaderName:
"_get_mcp_client_side_auth_header_name",
return_value="custom-auth-header",
):
# Create ASGI scope with custom header
scope = {
"type": "http",
@ -725,7 +716,6 @@ class TestMCPCustomHeaderName:
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
mcp_protocol_version,
) = await MCPRequestHandler.process_mcp_request(scope)
# Assert the results
@ -733,7 +723,6 @@ class TestMCPCustomHeaderName:
assert mcp_auth_header == "custom-auth-token"
assert mcp_servers is None
assert mcp_server_auth_headers == {}
assert mcp_protocol_version is None
# Verify the mock was called
mock_auth.assert_called_once()
@ -897,7 +886,6 @@ class TestMCPAccessGroupsE2E:
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
mcp_protocol_version,
) = await MCPRequestHandler.process_mcp_request(scope)
# Assert the results
@ -907,7 +895,6 @@ class TestMCPAccessGroupsE2E:
mcp_servers is None
) # x-mcp-access-groups is not parsed as mcp_servers
assert mcp_server_auth_headers == {}
assert mcp_protocol_version is None
# Verify the mock was called
mock_auth.assert_called_once()
@ -948,7 +935,6 @@ class TestMCPAccessGroupsE2E:
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
mcp_protocol_version,
) = await MCPRequestHandler.process_mcp_request(scope)
# Assert the results
@ -956,7 +942,6 @@ class TestMCPAccessGroupsE2E:
assert mcp_auth_header is None
assert mcp_servers == ["server1", "dev_group", "server2"]
assert mcp_server_auth_headers == {}
assert mcp_protocol_version is None
# Verify the mock was called
mock_auth.assert_called_once()
@ -978,7 +963,6 @@ def test_mcp_path_based_server_segregation(monkeypatch):
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
mcp_protocol_version,
) = get_auth_context()
# Capture the MCP servers for testing

View file

@ -88,19 +88,21 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
working_server = MagicMock()
working_server.name = "working_server"
working_server.alias = "working"
failing_server = MagicMock()
failing_server.name = "failing_server"
failing_server.alias = "failing"
# Mock global_mcp_server_manager
mock_manager = MagicMock()
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["working_server", "failing_server"])
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["working_server", "failing_server"]
)
mock_manager.get_mcp_server_by_id = lambda server_id: (
working_server if server_id == "working_server" else failing_server
)
async def mock_get_tools_from_server(server, mcp_auth_header=None, mcp_protocol_version=None):
async def mock_get_tools_from_server(server, mcp_auth_header=None):
if server.name == "working_server":
# Working server returns tools
tool1 = MagicMock()
@ -111,7 +113,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
else:
# Failing server raises an exception
raise Exception("Server connection failed")
mock_manager._get_tools_from_server = mock_get_tools_from_server
with patch(
@ -124,26 +126,29 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
# Test with server-specific auth headers
mcp_server_auth_headers = {
"working": "Bearer working-token",
"failing": "Bearer failing-token"
"failing": "Bearer failing-token",
}
result = await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=None,
mcp_servers=None,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_protocol_version=None
)
# Verify that tools from the working server are returned
assert len(result) == 1
assert result[0].name == "working_tool_1"
# Verify failure logging
mock_logger.exception.assert_any_call("Error getting tools from server failing_server: Server connection failed")
mock_logger.exception.assert_any_call(
"Error getting tools from server failing_server: Server connection failed"
)
# Verify success logging
mock_logger.info.assert_any_call("Successfully fetched 1 tools total from all MCP servers")
mock_logger.info.assert_any_call(
"Successfully fetched 1 tools total from all MCP servers"
)
@pytest.mark.asyncio
@ -165,22 +170,24 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
failing_server1 = MagicMock()
failing_server1.name = "failing_server1"
failing_server1.alias = "failing1"
failing_server2 = MagicMock()
failing_server2.name = "failing_server2"
failing_server2.alias = "failing2"
# Mock global_mcp_server_manager
mock_manager = MagicMock()
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["failing_server1", "failing_server2"])
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["failing_server1", "failing_server2"]
)
mock_manager.get_mcp_server_by_id = lambda server_id: (
failing_server1 if server_id == "failing_server1" else failing_server2
)
async def mock_get_tools_from_server(server, mcp_auth_header=None, mcp_protocol_version=None):
async def mock_get_tools_from_server(server, mcp_auth_header=None):
# All servers fail
raise Exception(f"Server {server.name} connection failed")
mock_manager._get_tools_from_server = mock_get_tools_from_server
with patch(
@ -193,26 +200,31 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
# Test with server-specific auth headers
mcp_server_auth_headers = {
"failing1": "Bearer failing1-token",
"failing2": "Bearer failing2-token"
"failing2": "Bearer failing2-token",
}
result = await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=None,
mcp_servers=None,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_protocol_version=None
)
# Verify that empty list is returned
assert len(result) == 0
# Verify failure logging for both servers
mock_logger.exception.assert_any_call("Error getting tools from server failing_server1: Server failing_server1 connection failed")
mock_logger.exception.assert_any_call("Error getting tools from server failing_server2: Server failing_server2 connection failed")
mock_logger.exception.assert_any_call(
"Error getting tools from server failing_server1: Server failing_server1 connection failed"
)
mock_logger.exception.assert_any_call(
"Error getting tools from server failing_server2: Server failing_server2 connection failed"
)
# Verify total logging
mock_logger.info.assert_any_call("Successfully fetched 0 tools total from all MCP servers")
mock_logger.info.assert_any_call(
"Successfully fetched 0 tools total from all MCP servers"
)
@pytest.mark.asyncio
@ -288,55 +300,66 @@ async def test_concurrent_initialize_session_managers():
)
except ImportError:
pytest.skip("MCP server not available")
# Import the module to reset state
import litellm.proxy._experimental.mcp_server.server as mcp_server
# Reset state before test
original_initialized = mcp_server._SESSION_MANAGERS_INITIALIZED
original_session_cm = mcp_server._session_manager_cm
original_sse_session_cm = mcp_server._sse_session_manager_cm
try:
mcp_server._SESSION_MANAGERS_INITIALIZED = False
mcp_server._session_manager_cm = None
mcp_server._sse_session_manager_cm = None
# Mock the session managers to avoid actual MCP initialization
with patch('litellm.proxy._experimental.mcp_server.server.session_manager') as mock_session_manager, \
patch('litellm.proxy._experimental.mcp_server.server.sse_session_manager') as mock_sse_session_manager, \
patch('litellm.proxy._experimental.mcp_server.server.verbose_logger'):
with patch(
"litellm.proxy._experimental.mcp_server.server.session_manager"
) as mock_session_manager, patch(
"litellm.proxy._experimental.mcp_server.server.sse_session_manager"
) as mock_sse_session_manager, patch(
"litellm.proxy._experimental.mcp_server.server.verbose_logger"
):
# Mock the run() method to return a mock context manager
mock_cm = AsyncMock()
mock_cm.__aenter__ = AsyncMock()
mock_cm.__aexit__ = AsyncMock()
mock_session_manager.run.return_value = mock_cm
mock_sse_session_manager.run.return_value = mock_cm
# Create multiple concurrent tasks that call initialize_session_managers
async def init_task():
await initialize_session_managers()
return "success"
# Run 10 concurrent initialization attempts
tasks = [init_task() for _ in range(10)]
results = await asyncio.gather(*tasks, return_exceptions=True)
# All tasks should complete successfully (no exceptions)
assert all(result == "success" for result in results), f"Some tasks failed: {results}"
assert all(
result == "success" for result in results
), f"Some tasks failed: {results}"
# session_manager.run() should only be called once due to the lock
assert mock_session_manager.run.call_count == 1, f"Expected 1 call to session_manager.run(), got {mock_session_manager.run.call_count}"
assert mock_sse_session_manager.run.call_count == 1, f"Expected 1 call to sse_session_manager.run(), got {mock_sse_session_manager.run.call_count}"
assert (
mock_session_manager.run.call_count == 1
), f"Expected 1 call to session_manager.run(), got {mock_session_manager.run.call_count}"
assert (
mock_sse_session_manager.run.call_count == 1
), f"Expected 1 call to sse_session_manager.run(), got {mock_sse_session_manager.run.call_count}"
# The context managers should only be entered once each
assert mock_cm.__aenter__.call_count == 2, f"Expected 2 calls to __aenter__ (one for each session manager), got {mock_cm.__aenter__.call_count}"
assert (
mock_cm.__aenter__.call_count == 2
), f"Expected 2 calls to __aenter__ (one for each session manager), got {mock_cm.__aenter__.call_count}"
# State should be properly set
assert mcp_server._SESSION_MANAGERS_INITIALIZED is True
finally:
# Restore original state
mcp_server._SESSION_MANAGERS_INITIALIZED = original_initialized
@ -371,14 +394,12 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name():
name="custom_solutions/user_123",
alias="custom_solutions/user_123",
transport=MCPTransport.http,
spec_version=MCPSpecVersion.jun_2025,
)
other_server = MCPServer(
server_id="other_server_in_group_id",
name="custom_solutions/another_user_456",
alias="custom_solutions/another_user_456",
transport=MCPTransport.http,
spec_version=MCPSpecVersion.jun_2025,
)
global_mcp_server_manager.registry[specific_server.server_id] = specific_server
global_mcp_server_manager.registry[other_server.server_id] = other_server
@ -392,9 +413,13 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name():
mock_get_tools_spy = AsyncMock(return_value=[])
# Mock the function that checks DB for an access group named "custom_solutions"
mock_db_lookup = AsyncMock(return_value=[specific_server.server_id, other_server.server_id])
mock_db_lookup = AsyncMock(
return_value=[specific_server.server_id, other_server.server_id]
)
mock_get_allowed = AsyncMock(return_value=[specific_server.server_id, other_server.server_id])
mock_get_allowed = AsyncMock(
return_value=[specific_server.server_id, other_server.server_id]
)
with patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers",
@ -415,7 +440,9 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name():
)
# Get the list of actual server objects that the orchestrator tried to contact
called_servers = [call.kwargs["server"] for call in mock_get_tools_spy.call_args_list]
called_servers = [
call.kwargs["server"] for call in mock_get_tools_spy.call_args_list
]
assert len(called_servers) == 1, "Should have resolved to exactly one server."
assert (

View file

@ -5,13 +5,13 @@ from unittest.mock import MagicMock, AsyncMock
import pytest
# Add the parent directory to the path so we can import litellm
sys.path.insert(0, '../../../../../')
sys.path.insert(0, "../../../../../")
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
_deserialize_env_dict,
)
from litellm.proxy._types import LiteLLM_MCPServerTable, MCPSpecVersion, MCPTransport
from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -24,12 +24,12 @@ class TestMCPServerManager:
env_json = '{"PATH": "/usr/bin", "DEBUG": "1"}'
result = _deserialize_env_dict(env_json)
assert result == {"PATH": "/usr/bin", "DEBUG": "1"}
# Test already dict
env_dict = {"PATH": "/usr/bin", "DEBUG": "1"}
result = _deserialize_env_dict(env_dict)
assert result == {"PATH": "/usr/bin", "DEBUG": "1"}
# Test invalid JSON
invalid_json = '{"PATH": "/usr/bin", "DEBUG": 1'
result = _deserialize_env_dict(invalid_json)
@ -38,27 +38,26 @@ class TestMCPServerManager:
def test_add_update_server_stdio(self):
"""Test adding stdio MCP server"""
manager = MCPServerManager()
stdio_server = LiteLLM_MCPServerTable(
server_id="stdio-server-1",
alias="test_stdio_server",
description="Test stdio server",
url=None,
transport=MCPTransport.stdio,
spec_version=MCPSpecVersion.mar_2025,
command="python",
args=["-m", "server"],
env={"DEBUG": "1", "TEST": "1"},
created_at=datetime.now(),
updated_at=datetime.now()
updated_at=datetime.now(),
)
manager.add_update_server(stdio_server)
# Verify server was added
assert "stdio-server-1" in manager.registry
added_server = manager.registry["stdio-server-1"]
assert added_server.server_id == "stdio-server-1"
assert added_server.name == "test_stdio_server"
assert added_server.transport == MCPTransport.stdio
@ -69,20 +68,19 @@ class TestMCPServerManager:
def test_create_mcp_client_stdio(self):
"""Test creating MCP client for stdio transport"""
manager = MCPServerManager()
stdio_server = MCPServer(
server_id="stdio-server-2",
name="test_stdio_server",
url=None,
transport=MCPTransport.stdio,
spec_version=MCPSpecVersion.mar_2025,
command="node",
args=["server.js"],
env={"NODE_ENV": "test"}
env={"NODE_ENV": "test"},
)
client = manager._create_mcp_client(stdio_server)
assert client.transport_type == MCPTransport.stdio
assert client.stdio_config is not None
assert client.stdio_config["command"] == "node"
@ -93,24 +91,28 @@ class TestMCPServerManager:
async def test_list_tools_with_server_specific_auth_headers(self):
"""Test list_tools method with server-specific auth headers"""
manager = MCPServerManager()
# Mock servers
server1 = MagicMock()
server1.name = "github"
server1.alias = "github"
server1.server_name = "github"
server2 = MagicMock()
server2.name = "zapier"
server2.alias = "zapier"
server2.server_name = "zapier"
# Mock get_allowed_mcp_servers to return our test servers
manager.get_allowed_mcp_servers = AsyncMock(return_value=["github", "zapier"])
manager.get_mcp_server_by_id = MagicMock(side_effect=lambda x: server1 if x == "github" else server2)
manager.get_mcp_server_by_id = MagicMock(
side_effect=lambda x: server1 if x == "github" else server2
)
# Mock _get_tools_from_server to return different results
async def mock_get_tools_from_server(server, mcp_auth_header=None, mcp_protocol_version=None):
async def mock_get_tools_from_server(
server, mcp_auth_header=None, mcp_protocol_version=None
):
if server.name == "github":
tool1 = MagicMock()
tool1.name = "github_tool_1"
@ -121,20 +123,22 @@ class TestMCPServerManager:
tool1 = MagicMock()
tool1.name = "zapier_tool_1"
return [tool1]
manager._get_tools_from_server = mock_get_tools_from_server
# Test with server-specific auth headers
mcp_server_auth_headers = {
"github": "Bearer github-token",
"zapier": "zapier-api-key"
"zapier": "zapier-api-key",
}
result = await manager.list_tools(mcp_server_auth_headers=mcp_server_auth_headers)
result = await manager.list_tools(
mcp_server_auth_headers=mcp_server_auth_headers
)
# Verify that both servers were called with their specific auth headers
assert len(result) == 3 # 2 from github + 1 from zapier
# Verify the tools have the expected names
tool_names = [tool.name for tool in result]
assert "github_tool_1" in tool_names
@ -145,32 +149,34 @@ class TestMCPServerManager:
async def test_list_tools_fallback_to_legacy_auth_header(self):
"""Test that list_tools falls back to legacy auth header when server-specific not available"""
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.name = "github"
server.alias = "github"
server.server_name = "github"
# Mock get_allowed_mcp_servers
manager.get_allowed_mcp_servers = AsyncMock(return_value=["github"])
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock _get_tools_from_server
async def mock_get_tools_from_server(server, mcp_auth_header=None, mcp_protocol_version=None):
async def mock_get_tools_from_server(
server, mcp_auth_header=None, mcp_protocol_version=None
):
assert mcp_auth_header == "legacy-token" # Should use legacy header
tool = MagicMock()
tool.name = "github_tool_1"
return [tool]
manager._get_tools_from_server = mock_get_tools_from_server
# Test with only legacy auth header (no server-specific headers)
result = await manager.list_tools(
mcp_auth_header="legacy-token",
mcp_server_auth_headers={} # Empty server-specific headers
mcp_server_auth_headers={}, # Empty server-specific headers
)
assert len(result) == 1
assert result[0].name == "github_tool_1"
@ -178,32 +184,36 @@ class TestMCPServerManager:
async def test_list_tools_prioritizes_server_specific_over_legacy(self):
"""Test that server-specific auth headers take priority over legacy header"""
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.name = "github"
server.alias = "github"
server.server_name = "github"
# Mock get_allowed_mcp_servers
manager.get_allowed_mcp_servers = AsyncMock(return_value=["github"])
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock _get_tools_from_server
async def mock_get_tools_from_server(server, mcp_auth_header=None, mcp_protocol_version=None):
assert mcp_auth_header == "server-specific-token" # Should use server-specific header
async def mock_get_tools_from_server(
server, mcp_auth_header=None, mcp_protocol_version=None
):
assert (
mcp_auth_header == "server-specific-token"
) # Should use server-specific header
tool = MagicMock()
tool.name = "github_tool_1"
return [tool]
manager._get_tools_from_server = mock_get_tools_from_server
# Test with both legacy and server-specific headers
result = await manager.list_tools(
mcp_auth_header="legacy-token",
mcp_server_auth_headers={"github": "server-specific-token"}
mcp_server_auth_headers={"github": "server-specific-token"},
)
assert len(result) == 1
assert result[0].name == "github_tool_1"
@ -211,32 +221,36 @@ class TestMCPServerManager:
async def test_list_tools_handles_missing_server_alias(self):
"""Test that list_tools handles servers without alias gracefully"""
manager = MCPServerManager()
# Mock server without alias
server = MagicMock()
server.name = "github"
server.alias = None # No alias
server.server_name = "github"
# Mock get_allowed_mcp_servers
manager.get_allowed_mcp_servers = AsyncMock(return_value=["github"])
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock _get_tools_from_server
async def mock_get_tools_from_server(server, mcp_auth_header=None, mcp_protocol_version=None):
assert mcp_auth_header == "server-specific-token" # Should use server-specific header via server_name
async def mock_get_tools_from_server(
server, mcp_auth_header=None, mcp_protocol_version=None
):
assert (
mcp_auth_header == "server-specific-token"
) # Should use server-specific header via server_name
tool = MagicMock()
tool.name = "github_tool_1"
return [tool]
manager._get_tools_from_server = mock_get_tools_from_server
# Test with server-specific headers that match server_name (even without alias)
result = await manager.list_tools(
mcp_auth_header="legacy-token",
mcp_server_auth_headers={"github": "server-specific-token"}
mcp_server_auth_headers={"github": "server-specific-token"},
)
assert len(result) == 1
assert result[0].name == "github_tool_1"
@ -244,14 +258,14 @@ class TestMCPServerManager:
async def test_health_check_server_healthy(self):
"""Test health check for a healthy server"""
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.server_id = "test-server"
server.name = "test-server"
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock successful _get_tools_from_server
async def mock_get_tools_from_server(server, mcp_auth_header=None):
tool1 = MagicMock()
@ -259,12 +273,12 @@ class TestMCPServerManager:
tool2 = MagicMock()
tool2.name = "tool2"
return [tool1, tool2]
manager._get_tools_from_server = mock_get_tools_from_server
# Perform health check
result = await manager.health_check_server("test-server")
# Verify results
assert result["server_id"] == "test-server"
assert result["status"] == "healthy"
@ -278,23 +292,23 @@ class TestMCPServerManager:
async def test_health_check_server_unhealthy(self):
"""Test health check for an unhealthy server"""
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.server_id = "test-server"
server.name = "test-server"
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock failed _get_tools_from_server
async def mock_get_tools_from_server(server, mcp_auth_header=None):
raise Exception("Connection timeout")
manager._get_tools_from_server = mock_get_tools_from_server
# Perform health check
result = await manager.health_check_server("test-server")
# Verify results
assert result["server_id"] == "test-server"
assert result["status"] == "unhealthy"
@ -307,13 +321,13 @@ class TestMCPServerManager:
async def test_health_check_server_not_found(self):
"""Test health check for a server that doesn't exist"""
manager = MCPServerManager()
# Mock server not found
manager.get_mcp_server_by_id = MagicMock(return_value=None)
# Perform health check
result = await manager.health_check_server("non-existent-server")
# Verify results
assert result["server_id"] == "non-existent-server"
assert result["status"] == "unknown"
@ -325,22 +339,19 @@ class TestMCPServerManager:
async def test_health_check_all_servers(self):
"""Test health check for all servers"""
manager = MCPServerManager()
# Mock servers
server1 = MagicMock()
server1.server_id = "server1"
server1.name = "server1"
server2 = MagicMock()
server2.server_id = "server2"
server2.name = "server2"
# Mock registry
manager.registry = {
"server1": server1,
"server2": server2
}
manager.registry = {"server1": server1, "server2": server2}
# Mock get_mcp_server_by_id
def mock_get_server_by_id(server_id):
if server_id == "server1":
@ -348,9 +359,9 @@ class TestMCPServerManager:
elif server_id == "server2":
return server2
return None
manager.get_mcp_server_by_id = mock_get_server_by_id
# Mock _get_tools_from_server with different results
async def mock_get_tools_from_server(server, mcp_auth_header=None):
if server.server_id == "server1":
@ -360,22 +371,22 @@ class TestMCPServerManager:
elif server.server_id == "server2":
raise Exception("Connection failed")
return []
manager._get_tools_from_server = mock_get_tools_from_server
# Perform health check for all servers
result = await manager.health_check_all_servers()
# Verify results
assert len(result) == 2
assert "server1" in result
assert "server2" in result
# Check server1 (healthy)
assert result["server1"]["status"] == "healthy"
assert result["server1"]["tools_count"] == 1
assert result["server1"]["error"] is None
# Check server2 (unhealthy)
assert result["server2"]["status"] == "unhealthy"
assert result["server2"]["error"] == "Connection failed"
@ -384,26 +395,26 @@ class TestMCPServerManager:
async def test_health_check_server_with_auth_header(self):
"""Test health check with authentication header"""
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.server_id = "test-server"
server.name = "test-server"
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock _get_tools_from_server to verify auth header is passed
async def mock_get_tools_from_server(server, mcp_auth_header=None):
assert mcp_auth_header == "test-token"
tool = MagicMock()
tool.name = "tool1"
return [tool]
manager._get_tools_from_server = mock_get_tools_from_server
# Perform health check with auth header
result = await manager.health_check_server("test-server", "test-token")
# Verify results
assert result["server_id"] == "test-server"
assert result["status"] == "healthy"
@ -411,4 +422,4 @@ class TestMCPServerManager:
if __name__ == "__main__":
pytest.main([__file__])
pytest.main([__file__])

View file

@ -1,12 +1,13 @@
"""
Unit tests for Bedrock Guardrails
"""
import json
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
sys.path.insert(0, os.path.abspath("../../../../../.."))
@ -860,7 +861,6 @@ async def test__redact_pii_matches_comprehensive_coverage():
print("Comprehensive coverage redaction test passed")
@pytest.mark.asyncio
async def test_bedrock_guardrail_respects_custom_runtime_endpoint(monkeypatch):
"""Test that BedrockGuardrail respects aws_bedrock_runtime_endpoint when set"""

View file

@ -18,7 +18,6 @@ from typing import Optional
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
LitellmUserRoles,
MCPSpecVersion,
MCPTransport,
UserAPIKeyAuth,
)
@ -31,7 +30,6 @@ def generate_mock_mcp_server_db_record(
alias: str = "Test DB Server",
url: str = "https://db-server.example.com/mcp",
transport: str = "sse",
spec_version: str = "2025-03-26",
auth_type: Optional[str] = None,
) -> LiteLLM_MCPServerTable:
"""Generate a mock MCP server record from database"""
@ -41,11 +39,6 @@ def generate_mock_mcp_server_db_record(
alias=alias,
url=url,
transport=MCPTransport.sse if transport == "sse" else MCPTransport.http,
spec_version=(
MCPSpecVersion.mar_2025
if spec_version == "2025-03-26"
else MCPSpecVersion.nov_2024
),
auth_type=MCPAuth.api_key if auth_type == "api_key" else None,
created_at=now,
updated_at=now,
@ -59,7 +52,6 @@ def generate_mock_mcp_server_config_record(
name: str = "Test Config Server",
url: str = "https://config-server.example.com/mcp",
transport: str = "http",
spec_version: str = "2025-03-26",
auth_type: Optional[str] = None,
) -> MCPServer:
"""Generate a mock MCP server record from config.yaml"""
@ -70,11 +62,6 @@ def generate_mock_mcp_server_config_record(
server_name=name,
url=url,
transport=MCPTransport.http if transport == "http" else MCPTransport.sse,
spec_version=(
MCPSpecVersion.mar_2025
if spec_version == "2025-03-26"
else MCPSpecVersion.nov_2024
),
auth_type=MCPAuth.api_key if auth_type == "api_key" else None,
mcp_info=MCPInfo(
server_name=name,
@ -98,24 +85,34 @@ def generate_mock_user_api_key_auth(
)
def generate_mock_team_record(team_id: str, team_alias: str, organization_id: str, mcp_servers: List[str]):
def generate_mock_team_record(
team_id: str, team_alias: str, organization_id: str, mcp_servers: List[str]
):
"""Generate a mock team record with object permissions"""
return MagicMock(
team_id=team_id,
team_alias=team_alias,
organization_id=organization_id,
members_with_roles=[{"user_id": "test_user_id"}],
object_permission=MagicMock(mcp_servers=mcp_servers)
object_permission=MagicMock(mcp_servers=mcp_servers),
)
def setup_mock_prisma_client(mock_prisma_client: MagicMock, team_records: List[MagicMock], mcp_servers: List[LiteLLM_MCPServerTable]):
def setup_mock_prisma_client(
mock_prisma_client: MagicMock,
team_records: List[MagicMock],
mcp_servers: List[LiteLLM_MCPServerTable],
):
"""Helper to set up a mock prisma client with proper async behavior"""
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_teamtable = AsyncMock()
mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=team_records)
mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(
return_value=team_records
)
mock_prisma_client.db.litellm_mcpservertable = AsyncMock()
mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=mcp_servers)
mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=mcp_servers
)
return mock_prisma_client
@ -126,7 +123,7 @@ class TestListMCPServers:
async def test_list_mcp_servers_config_yaml_only(self):
"""
Test 1: Returns MCPs defined on the config.yaml only
Scenario: No DB MCPs, only config.yaml MCPs
Expected: Should return only config.yaml MCPs
"""
@ -139,13 +136,13 @@ class TestListMCPServers:
team_id="team1",
team_alias="Team 1",
organization_id="org1",
mcp_servers=["config_server_1", "config_server_2"]
mcp_servers=["config_server_1", "config_server_2"],
)
],
mcp_servers=[] # No DB servers in this test
mcp_servers=[], # No DB servers in this test
)
mock_user_auth = generate_mock_user_api_key_auth()
# Mock config MCPs
config_server_1 = generate_mock_mcp_server_config_record(
server_id="config_server_1",
@ -159,7 +156,7 @@ class TestListMCPServers:
url="https://mcp.deepwiki.com/mcp",
transport="http",
)
# Mock global MCP server manager
mock_manager = MagicMock()
mock_manager.config_mcp_servers = {
@ -169,7 +166,7 @@ class TestListMCPServers:
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["config_server_1", "config_server_2"]
)
# Mock the new method that returns servers with health and team data
mock_servers_with_health = [
generate_mock_mcp_server_db_record(
@ -183,12 +180,12 @@ class TestListMCPServers:
alias="DeepWiki MCP",
url="https://mcp.deepwiki.com/mcp",
transport="http",
)
),
]
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(
return_value=mock_servers_with_health
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
@ -199,22 +196,21 @@ class TestListMCPServers:
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_all_mcp_servers,
)
result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth)
# Verify results
assert len(result) == 2
# Check that both config servers are returned
server_ids = [server.server_id for server in result]
assert "config_server_1" in server_ids
assert "config_server_2" in server_ids
# Check server details
for server in result:
if server.server_id == "config_server_1":
@ -230,7 +226,7 @@ class TestListMCPServers:
async def test_list_mcp_servers_combined_config_and_db(self):
"""
Test 2: If both config.yaml and DB then combines both and returns the result
Scenario: Both DB and config.yaml have MCPs
Expected: Should return combined list from both sources without duplicates
"""
@ -247,7 +243,7 @@ class TestListMCPServers:
url="https://slack-mcp.example.com/mcp",
transport="http",
)
# Mock dependencies
mock_prisma_client = MagicMock()
mock_prisma_client = setup_mock_prisma_client(
@ -257,13 +253,18 @@ class TestListMCPServers:
team_id="team1",
team_alias="Team 1",
organization_id="org1",
mcp_servers=["db_server_1", "db_server_2", "config_server_1", "config_server_2"]
mcp_servers=[
"db_server_1",
"db_server_2",
"config_server_1",
"config_server_2",
],
)
],
mcp_servers=[db_server_1, db_server_2] # DB servers for this test
mcp_servers=[db_server_1, db_server_2], # DB servers for this test
)
mock_user_auth = generate_mock_user_api_key_auth()
# Mock config MCPs
config_server_1 = generate_mock_mcp_server_config_record(
server_id="config_server_1",
@ -277,7 +278,7 @@ class TestListMCPServers:
url="https://mcp.deepwiki.com/mcp",
transport="http",
)
# Mock global MCP server manager
mock_manager = MagicMock()
mock_manager.config_mcp_servers = {
@ -292,7 +293,7 @@ class TestListMCPServers:
"config_server_2",
]
)
# Mock the new method that returns servers with health and team data
mock_servers_with_health = [
db_server_1,
@ -308,12 +309,12 @@ class TestListMCPServers:
alias="DeepWiki MCP",
url="https://mcp.deepwiki.com/mcp",
transport="http",
)
),
]
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(
return_value=mock_servers_with_health
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
@ -324,24 +325,23 @@ class TestListMCPServers:
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_all_mcp_servers,
)
result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth)
# Verify results
assert len(result) == 4
# Check that both DB and config servers are returned
server_ids = [server.server_id for server in result]
assert "db_server_1" in server_ids
assert "db_server_2" in server_ids
assert "config_server_1" in server_ids
assert "config_server_2" in server_ids
# Check server details
for server in result:
if server.server_id == "db_server_1":
@ -365,7 +365,7 @@ class TestListMCPServers:
async def test_list_mcp_servers_non_admin_user_filtered(self):
"""
Test 3: Non-admin users only see MCPs they have access to
Scenario: Non-admin user with limited access
Expected: Should return only MCPs the user has access to
"""
@ -375,7 +375,7 @@ class TestListMCPServers:
alias="Allowed Gmail MCP",
url="https://gmail-mcp.example.com/mcp",
)
# Mock dependencies
mock_prisma_client = MagicMock()
mock_prisma_client = setup_mock_prisma_client(
@ -385,23 +385,23 @@ class TestListMCPServers:
team_id="team1",
team_alias="Team 1",
organization_id="org1",
mcp_servers=["db_server_allowed", "config_server_allowed"]
mcp_servers=["db_server_allowed", "config_server_allowed"],
)
],
mcp_servers=[db_server_allowed] # Only the allowed DB server
mcp_servers=[db_server_allowed], # Only the allowed DB server
)
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.INTERNAL_USER, # Non-admin user
team_id="team_123",
)
# Mock config MCPs - user has access to one
config_server_allowed = generate_mock_mcp_server_config_record(
server_id="config_server_allowed",
name="Allowed Zapier MCP",
url="https://actions.zapier.com/mcp/sse",
)
# Mock global MCP server manager
mock_manager = MagicMock()
mock_manager.config_mcp_servers = {
@ -414,7 +414,7 @@ class TestListMCPServers:
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["db_server_allowed", "config_server_allowed"]
)
# Mock the new method that returns servers with health and team data
mock_servers_with_health = [
db_server_allowed,
@ -422,12 +422,12 @@ class TestListMCPServers:
server_id="config_server_allowed",
alias="Allowed Zapier MCP",
url="https://actions.zapier.com/mcp/sse",
)
),
]
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(
return_value=mock_servers_with_health
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
@ -438,23 +438,22 @@ class TestListMCPServers:
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_all_mcp_servers,
)
result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth)
# Verify results - should only return servers user has access to
assert len(result) == 2
# Check that only allowed servers are returned
server_ids = [server.server_id for server in result]
assert "db_server_allowed" in server_ids
assert "config_server_allowed" in server_ids
assert "config_server_not_allowed" not in server_ids
# Check server details
for server in result:
if server.server_id == "db_server_allowed":
@ -473,28 +472,29 @@ class TestMCPHealthCheckEndpoints:
"""Test successful health check for a specific MCP server"""
# Mock server
mock_server = generate_mock_mcp_server_db_record(
server_id="test-server",
alias="Test Server"
server_id="test-server", alias="Test Server"
)
# Mock dependencies
mock_prisma_client = MagicMock()
# Mock global MCP server manager
mock_manager = MagicMock()
mock_manager.health_check_server = AsyncMock(return_value={
"server_id": "test-server",
"status": "healthy",
"tools_count": 3,
"last_health_check": "2024-01-01T12:00:00",
"response_time_ms": 150.5,
"error": None
})
mock_manager.health_check_server = AsyncMock(
return_value={
"server_id": "test-server",
"status": "healthy",
"tools_count": 3,
"last_health_check": "2024-01-01T12:00:00",
"response_time_ms": 150.5,
"error": None,
}
)
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
@ -508,17 +508,15 @@ class TestMCPHealthCheckEndpoints:
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
AsyncMock(return_value=mock_server),
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
health_check_mcp_server,
)
result = await health_check_mcp_server(
server_id="test-server",
user_api_key_dict=mock_user_auth
server_id="test-server", user_api_key_dict=mock_user_auth
)
# Verify results
assert result["server_id"] == "test-server"
assert result["status"] == "healthy"
@ -531,11 +529,11 @@ class TestMCPHealthCheckEndpoints:
"""Test health check for a server that doesn't exist"""
# Mock dependencies
mock_prisma_client = MagicMock()
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
@ -543,19 +541,17 @@ class TestMCPHealthCheckEndpoints:
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
AsyncMock(return_value=None),
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
health_check_mcp_server,
)
# Should raise HTTPException
with pytest.raises(Exception) as exc_info:
await health_check_mcp_server(
server_id="non-existent-server",
user_api_key_dict=mock_user_auth
server_id="non-existent-server", user_api_key_dict=mock_user_auth
)
assert "not found" in str(exc_info.value)
@pytest.mark.asyncio
@ -563,20 +559,19 @@ class TestMCPHealthCheckEndpoints:
"""Test health check for a server user doesn't have access to"""
# Mock server
mock_server = generate_mock_mcp_server_db_record(
server_id="test-server",
alias="Test Server"
server_id="test-server", alias="Test Server"
)
# Mock dependencies
mock_prisma_client = MagicMock()
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.INTERNAL_USER # Non-admin user
)
# Mock user doesn't have access to this server
mock_user_servers = []
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
@ -590,19 +585,17 @@ class TestMCPHealthCheckEndpoints:
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
AsyncMock(return_value=mock_server),
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
health_check_mcp_server,
)
# Should raise HTTPException
with pytest.raises(Exception) as exc_info:
await health_check_mcp_server(
server_id="test-server",
user_api_key_dict=mock_user_auth
server_id="test-server", user_api_key_dict=mock_user_auth
)
assert "permission" in str(exc_info.value)
@pytest.mark.asyncio
@ -614,49 +607,53 @@ class TestMCPHealthCheckEndpoints:
team_id="team1",
team_alias="Team 1",
organization_id="org1",
mcp_servers=["server1", "server2"]
mcp_servers=["server1", "server2"],
)
]
# Mock DB servers
db_servers = [
generate_mock_mcp_server_db_record(server_id="server1"),
generate_mock_mcp_server_db_record(server_id="server2")
generate_mock_mcp_server_db_record(server_id="server2"),
]
# Mock dependencies
mock_prisma_client = MagicMock()
mock_prisma_client = setup_mock_prisma_client(
mock_prisma_client=mock_prisma_client,
team_records=team_records,
mcp_servers=db_servers
mcp_servers=db_servers,
)
# Mock global MCP server manager
mock_manager = MagicMock()
mock_manager.health_check_allowed_servers = AsyncMock(return_value={
"server1": {
"server_id": "server1",
"status": "healthy",
"tools_count": 2,
"last_health_check": "2024-01-01T12:00:00",
"response_time_ms": 100.0,
"error": None
},
"server2": {
"server_id": "server2",
"status": "unhealthy",
"last_health_check": "2024-01-01T12:00:00",
"response_time_ms": 5000.0,
"error": "Connection timeout"
mock_manager.health_check_allowed_servers = AsyncMock(
return_value={
"server1": {
"server_id": "server1",
"status": "healthy",
"tools_count": 2,
"last_health_check": "2024-01-01T12:00:00",
"response_time_ms": 100.0,
"error": None,
},
"server2": {
"server_id": "server2",
"status": "unhealthy",
"last_health_check": "2024-01-01T12:00:00",
"response_time_ms": 5000.0,
"error": "Connection timeout",
},
}
})
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1", "server2"])
)
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["server1", "server2"]
)
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.INTERNAL_USER
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
@ -667,14 +664,15 @@ class TestMCPHealthCheckEndpoints:
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
health_check_all_mcp_servers,
)
result = await health_check_all_mcp_servers(user_api_key_dict=mock_user_auth)
result = await health_check_all_mcp_servers(
user_api_key_dict=mock_user_auth
)
# Verify results
assert result["total_servers"] == 2
assert result["healthy_count"] == 1
@ -682,7 +680,7 @@ class TestMCPHealthCheckEndpoints:
assert result["unknown_count"] == 0
assert "server1" in result["servers"]
assert "server2" in result["servers"]
# Check individual server results
assert result["servers"]["server1"]["status"] == "healthy"
assert result["servers"]["server1"]["tools_count"] == 2
@ -694,32 +692,33 @@ class TestMCPHealthCheckEndpoints:
"""Test that fetch_all_mcp_servers includes health check status"""
# Mock server with health status
mock_server = generate_mock_mcp_server_db_record(
server_id="test-server",
alias="Test Server"
server_id="test-server", alias="Test Server"
)
# Add health status to the mock server
mock_server.status = "healthy"
mock_server.last_health_check = datetime.now()
mock_server.health_check_error = None
# Mock dependencies
mock_prisma_client = MagicMock()
mock_prisma_client = setup_mock_prisma_client(
mock_prisma_client=mock_prisma_client,
team_records=[],
mcp_servers=[] # Don't add servers here since we're mocking get_all_mcp_servers
mcp_servers=[], # Don't add servers here since we're mocking get_all_mcp_servers
)
# Mock global MCP server manager
mock_manager = MagicMock()
mock_manager.config_mcp_servers = {}
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=[])
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(return_value=[mock_server])
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(
return_value=[mock_server]
)
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
@ -730,14 +729,13 @@ class TestMCPHealthCheckEndpoints:
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_all_mcp_servers,
)
result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth)
# Verify health check status is included
assert len(result) == 1
server = result[0]

View file

@ -168,6 +168,7 @@ class TestResponseAPILoggingUtils:
"input_tokens": 10,
"output_tokens": 20,
"total_tokens": 30,
"input_tokens_details": {"cached_tokens": 2},
"output_tokens_details": {"reasoning_tokens": 5},
}
@ -181,6 +182,7 @@ class TestResponseAPILoggingUtils:
assert result.prompt_tokens == 10
assert result.completion_tokens == 20
assert result.total_tokens == 30
assert result.prompt_tokens_details and result.prompt_tokens_details.cached_tokens == 2
def test_transform_response_api_usage_with_none_values(self):
"""Test transformation handles None values properly"""

View file

@ -11,11 +11,11 @@ if ! command -v nvm &> /dev/null; then
fi
# Use nvm to set the required Node.js version
nvm use v18.17.0
nvm use v20
# Check if nvm use was successful
if [ $? -ne 0 ]; then
echo "Error: Failed to switch to Node.js v18.17.0. Deployment aborted."
echo "Error: Failed to switch to Node.js v20. Deployment aborted."
exit 1
fi

File diff suppressed because it is too large Load diff

View file

@ -6,7 +6,9 @@
"dev": "next dev",
"build": "next build",
"start": "next start",
"lint": "next lint"
"lint": "next lint",
"test": "vitest",
"test:watch": "vitest -w"
},
"dependencies": {
"@anthropic-ai/sdk": "^0.54.0",
@ -40,6 +42,10 @@
},
"devDependencies": {
"@tailwindcss/forms": "^0.5.7",
"@testing-library/jest-dom": "^6.8.0",
"@testing-library/react": "^16.3.0",
"@testing-library/user-event": "^14.6.1",
"@types/babel__traverse": "^7.28.0",
"@types/lodash": "^4.17.15",
"@types/node": "^20",
"@types/react": "18.2.48",
@ -47,13 +53,18 @@
"@types/react-dom": "^18",
"@types/react-syntax-highlighter": "^15.5.11",
"@types/uuid": "^10.0.0",
"@vitest/coverage-v8": "^3.2.4",
"@vitest/ui": "^3.2.4",
"autoprefixer": "^10.4.17",
"eslint": "^8",
"eslint-config-next": "14.2.32",
"jsdom": "^27.0.0",
"postcss": "^8.4.33",
"prettier": "3.2.5",
"tailwindcss": "^3.4.1",
"typescript": "5.3.3"
"typescript": "5.3.3",
"vite": "^5.4.20",
"vitest": "^3.2.4"
},
"overrides": {
"prismjs": ">=1.30.0",

View file

@ -48,7 +48,6 @@ const MCPConnectionTest: React.FC<MCPConnectionTestProps> = ({
alias: formValues.alias || "",
url: formValues.url,
transport: formValues.transport,
spec_version: formValues.spec_version,
auth_type: formValues.auth_type,
mcp_info: formValues.mcp_info,
};

View file

@ -369,25 +369,6 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
{/* Stdio Configuration - only show for stdio transport */}
<StdioConfiguration isVisible={transportType === "stdio"} />
<Form.Item
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
MCP Version
<Tooltip title="Select the MCP specification version your server supports">
<InfoCircleOutlined className="ml-2 text-gray-400 hover:text-gray-600" />
</Tooltip>
</span>
}
name="spec_version"
rules={[{ required: true, message: "Please select a spec version" }]}
>
<Select placeholder="Select MCP version" className="rounded-lg" size="large">
<Select.Option value="2025-06-18">2025-06-18 (Latest)</Select.Option>
<Select.Option value="2025-03-26">2025-03-26</Select.Option>
<Select.Option value="2024-11-05">2024-11-05</Select.Option>
</Select>
</Form.Item>
<Form.Item
label={
<span className="text-sm font-medium text-gray-700 flex items-center">

View file

@ -40,7 +40,6 @@ const MCPConnectionStatus: React.FC<MCPConnectionStatusProps> = ({
server_name: formValues.server_name || "",
url: formValues.url,
transport: formValues.transport,
spec_version: formValues.spec_version,
auth_type: formValues.auth_type,
mcp_info: formValues.mcp_info,
};
@ -83,7 +82,7 @@ const MCPConnectionStatus: React.FC<MCPConnectionStatusProps> = ({
setHasShownSuccessMessage(false);
onToolsLoaded?.([]);
}
}, [formValues.url, formValues.transport, formValues.auth_type, formValues.spec_version, accessToken]);
}, [formValues.url, formValues.transport, formValues.auth_type, accessToken]);
// Don't show anything if required fields aren't filled
if (!canFetchTools && !formValues.url) {

View file

@ -59,7 +59,6 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({ mcpServer, accessToken, o
server_name: mcpServer.server_name,
url: mcpServer.url,
transport: mcpServer.transport,
spec_version: mcpServer.spec_version,
auth_type: mcpServer.auth_type,
mcp_info: mcpServer.mcp_info,
};
@ -179,13 +178,6 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({ mcpServer, accessToken, o
<Select.Option value="basic">Basic Auth</Select.Option>
</Select>
</Form.Item>
<Form.Item label="MCP Version" name="spec_version" rules={[{ required: true }]}>
<Select>
<Select.Option value="2025-06-18">2025-06-18 (Latest)</Select.Option>
<Select.Option value="2025-03-26">2025-03-26</Select.Option>
<Select.Option value="2024-11-05">2024-11-05</Select.Option>
</Select>
</Form.Item>
<Form.Item
label={

View file

@ -225,10 +225,6 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
<Text className="font-medium">Auth Type</Text>
<div>{handleAuth(mcpServer.auth_type)}</div>
</div>
<div>
<Text className="font-medium">Spec Version</Text>
<div>{mcpServer.spec_version}</div>
</div>
<div>
<Text className="font-medium">Access Groups</Text>
<div>

View file

@ -185,7 +185,6 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
alias: "",
url: "",
transport: "",
spec_version: "",
auth_type: "",
created_at: "",
created_by: "",

View file

@ -132,7 +132,6 @@ export interface MCPServer {
description?: string | null;
url: string;
transport?: string | null;
spec_version?: string | null;
auth_type?: string | null;
mcp_info?: MCPInfo | null;
created_at: string;

View file

@ -1,257 +0,0 @@
import React, { useState } from "react";
interface RecognitionMetadata {
recognizer_name: string;
recognizer_identifier: string;
}
interface GuardrailEntity {
end: number;
score: number;
start: number;
entity_type: string;
analysis_explanation: string | null;
recognition_metadata: RecognitionMetadata;
}
interface MaskedEntityCount {
[key: string]: number;
}
interface GuardrailInformation {
duration: number;
end_time: number;
start_time: number;
guardrail_mode: string;
guardrail_name: string;
guardrail_status: string;
guardrail_response: GuardrailEntity[];
masked_entity_count: MaskedEntityCount;
}
interface GuardrailViewerProps {
data: GuardrailInformation;
}
export function GuardrailViewer({ data }: GuardrailViewerProps) {
const [sectionExpanded, setSectionExpanded] = useState(true);
const [entityListExpanded, setEntityListExpanded] = useState(true);
const [expandedEntities, setExpandedEntities] = useState<Record<number, boolean>>({});
if (!data) {
return null;
}
// Calculate total masked entities
const totalMaskedEntities = data.masked_entity_count ?
Object.values(data.masked_entity_count).reduce((sum, count) => sum + count, 0) : 0;
const formatTime = (timestamp: number): string => {
const date = new Date(timestamp * 1000);
return date.toLocaleString();
};
const toggleEntity = (index: number) => {
setExpandedEntities(prev => ({
...prev,
[index]: !prev[index]
}));
};
const getScoreColor = (score: number): string => {
if (score >= 0.8) return "text-green-600";
return "text-yellow-600";
};
return (
<div className="bg-white rounded-lg shadow mb-6">
<div
className="flex justify-between items-center p-4 border-b cursor-pointer hover:bg-gray-50"
onClick={() => setSectionExpanded(!sectionExpanded)}
>
<div className="flex items-center">
<svg
className={`w-5 h-5 mr-2 text-gray-600 transition-transform ${sectionExpanded ? 'transform rotate-90' : ''}`}
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M9 5l7 7-7 7" />
</svg>
<h3 className="text-lg font-medium">Guardrail Information</h3>
<span className={`ml-3 px-2 py-1 rounded-md text-xs font-medium inline-block ${
data.guardrail_status === "success"
? 'bg-green-100 text-green-800'
: 'bg-red-100 text-red-800'
}`}>
{data.guardrail_status}
</span>
{totalMaskedEntities > 0 && (
<span className="ml-3 px-2 py-1 bg-blue-50 text-blue-700 rounded-md text-xs font-medium">
{totalMaskedEntities} masked {totalMaskedEntities === 1 ? 'entity' : 'entities'}
</span>
)}
</div>
<span className="text-sm text-gray-500">{sectionExpanded ? 'Click to collapse' : 'Click to expand'}</span>
</div>
{sectionExpanded && (
<div className="p-4">
<div className="bg-white rounded-lg border p-4 mb-4">
<div className="grid grid-cols-2 gap-4">
<div className="space-y-2">
<div className="flex">
<span className="font-medium w-1/3">Guardrail Name:</span>
<span className="font-mono">{data.guardrail_name}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Mode:</span>
<span className="font-mono">{data.guardrail_mode}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Status:</span>
<span className={`px-2 py-1 rounded-md text-xs font-medium inline-block ${
data.guardrail_status === "success"
? 'bg-green-100 text-green-800'
: 'bg-red-100 text-red-800'
}`}>
{data.guardrail_status}
</span>
</div>
</div>
<div className="space-y-2">
<div className="flex">
<span className="font-medium w-1/3">Start Time:</span>
<span>{formatTime(data.start_time)}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">End Time:</span>
<span>{formatTime(data.end_time)}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Duration:</span>
<span>{data.duration.toFixed(4)}s</span>
</div>
</div>
</div>
{/* Masked Entity Summary */}
{data.masked_entity_count && Object.keys(data.masked_entity_count).length > 0 && (
<div className="mt-4 pt-4 border-t">
<h4 className="font-medium mb-2">Masked Entity Summary</h4>
<div className="flex flex-wrap gap-2">
{Object.entries(data.masked_entity_count).map(([entityType, count]) => (
<span key={entityType} className="px-3 py-1.5 bg-blue-50 text-blue-700 rounded-md text-xs font-medium">
{entityType}: {count}
</span>
))}
</div>
</div>
)}
</div>
{/* Detected Entities Section */}
{data.guardrail_response && data.guardrail_response.length > 0 && (
<div className="mt-4">
<div
className="flex items-center mb-2 cursor-pointer"
onClick={() => setEntityListExpanded(!entityListExpanded)}
>
<svg
className={`w-5 h-5 mr-2 transition-transform ${entityListExpanded ? 'transform rotate-90' : ''}`}
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M9 5l7 7-7 7" />
</svg>
<h4 className="font-medium">Detected Entities ({data.guardrail_response.length})</h4>
</div>
{entityListExpanded && (
<div className="space-y-2">
{data.guardrail_response.map((entity, index) => {
const isExpanded = expandedEntities[index] || false;
return (
<div key={index} className="border rounded-lg overflow-hidden">
<div
className="flex items-center justify-between p-3 bg-gray-50 cursor-pointer hover:bg-gray-100"
onClick={() => toggleEntity(index)}
>
<div className="flex items-center">
<svg
className={`w-5 h-5 mr-2 transition-transform ${isExpanded ? 'transform rotate-90' : ''}`}
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M9 5l7 7-7 7" />
</svg>
<span className="font-medium mr-2">{entity.entity_type}</span>
<span className={`font-mono ${getScoreColor(entity.score)}`}>
Score: {entity.score.toFixed(2)}
</span>
</div>
<span className="text-xs text-gray-500">
Position: {entity.start}-{entity.end}
</span>
</div>
{isExpanded && (
<div className="p-3 border-t bg-white">
<div className="grid grid-cols-2 gap-4 mb-2">
<div className="space-y-2">
<div className="flex">
<span className="font-medium w-1/3">Entity Type:</span>
<span>{entity.entity_type}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Position:</span>
<span>Characters {entity.start}-{entity.end}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Confidence:</span>
<span className={getScoreColor(entity.score)}>
{entity.score.toFixed(2)}
</span>
</div>
</div>
<div className="space-y-2">
{entity.recognition_metadata && (
<>
<div className="flex">
<span className="font-medium w-1/3">Recognizer:</span>
<span>{entity.recognition_metadata.recognizer_name}</span>
</div>
<div className="flex overflow-hidden">
<span className="font-medium w-1/3">Identifier:</span>
<span className="truncate text-xs font-mono">
{entity.recognition_metadata.recognizer_identifier}
</span>
</div>
</>
)}
{entity.analysis_explanation && (
<div className="flex">
<span className="font-medium w-1/3">Explanation:</span>
<span>{entity.analysis_explanation}</span>
</div>
)}
</div>
</div>
</div>
)}
</div>
);
})}
</div>
)}
</div>
)}
</div>
)}
</div>
);
}

View file

@ -0,0 +1,118 @@
import React from 'react';
import { describe, it, expect } from 'vitest';
import BedrockGuardrailDetails, {
BedrockGuardrailResponse,
} from '@/components/view_logs/GuardrailViewer/BedrockGuardrailDetails';
import { renderWithProviders, screen } from "../../../../tests/test-utils"
import {
makeAssessment,
makeBedrockCoverage,
makeBedrockResponse,
makeBedrockUsage,
} from "@/components/view_logs/GuardrailViewer/__tests__/fixtures"
describe('BedrockGuardrailDetails', () => {
it('returns null when response is falsy', () => {
// @ts-expect-error testing nullish handling
const { container } = renderWithProviders(<BedrockGuardrailDetails response={undefined} />);
expect(container).toBeEmptyDOMElement();
});
it('renders top summary: action chip, reason, blocked response', () => {
const resp: BedrockGuardrailResponse = makeBedrockResponse({
action: 'GUARDRAIL_INTERVENED',
actionReason: 'Policy violation',
blockedResponse: '[blocked]',
});
renderWithProviders(<BedrockGuardrailDetails response={resp} />);
expect(screen.getByText('Action:')).toBeInTheDocument();
expect(screen.getByText('Policy violation')).toBeInTheDocument();
expect(screen.getByText('[blocked]')).toBeInTheDocument();
});
it('renders coverage and usage pills', () => {
const resp = makeBedrockResponse({
guardrailCoverage: makeBedrockCoverage(),
usage: makeBedrockUsage({ contentPolicyUnits: 7, wordPolicyUnits: 1 }),
});
renderWithProviders(<BedrockGuardrailDetails response={resp} />);
expect(screen.getByText(/text guarded 27\/100/)).toBeInTheDocument();
expect(screen.getByText(/images guarded 1\/3/)).toBeInTheDocument();
expect(screen.getByText(/contentPolicyUnits: 7/)).toBeInTheDocument();
expect(screen.getByText(/wordPolicyUnits: 1/)).toBeInTheDocument();
});
it('renders outputs when present (prefers `outputs`, falls back to `output`)', () => {
// Using outputs
let resp = makeBedrockResponse({ outputs: [{ text: 'hello' }] });
const { rerender } = renderWithProviders(<BedrockGuardrailDetails response={resp} />);
expect(screen.getByText('Outputs')).toBeInTheDocument();
expect(screen.getByText('hello')).toBeInTheDocument();
// Using output
resp = makeBedrockResponse({ outputs: undefined, output: [{ text: 'world' }] });
rerender(<BedrockGuardrailDetails response={resp} />);
expect(screen.getByText('world')).toBeInTheDocument();
});
it('renders assessments with all policy sections and metrics', () => {
const resp = makeBedrockResponse({
assessments: [makeAssessment()],
});
renderWithProviders(<BedrockGuardrailDetails response={resp} />);
// Assessment section present
expect(screen.getByText('Assessment #1')).toBeInTheDocument();
// Word policy sections
expect(screen.getByText('Word Policy')).toBeInTheDocument();
expect(screen.getByText('Custom Words')).toBeInTheDocument();
expect(screen.getByText('Managed Word Lists')).toBeInTheDocument();
// Contextual grounding table headers
expect(screen.getByText('Contextual Grounding')).toBeInTheDocument();
expect(screen.getAllByText('Score').length).toBeGreaterThan(0);
expect(screen.getAllByText('Threshold').length).toBeGreaterThan(0);
// Sensitive Info sections
expect(screen.getByText('Sensitive Information')).toBeInTheDocument();
expect(screen.getByText('PII Entities')).toBeInTheDocument();
expect(screen.getByText('Custom Regexes')).toBeInTheDocument();
// Topic Policy
expect(screen.getByText('Topic Policy')).toBeInTheDocument();
expect(screen.getByText('weapons')).toBeInTheDocument();
// Invocation Metrics
expect(screen.getByText('Invocation Metrics')).toBeInTheDocument();
// Raw JSON section exists (closed by default)
expect(screen.getByText('Raw Bedrock Guardrail Response')).toBeInTheDocument();
});
it('handles non-text outputs gracefully', () => {
const resp = makeBedrockResponse({ outputs: [{}, { text: 'texty' }] });
renderWithProviders(<BedrockGuardrailDetails response={resp} />);
expect(screen.getByText('(non-text output)')).toBeInTheDocument();
expect(screen.getByText('texty')).toBeInTheDocument();
});
it('gracefully handles missing optional sections', () => {
const resp = makeBedrockResponse({
assessments: [
{
// only include minimal fields; others omitted
invocationMetrics: { guardrailProcessingLatency: 5 },
} as any,
],
usage: undefined,
guardrailCoverage: undefined,
outputs: [],
});
renderWithProviders(<BedrockGuardrailDetails response={resp} />);
// No crash, minimal render: Assessment + Invocation Metrics present, but no usage/coverage chips at top
expect(screen.getByText('Assessment #1')).toBeInTheDocument();
});
});

View file

@ -0,0 +1,508 @@
import React, { useState } from "react";
export type BedrockGuardrailAction = "NONE" | "GUARDRAIL_INTERVENED";
export interface BedrockGuardrailUsage {
wordPolicyUnits?: number;
topicPolicyUnits?: number;
contentPolicyUnits?: number;
contentPolicyImageUnits?: number;
automatedReasoningPolicies?: number;
automatedReasoningPolicyUnits?: number;
contextualGroundingPolicyUnits?: number;
sensitiveInformationPolicyUnits?: number;
sensitiveInformationPolicyFreeUnits?: number;
}
export interface BedrockGuardrailCoverageDim {
total?: number;
guarded?: number;
}
export interface BedrockGuardrailCoverage {
images?: BedrockGuardrailCoverageDim;
textCharacters?: BedrockGuardrailCoverageDim;
}
export interface BedrockWordItem {
action?: string; // e.g., "BLOCKED" or "ALLOWED"
detected?: boolean;
match?: string;
type?: string; // present on managedWordLists entries
}
export interface BedrockContentFilter {
type?: string; // e.g. "HATE", etc.
action?: string; // "BLOCKED" | "NONE" (varies by config)
detected?: boolean;
filterStrength?: string; // e.g., "HIGH"/"MEDIUM"/"LOW"
confidence?: string; // e.g., "HIGH"/"MEDIUM"/"LOW"
}
export interface BedrockContextualGroundingFilter {
type?: string; // "GROUNDING" | "RELEVANCE"
action?: string; // "BLOCKED" | ...
detected?: boolean;
score?: number; // model score
threshold?: number; // configured threshold
}
export interface BedrockPiiEntity {
type?: string; // e.g., "ADDRESS", "EMAIL", etc.
match?: string; // the matched text (may be masked)
detected?: boolean;
action?: string; // "BLOCKED" | "ANONYMIZED" | ...
}
export interface BedrockRegexFinding {
name?: string;
regex?: string;
match?: string;
detected?: boolean;
action?: string;
}
export interface BedrockTopic {
name?: string; // topic list entry name
type?: string; // "DENY" | "ALLOW" depending on list
detected?: boolean;
action?: string; // "BLOCKED" | "NONE"
}
export interface BedrockInvocationMetrics {
usage?: BedrockGuardrailUsage;
guardrailCoverage?: BedrockGuardrailCoverage;
guardrailProcessingLatency?: number; // in ms
}
export interface BedrockAssessment {
wordPolicy?: { customWords?: BedrockWordItem[]; managedWordLists?: BedrockWordItem[] };
contentPolicy?: { filters?: BedrockContentFilter[] };
topicPolicy?: { topics?: BedrockTopic[] };
sensitiveInformationPolicy?: { piiEntities?: BedrockPiiEntity[]; regexes?: BedrockRegexFinding[] };
contextualGroundingPolicy?: { filters?: BedrockContextualGroundingFilter[] };
automatedReasoningPolicy?: { findings?: any[] };
invocationMetrics?: BedrockInvocationMetrics;
}
export interface BedrockOutputContent {
text?: string;
}
export interface BedrockGuardrailResponse {
action?: BedrockGuardrailAction;
actionReason?: string | null;
outputs?: BedrockOutputContent[];
output?: BedrockOutputContent[];
usage?: BedrockGuardrailUsage;
guardrailCoverage?: BedrockGuardrailCoverage;
assessments?: BedrockAssessment[];
blockedResponse?: string;
}
/** ====== UI helpers ====== */
type ChipTone = "green" | "red" | "blue" | "slate" | "amber";
const chip = (
text: React.ReactNode,
tone: ChipTone = "slate"
) => {
const map: Record<ChipTone, string> = {
green: "bg-green-100 text-green-800",
red: "bg-red-100 text-red-800",
blue: "bg-blue-50 text-blue-700",
slate: "bg-slate-100 text-slate-800",
amber: "bg-amber-100 text-amber-800",
};
return <span className={`px-2 py-1 rounded-md text-xs font-medium inline-block ${map[tone]}`}>{text}</span>;
};
const boolPill = (b?: boolean) => (b ? chip("detected", "red") : chip("not detected", "slate"));
interface SectionProps {
title: string;
count?: number;
defaultOpen?: boolean;
right?: React.ReactNode;
children?: React.ReactNode;
}
const Section: React.FC<SectionProps> = ({
title,
count,
defaultOpen = true,
right,
children,
}) => {
const [open, setOpen] = useState(defaultOpen);
return (
<div className="border rounded-lg overflow-hidden">
<div
className="flex items-center justify-between p-3 bg-gray-50 cursor-pointer hover:bg-gray-100"
onClick={() => setOpen((v) => !v)}
>
<div className="flex items-center">
<svg className={`w-5 h-5 mr-2 transition-transform ${open ? "transform rotate-90" : ""}`} fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M9 5l7 7-7 7" />
</svg>
<h5 className="font-medium">
{title} {typeof count === "number" && <span className="text-gray-500 font-normal">({count})</span>}
</h5>
</div>
<div>{right}</div>
</div>
{open && <div className="p-3 border-t bg-white">{children}</div>}
</div>
);
};
interface KVProps {
label: string;
children?: React.ReactNode;
mono?: boolean;
}
const KV: React.FC<KVProps> = ({ label, children, mono }) => (
<div className="flex">
<span className="font-medium w-1/3">{label}</span>
<span className={mono ? "font-mono text-sm break-all" : ""}>{children}</span>
</div>
);
const Divider: React.FC = () => <div className="my-3 border-t" />;
/** ====== Main component ====== */
export const BedrockGuardrailDetails: React.FC<{ response: BedrockGuardrailResponse }> = ({ response }) => {
if (!response) return null;
const outputs: BedrockOutputContent[] = (response.outputs ?? response.output ?? []) as BedrockOutputContent[];
const actionTone: "green" | "red" = response.action === "GUARDRAIL_INTERVENED" ? "red" : "green";
const coverageChips = (
<div className="flex flex-wrap gap-2">
{response.guardrailCoverage?.textCharacters && (
chip(
`text guarded ${response.guardrailCoverage.textCharacters.guarded ?? 0}/${response.guardrailCoverage.textCharacters.total ?? 0}`,
"blue"
)
)}
{response.guardrailCoverage?.images && (
chip(
`images guarded ${response.guardrailCoverage.images.guarded ?? 0}/${response.guardrailCoverage.images.total ?? 0}`,
"blue"
)
)}
</div>
);
const usagePills = response.usage && (
<div className="flex flex-wrap gap-2">
{Object.entries(response.usage).map(([k, v]) =>
typeof v === "number" ? (
<span key={k} className="px-2 py-1 bg-slate-100 text-slate-800 rounded-md text-xs font-medium">
{k}: {v}
</span>
) : null
)}
</div>
);
return (
<div className="space-y-3">
{/* Top summary card */}
<div className="border rounded-lg p-4">
<div className="grid grid-cols-2 gap-4">
<div className="space-y-2">
<KV label="Action:">
{chip(response.action ?? "N/A", actionTone)}
</KV>
{response.actionReason && <KV label="Action Reason:">{response.actionReason}</KV>}
{response.blockedResponse && (
<KV label="Blocked Response:">
<span className="italic">{response.blockedResponse}</span>
</KV>
)}
</div>
<div className="space-y-2">
<KV label="Coverage:">{coverageChips}</KV>
<KV label="Usage:">{usagePills}</KV>
</div>
</div>
{/* Outputs */}
{outputs.length > 0 && (
<>
<Divider />
<h4 className="font-medium mb-2">Outputs</h4>
<div className="space-y-2">
{outputs.map((o, i) => (
<div key={i} className="p-3 bg-gray-50 rounded-md">
<div className="text-sm whitespace-pre-wrap">{o.text ?? <em>(non-text output)</em>}</div>
</div>
))}
</div>
</>
)}
</div>
{/* Assessments */}
{response.assessments?.length ? (
<div className="space-y-3">
{response.assessments.map((assess, idx) => {
const policyBadges = (
<div className="flex flex-wrap gap-1">
{assess.wordPolicy && chip("word", "slate")}
{assess.contentPolicy && chip("content", "slate")}
{assess.topicPolicy && chip("topic", "slate")}
{assess.sensitiveInformationPolicy && chip("sensitive-info", "slate")}
{assess.contextualGroundingPolicy && chip("contextual-grounding", "slate")}
{assess.automatedReasoningPolicy && chip("automated-reasoning", "slate")}
</div>
);
return (
<Section
key={idx}
title={`Assessment #${idx + 1}`}
defaultOpen
right={
<div className="flex items-center gap-3">
{assess.invocationMetrics?.guardrailProcessingLatency != null &&
chip(`${assess.invocationMetrics.guardrailProcessingLatency} ms`, "amber")}
{policyBadges}
</div>
}
>
{/* Word policy */}
{assess.wordPolicy && (
<div className="mb-3">
<h6 className="font-medium mb-2">Word Policy</h6>
{(assess.wordPolicy.customWords?.length ?? 0) > 0 && (
<Section title="Custom Words" defaultOpen>
<div className="space-y-2">
{assess.wordPolicy.customWords!.map((w, i) => (
<div key={i} className="flex justify-between items-center p-2 bg-gray-50 rounded">
<div className="flex items-center gap-2">
{chip(w.action ?? "N/A", w.detected ? "red" : "slate")}
<span className="font-mono text-sm break-all">{w.match}</span>
</div>
{boolPill(w.detected)}
</div>
))}
</div>
</Section>
)}
{(assess.wordPolicy.managedWordLists?.length ?? 0) > 0 && (
<Section title="Managed Word Lists" defaultOpen={false}>
<div className="space-y-2">
{assess.wordPolicy.managedWordLists!.map((w, i) => (
<div key={i} className="flex justify-between items-center p-2 bg-gray-50 rounded">
<div className="flex items-center gap-2">
{chip(w.action ?? "N/A", w.detected ? "red" : "slate")}
<span className="font-mono text-sm break-all">{w.match}</span>
{w.type && chip(w.type, "slate")}
</div>
{boolPill(w.detected)}
</div>
))}
</div>
</Section>
)}
</div>
)}
{/* Content policy */}
{assess.contentPolicy?.filters?.length ? (
<div className="mb-3">
<h6 className="font-medium mb-2">Content Policy</h6>
<div className="overflow-x-auto">
<table className="min-w-full text-sm">
<thead>
<tr className="text-left text-gray-600">
<th className="py-1 pr-4">Type</th>
<th className="py-1 pr-4">Action</th>
<th className="py-1 pr-4">Detected</th>
<th className="py-1 pr-4">Strength</th>
<th className="py-1 pr-4">Confidence</th>
</tr>
</thead>
<tbody>
{assess.contentPolicy.filters!.map((f, i) => (
<tr key={i} className="border-t">
<td className="py-1 pr-4">{f.type ?? "—"}</td>
<td className="py-1 pr-4">{chip(f.action ?? "—", f.detected ? "red" : "slate")}</td>
<td className="py-1 pr-4">{boolPill(f.detected)}</td>
<td className="py-1 pr-4">{f.filterStrength ?? "—"}</td>
<td className="py-1 pr-4">{f.confidence ?? "—"}</td>
</tr>
))}
</tbody>
</table>
</div>
</div>
) : null}
{/* Contextual grounding */}
{assess.contextualGroundingPolicy?.filters?.length ? (
<div className="mb-3">
<h6 className="font-medium mb-2">Contextual Grounding</h6>
<div className="overflow-x-auto">
<table className="min-w-full text-sm">
<thead>
<tr className="text-left text-gray-600">
<th className="py-1 pr-4">Type</th>
<th className="py-1 pr-4">Action</th>
<th className="py-1 pr-4">Detected</th>
<th className="py-1 pr-4">Score</th>
<th className="py-1 pr-4">Threshold</th>
</tr>
</thead>
<tbody>
{assess.contextualGroundingPolicy.filters!.map((f, i) => (
<tr key={i} className="border-t">
<td className="py-1 pr-4">{f.type ?? "—"}</td>
<td className="py-1 pr-4">{chip(f.action ?? "—", f.detected ? "red" : "slate")}</td>
<td className="py-1 pr-4">{boolPill(f.detected)}</td>
<td className="py-1 pr-4">{f.score ?? "—"}</td>
<td className="py-1 pr-4">{f.threshold ?? "—"}</td>
</tr>
))}
</tbody>
</table>
</div>
</div>
) : null}
{/* Sensitive Information */}
{assess.sensitiveInformationPolicy && (
<div className="mb-3">
<h6 className="font-medium mb-2">Sensitive Information</h6>
{(assess.sensitiveInformationPolicy.piiEntities?.length ?? 0) > 0 && (
<Section title="PII Entities" defaultOpen>
<div className="space-y-2">
{assess.sensitiveInformationPolicy.piiEntities!.map((p, i) => (
<div key={i} className="flex justify-between items-center p-2 bg-gray-50 rounded">
<div className="flex items-center gap-2">
{chip(p.action ?? "N/A", p.detected ? "red" : "slate")}
{p.type && chip(p.type, "slate")}
<span className="font-mono text-xs break-all">{p.match}</span>
</div>
{boolPill(p.detected)}
</div>
))}
</div>
</Section>
)}
{(assess.sensitiveInformationPolicy.regexes?.length ?? 0) > 0 && (
<Section title="Custom Regexes" defaultOpen={false}>
<div className="space-y-2">
{assess.sensitiveInformationPolicy.regexes!.map((r, i) => (
<div key={i} className="flex flex-col sm:flex-row sm:items-center sm:justify-between p-2 bg-gray-50 rounded gap-1">
<div className="flex items-center gap-2">
{chip(r.action ?? "N/A", r.detected ? "red" : "slate")}
<span className="font-medium">{r.name ?? "regex"}</span>
<span className="font-mono text-xs break-all">{r.regex}</span>
</div>
<div className="flex items-center gap-2">
{boolPill(r.detected)}
{r.match && <span className="font-mono text-xs break-all">{r.match}</span>}
</div>
</div>
))}
</div>
</Section>
)}
</div>
)}
{/* Topic policy */}
{assess.topicPolicy?.topics?.length ? (
<div className="mb-3">
<h6 className="font-medium mb-2">Topic Policy</h6>
<div className="flex flex-wrap gap-2">
{assess.topicPolicy.topics!.map((t, i) => (
<div key={i} className="px-3 py-1.5 bg-gray-50 rounded-md text-xs">
<div className="flex items-center gap-2">
{chip(t.action ?? "N/A", t.detected ? "red" : "slate")}
<span className="font-medium">{t.name ?? "topic"}</span>
{t.type && chip(t.type, "slate")}
{boolPill(t.detected)}
</div>
</div>
))}
</div>
</div>
) : null}
{/* Invocation metrics */}
{assess.invocationMetrics && (
<Section title="Invocation Metrics" defaultOpen={false}>
<div className="grid grid-cols-2 gap-4">
<div className="space-y-2">
<KV label="Latency (ms)">{assess.invocationMetrics.guardrailProcessingLatency ?? "—"}</KV>
<KV label="Coverage:">
<div className="flex flex-wrap gap-2">
{assess.invocationMetrics.guardrailCoverage?.textCharacters &&
chip(
`text ${assess.invocationMetrics.guardrailCoverage.textCharacters.guarded ?? 0}/${
assess.invocationMetrics.guardrailCoverage.textCharacters.total ?? 0
}`,
"blue"
)}
{assess.invocationMetrics.guardrailCoverage?.images &&
chip(
`images ${assess.invocationMetrics.guardrailCoverage.images.guarded ?? 0}/${
assess.invocationMetrics.guardrailCoverage.images.total ?? 0
}`,
"blue"
)}
</div>
</KV>
</div>
<div className="space-y-2">
<KV label="Usage:">
<div className="flex flex-wrap gap-2">
{assess.invocationMetrics.usage &&
Object.entries(assess.invocationMetrics.usage).map(([k, v]) =>
typeof v === "number" ? (
<span key={k} className="px-2 py-1 bg-slate-100 text-slate-800 rounded-md text-xs font-medium">
{k}: {v}
</span>
) : null
)}
</div>
</KV>
</div>
</div>
</Section>
)}
{/* Automated reasoning (fallback render) */}
{assess.automatedReasoningPolicy?.findings?.length ? (
<Section title="Automated Reasoning Findings" defaultOpen={false}>
<div className="space-y-2">
{assess.automatedReasoningPolicy.findings!.map((f, i) => (
<pre key={i} className="bg-gray-50 rounded p-2 text-xs overflow-x-auto">
{JSON.stringify(f, null, 2)}
</pre>
))}
</div>
</Section>
) : null}
</Section>
);
})}
</div>
) : null}
{/* Raw JSON (for debugging / completeness) */}
<Section title="Raw Bedrock Guardrail Response" defaultOpen={false}>
<pre className="bg-gray-50 rounded p-3 text-xs overflow-x-auto">{JSON.stringify(response, null, 2)}</pre>
</Section>
</div>
);
};
export default BedrockGuardrailDetails;

View file

@ -0,0 +1,154 @@
import React from 'react';
import { describe, it, expect, vi, beforeEach } from 'vitest';
import userEvent from '@testing-library/user-event';
import { renderWithProviders, screen } from "../../../../tests/test-utils"
import {
makeBedrockResponse, makeEntity,
makeGuardrailInformation,
} from "@/components/view_logs/GuardrailViewer/__tests__/fixtures"
import GuardrailViewer from "@/components/view_logs/GuardrailViewer/GuardrailViewer"
// We will mock child components selectively for some tests to assert prop passthrough,
// but also run an integration-style render without mocks.
const PresidioPath = '@/components/view_logs/GuardrailViewer/PresidioDetectedEntities';
const BedrockPath = '@/components/view_logs/GuardrailViewer/BedrockGuardrailDetails';
describe('GuardrailViewer', () => {
beforeEach(() => {
vi.resetModules();
});
it('shows header, status pill color, duration rounding, and time labels', () => {
const data = makeGuardrailInformation({ duration: 1.23456, guardrail_status: 'success' });
renderWithProviders(<GuardrailViewer data={data} />);
expect(screen.getByText('Guardrail Information')).toBeInTheDocument();
// header status pill (success => green)
const statusBadges = screen.getAllByText('success');
// there are two status locations: header chip and grid "Status"
expect(statusBadges.length).toBeGreaterThanOrEqual(1);
// Quick class assertion for at least one of them
expect(statusBadges[0].className).toMatch(/bg-green-100/);
// duration displays with 4 decimals
expect(screen.getByText(/1\.2346s/)).toBeInTheDocument();
// time labels exist
expect(screen.getByText('Start Time:')).toBeInTheDocument();
expect(screen.getByText('End Time:')).toBeInTheDocument();
});
it('calculates and displays masked entity totals with pluralization', () => {
const data = makeGuardrailInformation({
masked_entity_count: { EMAIL_ADDRESS: 2, PHONE_NUMBER: 1 },
});
renderWithProviders(<GuardrailViewer data={data} />);
expect(screen.getByText('3 masked entities')).toBeInTheDocument();
// summary chips for each entry
expect(screen.getByText('EMAIL_ADDRESS: 2')).toBeInTheDocument();
expect(screen.getByText('PHONE_NUMBER: 1')).toBeInTheDocument();
});
it('hides masked badge & summary when count is zero/empty', () => {
const data = makeGuardrailInformation({ masked_entity_count: {} });
renderWithProviders(<GuardrailViewer data={data} />);
expect(screen.queryByText(/masked entity/)).not.toBeInTheDocument();
expect(screen.queryByText('Masked Entity Summary')).not.toBeInTheDocument();
});
it('toggles main section open/closed and chevron rotation class', async () => {
const user = userEvent.setup();
const data = makeGuardrailInformation();
renderWithProviders(<GuardrailViewer data={data} />);
const header = screen.getByText('Guardrail Information').closest('div')!;
// Initially expanded
expect(screen.getByText('Click to collapse')).toBeInTheDocument();
// Click to collapse
await user.click(header);
expect(screen.getByText('Click to expand')).toBeInTheDocument();
// Details gone
expect(screen.queryByText('Masked Entity Summary')).not.toBeInTheDocument();
// Click to expand again
await user.click(header);
expect(screen.getByText('Click to collapse')).toBeInTheDocument();
});
it('defaults to presidio provider when guardrail_provider is undefined', async () => {
vi.doMock(PresidioPath, () => ({
__esModule: true,
default: ({ entities }: any) => <div data-testid="presidio-mock">presidio {entities?.length}</div>,
}));
const { default: Component } = await import('@/components/view_logs/GuardrailViewer/GuardrailViewer');
const data = makeGuardrailInformation({
guardrail_provider: undefined,
guardrail_response: [makeEntity(), makeEntity()],
});
renderWithProviders(<Component data={data} />);
expect(screen.getByTestId('presidio-mock')).toHaveTextContent('presidio 2');
});
it('renders PresidioDetectedEntities when provider="presidio" and response has entities', async () => {
vi.doMock(PresidioPath, () => ({
__esModule: true,
default: ({ entities }: any) => <div data-testid="presidio-mock">count:{entities?.length}</div>,
}));
const { default: Component } = await import('@/components/view_logs/GuardrailViewer/GuardrailViewer');
const data = makeGuardrailInformation({
guardrail_provider: 'presidio',
guardrail_response: [makeEntity()],
});
renderWithProviders(<Component data={data} />);
expect(screen.getByTestId('presidio-mock')).toHaveTextContent('count:1');
});
it('renders BedrockGuardrailDetails when provider="bedrock"', async () => {
vi.doMock(BedrockPath, () => ({
__esModule: true,
default: ({ response }: any) => (
<div data-testid="bedrock-mock">{response?.action ?? 'no-action'}</div>
),
}));
const { default: Component } = await import('@/components/view_logs/GuardrailViewer/GuardrailViewer');
const data = makeGuardrailInformation({
guardrail_provider: 'bedrock',
guardrail_response: makeBedrockResponse({ action: 'GUARDRAIL_INTERVENED' }),
});
renderWithProviders(<Component data={data} />);
expect(screen.getByTestId('bedrock-mock')).toHaveTextContent('GUARDRAIL_INTERVENED');
});
it('unknown provider renders neither Presidio nor Bedrock details', () => {
const data = makeGuardrailInformation({
guardrail_provider: 'unknown',
});
renderWithProviders(<GuardrailViewer data={data} />);
// Summary still present
expect(screen.getByText('Guardrail Information')).toBeInTheDocument();
// No provider sections
expect(screen.queryByText(/Detected Entities/)).not.toBeInTheDocument();
expect(screen.queryByText(/Raw Bedrock Guardrail Response/)).not.toBeInTheDocument();
});
it('integration: renders with real Bedrock details without mocks', () => {
const data = makeGuardrailInformation({
guardrail_provider: 'bedrock',
guardrail_response: makeBedrockResponse({
action: 'NONE',
outputs: [{ text: 'ok' }],
}),
});
renderWithProviders(<GuardrailViewer data={data} />);
// Bedrock summary bits
expect(screen.getByText('Outputs')).toBeInTheDocument();
expect(screen.getByText('ok')).toBeInTheDocument();
});
});

View file

@ -0,0 +1,176 @@
import React, { useState } from "react";
import { Tooltip } from "antd";
import PresidioDetectedEntities from "./PresidioDetectedEntities";
import BedrockGuardrailDetails, {
BedrockGuardrailResponse,
} from "@/components/view_logs/GuardrailViewer/BedrockGuardrailDetails";
interface RecognitionMetadata {
recognizer_name: string;
recognizer_identifier: string;
}
interface GuardrailEntity {
end: number;
score: number;
start: number;
entity_type: string;
analysis_explanation: string | null;
recognition_metadata: RecognitionMetadata;
}
interface MaskedEntityCount {
[key: string]: number;
}
interface GuardrailInformation {
duration: number;
end_time: number;
start_time: number;
guardrail_mode: string;
guardrail_name: string;
guardrail_status: string;
guardrail_response: GuardrailEntity[] | BedrockGuardrailResponse;
masked_entity_count: MaskedEntityCount;
guardrail_provider?: string; // "presidio" | other providers
}
interface GuardrailViewerProps {
data: GuardrailInformation;
}
const GuardrailViewer = ({ data }: GuardrailViewerProps) => {
const [sectionExpanded, setSectionExpanded] = useState(true);
// Default to presidio for backwards compatibility
const guardrailProvider = data.guardrail_provider ?? "presidio";
if (!data) return null;
const isSuccess =
typeof data.guardrail_status === "string" &&
data.guardrail_status.toLowerCase() === "success";
const tooltipTitle = isSuccess ? null : "Guardrail failed to run.";
// Calculate total masked entities
const totalMaskedEntities = data.masked_entity_count ?
Object.values(data.masked_entity_count).reduce((sum, count) => sum + count, 0) : 0;
const formatTime = (timestamp: number): string => {
const date = new Date(timestamp * 1000);
return date.toLocaleString();
};
return (
<div className="bg-white rounded-lg shadow mb-6">
<div
className="flex justify-between items-center p-4 border-b cursor-pointer hover:bg-gray-50"
onClick={() => setSectionExpanded(!sectionExpanded)}
>
<div className="flex items-center">
<svg
className={`w-5 h-5 mr-2 text-gray-600 transition-transform ${sectionExpanded ? 'transform rotate-90' : ''}`}
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M9 5l7 7-7 7" />
</svg>
<h3 className="text-lg font-medium">Guardrail Information</h3>
{/* Header status chip with tooltip */}
<Tooltip title={tooltipTitle} placement="top" arrow destroyTooltipOnHide>
<span
className={`ml-3 px-2 py-1 rounded-md text-xs font-medium inline-block ${
isSuccess ? "bg-green-100 text-green-800" : "bg-red-100 text-red-800 cursor-help"
}`}
>
{data.guardrail_status}
</span>
</Tooltip>
{totalMaskedEntities > 0 && (
<span className="ml-3 px-2 py-1 bg-blue-50 text-blue-700 rounded-md text-xs font-medium">
{totalMaskedEntities} masked {totalMaskedEntities === 1 ? 'entity' : 'entities'}
</span>
)}
</div>
<span className="text-sm text-gray-500">{sectionExpanded ? 'Click to collapse' : 'Click to expand'}</span>
</div>
{sectionExpanded && (
<div className="p-4">
<div className="bg-white rounded-lg border p-4 mb-4">
<div className="grid grid-cols-2 gap-4">
<div className="space-y-2">
<div className="flex">
<span className="font-medium w-1/3">Guardrail Name:</span>
<span className="font-mono">{data.guardrail_name}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Mode:</span>
<span className="font-mono">{data.guardrail_mode}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Status:</span>
<Tooltip title={tooltipTitle} placement="top" arrow destroyTooltipOnHide>
<span
className={`px-2 py-1 rounded-md text-xs font-medium inline-block ${
isSuccess ? "bg-green-100 text-green-800" : "bg-red-100 text-red-800 cursor-help"
}`}
>
{data.guardrail_status}
</span>
</Tooltip>
</div>
</div>
<div className="space-y-2">
<div className="flex">
<span className="font-medium w-1/3">Start Time:</span>
<span>{formatTime(data.start_time)}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">End Time:</span>
<span>{formatTime(data.end_time)}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Duration:</span>
<span>{data.duration.toFixed(4)}s</span>
</div>
</div>
</div>
{/* Masked Entity Summary */}
{data.masked_entity_count && Object.keys(data.masked_entity_count).length > 0 && (
<div className="mt-4 pt-4 border-t">
<h4 className="font-medium mb-2">Masked Entity Summary</h4>
<div className="flex flex-wrap gap-2">
{Object.entries(data.masked_entity_count).map(([entityType, count]) => (
<span key={entityType} className="px-3 py-1.5 bg-blue-50 text-blue-700 rounded-md text-xs font-medium">
{entityType}: {count}
</span>
))}
</div>
</div>
)}
</div>
{/* Provider-specific Detected Entities */}
{guardrailProvider === "presidio" && (data.guardrail_response as GuardrailEntity[])?.length > 0 && (
<PresidioDetectedEntities entities={data.guardrail_response as GuardrailEntity[]} />
)}
{guardrailProvider === "bedrock" && data.guardrail_response && (
<div className="mt-4">
<BedrockGuardrailDetails response={data.guardrail_response as BedrockGuardrailResponse} />
</div>
)}
</div>
)}
</div>
);
};
export default GuardrailViewer;

View file

@ -0,0 +1,55 @@
import React from 'react';
import { describe, it, expect } from 'vitest';
import userEvent from '@testing-library/user-event';
import PresidioDetectedEntities from '@/components/view_logs/GuardrailViewer/PresidioDetectedEntities';
import { renderWithProviders, screen } from "../../../../tests/test-utils"
import { makeEntity } from "@/components/view_logs/GuardrailViewer/__tests__/fixtures"
describe('PresidioDetectedEntities', () => {
it('renders null when entities empty', () => {
const { container } = renderWithProviders(<PresidioDetectedEntities entities={[]} />);
expect(container).toBeEmptyDOMElement();
});
it('renders per-entity header info including score color and position', async () => {
const user = userEvent.setup();
const e = makeEntity({ start: 10, end: 20, score: 0.92, entity_type: 'EMAIL_ADDRESS' });
renderWithProviders(<PresidioDetectedEntities entities={[e]} />);
// Header row values
expect(screen.getByText('EMAIL_ADDRESS')).toBeInTheDocument();
expect(screen.getByText(/Score: 0\.92/)).toBeInTheDocument();
expect(screen.getByText('Position: 10-20')).toBeInTheDocument();
// Expand details
await user.click(screen.getByText('EMAIL_ADDRESS'));
expect(screen.getByText('Entity Type:')).toBeInTheDocument();
expect(screen.getByText('Characters 10-20')).toBeInTheDocument();
expect(screen.getByText('Confidence:')).toBeInTheDocument();
// Recognizer details
expect(screen.getByText('EmailRecognizer')).toBeInTheDocument();
expect(screen.getByText('email_v1')).toBeInTheDocument();
// Explanation
expect(screen.getByText('Matched via pattern')).toBeInTheDocument();
});
it('handles missing metadata & low scores gracefully', async () => {
const user = userEvent.setup();
const e = makeEntity({
score: 0.3,
recognition_metadata: undefined as any,
analysis_explanation: null,
entity_type: 'NAME',
start: 0,
end: 0,
});
renderWithProviders(<PresidioDetectedEntities entities={[e]} />);
await user.click(screen.getByText('NAME'));
// No recognizer/explanation rows
expect(screen.queryByText('Recognizer:')).not.toBeInTheDocument();
expect(screen.queryByText('Explanation:')).not.toBeInTheDocument();
// Position still renders
expect(screen.getByText('Characters 0-0')).toBeInTheDocument();
});
});

View file

@ -0,0 +1,136 @@
import React, { useState } from "react";
interface RecognitionMetadata {
recognizer_name: string;
recognizer_identifier: string;
}
export interface GuardrailEntity {
end: number;
score: number;
start: number;
entity_type: string;
analysis_explanation: string | null;
recognition_metadata: RecognitionMetadata;
}
interface PresidioDetectedEntitiesProps {
entities: GuardrailEntity[];
}
const getScoreColor = (score: number): string => {
if (score >= 0.8) return "text-green-600";
return "text-yellow-600";
};
const PresidioDetectedEntities = ({ entities }: PresidioDetectedEntitiesProps) => {
const [entityListExpanded, setEntityListExpanded] = useState(true);
const [expandedEntities, setExpandedEntities] = useState<Record<number, boolean>>({});
const toggleEntity = (index: number) => {
setExpandedEntities((prev) => ({
...prev,
[index]: !prev[index],
}));
};
if (!entities || entities.length === 0) return null;
return (
<div className="mt-4">
<div
className="flex items-center mb-2 cursor-pointer"
onClick={() => setEntityListExpanded(!entityListExpanded)}
>
<svg
className={`w-5 h-5 mr-2 transition-transform ${entityListExpanded ? "transform rotate-90" : ""}`}
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M9 5l7 7-7 7" />
</svg>
<h4 className="font-medium">Detected Entities ({entities.length})</h4>
</div>
{entityListExpanded && (
<div className="space-y-2">
{entities.map((entity, index) => {
const isExpanded = expandedEntities[index] || false;
return (
<div key={index} className="border rounded-lg overflow-hidden">
<div
className="flex items-center justify-between p-3 bg-gray-50 cursor-pointer hover:bg-gray-100"
onClick={() => toggleEntity(index)}
>
<div className="flex items-center">
<svg
className={`w-5 h-5 mr-2 transition-transform ${isExpanded ? "transform rotate-90" : ""}`}
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M9 5l7 7-7 7" />
</svg>
<span className="font-medium mr-2">{entity.entity_type}</span>
<span className={`font-mono ${getScoreColor(entity.score)}`}>
Score: {entity.score.toFixed(2)}
</span>
</div>
<span className="text-xs text-gray-500">Position: {entity.start}-{entity.end}</span>
</div>
{isExpanded && (
<div className="p-3 border-t bg-white">
<div className="grid grid-cols-2 gap-4 mb-2">
<div className="space-y-2">
<div className="flex">
<span className="font-medium w-1/3">Entity Type:</span>
<span>{entity.entity_type}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Position:</span>
<span>Characters {entity.start}-{entity.end}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Confidence:</span>
<span className={getScoreColor(entity.score)}>{entity.score.toFixed(2)}</span>
</div>
</div>
<div className="space-y-2">
{entity.recognition_metadata && (
<>
<div className="flex">
<span className="font-medium w-1/3">Recognizer:</span>
<span>{entity.recognition_metadata.recognizer_name}</span>
</div>
<div className="flex overflow-hidden">
<span className="font-medium w-1/3">Identifier:</span>
<span className="truncate text-xs font-mono">
{entity.recognition_metadata.recognizer_identifier}
</span>
</div>
</>
)}
{entity.analysis_explanation && (
<div className="flex">
<span className="font-medium w-1/3">Explanation:</span>
<span>{entity.analysis_explanation}</span>
</div>
)}
</div>
</div>
</div>
)}
</div>
);
})}
</div>
)}
</div>
);
}
export default PresidioDetectedEntities;

View file

@ -0,0 +1,119 @@
import type {
BedrockGuardrailResponse,
BedrockAssessment,
BedrockGuardrailCoverage,
BedrockGuardrailUsage,
} from '@/components/view_logs/GuardrailViewer/BedrockGuardrailDetails';
export interface RecognitionMetadata {
recognizer_name: string;
recognizer_identifier: string;
}
export interface GuardrailEntity {
end: number;
score: number;
start: number;
entity_type: string;
analysis_explanation: string | null;
recognition_metadata: RecognitionMetadata;
}
export interface GuardrailInformation {
duration: number;
end_time: number;
start_time: number;
guardrail_mode: string;
guardrail_name: string;
guardrail_status: string;
guardrail_response: GuardrailEntity[] | BedrockGuardrailResponse;
masked_entity_count: Record<string, number>;
guardrail_provider?: string;
}
// ===== Builders =====
export const makeEntity = (overrides: Partial<GuardrailEntity> = {}): GuardrailEntity => ({
end: 18,
start: 5,
score: 0.92,
entity_type: 'EMAIL_ADDRESS',
analysis_explanation: 'Matched via pattern',
recognition_metadata: {
recognizer_name: 'EmailRecognizer',
recognizer_identifier: 'email_v1',
},
...overrides,
});
export const makeMaskedCounts = (overrides: Record<string, number> = {}) => ({
EMAIL_ADDRESS: 2,
PHONE_NUMBER: 1,
...overrides,
});
export const makeBedrockUsage = (overrides: Partial<BedrockGuardrailUsage> = {}): BedrockGuardrailUsage => ({
contentPolicyUnits: 4,
topicPolicyUnits: 2,
...overrides,
});
export const makeBedrockCoverage = (
overrides: Partial<BedrockGuardrailCoverage> = {}
): BedrockGuardrailCoverage => ({
textCharacters: { guarded: 27, total: 100 },
images: { guarded: 1, total: 3 },
...overrides,
});
export const makeAssessment = (overrides: Partial<BedrockAssessment> = {}): BedrockAssessment => ({
wordPolicy: {
customWords: [{ action: 'BLOCKED', detected: true, match: 'badword' }],
managedWordLists: [{ action: 'ALLOWED', detected: false, match: 'ok', type: 'PROFANITY' }],
},
contentPolicy: {
filters: [
{ type: 'HATE', action: 'BLOCKED', detected: true, filterStrength: 'HIGH', confidence: 'MEDIUM' },
{ type: 'VIOLENCE', action: 'NONE', detected: false, filterStrength: 'LOW', confidence: 'LOW' },
],
},
topicPolicy: { topics: [{ name: 'weapons', type: 'DENY', detected: true, action: 'BLOCKED' }] },
sensitiveInformationPolicy: {
piiEntities: [{ type: 'EMAIL', match: 'x@y.com', detected: true, action: 'ANONYMIZED' }],
regexes: [{ name: 'ticket', regex: '#[0-9]+', match: '#123', detected: true, action: 'BLOCKED' }],
},
contextualGroundingPolicy: {
filters: [{ type: 'GROUNDING', action: 'BLOCKED', detected: true, score: 0.2, threshold: 0.5 }],
},
automatedReasoningPolicy: { findings: [{ foo: 'bar' }] },
invocationMetrics: {
guardrailProcessingLatency: 42,
usage: makeBedrockUsage(),
guardrailCoverage: makeBedrockCoverage(),
},
...overrides,
});
export const makeBedrockResponse = (
overrides: Partial<BedrockGuardrailResponse> = {}
): BedrockGuardrailResponse => ({
action: 'NONE',
outputs: [{ text: 'ok' }],
usage: makeBedrockUsage(),
guardrailCoverage: makeBedrockCoverage(),
assessments: [makeAssessment()],
...overrides,
});
export const makeGuardrailInformation = (
overrides: Partial<GuardrailInformation> = {}
): GuardrailInformation => ({
guardrail_name: 'pii-rail',
guardrail_mode: 'post',
guardrail_status: 'success',
start_time: 1_700_000_000,
end_time: 1_700_000_123,
duration: 0.123456,
guardrail_response: [makeEntity()],
masked_entity_count: makeMaskedCounts(),
...overrides,
});

Some files were not shown because too many files have changed in this diff Show more