mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge branch 'BerriAI:main' into LangfuseUsageDetails
This commit is contained in:
commit
f0490ab60d
104 changed files with 8386 additions and 3387 deletions
|
|
@ -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 .
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
```
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
95
docs/my-website/docs/providers/bedrock_embedding.md
Normal file
95
docs/my-website/docs/providers/bedrock_embedding.md
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -411,6 +411,7 @@ const sidebars = {
|
|||
label: "Bedrock",
|
||||
items: [
|
||||
"providers/bedrock",
|
||||
"providers/bedrock_embedding",
|
||||
"providers/bedrock_agents",
|
||||
"providers/bedrock_batches",
|
||||
"providers/bedrock_vector_store",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
56
litellm/litellm_core_utils/cached_imports.py
Normal file
56
litellm/litellm_core_utils/cached_imports.py
Normal 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
|
||||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
137
litellm/litellm_core_utils/object_pooling.py
Normal file
137
litellm/litellm_core_utils/object_pooling.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 #############
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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
1683
poetry.lock
generated
File diff suppressed because it is too large
Load diff
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
43
tests/code_coverage_tests/test_chat_completion_imports.py
Normal file
43
tests/code_coverage_tests/test_chat_completion_imports.py
Normal 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")
|
||||
106
tests/litellm_utils_tests/test_object_pooling.py
Normal file
106
tests/litellm_utils_tests/test_object_pooling.py
Normal 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__])
|
||||
|
|
@ -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}")
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
134
tests/test_litellm/integrations/test_langsmith_init.py
Normal file
134
tests/test_litellm/integrations/test_langsmith_init.py
Normal 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}"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
2285
ui/litellm-dashboard/package-lock.json
generated
2285
ui/litellm-dashboard/package-lock.json
generated
File diff suppressed because it is too large
Load diff
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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={
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -185,7 +185,6 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
|
|||
alias: "",
|
||||
url: "",
|
||||
transport: "",
|
||||
spec_version: "",
|
||||
auth_type: "",
|
||||
created_at: "",
|
||||
created_by: "",
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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;
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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;
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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;
|
||||
|
|
@ -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
Loading…
Add table
Reference in a new issue