mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge branch 'main' into litellm_dev_10_06_2025_p1
This commit is contained in:
commit
94a34dd53a
86 changed files with 3540 additions and 502 deletions
0
MCP_SSL_CHANGES_SUMMARY.md
Normal file
0
MCP_SSL_CHANGES_SUMMARY.md
Normal file
|
|
@ -23,6 +23,43 @@ LiteLLM Proxy provides an MCP Gateway that allows you to use a fixed endpoint fo
|
|||
|
||||
## Adding your MCP
|
||||
|
||||
### Prerequisites
|
||||
|
||||
To store MCP servers in the database, you need to enable database storage:
|
||||
|
||||
**Environment Variable:**
|
||||
```bash
|
||||
export STORE_MODEL_IN_DB=True
|
||||
```
|
||||
|
||||
**OR in config.yaml:**
|
||||
```yaml
|
||||
general_settings:
|
||||
store_model_in_db: true
|
||||
```
|
||||
|
||||
#### Fine-grained Database Storage Control
|
||||
|
||||
By default, when `store_model_in_db` is `true`, all object types (models, MCPs, guardrails, vector stores, etc.) are stored in the database. If you want to store only specific object types, use the `supported_db_objects` setting.
|
||||
|
||||
**Example: Store only MCP servers in the database**
|
||||
|
||||
```yaml title="config.yaml" showLineNumbers
|
||||
general_settings:
|
||||
store_model_in_db: true
|
||||
supported_db_objects: ["mcp"] # Only store MCP servers in DB
|
||||
|
||||
model_list:
|
||||
- model_name: gpt-4o
|
||||
litellm_params:
|
||||
model: openai/gpt-4o
|
||||
api_key: sk-xxxxxxx
|
||||
```
|
||||
|
||||
**See all available object types:** [Config Settings - supported_db_objects](./proxy/config_settings.md#general_settings---reference)
|
||||
|
||||
If `supported_db_objects` is not set, all object types are loaded from the database (default behavior).
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="ui" label="LiteLLM UI">
|
||||
|
||||
|
|
|
|||
|
|
@ -171,6 +171,7 @@ os.environ["OPENAI_BASE_URL"] = "https://your_host/v1" # OPTIONAL
|
|||
| gpt-5-2025-08-07 | `response = completion(model="gpt-5-2025-08-07", messages=messages)` |
|
||||
| gpt-5-mini-2025-08-07 | `response = completion(model="gpt-5-mini-2025-08-07", messages=messages)` |
|
||||
| gpt-5-nano-2025-08-07 | `response = completion(model="gpt-5-nano-2025-08-07", messages=messages)` |
|
||||
| gpt-5-pro | `response = completion(model="gpt-5-pro", messages=messages)` |
|
||||
| gpt-4.1 | `response = completion(model="gpt-4.1", messages=messages)` |
|
||||
| gpt-4.1-mini | `response = completion(model="gpt-4.1-mini", messages=messages)` |
|
||||
| gpt-4.1-nano | `response = completion(model="gpt-4.1-nano", messages=messages)` |
|
||||
|
|
@ -749,4 +750,24 @@ In your logs you should see the forwarded org id
|
|||
```bash
|
||||
LiteLLM:DEBUG: utils.py:255 - Request to litellm:
|
||||
LiteLLM:DEBUG: utils.py:255 - litellm.acompletion(... organization='my-special-org',)
|
||||
```
|
||||
|
||||
## GPT-5 Pro Special Notes
|
||||
|
||||
GPT-5 Pro is OpenAI's most advanced reasoning model with unique characteristics:
|
||||
|
||||
- **Responses API Only**: GPT-5 Pro is only available through the `/v1/responses` endpoint
|
||||
- **No Streaming**: Does not support streaming responses
|
||||
- **High Reasoning**: Designed for complex reasoning tasks with highest effort reasoning
|
||||
- **Context Window**: 400,000 tokens input, 272,000 tokens output
|
||||
- **Pricing**: $15.00 input / $120.00 output per 1M tokens (Standard), $7.50 input / $60.00 output (Batch)
|
||||
- **Tools**: Supports Web Search, File Search, Image Generation, MCP (but not Code Interpreter or Computer Use)
|
||||
- **Modalities**: Text and Image input, Text output only
|
||||
|
||||
```python
|
||||
# GPT-5 Pro usage example
|
||||
response = completion(
|
||||
model="gpt-5-pro",
|
||||
messages=[{"role": "user", "content": "Solve this complex reasoning problem..."}]
|
||||
)
|
||||
```
|
||||
|
|
@ -191,7 +191,7 @@ print(json.loads(completion.choices[0].message.content))
|
|||
model_list:
|
||||
- model_name: gemini-2.5-pro
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-1.5-pro
|
||||
model: vertex_ai/gemini-2.5-pro
|
||||
vertex_project: "project-id"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: "/path/to/service_account.json" # [OPTIONAL] Do this OR `!gcloud auth application-default login` - run this to add vertex credentials to your env
|
||||
|
|
@ -277,7 +277,7 @@ except JSONSchemaValidationError as e:
|
|||
model_list:
|
||||
- model_name: gemini-2.5-pro
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-1.5-pro
|
||||
model: vertex_ai/gemini-2.5-pro
|
||||
vertex_project: "project-id"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: "/path/to/service_account.json" # [OPTIONAL] Do this OR `!gcloud auth application-default login` - run this to add vertex credentials to your env
|
||||
|
|
@ -981,11 +981,155 @@ curl http://0.0.0.0:4000/v1/chat/completions \
|
|||
|
||||
### **Context Caching**
|
||||
|
||||
Use Vertex AI context caching is supported by calling provider api directly. (Unified Endpoint support coming soon.).
|
||||
#### Unified Endpoint
|
||||
|
||||
Use Vertex AI context caching in the same way as [**Google AI Studio - Context Caching**](../providers/gemini.md#context-caching)
|
||||
|
||||
|
||||
##### Example usage
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
for _ in range(2):
|
||||
resp = completion(
|
||||
model="vertex_ai/gemini-2.5-pro",
|
||||
messages=[
|
||||
# System Message
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Here is the full text of a complex legal agreement" * 4000,
|
||||
"cache_control": {"type": "ephemeral"}, # 👈 KEY CHANGE
|
||||
}
|
||||
],
|
||||
},
|
||||
# marked for caching with the cache_control parameter, so that this checkpoint can read from the previous cache.
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "What are the key terms and conditions in this agreement?",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
}]
|
||||
)
|
||||
|
||||
print(resp.usage) # 👈 2nd usage block will be less, since cached tokens used
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="sdk-ttl" label="SDK with Custom TTL">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
# Cache for 2 hours (7200 seconds)
|
||||
resp = completion(
|
||||
model="vertex_ai/gemini-2.5-pro",
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Here is the full text of a complex legal agreement" * 4000,
|
||||
"cache_control": {
|
||||
"type": "ephemeral",
|
||||
"ttl": "7200s" # 👈 Cache for 2 hours
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "What are the key terms and conditions in this agreement?",
|
||||
"cache_control": {
|
||||
"type": "ephemeral",
|
||||
"ttl": "3600s" # 👈 This TTL will be ignored (first one is used)
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
print(resp.usage)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
1. Setup config.yaml
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gemini-2.5-pro
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-2.5-pro
|
||||
vertex_project: "project-id"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: "/path/to/service_account.json"
|
||||
```
|
||||
|
||||
2. Start proxy
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
3. Test it!
|
||||
|
||||
```bash
|
||||
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-d '{
|
||||
"model": "gemini-2.5-flash",
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Long cache message (must be >= 1024 tokens)",
|
||||
"cache_control": {
|
||||
"type": "ephemeral",
|
||||
"ttl": "7200s"
|
||||
}
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "What is the text about?"
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}'
|
||||
|
||||
```
|
||||
|
||||
#### Calling provider api directly
|
||||
|
||||
[**Go straight to provider**](../pass_through/vertex_ai.md#context-caching)
|
||||
|
||||
#### 1. Create the Cache
|
||||
##### 1. Create the Cache
|
||||
|
||||
First, create the cache by sending a `POST` request to the `cachedContents` endpoint via the LiteLLM proxy.
|
||||
|
||||
|
|
@ -1011,7 +1155,7 @@ curl http://0.0.0.0:4000/vertex_ai/v1/projects/{project_id}/locations/{location}
|
|||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
#### 2. Get the Cache Name from the Response
|
||||
##### 2. Get the Cache Name from the Response
|
||||
|
||||
Vertex AI will return a response containing the `name` of the cached content. This name is the identifier for your cached data.
|
||||
|
||||
|
|
@ -1030,7 +1174,7 @@ Vertex AI will return a response containing the `name` of the cached content. Th
|
|||
}
|
||||
```
|
||||
|
||||
#### 3. Use the Cached Content
|
||||
##### 3. Use the Cached Content
|
||||
|
||||
Use the `name` from the response as `cachedContent` or `cached_content` in subsequent API calls to reuse the cached information. This is passed in the body of your request to `/chat/completions`.
|
||||
|
||||
|
|
|
|||
|
|
@ -224,6 +224,7 @@ router_settings:
|
|||
| service_account_settings | List[Dict[str, Any]] | Set `service_account_settings` if you want to create settings that only apply to service account keys (Doc on service accounts)[./service_accounts.md] |
|
||||
| image_generation_model | str | The default model to use for image generation - ignores model set in request |
|
||||
| store_model_in_db | boolean | If true, enables storing model + credential information in the DB. |
|
||||
| supported_db_objects | List[str] | Fine-grained control over which object types to load from the database when `store_model_in_db` is True. Available types: `"models"`, `"mcp"`, `"guardrails"`, `"vector_stores"`, `"pass_through_endpoints"`, `"prompts"`, `"model_cost_map"`. If not set, all object types are loaded (default behavior). Example: `supported_db_objects: ["mcp"]` to only load MCP servers from DB. |
|
||||
| store_prompts_in_spend_logs | boolean | If true, allows prompts and responses to be stored in the spend logs table. |
|
||||
| max_request_size_mb | int | The maximum size for requests in MB. Requests above this size will be rejected. |
|
||||
| max_response_size_mb | int | The maximum size for responses in MB. LLM Responses above this size will not be sent. |
|
||||
|
|
|
|||
354
docs/my-website/docs/proxy/security_encryption_faq.md
Normal file
354
docs/my-website/docs/proxy/security_encryption_faq.md
Normal file
|
|
@ -0,0 +1,354 @@
|
|||
# LiteLLM Self-Hosted Security & Encryption FAQ
|
||||
|
||||
## Data in Transit Encryption
|
||||
|
||||
### Does the product encrypt data in transit?
|
||||
|
||||
**Yes**, LiteLLM encrypts data in transit using TLS/SSL.
|
||||
|
||||
### Available in both OSS and Enterprise?
|
||||
|
||||
**Yes**, TLS encryption is available in both Open Source and Enterprise versions.
|
||||
|
||||
### In transit between the calling client and the product?
|
||||
|
||||
**Yes**, HTTPS/TLS is supported through SSL certificate configuration.
|
||||
|
||||
**Configuration:**
|
||||
```bash
|
||||
# CLI
|
||||
litellm --ssl_keyfile_path /path/to/key.pem --ssl_certfile_path /path/to/cert.pem
|
||||
|
||||
# Environment Variables
|
||||
export SSL_KEYFILE_PATH="/path/to/key.pem"
|
||||
export SSL_CERTFILE_PATH="/path/to/cert.pem"
|
||||
```
|
||||
|
||||
**Documentation Reference:** `docs/my-website/docs/guides/security_settings.md`
|
||||
|
||||
### In transit between the product and the LLM providers?
|
||||
|
||||
**Yes**, all connections to LLM providers use TLS encryption by default.
|
||||
|
||||
**Implementation Details:**
|
||||
- Uses Python's `ssl.create_default_context()`
|
||||
- Leverages HTTPX and aiohttp libraries with SSL/TLS enabled
|
||||
- Uses certifi CA bundle by default for SSL verification
|
||||
|
||||
**Code Reference:** `litellm/llms/custom_httpx/http_handler.py` (lines 43-105)
|
||||
|
||||
### Are TCP sessions to the LLM providers shared?
|
||||
|
||||
**Yes**, TCP connections are pooled and reused.
|
||||
|
||||
**Details:**
|
||||
- Connection pooling is enabled by default
|
||||
- Default: 1000 max concurrent connections with keepalive
|
||||
- Sessions are maintained across requests to the same provider
|
||||
- Reduces overhead of TLS handshakes
|
||||
|
||||
**Code Reference:** `litellm/llms/custom_httpx/http_handler.py` (lines 704-712)
|
||||
|
||||
### Or does the product negotiate a new TLS session with the same LLM provider for every sequential call?
|
||||
|
||||
**No**, TLS sessions are reused through connection pooling. New TLS handshakes are not performed for every request.
|
||||
|
||||
### How is it encrypted?
|
||||
|
||||
**TLS 1.2 and TLS 1.3**
|
||||
|
||||
Uses Python's default SSL context which supports both TLS 1.2 and TLS 1.3. The specific version negotiated depends on:
|
||||
- Python version
|
||||
- System SSL library (typically OpenSSL)
|
||||
- Server capabilities
|
||||
|
||||
**Implementation:** `ssl.create_default_context()` in Python
|
||||
|
||||
### How are these added to the product's configuration?
|
||||
|
||||
#### x.509 Certificate
|
||||
|
||||
**Method 1: CLI Arguments**
|
||||
```bash
|
||||
litellm --ssl_certfile_path /path/to/certificate.pem
|
||||
```
|
||||
|
||||
**Method 2: Environment Variable**
|
||||
```bash
|
||||
export SSL_CERTFILE_PATH="/path/to/certificate.pem"
|
||||
```
|
||||
|
||||
#### Private Key
|
||||
|
||||
**Method 1: CLI Arguments**
|
||||
```bash
|
||||
litellm --ssl_keyfile_path /path/to/private_key.pem
|
||||
```
|
||||
|
||||
**Method 2: Environment Variable**
|
||||
```bash
|
||||
export SSL_KEYFILE_PATH="/path/to/private_key.pem"
|
||||
```
|
||||
|
||||
#### Certificate Bundle/Chain
|
||||
|
||||
**For client-to-proxy connections:**
|
||||
Use standard SSL certificate setup with intermediate certificates bundled in the certfile.
|
||||
|
||||
**For proxy-to-LLM provider connections:**
|
||||
|
||||
**Method 1: Config YAML**
|
||||
```yaml
|
||||
litellm_settings:
|
||||
ssl_verify: "/path/to/ca_bundle.pem"
|
||||
```
|
||||
|
||||
**Method 2: Environment Variable**
|
||||
```bash
|
||||
export SSL_CERT_FILE="/path/to/ca_bundle.pem"
|
||||
```
|
||||
|
||||
**Method 3: Client Certificate Authentication**
|
||||
```yaml
|
||||
litellm_settings:
|
||||
ssl_certificate: "/path/to/client_certificate.pem"
|
||||
```
|
||||
|
||||
or
|
||||
|
||||
```bash
|
||||
export SSL_CERTIFICATE="/path/to/client_certificate.pem"
|
||||
```
|
||||
|
||||
### Documentation Coverage
|
||||
|
||||
**Primary Documentation:**
|
||||
- `docs/my-website/docs/guides/security_settings.md` - SSL/TLS configuration guide
|
||||
|
||||
**Additional References:**
|
||||
- `litellm/proxy/proxy_cli.py` (lines 455-467) - CLI options
|
||||
- `docs/my-website/docs/completion/http_handler_config.md` - Custom HTTP handler configuration
|
||||
|
||||
---
|
||||
|
||||
## Data at Rest Encryption
|
||||
|
||||
### Does the product encrypt data at rest?
|
||||
|
||||
**Partially**. Only specific sensitive data is encrypted at rest.
|
||||
|
||||
### What data is stored in encrypted form?
|
||||
|
||||
#### Encrypted Data:
|
||||
1. **LLM API Keys** - Model credentials in `LiteLLM_ProxyModelTable.litellm_params`
|
||||
2. **Provider Credentials** - Stored in `LiteLLM_CredentialsTable.credential_values`
|
||||
3. **Configuration Secrets** - Sensitive config values in `LiteLLM_Config` table
|
||||
4. **Virtual Keys** - When using secret managers (optional feature)
|
||||
|
||||
#### NOT Encrypted:
|
||||
1. **Spend Logs** - Request/response data in `LiteLLM_SpendLogs`
|
||||
2. **Audit Logs** - Change history in `LiteLLM_AuditLog`
|
||||
3. **User/Team/Organization Data** - Metadata and configuration
|
||||
4. **Cached Prompts and Completions** - Cache data is stored in plaintext
|
||||
|
||||
### Cached prompts and completions?
|
||||
|
||||
**No**, cached prompts and completions are **NOT encrypted**.
|
||||
|
||||
Cache backends (Redis, S3, local disk) store data as plaintext JSON.
|
||||
|
||||
**Code References:**
|
||||
- `litellm/caching/redis_cache.py`
|
||||
- `litellm/caching/s3_cache.py`
|
||||
- `litellm/caching/caching.py`
|
||||
|
||||
### Configuration data?
|
||||
|
||||
**Partially encrypted**.
|
||||
|
||||
#### What IS Encrypted:
|
||||
- LLM API keys and credentials in model configurations
|
||||
- Sensitive values in `LiteLLM_Config` table
|
||||
- Credential values in `LiteLLM_CredentialsTable`
|
||||
|
||||
#### What is NOT Encrypted:
|
||||
- Model names and aliases
|
||||
- Rate limits and budget settings
|
||||
- User/team/organization metadata
|
||||
- Non-sensitive configuration parameters
|
||||
|
||||
**Code Reference:** `litellm/proxy/management_endpoints/model_management_endpoints.py` (lines 275-308)
|
||||
|
||||
### Log data?
|
||||
|
||||
**No**, log data is **NOT encrypted**.
|
||||
|
||||
Log data stored in database tables is in plaintext:
|
||||
- `LiteLLM_SpendLogs` - Contains request/response data, tokens, spend
|
||||
- `LiteLLM_ErrorLogs` - Error information
|
||||
- `LiteLLM_AuditLog` - Audit trail of changes
|
||||
|
||||
**Note:** You can disable logging to avoid storing sensitive data:
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
disable_spend_logs: True # Disable writing spend logs to DB
|
||||
disable_error_logs: True # Disable writing error logs to DB
|
||||
```
|
||||
|
||||
**Documentation:** `docs/my-website/docs/proxy/db_info.md` (lines 52-60)
|
||||
|
||||
### Where is it stored?
|
||||
|
||||
#### In the DB?
|
||||
|
||||
**Yes**, encrypted data is stored in PostgreSQL database.
|
||||
|
||||
**Key Tables with Encrypted Data:**
|
||||
- `LiteLLM_ProxyModelTable` - Model configurations with encrypted API keys
|
||||
- `LiteLLM_CredentialsTable` - Credential values
|
||||
- `LiteLLM_Config` - Configuration secrets
|
||||
|
||||
**Schema Reference:** `schema.prisma`
|
||||
|
||||
#### In the filesystem?
|
||||
|
||||
**No**, encrypted data is not stored in the filesystem by default.
|
||||
|
||||
**Note:** If using disk cache (`disk_cache_dir`), cached data is stored unencrypted.
|
||||
|
||||
#### Somewhere else?
|
||||
|
||||
**Optional:** When using secret managers (AWS Secrets Manager, Azure Key Vault, HashiCorp Vault), encrypted data can be stored externally.
|
||||
|
||||
**Configuration:**
|
||||
```yaml
|
||||
general_settings:
|
||||
key_management_system: "aws_secret_manager" # or "azure_key_vault", "hashicorp_vault"
|
||||
```
|
||||
|
||||
**Documentation:** `docs/my-website/docs/secret.md`
|
||||
|
||||
### How is it encrypted?
|
||||
|
||||
**Algorithm:** NaCl SecretBox (XSalsa20-Poly1305 AEAD)
|
||||
|
||||
**NOT AES-256** - LiteLLM uses NaCl (Networking and Cryptography Library) which provides:
|
||||
- XSalsa20 stream cipher
|
||||
- Poly1305 MAC for authentication
|
||||
- Equivalent security to AES-256
|
||||
|
||||
**Key Derivation:**
|
||||
1. Takes `LITELLM_SALT_KEY` (or `LITELLM_MASTER_KEY` if salt key not set)
|
||||
2. Hashes with SHA-256 to derive 256-bit encryption key
|
||||
3. Uses NaCl SecretBox for authenticated encryption
|
||||
|
||||
**Code Reference:** `litellm/proxy/common_utils/encrypt_decrypt_utils.py` (lines 69-112)
|
||||
|
||||
**Implementation:**
|
||||
```python
|
||||
import hashlib
|
||||
import nacl.secret
|
||||
|
||||
# Derive 256-bit key from salt
|
||||
hash_object = hashlib.sha256(signing_key.encode())
|
||||
hash_bytes = hash_object.digest()
|
||||
|
||||
# Create SecretBox and encrypt
|
||||
box = nacl.secret.SecretBox(hash_bytes)
|
||||
encrypted = box.encrypt(value_bytes)
|
||||
```
|
||||
|
||||
### Setting the Encryption Key
|
||||
|
||||
**Required Environment Variable:**
|
||||
```bash
|
||||
export LITELLM_SALT_KEY="your-strong-random-key-here"
|
||||
```
|
||||
|
||||
**Important Notes:**
|
||||
- ⚠️ **Must be set before adding any models**
|
||||
- ⚠️ **Never change this key** - encrypted data becomes unrecoverable
|
||||
- ⚠️ Use a strong random key (recommended: https://1password.com/password-generator/)
|
||||
- If not set, falls back to `LITELLM_MASTER_KEY`
|
||||
|
||||
**Documentation:** `docs/my-website/docs/proxy/prod.md` (section 8, lines 184-196)
|
||||
|
||||
### Documentation Coverage
|
||||
|
||||
**Primary Documentation:**
|
||||
- `docs/my-website/docs/proxy/prod.md` (section 8) - LITELLM_SALT_KEY setup
|
||||
- `docs/my-website/docs/secret.md` - Secret management systems
|
||||
- `docs/my-website/docs/proxy/db_info.md` - Database information
|
||||
|
||||
**Additional References:**
|
||||
- `security.md` - General security measures
|
||||
- `docs/my-website/docs/data_security.md` - Data privacy overview
|
||||
- `schema.prisma` - Database schema with encrypted fields
|
||||
|
||||
---
|
||||
|
||||
## Summary of Security Features
|
||||
|
||||
### ✅ Provided Out of the Box
|
||||
|
||||
1. **TLS/SSL encryption** for client-to-proxy connections
|
||||
2. **TLS encryption** for proxy-to-LLM provider connections (with connection pooling)
|
||||
3. **Encrypted storage** of LLM API keys and credentials
|
||||
4. **Support for TLS 1.2 and TLS 1.3**
|
||||
5. **Connection pooling** to reduce TLS handshake overhead
|
||||
|
||||
### ⚠️ Important Limitations
|
||||
|
||||
1. **Cached data is NOT encrypted** (Redis, S3, disk cache)
|
||||
2. **Log data is NOT encrypted** (spend logs, audit logs)
|
||||
3. **Request/response payloads in logs are NOT encrypted**
|
||||
4. **Uses NaCl SecretBox, NOT AES-256** (equivalent security)
|
||||
5. **TLS version not explicitly configured** - uses Python/system defaults
|
||||
|
||||
### 🔧 Configuration Requirements
|
||||
|
||||
**For Production Deployments:**
|
||||
|
||||
1. **Set LITELLM_SALT_KEY** before adding any models
|
||||
2. **Configure SSL certificates** for HTTPS client connections
|
||||
3. **Consider disabling logs** if they contain sensitive data
|
||||
4. **Use secret managers** for enhanced security (optional)
|
||||
5. **Configure CA bundles** if using custom certificates
|
||||
|
||||
---
|
||||
|
||||
## Quick Start Security Checklist
|
||||
|
||||
```bash
|
||||
# 1. Generate a strong salt key
|
||||
export LITELLM_SALT_KEY="$(openssl rand -base64 32)"
|
||||
|
||||
# 2. Set up SSL certificates (for HTTPS)
|
||||
export SSL_KEYFILE_PATH="/path/to/private_key.pem"
|
||||
export SSL_CERTFILE_PATH="/path/to/certificate.pem"
|
||||
|
||||
# 3. Configure database
|
||||
export DATABASE_URL="postgresql://user:password@host:port/dbname"
|
||||
|
||||
# 4. (Optional) Disable logs if they contain sensitive data
|
||||
# Add to config.yaml:
|
||||
# general_settings:
|
||||
# disable_spend_logs: True
|
||||
# disable_error_logs: True
|
||||
|
||||
# 5. Start LiteLLM Proxy
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Additional Resources
|
||||
|
||||
- **LiteLLM Documentation:** https://docs.litellm.ai/
|
||||
- **Security Settings Guide:** https://docs.litellm.ai/docs/guides/security_settings
|
||||
- **Production Deployment:** https://docs.litellm.ai/docs/proxy/prod
|
||||
- **Secret Management:** https://docs.litellm.ai/docs/secret
|
||||
|
||||
For security inquiries: support@berri.ai
|
||||
|
||||
|
|
@ -674,6 +674,7 @@ const sidebars = {
|
|||
items: [
|
||||
"data_security",
|
||||
"data_retention",
|
||||
"proxy/security_encryption_faq",
|
||||
"migration_policy",
|
||||
{
|
||||
type: "category",
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@
|
|||
| gpt-3.5-turbo-16k | `completion('gpt-3.5-turbo-16k', messages)` | `os.environ['OPENAI_API_KEY']` |
|
||||
| gpt-3.5-turbo-16k-0613 | `completion('gpt-3.5-turbo-16k-0613', messages)` | `os.environ['OPENAI_API_KEY']` |
|
||||
| gpt-4 | `completion('gpt-4', messages)` | `os.environ['OPENAI_API_KEY']` |
|
||||
| gpt-5-pro | `completion('gpt-5-pro', messages)` | `os.environ['OPENAI_API_KEY']` |
|
||||
|
||||
## Azure OpenAI Chat Completion Models
|
||||
For Azure calls add the `azure/` prefix to `model`. If your azure deployment name is `gpt-v-2` set `model` = `azure/gpt-v-2`
|
||||
|
|
|
|||
|
|
@ -0,0 +1,3 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "mcp_tool_permissions" JSONB;
|
||||
|
||||
|
|
@ -156,6 +156,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
object_permission_id String @id @default(uuid())
|
||||
mcp_servers String[] @default([])
|
||||
mcp_access_groups String[] @default([])
|
||||
mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]}
|
||||
vector_stores String[] @default([])
|
||||
teams LiteLLM_TeamTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
|
|
|
|||
|
|
@ -5,8 +5,9 @@ LiteLLM Proxy uses this MCP Client to connnect to other MCP servers.
|
|||
import asyncio
|
||||
import base64
|
||||
from datetime import timedelta
|
||||
from typing import Dict, List, Optional, Union
|
||||
from typing import Callable, Dict, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
from mcp import ClientSession, StdioServerParameters
|
||||
from mcp.client.sse import sse_client
|
||||
from mcp.client.stdio import stdio_client
|
||||
|
|
@ -17,6 +18,8 @@ from mcp.types import TextContent
|
|||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
|
||||
from litellm.types.llms.custom_http import VerifyTypes
|
||||
from litellm.types.mcp import (
|
||||
MCPAuth,
|
||||
MCPAuthType,
|
||||
|
|
@ -48,6 +51,7 @@ class MCPClient:
|
|||
timeout: float = 60.0,
|
||||
stdio_config: Optional[MCPStdioConfig] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
ssl_verify: Optional[VerifyTypes] = None,
|
||||
):
|
||||
self.server_url: str = server_url
|
||||
self.transport_type: MCPTransport = transport_type
|
||||
|
|
@ -62,6 +66,7 @@ class MCPClient:
|
|||
self._task: Optional[asyncio.Task] = None
|
||||
self.stdio_config: Optional[MCPStdioConfig] = stdio_config
|
||||
self.extra_headers: Optional[Dict[str, str]] = extra_headers
|
||||
self.ssl_verify: Optional[VerifyTypes] = ssl_verify
|
||||
# handle the basic auth value if provided
|
||||
if auth_value:
|
||||
self.update_auth_value(auth_value)
|
||||
|
|
@ -104,10 +109,12 @@ class MCPClient:
|
|||
await self._session.initialize()
|
||||
elif self.transport_type == MCPTransport.sse:
|
||||
headers = self._get_auth_headers()
|
||||
httpx_client_factory = self._create_httpx_client_factory()
|
||||
self._transport_ctx = sse_client(
|
||||
url=self.server_url,
|
||||
timeout=self.timeout,
|
||||
headers=headers,
|
||||
httpx_client_factory=httpx_client_factory,
|
||||
)
|
||||
self._transport = await self._transport_ctx.__aenter__()
|
||||
self._session_ctx = ClientSession(
|
||||
|
|
@ -117,13 +124,15 @@ class MCPClient:
|
|||
await self._session.initialize()
|
||||
else: # http
|
||||
headers = self._get_auth_headers()
|
||||
httpx_client_factory = self._create_httpx_client_factory()
|
||||
verbose_logger.debug(
|
||||
"litellm headers for streamablehttp_client: ", headers
|
||||
"litellm headers for streamablehttp_client: %s", headers
|
||||
)
|
||||
self._transport_ctx = streamablehttp_client(
|
||||
url=self.server_url,
|
||||
timeout=timedelta(seconds=self.timeout),
|
||||
headers=headers,
|
||||
httpx_client_factory=httpx_client_factory,
|
||||
)
|
||||
self._transport = await self._transport_ctx.__aenter__()
|
||||
self._session_ctx = ClientSession(
|
||||
|
|
@ -215,6 +224,41 @@ class MCPClient:
|
|||
|
||||
return headers
|
||||
|
||||
def _create_httpx_client_factory(self) -> Callable[..., httpx.AsyncClient]:
|
||||
"""
|
||||
Create a custom httpx client factory that uses LiteLLM's SSL configuration.
|
||||
|
||||
This factory follows the same CA bundle path logic as http_handler.py:
|
||||
1. Check ssl_verify parameter (can be SSLContext, bool, or path to CA bundle)
|
||||
2. Check SSL_VERIFY environment variable
|
||||
3. Check SSL_CERT_FILE environment variable
|
||||
4. Fall back to certifi CA bundle
|
||||
"""
|
||||
|
||||
def factory(
|
||||
*,
|
||||
headers: Optional[Dict[str, str]] = None,
|
||||
timeout: Optional[httpx.Timeout] = None,
|
||||
auth: Optional[httpx.Auth] = None,
|
||||
) -> httpx.AsyncClient:
|
||||
"""Create an httpx.AsyncClient with LiteLLM's SSL configuration."""
|
||||
# Get unified SSL configuration using the same logic as http_handler.py
|
||||
ssl_config = get_ssl_configuration(self.ssl_verify)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"MCP client using SSL configuration: {type(ssl_config).__name__}"
|
||||
)
|
||||
|
||||
return httpx.AsyncClient(
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
auth=auth,
|
||||
verify=ssl_config,
|
||||
follow_redirects=True,
|
||||
)
|
||||
|
||||
return factory
|
||||
|
||||
async def list_tools(self) -> List[MCPTool]:
|
||||
"""List available tools from the server."""
|
||||
if not self._session:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
|
||||
|
||||
import httpx
|
||||
from openai.types.responses import ResponseReasoningItem
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.azure.common_utils import BaseAzureLLM
|
||||
|
|
@ -38,6 +39,50 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
model = model.replace("o_series/", "")
|
||||
return model
|
||||
|
||||
def _handle_reasoning_item(self, item: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Handle reasoning items specifically to filter out status=None using OpenAI's model.
|
||||
Issue: https://github.com/BerriAI/litellm/issues/13484
|
||||
OpenAI API does not accept ReasoningItem(status=None), so we need to:
|
||||
1. Check if the item is a reasoning type
|
||||
2. Create a ResponseReasoningItem object with the item data
|
||||
3. Convert it back to dict with exclude_none=True to filter None values
|
||||
"""
|
||||
if item.get("type") == "reasoning":
|
||||
try:
|
||||
# Ensure required fields are present for ResponseReasoningItem
|
||||
item_data = dict(item)
|
||||
if "id" not in item_data:
|
||||
item_data["id"] = f"reasoning_{hash(str(item_data))}"
|
||||
if "summary" not in item_data:
|
||||
item_data["summary"] = (
|
||||
item_data.get("reasoning_content", "")[:100] + "..."
|
||||
if len(item_data.get("reasoning_content", "")) > 100
|
||||
else item_data.get("reasoning_content", "")
|
||||
)
|
||||
|
||||
# Create ResponseReasoningItem object from the item data
|
||||
reasoning_item = ResponseReasoningItem(**item_data)
|
||||
|
||||
# Convert back to dict with exclude_none=True to exclude None fields
|
||||
dict_reasoning_item = reasoning_item.model_dump(exclude_none=True)
|
||||
dict_reasoning_item.pop("status", None)
|
||||
|
||||
return dict_reasoning_item
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
f"Failed to create ResponseReasoningItem, falling back to manual filtering: {e}"
|
||||
)
|
||||
# Fallback: manually filter out known None fields
|
||||
filtered_item = {
|
||||
k: v
|
||||
for k, v in item.items()
|
||||
if v is not None
|
||||
or k not in {"status", "content", "encrypted_content"}
|
||||
}
|
||||
return filtered_item
|
||||
return item
|
||||
|
||||
def transform_responses_api_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -48,12 +93,13 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
) -> Dict:
|
||||
"""No transform applied since inputs are in OpenAI spec already"""
|
||||
stripped_model_name = self.get_stripped_model_name(model)
|
||||
return dict(
|
||||
ResponsesAPIRequestParams(
|
||||
model=stripped_model_name,
|
||||
input=input,
|
||||
**response_api_optional_request_params,
|
||||
)
|
||||
|
||||
return super().transform_responses_api_request(
|
||||
model=stripped_model_name,
|
||||
input=input,
|
||||
response_api_optional_request_params=response_api_optional_request_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
def get_complete_url(
|
||||
|
|
@ -217,15 +263,15 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
at the correct location (before any query parameters).
|
||||
"""
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
|
||||
# Parse the URL to separate its components
|
||||
parsed_url = urlparse(api_base)
|
||||
|
||||
|
||||
# Insert the response_id and /cancel at the end of the path component
|
||||
# Remove trailing slash if present to avoid double slashes
|
||||
path = parsed_url.path.rstrip("/")
|
||||
new_path = f"{path}/{response_id}/cancel"
|
||||
|
||||
|
||||
# Reconstruct the URL with all original components but with the modified path
|
||||
cancel_url = urlunparse(
|
||||
(
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
|
||||
from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error
|
||||
from litellm.llms.openai.openai import OpenAIConfig
|
||||
from litellm.llms.xai.chat.transformation import XAIChatConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse, ProviderField
|
||||
|
|
@ -35,9 +36,24 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
for param in supported_params:
|
||||
if param != "tool_choice":
|
||||
filtered_supported_params.append(param)
|
||||
return filtered_supported_params
|
||||
supported_params = filtered_supported_params
|
||||
|
||||
# Filter out unsupported parameters for specific models
|
||||
if not self._supports_stop_reason(model):
|
||||
supported_params = [param for param in supported_params if param != "stop"]
|
||||
|
||||
return supported_params
|
||||
|
||||
def _supports_stop_reason(self, model: str) -> bool:
|
||||
"""
|
||||
Check if the model supports stop tokens.
|
||||
"""
|
||||
if "grok" in model:
|
||||
# Reuse Xai method for Grok model
|
||||
xai_config = XAIChatConfig()
|
||||
return xai_config._supports_stop_reason(model)
|
||||
return True
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
@ -53,9 +69,7 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
else:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
|
||||
headers["Content-Type"] = (
|
||||
"application/json" # tell Azure AI Studio to expect JSON
|
||||
)
|
||||
headers["Content-Type"] = "application/json" # tell Azure AI Studio to expect JSON
|
||||
|
||||
return headers
|
||||
|
||||
|
|
@ -65,10 +79,7 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
"""
|
||||
parsed_url = urlparse(api_base)
|
||||
host = parsed_url.hostname
|
||||
if host and (
|
||||
host.endswith(".services.ai.azure.com")
|
||||
or host.endswith(".openai.azure.com")
|
||||
):
|
||||
if host and (host.endswith(".services.ai.azure.com") or host.endswith(".openai.azure.com")):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
|
@ -115,13 +126,9 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
|
||||
# Add the path to the base URL
|
||||
if "services.ai.azure.com" in api_base:
|
||||
new_url = _add_path_to_api_base(
|
||||
api_base=api_base, ending_path="/models/chat/completions"
|
||||
)
|
||||
new_url = _add_path_to_api_base(api_base=api_base, ending_path="/models/chat/completions")
|
||||
else:
|
||||
new_url = _add_path_to_api_base(
|
||||
api_base=api_base, ending_path="/chat/completions"
|
||||
)
|
||||
new_url = _add_path_to_api_base(api_base=api_base, ending_path="/chat/completions")
|
||||
|
||||
# Use the new query_params dictionary
|
||||
final_url = httpx.URL(new_url).copy_with(params=query_params)
|
||||
|
|
@ -191,11 +198,7 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
dynamic_api_key = api_key or get_secret_str("AZURE_AI_API_KEY")
|
||||
|
||||
if self._is_azure_openai_model(model=model, api_base=api_base):
|
||||
verbose_logger.debug(
|
||||
"Model={} is Azure OpenAI model. Setting custom_llm_provider='azure'.".format(
|
||||
model
|
||||
)
|
||||
)
|
||||
verbose_logger.debug("Model={} is Azure OpenAI model. Setting custom_llm_provider='azure'.".format(model))
|
||||
custom_llm_provider = "azure"
|
||||
return api_base, dynamic_api_key, custom_llm_provider
|
||||
|
||||
|
|
@ -211,9 +214,7 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
if extra_body and isinstance(extra_body, dict):
|
||||
optional_params.update(extra_body)
|
||||
optional_params.pop("max_retries", None)
|
||||
return super().transform_request(
|
||||
model, messages, optional_params, litellm_params, headers
|
||||
)
|
||||
return super().transform_request(model, messages, optional_params, litellm_params, headers)
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
|
|
@ -252,47 +253,30 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
|
||||
if should_drop_params and "Extra inputs are not permitted" in error_text:
|
||||
return True
|
||||
elif (
|
||||
"unknown field: parameter index is not a valid field" in error_text
|
||||
): # remove index from tool calls
|
||||
elif "unknown field: parameter index is not a valid field" in error_text: # remove index from tool calls
|
||||
return True
|
||||
elif (
|
||||
AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value
|
||||
in error_text
|
||||
AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value in error_text
|
||||
): # remove extra-parameters from tool calls
|
||||
return True
|
||||
return super().should_retry_llm_api_inside_llm_translation_on_http_error(
|
||||
e=e, litellm_params=litellm_params
|
||||
)
|
||||
return super().should_retry_llm_api_inside_llm_translation_on_http_error(e=e, litellm_params=litellm_params)
|
||||
|
||||
@property
|
||||
def max_retry_on_unprocessable_entity_error(self) -> int:
|
||||
return 2
|
||||
|
||||
def transform_request_on_unprocessable_entity_error(
|
||||
self, e: httpx.HTTPStatusError, request_data: dict
|
||||
) -> dict:
|
||||
def transform_request_on_unprocessable_entity_error(self, e: httpx.HTTPStatusError, request_data: dict) -> dict:
|
||||
_messages = cast(Optional[List[AllMessageValues]], request_data.get("messages"))
|
||||
if (
|
||||
"unknown field: parameter index is not a valid field" in e.response.text
|
||||
and _messages is not None
|
||||
):
|
||||
if "unknown field: parameter index is not a valid field" in e.response.text and _messages is not None:
|
||||
litellm.remove_index_from_tool_calls(
|
||||
messages=_messages,
|
||||
)
|
||||
elif (
|
||||
AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value
|
||||
in e.response.text
|
||||
):
|
||||
request_data = self._drop_extra_params_from_request_data(
|
||||
request_data, e.response.text
|
||||
)
|
||||
elif AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value in e.response.text:
|
||||
request_data = self._drop_extra_params_from_request_data(request_data, e.response.text)
|
||||
data = drop_params_from_unprocessable_entity_error(e=e, data=request_data)
|
||||
return data
|
||||
|
||||
def _drop_extra_params_from_request_data(
|
||||
self, request_data: dict, error_text: str
|
||||
) -> dict:
|
||||
def _drop_extra_params_from_request_data(self, request_data: dict, error_text: str) -> dict:
|
||||
params_to_drop = self._extract_params_to_drop_from_error_text(error_text)
|
||||
if params_to_drop:
|
||||
for param in params_to_drop:
|
||||
|
|
@ -300,9 +284,7 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
request_data.pop(param, None)
|
||||
return request_data
|
||||
|
||||
def _extract_params_to_drop_from_error_text(
|
||||
self, error_text: str
|
||||
) -> Optional[List[str]]:
|
||||
def _extract_params_to_drop_from_error_text(self, error_text: str) -> Optional[List[str]]:
|
||||
"""
|
||||
Error text looks like this"
|
||||
"Extra parameters ['stream_options', 'extra-parameters'] are not allowed when extra-parameters is not set or set to be 'error'.
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
|
|||
"presence_penalty",
|
||||
"frequency_penalty",
|
||||
"top_logprobs",
|
||||
"stop",
|
||||
]
|
||||
|
||||
return [
|
||||
|
|
|
|||
|
|
@ -1,12 +1,4 @@
|
|||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Dict,
|
||||
Optional,
|
||||
Union,
|
||||
cast,
|
||||
get_type_hints,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional, Union, cast, get_type_hints
|
||||
|
||||
import httpx
|
||||
from openai.types.responses import ResponseReasoningItem
|
||||
|
|
@ -127,7 +119,6 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
2. Create a ResponseReasoningItem object with the item data
|
||||
3. Convert it back to dict with exclude_none=True to filter None values
|
||||
"""
|
||||
verbose_logger.debug(f"Handling reasoning item: {item}")
|
||||
if item.get("type") == "reasoning":
|
||||
try:
|
||||
# Ensure required fields are present for ResponseReasoningItem
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
|
|||
"""
|
||||
|
||||
import re
|
||||
from typing import List, Optional, Tuple
|
||||
from typing import List, Optional, Tuple, Literal
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.llms.vertex_ai import CachedContentRequestBody
|
||||
|
|
@ -155,13 +155,18 @@ def separate_cached_messages(
|
|||
|
||||
|
||||
def transform_openai_messages_to_gemini_context_caching(
|
||||
model: str, messages: List[AllMessageValues], cache_key: str
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
cache_key: str,
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
) -> CachedContentRequestBody:
|
||||
# Extract TTL from cached messages BEFORE system message transformation
|
||||
ttl = extract_ttl_from_cached_messages(messages)
|
||||
|
||||
supports_system_message = get_supports_system_message(
|
||||
model=model, custom_llm_provider="gemini"
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
|
||||
transformed_system_messages, new_messages = _transform_system_message(
|
||||
|
|
@ -170,9 +175,14 @@ def transform_openai_messages_to_gemini_context_caching(
|
|||
|
||||
transformed_messages = _gemini_convert_messages_with_history(messages=new_messages)
|
||||
|
||||
model_name = "models/{}".format(model)
|
||||
|
||||
if custom_llm_provider == "vertex_ai" or custom_llm_provider == "vertex_ai_beta":
|
||||
model_name = f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/{model_name}"
|
||||
|
||||
data = CachedContentRequestBody(
|
||||
contents=transformed_messages,
|
||||
model="models/{}".format(model),
|
||||
model=model_name,
|
||||
displayName=cache_key,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -41,8 +41,11 @@ class ContextCachingEndpoints(VertexBase):
|
|||
def _get_token_and_url_context_caching(
|
||||
self,
|
||||
gemini_api_key: Optional[str],
|
||||
custom_llm_provider: Literal["gemini"],
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
api_base: Optional[str],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
) -> Tuple[Optional[str], str]:
|
||||
"""
|
||||
Internal function. Returns the token and url for the call.
|
||||
|
|
@ -58,9 +61,15 @@ class ContextCachingEndpoints(VertexBase):
|
|||
url = "https://generativelanguage.googleapis.com/v1beta/{}?key={}".format(
|
||||
endpoint, gemini_api_key
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
auth_header = vertex_auth_header
|
||||
endpoint = "cachedContents"
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}"
|
||||
else:
|
||||
raise NotImplementedError
|
||||
auth_header = vertex_auth_header
|
||||
endpoint = "cachedContents"
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1beta1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}"
|
||||
|
||||
|
||||
return self._check_custom_proxy(
|
||||
api_base=api_base,
|
||||
|
|
@ -80,6 +89,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
api_key: str,
|
||||
api_base: Optional[str],
|
||||
logging_obj: Logging,
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Checks if content already cached.
|
||||
|
|
@ -94,8 +107,11 @@ class ContextCachingEndpoints(VertexBase):
|
|||
|
||||
_, url = self._get_token_and_url_context_caching(
|
||||
gemini_api_key=api_key,
|
||||
custom_llm_provider="gemini",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
try:
|
||||
## LOGGING
|
||||
|
|
@ -145,6 +161,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
api_key: str,
|
||||
api_base: Optional[str],
|
||||
logging_obj: Logging,
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str]
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Checks if content already cached.
|
||||
|
|
@ -159,8 +179,11 @@ class ContextCachingEndpoints(VertexBase):
|
|||
|
||||
_, url = self._get_token_and_url_context_caching(
|
||||
gemini_api_key=api_key,
|
||||
custom_llm_provider="gemini",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
try:
|
||||
## LOGGING
|
||||
|
|
@ -212,6 +235,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
client: Optional[HTTPHandler],
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
logging_obj: Logging,
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
extra_headers: Optional[dict] = None,
|
||||
cached_content: Optional[str] = None,
|
||||
) -> Tuple[List[AllMessageValues], dict, Optional[str]]:
|
||||
|
|
@ -240,8 +267,11 @@ class ContextCachingEndpoints(VertexBase):
|
|||
## AUTHORIZATION ##
|
||||
token, url = self._get_token_and_url_context_caching(
|
||||
gemini_api_key=api_key,
|
||||
custom_llm_provider="gemini",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
|
||||
headers = {
|
||||
|
|
@ -273,6 +303,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
if google_cache_name:
|
||||
return non_cached_messages, optional_params, google_cache_name
|
||||
|
|
@ -280,7 +314,12 @@ class ContextCachingEndpoints(VertexBase):
|
|||
## TRANSFORM REQUEST
|
||||
cached_content_request_body = (
|
||||
transform_openai_messages_to_gemini_context_caching(
|
||||
model=model, messages=cached_messages, cache_key=generated_cache_key
|
||||
model=model,
|
||||
messages=cached_messages,
|
||||
cache_key=generated_cache_key,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -328,6 +367,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
client: Optional[AsyncHTTPHandler],
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
logging_obj: Logging,
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
extra_headers: Optional[dict] = None,
|
||||
cached_content: Optional[str] = None,
|
||||
) -> Tuple[List[AllMessageValues], dict, Optional[str]]:
|
||||
|
|
@ -356,8 +399,11 @@ class ContextCachingEndpoints(VertexBase):
|
|||
## AUTHORIZATION ##
|
||||
token, url = self._get_token_and_url_context_caching(
|
||||
gemini_api_key=api_key,
|
||||
custom_llm_provider="gemini",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
|
||||
headers = {
|
||||
|
|
@ -386,6 +432,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
|
||||
if google_cache_name:
|
||||
|
|
@ -394,7 +444,12 @@ class ContextCachingEndpoints(VertexBase):
|
|||
## TRANSFORM REQUEST
|
||||
cached_content_request_body = (
|
||||
transform_openai_messages_to_gemini_context_caching(
|
||||
model=model, messages=cached_messages, cache_key=generated_cache_key
|
||||
model=model,
|
||||
messages=cached_messages,
|
||||
cache_key=generated_cache_key,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -514,34 +514,35 @@ def sync_transform_request_body(
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
litellm_params: dict,
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
) -> RequestBody:
|
||||
from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints
|
||||
|
||||
context_caching_endpoints = ContextCachingEndpoints()
|
||||
|
||||
if gemini_api_key is not None:
|
||||
(
|
||||
messages,
|
||||
optional_params,
|
||||
cached_content,
|
||||
) = context_caching_endpoints.check_and_create_cache(
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=gemini_api_key,
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
extra_headers=extra_headers,
|
||||
cached_content=optional_params.pop("cached_content", None),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
else: # [TODO] implement context caching for gemini as well
|
||||
cached_content = None
|
||||
if "cached_content" in optional_params:
|
||||
cached_content = optional_params.pop("cached_content")
|
||||
elif "cachedContent" in optional_params:
|
||||
cached_content = optional_params.pop("cachedContent")
|
||||
(
|
||||
messages,
|
||||
optional_params,
|
||||
cached_content,
|
||||
) = context_caching_endpoints.check_and_create_cache(
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=gemini_api_key or "dummy",
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
extra_headers=extra_headers,
|
||||
cached_content=optional_params.pop("cached_content", None),
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header,
|
||||
)
|
||||
|
||||
|
||||
return _transform_request_body(
|
||||
messages=messages,
|
||||
|
|
@ -565,34 +566,34 @@ async def async_transform_request_body(
|
|||
logging_obj: litellm.litellm_core_utils.litellm_logging.Logging, # type: ignore
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
litellm_params: dict,
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
) -> RequestBody:
|
||||
from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints
|
||||
|
||||
context_caching_endpoints = ContextCachingEndpoints()
|
||||
|
||||
if gemini_api_key is not None:
|
||||
(
|
||||
messages,
|
||||
optional_params,
|
||||
cached_content,
|
||||
) = await context_caching_endpoints.async_check_and_create_cache(
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=gemini_api_key,
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
extra_headers=extra_headers,
|
||||
cached_content=optional_params.pop("cached_content", None),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
else: # [TODO] implement context caching for gemini as well
|
||||
cached_content = None
|
||||
if "cached_content" in optional_params:
|
||||
cached_content = optional_params.pop("cached_content")
|
||||
elif "cachedContent" in optional_params:
|
||||
cached_content = optional_params.pop("cachedContent")
|
||||
(
|
||||
messages,
|
||||
optional_params,
|
||||
cached_content,
|
||||
) = await context_caching_endpoints.async_check_and_create_cache(
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=gemini_api_key or "dummy",
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
extra_headers=extra_headers,
|
||||
cached_content=optional_params.pop("cached_content", None),
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header,
|
||||
)
|
||||
|
||||
return _transform_request_body(
|
||||
messages=messages,
|
||||
|
|
|
|||
|
|
@ -1792,7 +1792,6 @@ class VertexLLM(VertexBase):
|
|||
gemini_api_key: Optional[str] = None,
|
||||
extra_headers: Optional[dict] = None,
|
||||
) -> CustomStreamWrapper:
|
||||
request_body = await async_transform_request_body(**data) # type: ignore
|
||||
|
||||
should_use_v1beta1_features = self.is_using_v1beta1_features(
|
||||
optional_params=optional_params
|
||||
|
|
@ -1826,6 +1825,13 @@ class VertexLLM(VertexBase):
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
request_body = await async_transform_request_body(
|
||||
**data,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=auth_header) # type: ignore
|
||||
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
|
|
@ -1913,7 +1919,12 @@ class VertexLLM(VertexBase):
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
request_body = await async_transform_request_body(**data) # type: ignore
|
||||
request_body = await async_transform_request_body(
|
||||
**data,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=auth_header) # type: ignore
|
||||
|
||||
_async_client_params = {}
|
||||
if timeout:
|
||||
_async_client_params["timeout"] = timeout
|
||||
|
|
@ -2088,7 +2099,11 @@ class VertexLLM(VertexBase):
|
|||
)
|
||||
|
||||
## TRANSFORMATION ##
|
||||
data = sync_transform_request_body(**transform_request_params)
|
||||
data = sync_transform_request_body(
|
||||
**transform_request_params,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=auth_header)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
|
|||
|
|
@ -2872,7 +2872,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
custom_llm_provider=custom_llm_provider, # type: ignore
|
||||
client=client,
|
||||
api_base=api_base,
|
||||
extra_headers=extra_headers,
|
||||
extra_headers=headers,
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
|
|
@ -2941,7 +2941,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
custom_llm_provider=custom_llm_provider, # type: ignore
|
||||
client=client,
|
||||
api_base=api_base,
|
||||
extra_headers=extra_headers,
|
||||
extra_headers=headers,
|
||||
)
|
||||
elif "openai" in model:
|
||||
# Vertex Model Garden - OpenAI compatible models
|
||||
|
|
|
|||
|
|
@ -333,6 +333,139 @@ class MCPRequestHandler:
|
|||
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
async def _get_key_object_permission(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
):
|
||||
"""Helper to get key object_permission from cache or DB."""
|
||||
from litellm.proxy.auth.auth_checks import get_object_permission
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if not user_api_key_auth:
|
||||
return None
|
||||
|
||||
# Already loaded
|
||||
if user_api_key_auth.object_permission:
|
||||
return user_api_key_auth.object_permission
|
||||
|
||||
# Need to fetch from DB
|
||||
if user_api_key_auth.object_permission_id and prisma_client:
|
||||
return await get_object_permission(
|
||||
object_permission_id=user_api_key_auth.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def _get_team_object_permission(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
):
|
||||
"""Helper to get team object_permission from cache or DB."""
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client:
|
||||
return None
|
||||
|
||||
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
|
||||
team_id=user_api_key_auth.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
return team_obj.object_permission if team_obj else None
|
||||
|
||||
@staticmethod
|
||||
async def get_allowed_tools_for_server(
|
||||
server_id: str,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> Optional[List[str]]:
|
||||
"""
|
||||
Get list of allowed tool names for a specific server based on key/team permissions.
|
||||
Follows same inheritance logic as get_allowed_mcp_servers.
|
||||
|
||||
Args:
|
||||
server_id: Server ID to check permissions for
|
||||
user_api_key_auth: User auth
|
||||
|
||||
Returns:
|
||||
List[str] if restrictions exist, None if no restrictions (allow all)
|
||||
"""
|
||||
if not user_api_key_auth:
|
||||
return None
|
||||
|
||||
try:
|
||||
# Get key and team object permissions
|
||||
key_obj_perm = await MCPRequestHandler._get_key_object_permission(user_api_key_auth)
|
||||
team_obj_perm = await MCPRequestHandler._get_team_object_permission(user_api_key_auth)
|
||||
|
||||
# Extract tool permissions for this server
|
||||
key_tools = key_obj_perm.mcp_tool_permissions.get(server_id) if key_obj_perm and key_obj_perm.mcp_tool_permissions else None
|
||||
team_tools = team_obj_perm.mcp_tool_permissions.get(server_id) if team_obj_perm and team_obj_perm.mcp_tool_permissions else None
|
||||
|
||||
# Apply same inheritance logic as get_allowed_mcp_servers
|
||||
if team_tools:
|
||||
if key_tools:
|
||||
# Both have restrictions → intersection
|
||||
return list(set(team_tools) & set(key_tools))
|
||||
else:
|
||||
# Only team has restrictions → inherit from team
|
||||
return team_tools
|
||||
else:
|
||||
# No team restrictions → use key restrictions
|
||||
return key_tools
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get allowed tools for server: {str(e)}")
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def is_tool_allowed_for_server(
|
||||
tool_name: str,
|
||||
server_id: str,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if a specific tool is allowed for a server based on key/team permissions.
|
||||
|
||||
Args:
|
||||
tool_name: Name of the tool to check
|
||||
server_id: Server ID
|
||||
user_api_key_auth: User auth
|
||||
|
||||
Returns:
|
||||
True if allowed, False if blocked
|
||||
"""
|
||||
allowed_tools = await MCPRequestHandler.get_allowed_tools_for_server(
|
||||
server_id=server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
# None means no restrictions (allow all)
|
||||
if allowed_tools is None:
|
||||
return True
|
||||
|
||||
# Empty list means no tools allowed
|
||||
if not allowed_tools:
|
||||
return False
|
||||
|
||||
# Check if tool is in allowed list
|
||||
return tool_name in allowed_tools
|
||||
|
||||
@staticmethod
|
||||
def is_tool_allowed(
|
||||
allowed_mcp_servers: List[str],
|
||||
|
|
|
|||
|
|
@ -602,6 +602,45 @@ class MCPServerManager:
|
|||
return tool_name not in server.disallowed_tools
|
||||
return True
|
||||
|
||||
async def check_tool_permission_for_key_team(
|
||||
self,
|
||||
tool_name: str,
|
||||
server: MCPServer,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> None:
|
||||
"""
|
||||
Check if a tool is allowed based on key/team object_permission.mcp_tool_permissions.
|
||||
Uses MCPRequestHandler.is_tool_allowed_for_server for consistent inheritance logic.
|
||||
Raises HTTPException if tool is not allowed.
|
||||
|
||||
Args:
|
||||
tool_name: Name of the tool to check
|
||||
server: MCPServer object
|
||||
user_api_key_auth: User authentication
|
||||
|
||||
Raises:
|
||||
HTTPException: If tool is not allowed for this key/team
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
if not user_api_key_auth:
|
||||
return
|
||||
|
||||
# Check if tool is allowed
|
||||
is_allowed = await MCPRequestHandler.is_tool_allowed_for_server(
|
||||
tool_name=tool_name,
|
||||
server_id=server.server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
if not is_allowed:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": f"Tool '{tool_name}' is not allowed for your key/team on server '{server.name}'. Contact proxy admin for access."
|
||||
},
|
||||
)
|
||||
|
||||
async def pre_call_tool_check(
|
||||
self,
|
||||
name: str,
|
||||
|
|
@ -621,6 +660,13 @@ class MCPServerManager:
|
|||
},
|
||||
)
|
||||
|
||||
## check tool-level permissions from object_permission
|
||||
await self.check_tool_permission_for_key_team(
|
||||
tool_name=name,
|
||||
server=server,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
pre_hook_kwargs = {
|
||||
"name": name,
|
||||
"arguments": arguments,
|
||||
|
|
|
|||
|
|
@ -499,6 +499,13 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
filtered_tools = await filter_tools_by_key_team_permissions(
|
||||
tools=filtered_tools,
|
||||
server_id=server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
all_tools.extend(filtered_tools)
|
||||
|
||||
verbose_logger.debug(
|
||||
|
|
@ -515,6 +522,25 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
return all_tools
|
||||
|
||||
async def filter_tools_by_key_team_permissions(
|
||||
tools: List[MCPTool],
|
||||
server_id: str,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> List[MCPTool]:
|
||||
"""Filter tools based on key/team mcp_tool_permissions."""
|
||||
# Filter by key/team tool-level permissions
|
||||
allowed_tool_names = await MCPRequestHandler.get_allowed_tools_for_server(
|
||||
server_id=server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
if allowed_tool_names is not None:
|
||||
filtered_tools = [t for t in tools if t.name in allowed_tool_names]
|
||||
else:
|
||||
# No restrictions, return all tools
|
||||
filtered_tools = tools
|
||||
|
||||
return filtered_tools
|
||||
|
||||
async def _list_mcp_tools(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
mcp_auth_header: Optional[str] = None,
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
|
|
@ -1 +1 @@
|
|||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[185],{77401:function(n,e,t){Promise.resolve().then(t.t.bind(t,39974,23)),Promise.resolve().then(t.t.bind(t,2778,23))},2778:function(){},39974:function(n){n.exports={style:{fontFamily:"'__Inter_1c856b', '__Inter_Fallback_1c856b'",fontStyle:"normal"},className:"__className_1c856b"}}},function(n){n.O(0,[919,986,971,117,744],function(){return n(n.s=77401)}),_N_E=n.O()}]);
|
||||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[185],{85210:function(n,e,t){Promise.resolve().then(t.t.bind(t,39974,23)),Promise.resolve().then(t.t.bind(t,2778,23))},2778:function(){},39974:function(n){n.exports={style:{fontFamily:"'__Inter_1c856b', '__Inter_Fallback_1c856b'",fontStyle:"normal"},className:"__className_1c856b"}}},function(n){n.O(0,[919,986,971,117,744],function(){return n(n.s=85210)}),_N_E=n.O()}]);
|
||||
|
|
@ -1 +1 @@
|
|||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[418],{96422:function(e,n,t){Promise.resolve().then(t.bind(t,52829))},52829:function(e,n,t){"use strict";t.r(n),t.d(n,{default:function(){return f}});var u=t(57437),s=t(2265),c=t(99376),r=t(72162);function f(){let e=(0,c.useSearchParams)().get("key"),[n,t]=(0,s.useState)(null);return(0,s.useEffect)(()=>{e&&t(e)},[e]),(0,u.jsx)(r.Z,{accessToken:n})}}},function(e){e.O(0,[50,521,154,162,971,117,744],function(){return e(e.s=96422)}),_N_E=e.O()}]);
|
||||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[418],{67355:function(e,n,t){Promise.resolve().then(t.bind(t,52829))},52829:function(e,n,t){"use strict";t.r(n),t.d(n,{default:function(){return f}});var u=t(57437),s=t(2265),c=t(99376),r=t(72162);function f(){let e=(0,c.useSearchParams)().get("key"),[n,t]=(0,s.useState)(null);return(0,s.useEffect)(()=>{e&&t(e)},[e]),(0,u.jsx)(r.Z,{accessToken:n})}}},function(e){e.O(0,[50,521,49,162,971,117,744],function(){return e(e.s=67355)}),_N_E=e.O()}]);
|
||||
|
|
@ -0,0 +1 @@
|
|||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[25],{38520:function(e,n,u){Promise.resolve().then(u.bind(u,22775))},22775:function(e,n,u){"use strict";u.r(n),u.d(n,{default:function(){return f}});var t=u(57437),s=u(2265),r=u(99376),c=u(97851);function f(){let e=(0,r.useSearchParams)().get("key"),[n,u]=(0,s.useState)(null);return(0,s.useEffect)(()=>{e&&u(e)},[e]),(0,t.jsx)(c.Z,{accessToken:n,publicPage:!0,premiumUser:!1,userRole:null})}}},function(e){e.O(0,[50,521,866,49,162,851,971,117,744],function(){return e(e.s=38520)}),_N_E=e.O()}]);
|
||||
|
|
@ -1 +0,0 @@
|
|||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[25],{9397:function(e,n,u){Promise.resolve().then(u.bind(u,22775))},22775:function(e,n,u){"use strict";u.r(n),u.d(n,{default:function(){return f}});var t=u(57437),s=u(2265),r=u(99376),c=u(97851);function f(){let e=(0,r.useSearchParams)().get("key"),[n,u]=(0,s.useState)(null);return(0,s.useEffect)(()=>{e&&u(e)},[e]),(0,t.jsx)(c.Z,{accessToken:n,publicPage:!0,premiumUser:!1,userRole:null})}}},function(e){e.O(0,[50,521,866,154,162,851,971,117,744],function(){return e(e.s=9397)}),_N_E=e.O()}]);
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
|
|
@ -1 +1 @@
|
|||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[744],{60400:function(e,n,t){Promise.resolve().then(t.t.bind(t,12846,23)),Promise.resolve().then(t.t.bind(t,19107,23)),Promise.resolve().then(t.t.bind(t,61060,23)),Promise.resolve().then(t.t.bind(t,4707,23)),Promise.resolve().then(t.t.bind(t,80,23)),Promise.resolve().then(t.t.bind(t,36423,23))}},function(e){var n=function(n){return e(e.s=n)};e.O(0,[971,117],function(){return n(54278),n(60400)}),_N_E=e.O()}]);
|
||||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[744],{78483:function(e,n,t){Promise.resolve().then(t.t.bind(t,12846,23)),Promise.resolve().then(t.t.bind(t,19107,23)),Promise.resolve().then(t.t.bind(t,61060,23)),Promise.resolve().then(t.t.bind(t,4707,23)),Promise.resolve().then(t.t.bind(t,80,23)),Promise.resolve().then(t.t.bind(t,36423,23))}},function(e){var n=function(n){return e(e.s=n)};e.O(0,[971,117],function(){return n(54278),n(78483)}),_N_E=e.O()}]);
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
|
|
@ -0,0 +1,9 @@
|
|||
<svg id="Layer_1" data-name="Layer 1" xmlns="http://www.w3.org/2000/svg" viewBox="0 0 146.36 139.16" xmlns:xlink="http://www.w3.org/1999/xlink">
|
||||
<defs>
|
||||
<style>
|
||||
.cls-1{fill:#29b5e8;fill-rule:evenodd;}
|
||||
</style>
|
||||
</defs>
|
||||
<path class="cls-1" d="M134.81,60.1l-16.47,9.49L134.81,79a8.65,8.65,0,1,1-8.67,15l-29.51-17a8.68,8.68,0,0,1-4.33-7.75,8.48,8.48,0,0,1,.31-2,8.68,8.68,0,0,1,4-5.19l29.51-16.94A8.69,8.69,0,0,1,138,48.31,8.58,8.58,0,0,1,134.81,60.1Zm-15.59,46L89.72,89.13a8.72,8.72,0,0,0-13.06,7.48v33.9a8.69,8.69,0,0,0,17.37,0v-19L110.54,121a8.66,8.66,0,1,0,8.68-15Zm-34-33.16L72.92,85.09a2.44,2.44,0,0,1-1.54.65H67.77a2.51,2.51,0,0,1-1.54-.65L54,72.9a2.45,2.45,0,0,1-.64-1.52v-3.6A2.5,2.5,0,0,1,54,66.25L66.23,54.06a2.5,2.5,0,0,1,1.54-.64h3.61a2.45,2.45,0,0,1,1.54.64L85.18,66.25a2.49,2.49,0,0,1,.63,1.53v3.6A2.44,2.44,0,0,1,85.18,72.9Zm-9.8-3.38A2.59,2.59,0,0,0,74.73,68l-3.55-3.51a2.51,2.51,0,0,0-1.54-.64h-.13a2.46,2.46,0,0,0-1.53.64L64.43,68a2.51,2.51,0,0,0-.63,1.55v.13a2.41,2.41,0,0,0,.63,1.52L68,74.7a2.48,2.48,0,0,0,1.53.64h.13a2.51,2.51,0,0,0,1.54-.64l3.55-3.53a2.49,2.49,0,0,0,.65-1.52ZM19.93,33.08,49.44,50a8.73,8.73,0,0,0,13.07-7.49V8.64a8.69,8.69,0,0,0-17.37,0v19l-16.53-9.5a8.65,8.65,0,1,0-8.68,15ZM84.69,51.16a8.64,8.64,0,0,0,5-1.13l29.5-17a8.65,8.65,0,1,0-8.68-15L94,27.61v-19a8.69,8.69,0,0,0-17.37,0v33.9A8.66,8.66,0,0,0,84.69,51.16ZM54.48,88a8.58,8.58,0,0,0-5,1.13L19.93,106.06a8.66,8.66,0,1,0,8.68,15l16.53-9.49v19a8.69,8.69,0,0,0,17.37,0V96.61A8.65,8.65,0,0,0,54.48,88Zm-8-15.87a8.61,8.61,0,0,0-4-10L13,45.14A8.69,8.69,0,0,0,1.17,48.31,8.59,8.59,0,0,0,4.35,60.1l16.47,9.49L4.35,79A8.65,8.65,0,1,0,13,94l29.48-17A8.59,8.59,0,0,0,46.47,72.13Zm93.15-56.22H138.3v1.63h1.32c.61,0,1-.28,1-.8S140.26,15.91,139.62,15.91Zm-2.94-1.5h3c1.62,0,2.7.89,2.7,2.27a2.16,2.16,0,0,1-1.08,1.9l1.17,1.68v.34h-1.69L139.62,19H138.3V20.6h-1.62Zm8.3,3.22a5.48,5.48,0,0,0-5.58-5.83c-3.31,0-5.51,2.39-5.51,5.83,0,3.28,2.2,5.82,5.51,5.82A5.47,5.47,0,0,0,145,17.63Zm1.38,0c0,3.89-2.6,7.14-7,7.14s-6.89-3.28-6.89-7.14,2.57-7.14,6.89-7.14S146.36,13.73,146.36,17.63Z">
|
||||
</path>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 2 KiB |
File diff suppressed because one or more lines are too long
|
|
@ -1,7 +1,7 @@
|
|||
2:I[19107,[],"ClientPageRoot"]
|
||||
3:I[55139,["665","static/chunks/3014691f-b7b79b78e27792f3.js","990","static/chunks/13b76428-ebdf3012af0e4489.js","50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","866","static/chunks/866-3523e0e07cf314f6.js","313","static/chunks/313-27c820a98e9413e5.js","154","static/chunks/154-f87cf692dcea3018.js","162","static/chunks/162-dd6427ff1a4ad9f4.js","851","static/chunks/851-bbe6d02cf41bb87a.js","931","static/chunks/app/page-f400068ac45ce482.js"],"default",1]
|
||||
3:I[73148,["665","static/chunks/3014691f-b7b79b78e27792f3.js","990","static/chunks/13b76428-ebdf3012af0e4489.js","50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","866","static/chunks/866-9e1803a09e9ae8da.js","313","static/chunks/313-0025fb08e386c4b8.js","49","static/chunks/49-b6f167418ea8dbf4.js","162","static/chunks/162-ffcd7d9fbb033bdf.js","851","static/chunks/851-0a73701a3a0b0187.js","931","static/chunks/app/page-5b9ff2d173a47e2c.js"],"default",1]
|
||||
4:I[4707,[],""]
|
||||
5:I[36423,[],""]
|
||||
0:["WkpkdsewrdPMuTzVGS_5j",[[["",{"children":["__PAGE__",{}]},"$undefined","$undefined",true],["",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/4103fa525703177b.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
0:["_S0i_Y-CCoYQWc9dIuLxF",[[["",{"children":["__PAGE__",{}]},"$undefined","$undefined",true],["",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/c200be8dd8638678.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
6:[["$","meta","0",{"name":"viewport","content":"width=device-width, initial-scale=1"}],["$","meta","1",{"charSet":"utf-8"}],["$","title","2",{"children":"LiteLLM Dashboard"}],["$","meta","3",{"name":"description","content":"LiteLLM Proxy Admin UI"}],["$","link","4",{"rel":"icon","href":"/favicon.ico","type":"image/x-icon","sizes":"16x16"}],["$","link","5",{"rel":"icon","href":"./favicon.ico"}],["$","meta","6",{"name":"next-size-adjust"}]]
|
||||
1:null
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
2:I[19107,[],"ClientPageRoot"]
|
||||
3:I[52829,["50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","154","static/chunks/154-f87cf692dcea3018.js","162","static/chunks/162-dd6427ff1a4ad9f4.js","418","static/chunks/app/model_hub/page-237d2973f13202c4.js"],"default",1]
|
||||
3:I[52829,["50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","49","static/chunks/49-b6f167418ea8dbf4.js","162","static/chunks/162-ffcd7d9fbb033bdf.js","418","static/chunks/app/model_hub/page-d7915f579b770030.js"],"default",1]
|
||||
4:I[4707,[],""]
|
||||
5:I[36423,[],""]
|
||||
0:["WkpkdsewrdPMuTzVGS_5j",[[["",{"children":["model_hub",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["model_hub",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","model_hub","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/4103fa525703177b.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
0:["_S0i_Y-CCoYQWc9dIuLxF",[[["",{"children":["model_hub",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["model_hub",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","model_hub","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/c200be8dd8638678.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
6:[["$","meta","0",{"name":"viewport","content":"width=device-width, initial-scale=1"}],["$","meta","1",{"charSet":"utf-8"}],["$","title","2",{"children":"LiteLLM Dashboard"}],["$","meta","3",{"name":"description","content":"LiteLLM Proxy Admin UI"}],["$","link","4",{"rel":"icon","href":"/favicon.ico","type":"image/x-icon","sizes":"16x16"}],["$","link","5",{"rel":"icon","href":"./favicon.ico"}],["$","meta","6",{"name":"next-size-adjust"}]]
|
||||
1:null
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
2:I[19107,[],"ClientPageRoot"]
|
||||
3:I[22775,["50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","866","static/chunks/866-3523e0e07cf314f6.js","154","static/chunks/154-f87cf692dcea3018.js","162","static/chunks/162-dd6427ff1a4ad9f4.js","851","static/chunks/851-bbe6d02cf41bb87a.js","25","static/chunks/app/model_hub_table/page-5d1aa98a47f9e9fd.js"],"default",1]
|
||||
3:I[22775,["50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","866","static/chunks/866-9e1803a09e9ae8da.js","49","static/chunks/49-b6f167418ea8dbf4.js","162","static/chunks/162-ffcd7d9fbb033bdf.js","851","static/chunks/851-0a73701a3a0b0187.js","25","static/chunks/app/model_hub_table/page-2f23c22a47b20607.js"],"default",1]
|
||||
4:I[4707,[],""]
|
||||
5:I[36423,[],""]
|
||||
0:["WkpkdsewrdPMuTzVGS_5j",[[["",{"children":["model_hub_table",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["model_hub_table",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","model_hub_table","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/4103fa525703177b.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
0:["_S0i_Y-CCoYQWc9dIuLxF",[[["",{"children":["model_hub_table",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["model_hub_table",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","model_hub_table","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/c200be8dd8638678.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
6:[["$","meta","0",{"name":"viewport","content":"width=device-width, initial-scale=1"}],["$","meta","1",{"charSet":"utf-8"}],["$","title","2",{"children":"LiteLLM Dashboard"}],["$","meta","3",{"name":"description","content":"LiteLLM Proxy Admin UI"}],["$","link","4",{"rel":"icon","href":"/favicon.ico","type":"image/x-icon","sizes":"16x16"}],["$","link","5",{"rel":"icon","href":"./favicon.ico"}],["$","meta","6",{"name":"next-size-adjust"}]]
|
||||
1:null
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
|
|
@ -1,7 +1,7 @@
|
|||
2:I[19107,[],"ClientPageRoot"]
|
||||
3:I[12011,["665","static/chunks/3014691f-b7b79b78e27792f3.js","50","static/chunks/50-d0da2dd7acce2eb9.js","154","static/chunks/154-f87cf692dcea3018.js","461","static/chunks/app/onboarding/page-099f7aa4c559d470.js"],"default",1]
|
||||
3:I[12011,["665","static/chunks/3014691f-b7b79b78e27792f3.js","50","static/chunks/50-d0da2dd7acce2eb9.js","49","static/chunks/49-b6f167418ea8dbf4.js","461","static/chunks/app/onboarding/page-1fd26064f407ad1d.js"],"default",1]
|
||||
4:I[4707,[],""]
|
||||
5:I[36423,[],""]
|
||||
0:["WkpkdsewrdPMuTzVGS_5j",[[["",{"children":["onboarding",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["onboarding",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","onboarding","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/4103fa525703177b.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
0:["_S0i_Y-CCoYQWc9dIuLxF",[[["",{"children":["onboarding",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["onboarding",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","onboarding","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/c200be8dd8638678.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
6:[["$","meta","0",{"name":"viewport","content":"width=device-width, initial-scale=1"}],["$","meta","1",{"charSet":"utf-8"}],["$","title","2",{"children":"LiteLLM Dashboard"}],["$","meta","3",{"name":"description","content":"LiteLLM Proxy Admin UI"}],["$","link","4",{"rel":"icon","href":"/favicon.ico","type":"image/x-icon","sizes":"16x16"}],["$","link","5",{"rel":"icon","href":"./favicon.ico"}],["$","meta","6",{"name":"next-size-adjust"}]]
|
||||
1:null
|
||||
|
|
|
|||
|
|
@ -1,39 +1,6 @@
|
|||
model_list:
|
||||
- model_name: openai/gpt-4o
|
||||
- model_name: gpt-5-mini
|
||||
litellm_params:
|
||||
model: openai/gpt-4o-mini
|
||||
api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5"
|
||||
api_key: dummy
|
||||
- model_name: "byok-wildcard/*"
|
||||
litellm_params:
|
||||
model: openai/*
|
||||
- model_name: xai-grok-3
|
||||
litellm_params:
|
||||
model: xai/grok-3
|
||||
- model_name: hosted_vllm/whisper-v3
|
||||
litellm_params:
|
||||
model: hosted_vllm/whisper-v3
|
||||
api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5"
|
||||
api_key: dummy
|
||||
- model_name: azure-hidden-model
|
||||
litellm_params:
|
||||
model: azure/gpt-4.1
|
||||
api_base: os.environ/AZURE_API_BASE_ALT
|
||||
api_version: "2023-05-15"
|
||||
model: azure/gpt-5-mini-2
|
||||
api_key: os.environ/AZURE_API_KEY_ALT
|
||||
|
||||
# mcp_servers:
|
||||
# github_mcp:
|
||||
# url: "https://api.githubcopilot.com/mcp"
|
||||
# auth_type: oauth2
|
||||
# authorization_url: https://github.com/login/oauth/authorize
|
||||
# token_url: https://github.com/login/oauth/access_token
|
||||
# client_id: os.environ/GITHUB_OAUTH_CLIENT_ID
|
||||
# client_secret: os.environ/GITHUB_OAUTH_CLIENT_SECRET
|
||||
# scopes: ["public_repo", "user:email"]
|
||||
# allowed_tools: ["list_tools"]
|
||||
# # disallowed_tools: ["repo_delete"]
|
||||
|
||||
litellm_settings:
|
||||
callbacks: ["prometheus"]
|
||||
custom_prometheus_metadata_labels: ["metadata.initiative", "metadata.business-unit"]
|
||||
api_base: os.environ/AZURE_API_BASE_ALT
|
||||
|
|
|
|||
|
|
@ -52,6 +52,24 @@ else:
|
|||
Span = Any
|
||||
|
||||
|
||||
class SupportedDBObjectType(str, enum.Enum):
|
||||
"""
|
||||
Supported database object types for fine-grained DB storage control.
|
||||
Use in general_settings.supported_db_objects to specify which objects to load from DB.
|
||||
"""
|
||||
|
||||
MODELS = "models"
|
||||
MCP = "mcp"
|
||||
GUARDRAILS = "guardrails"
|
||||
VECTOR_STORES = "vector_stores"
|
||||
PASS_THROUGH_ENDPOINTS = "pass_through_endpoints"
|
||||
PROMPTS = "prompts"
|
||||
MODEL_COST_MAP = "model_cost_map"
|
||||
|
||||
def __str__(self):
|
||||
return str(self.value)
|
||||
|
||||
|
||||
class LiteLLMTeamRoles(enum.Enum):
|
||||
# team admin
|
||||
TEAM_ADMIN = "admin"
|
||||
|
|
@ -712,6 +730,7 @@ class ModelParams(LiteLLMPydanticObjectBase):
|
|||
class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase):
|
||||
mcp_servers: Optional[List[str]] = None
|
||||
mcp_access_groups: Optional[List[str]] = None
|
||||
mcp_tool_permissions: Optional[Dict[str, List[str]]] = None
|
||||
vector_stores: Optional[List[str]] = None
|
||||
|
||||
|
||||
|
|
@ -1414,6 +1433,16 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase):
|
|||
object_permission_id: str
|
||||
mcp_servers: Optional[List[str]] = []
|
||||
mcp_access_groups: Optional[List[str]] = []
|
||||
mcp_tool_permissions: Optional[Dict[str, List[str]]] = None
|
||||
"""
|
||||
Mapping - server_id -> list of tools
|
||||
|
||||
Enforces allowed tools for a specific key/team/organization
|
||||
{
|
||||
"1234567890": ["tool_name_1", "tool_name_2"]
|
||||
}
|
||||
"""
|
||||
|
||||
vector_stores: Optional[List[str]] = []
|
||||
|
||||
|
||||
|
|
@ -1749,6 +1778,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
description="[DEPRECATED] Use 'user_header_mappings' instead. When set, the header value is treated as the end user id unless overridden by user_header_mappings.",
|
||||
)
|
||||
user_header_mappings: Optional[List[UserHeaderMapping]] = None
|
||||
supported_db_objects: Optional[List[SupportedDBObjectType]] = Field(
|
||||
None,
|
||||
description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map'. If not set, all objects are loaded (default behavior).",
|
||||
)
|
||||
|
||||
|
||||
class ConfigYAML(LiteLLMPydanticObjectBase):
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ import fastapi
|
|||
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching import DualCache
|
||||
|
|
@ -837,7 +838,7 @@ async def generate_key_fn(
|
|||
- enforced_params: Optional[List[str]] - List of enforced params for the key (Enterprise only). [Docs](https://docs.litellm.ai/docs/proxy/enterprise#enforce-required-params-for-llm-requests)
|
||||
- prompts: Optional[List[str]] - List of prompts that the key is allowed to use.
|
||||
- allowed_routes: Optional[list] - List of allowed routes for the key. Store the actual route or store a wildcard pattern for a set of routes. Example - ["/chat/completions", "/embeddings", "/keys/*"]
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission.
|
||||
- key_type: Optional[str] - Type of key that determines default allowed routes. Options: "llm_api" (can call LLM API routes), "management" (can call management routes), "read_only" (can only call info/read routes), "default" (uses default allowed routes). Defaults to "default".
|
||||
- prompts: Optional[List[str]] - List of allowed prompts for the key. If specified, the key will only be able to use these specific prompts.
|
||||
- auto_rotate: Optional[bool] - Whether this key should be automatically rotated (regenerated)
|
||||
|
|
@ -987,7 +988,7 @@ async def generate_service_account_key_fn(
|
|||
- tags: Optional[List[str]] - Tags for [tracking spend](https://litellm.vercel.app/docs/proxy/enterprise#tracking-spend-for-custom-tags) and/or doing [tag-based routing](https://litellm.vercel.app/docs/proxy/tag_routing).
|
||||
- enforced_params: Optional[List[str]] - List of enforced params for the key (Enterprise only). [Docs](https://docs.litellm.ai/docs/proxy/enterprise#enforce-required-params-for-llm-requests)
|
||||
- allowed_routes: Optional[list] - List of allowed routes for the key. Store the actual route or store a wildcard pattern for a set of routes. Example - ["/chat/completions", "/embeddings", "/keys/*"]
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission.
|
||||
Examples:
|
||||
|
||||
1. Allow users to turn on/off pii masking
|
||||
|
|
@ -1125,6 +1126,13 @@ async def _set_object_permission(
|
|||
return data_json
|
||||
|
||||
if "object_permission" in data_json:
|
||||
# Serialize mcp_tool_permissions JSON field to avoid GraphQL parsing issues
|
||||
# (e.g., server IDs starting with "3e64" being interpreted as floats)
|
||||
if "mcp_tool_permissions" in data_json["object_permission"]:
|
||||
data_json["object_permission"]["mcp_tool_permissions"] = safe_dumps(
|
||||
data_json["object_permission"]["mcp_tool_permissions"]
|
||||
)
|
||||
|
||||
created_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=data_json["object_permission"],
|
||||
|
|
@ -1300,7 +1308,7 @@ async def update_key_fn(
|
|||
- temp_budget_expiry: Optional[str] - Expiry time for the temporary budget increase (Enterprise only).
|
||||
- allowed_routes: Optional[list] - List of allowed routes for the key. Store the actual route or store a wildcard pattern for a set of routes. Example - ["/chat/completions", "/embeddings", "/keys/*"]
|
||||
- prompts: Optional[List[str]] - List of allowed prompts for the key. If specified, the key will only be able to use these specific prompts.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission.
|
||||
- auto_rotate: Optional[bool] - Whether this key should be automatically rotated
|
||||
- rotation_interval: Optional[str] - How often to rotate this key (e.g., '30d', '90d'). Required if auto_rotate=True
|
||||
Example:
|
||||
|
|
@ -2841,6 +2849,79 @@ async def list_keys(
|
|||
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/key/aliases",
|
||||
tags=["key management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def key_aliases() -> Dict[str, List[str]]:
|
||||
"""
|
||||
Lists all key aliases
|
||||
|
||||
Returns:
|
||||
{
|
||||
"aliases": List[str]
|
||||
}
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
verbose_proxy_logger.debug("Entering key_aliases function")
|
||||
|
||||
if prisma_client is None:
|
||||
verbose_proxy_logger.error("Database not connected")
|
||||
raise Exception("Database not connected")
|
||||
|
||||
where: Dict[str, Any] = {}
|
||||
try:
|
||||
where.update(_get_condition_to_filter_out_ui_session_tokens())
|
||||
except NameError:
|
||||
# Helper may not exist in some builds; ignore if missing
|
||||
pass
|
||||
|
||||
rows = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where=where,
|
||||
order=[{"key_alias": "asc"}],
|
||||
)
|
||||
|
||||
seen = set()
|
||||
aliases: List[str] = []
|
||||
for row in rows:
|
||||
alias = getattr(row, "key_alias", None)
|
||||
if alias is None and isinstance(row, dict):
|
||||
alias = row.get("key_alias")
|
||||
|
||||
if not alias:
|
||||
continue
|
||||
|
||||
alias_str = str(alias).strip()
|
||||
if alias_str and alias_str not in seen:
|
||||
seen.add(alias_str)
|
||||
aliases.append(alias_str)
|
||||
|
||||
verbose_proxy_logger.debug(f"Returning {len(aliases)} key aliases")
|
||||
|
||||
return {"aliases": aliases}
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error in key_aliases: {e}")
|
||||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "detail", f"error({str(e)})"),
|
||||
type=ProxyErrorTypes.internal_server_error,
|
||||
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="Authentication Error, " + str(e),
|
||||
type=ProxyErrorTypes.internal_server_error,
|
||||
param=getattr(e, "param", "None"),
|
||||
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
)
|
||||
|
||||
|
||||
def _validate_sort_params(
|
||||
sort_by: Optional[str], sort_order: str
|
||||
|
|
|
|||
|
|
@ -311,7 +311,7 @@ async def new_team( # noqa: PLR0915
|
|||
- model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias)
|
||||
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
|
||||
- prompts: Optional[List[str]] - List of prompts that the team is allowed to use.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission.
|
||||
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
|
||||
- team_member_rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for individual team members.
|
||||
- team_member_tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for individual team members.
|
||||
|
|
@ -769,7 +769,7 @@ async def update_team(
|
|||
- model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias)
|
||||
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
|
||||
- prompts: Optional[List[str]] - List of prompts that the team is allowed to use.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission.
|
||||
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
|
||||
- team_member_rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for individual team members.
|
||||
- team_member_tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for individual team members.
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ from typing import Dict, Optional, Union
|
|||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
|
||||
|
||||
async def attach_object_permission_to_dict(
|
||||
|
|
@ -114,6 +116,15 @@ async def handle_update_object_permission_common(
|
|||
if isinstance(new_object_permission, dict):
|
||||
existing_object_permissions_dict.update(new_object_permission)
|
||||
|
||||
#########################################################
|
||||
# Serialize mcp_tool_permissions JSON field to avoid GraphQL parsing issues
|
||||
# (e.g., server IDs starting with "3e64" being interpreted as floats)
|
||||
#########################################################
|
||||
if "mcp_tool_permissions" in existing_object_permissions_dict:
|
||||
existing_object_permissions_dict["mcp_tool_permissions"] = safe_dumps(
|
||||
existing_object_permissions_dict["mcp_tool_permissions"]
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Commit the update to the LiteLLM_ObjectPermissionTable
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -2957,6 +2957,40 @@ class ProxyConfig:
|
|||
|
||||
return config
|
||||
|
||||
def _should_load_db_object(
|
||||
self, object_type: Union[str, SupportedDBObjectType]
|
||||
) -> bool:
|
||||
"""
|
||||
Check if an object type should be loaded from the database based on general_settings.supported_db_objects.
|
||||
|
||||
Args:
|
||||
object_type: Type of object to check (e.g., SupportedDBObjectType.MODELS, "models", etc.)
|
||||
|
||||
Returns:
|
||||
True if the object should be loaded, False otherwise
|
||||
"""
|
||||
global general_settings
|
||||
|
||||
# Get the supported_db_objects configuration
|
||||
supported_db_objects = general_settings.get("supported_db_objects", None)
|
||||
|
||||
# If supported_db_objects is not set, load all objects (default behavior)
|
||||
if supported_db_objects is None:
|
||||
return True
|
||||
|
||||
# If supported_db_objects is set, only load specified objects
|
||||
if not isinstance(supported_db_objects, list):
|
||||
verbose_proxy_logger.warning(
|
||||
f"supported_db_objects is not a list, got {type(supported_db_objects)}. Loading all objects."
|
||||
)
|
||||
return True
|
||||
|
||||
# Convert object_type to string for comparison (handles both str and enum)
|
||||
object_type_str = str(object_type)
|
||||
|
||||
# Check if the object type is in the list (supports both str and enum values)
|
||||
return any(str(obj) == object_type_str for obj in supported_db_objects)
|
||||
|
||||
async def _get_models_from_db(self, prisma_client: PrismaClient) -> list:
|
||||
try:
|
||||
new_models = await prisma_client.db.litellm_proxymodeltable.find_many()
|
||||
|
|
@ -2988,12 +3022,14 @@ class ProxyConfig:
|
|||
f"Master key is not initialized or formatted. master_key={master_key}"
|
||||
)
|
||||
|
||||
new_models = await self._get_models_from_db(prisma_client=prisma_client)
|
||||
# Only load models from DB if "models" is in supported_db_objects (or if supported_db_objects is not set)
|
||||
if self._should_load_db_object(object_type="models"):
|
||||
new_models = await self._get_models_from_db(prisma_client=prisma_client)
|
||||
|
||||
# update llm router
|
||||
await self._update_llm_router(
|
||||
new_models=new_models, proxy_logging_obj=proxy_logging_obj
|
||||
)
|
||||
# update llm router
|
||||
await self._update_llm_router(
|
||||
new_models=new_models, proxy_logging_obj=proxy_logging_obj
|
||||
)
|
||||
|
||||
db_general_settings = await prisma_client.db.litellm_config.find_first(
|
||||
where={"param_name": "general_settings"}
|
||||
|
|
@ -3021,12 +3057,23 @@ class ProxyConfig:
|
|||
|
||||
ex. Vector Stores, Guardrails, MCP tools, etc.
|
||||
"""
|
||||
await self._init_guardrails_in_db(prisma_client=prisma_client)
|
||||
await self._init_vector_stores_in_db(prisma_client=prisma_client)
|
||||
await self._init_mcp_servers_in_db()
|
||||
await self._init_pass_through_endpoints_in_db()
|
||||
await self._init_prompts_in_db(prisma_client=prisma_client)
|
||||
await self._check_and_reload_model_cost_map(prisma_client=prisma_client)
|
||||
if self._should_load_db_object(object_type="guardrails"):
|
||||
await self._init_guardrails_in_db(prisma_client=prisma_client)
|
||||
|
||||
if self._should_load_db_object(object_type="vector_stores"):
|
||||
await self._init_vector_stores_in_db(prisma_client=prisma_client)
|
||||
|
||||
if self._should_load_db_object(object_type="mcp"):
|
||||
await self._init_mcp_servers_in_db()
|
||||
|
||||
if self._should_load_db_object(object_type="pass_through_endpoints"):
|
||||
await self._init_pass_through_endpoints_in_db()
|
||||
|
||||
if self._should_load_db_object(object_type="prompts"):
|
||||
await self._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
if self._should_load_db_object(object_type="model_cost_map"):
|
||||
await self._check_and_reload_model_cost_map(prisma_client=prisma_client)
|
||||
|
||||
async def _check_and_reload_model_cost_map(self, prisma_client: PrismaClient):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -156,6 +156,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
object_permission_id String @id @default(uuid())
|
||||
mcp_servers String[] @default([])
|
||||
mcp_access_groups String[] @default([])
|
||||
mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]}
|
||||
vector_stores String[] @default([])
|
||||
teams LiteLLM_TeamTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
|
|
|
|||
|
|
@ -6249,6 +6249,20 @@ class Router:
|
|||
returned_models.append(alias_model)
|
||||
else:
|
||||
returned_models.append(model)
|
||||
elif team_id is not None:
|
||||
# Fallback: if team_id is provided and model_name not in index,
|
||||
# check if model_name matches any team_public_model_name
|
||||
# O(n) scan but only when team_id lookup fails
|
||||
for idx, model in enumerate(self.model_list):
|
||||
if self.should_include_deployment(
|
||||
model_name=model_name, model=model, team_id=team_id
|
||||
):
|
||||
if model_alias is not None:
|
||||
alias_model = copy.deepcopy(model)
|
||||
alias_model["model_name"] = model_alias
|
||||
returned_models.append(alias_model)
|
||||
else:
|
||||
returned_models.append(model)
|
||||
|
||||
return returned_models
|
||||
|
||||
|
|
|
|||
316
litellm/utils.py
316
litellm/utils.py
|
|
@ -907,159 +907,169 @@ def _get_wrapper_timeout(
|
|||
|
||||
return timeout
|
||||
|
||||
def check_coroutine(value) -> bool:
|
||||
return get_coroutine_checker().is_async_callable(value)
|
||||
|
||||
|
||||
async def async_pre_call_deployment_hook(kwargs: Dict[str, Any], call_type: str):
|
||||
"""
|
||||
Allow modifying the request just before it's sent to the deployment.
|
||||
|
||||
Use this instead of 'async_pre_call_hook' when you need to modify the request AFTER a deployment is selected, but BEFORE the request is sent.
|
||||
"""
|
||||
try:
|
||||
typed_call_type = CallTypes(call_type)
|
||||
except ValueError:
|
||||
typed_call_type = None # unknown call type
|
||||
|
||||
modified_kwargs = kwargs.copy()
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
result = await callback.async_pre_call_deployment_hook(
|
||||
modified_kwargs, typed_call_type
|
||||
)
|
||||
if result is not None:
|
||||
modified_kwargs = result
|
||||
|
||||
return modified_kwargs
|
||||
|
||||
|
||||
async def async_post_call_success_deployment_hook(
|
||||
request_data: dict, response: Any, call_type: Optional[CallTypes]
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Allow modifying / reviewing the response just after it's received from the deployment.
|
||||
"""
|
||||
try:
|
||||
typed_call_type = CallTypes(call_type)
|
||||
except ValueError:
|
||||
typed_call_type = None # unknown call type
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
result = await callback.async_post_call_success_deployment_hook(
|
||||
request_data, cast(LLMResponseTypes, response), typed_call_type
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
return response
|
||||
|
||||
|
||||
def post_call_processing(
|
||||
original_response,
|
||||
model,
|
||||
optional_params: Optional[dict],
|
||||
original_function,
|
||||
rules_obj,
|
||||
):
|
||||
try:
|
||||
if original_response is None:
|
||||
pass
|
||||
else:
|
||||
call_type = original_function.__name__
|
||||
if (
|
||||
call_type == CallTypes.completion.value
|
||||
or call_type == CallTypes.acompletion.value
|
||||
):
|
||||
is_coroutine = check_coroutine(original_response)
|
||||
if is_coroutine is True:
|
||||
pass
|
||||
else:
|
||||
if (
|
||||
isinstance(original_response, ModelResponse)
|
||||
and len(original_response.choices) > 0
|
||||
):
|
||||
model_response: Optional[str] = original_response.choices[
|
||||
0
|
||||
].message.content # type: ignore
|
||||
if model_response is not None:
|
||||
### POST-CALL RULES ###
|
||||
rules_obj.post_call_rules(
|
||||
input=model_response, model=model
|
||||
)
|
||||
### JSON SCHEMA VALIDATION ###
|
||||
if litellm.enable_json_schema_validation is True:
|
||||
try:
|
||||
if (
|
||||
optional_params is not None
|
||||
and "response_format" in optional_params
|
||||
and optional_params["response_format"]
|
||||
is not None
|
||||
):
|
||||
json_response_format: Optional[dict] = None
|
||||
if (
|
||||
isinstance(
|
||||
optional_params["response_format"],
|
||||
dict,
|
||||
)
|
||||
and optional_params[
|
||||
"response_format"
|
||||
].get("json_schema")
|
||||
is not None
|
||||
):
|
||||
json_response_format = optional_params[
|
||||
"response_format"
|
||||
]
|
||||
elif _parsing._completions.is_basemodel_type(
|
||||
optional_params["response_format"] # type: ignore
|
||||
):
|
||||
json_response_format = (
|
||||
type_to_response_format_param(
|
||||
response_format=optional_params[
|
||||
"response_format"
|
||||
]
|
||||
)
|
||||
)
|
||||
if json_response_format is not None:
|
||||
litellm.litellm_core_utils.json_validation_rule.validate_schema(
|
||||
schema=json_response_format[
|
||||
"json_schema"
|
||||
]["schema"],
|
||||
response=model_response,
|
||||
)
|
||||
except TypeError:
|
||||
pass
|
||||
if (
|
||||
optional_params is not None
|
||||
and "response_format" in optional_params
|
||||
and isinstance(
|
||||
optional_params["response_format"], dict
|
||||
)
|
||||
and "type" in optional_params["response_format"]
|
||||
and optional_params["response_format"]["type"]
|
||||
== "json_object"
|
||||
and "response_schema"
|
||||
in optional_params["response_format"]
|
||||
and isinstance(
|
||||
optional_params["response_format"][
|
||||
"response_schema"
|
||||
],
|
||||
dict,
|
||||
)
|
||||
and "enforce_validation"
|
||||
in optional_params["response_format"]
|
||||
and optional_params["response_format"][
|
||||
"enforce_validation"
|
||||
]
|
||||
is True
|
||||
):
|
||||
# schema given, json response expected, and validation enforced
|
||||
litellm.litellm_core_utils.json_validation_rule.validate_schema(
|
||||
schema=optional_params["response_format"][
|
||||
"response_schema"
|
||||
],
|
||||
response=model_response,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
||||
def client(original_function): # noqa: PLR0915
|
||||
rules_obj = Rules()
|
||||
|
||||
def check_coroutine(value) -> bool:
|
||||
return get_coroutine_checker().is_async_callable(value)
|
||||
|
||||
async def async_pre_call_deployment_hook(kwargs: Dict[str, Any], call_type: str):
|
||||
"""
|
||||
Allow modifying the request just before it's sent to the deployment.
|
||||
|
||||
Use this instead of 'async_pre_call_hook' when you need to modify the request AFTER a deployment is selected, but BEFORE the request is sent.
|
||||
"""
|
||||
try:
|
||||
typed_call_type = CallTypes(call_type)
|
||||
except ValueError:
|
||||
typed_call_type = None # unknown call type
|
||||
|
||||
modified_kwargs = kwargs.copy()
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
result = await callback.async_pre_call_deployment_hook(
|
||||
modified_kwargs, typed_call_type
|
||||
)
|
||||
if result is not None:
|
||||
modified_kwargs = result
|
||||
|
||||
return modified_kwargs
|
||||
|
||||
async def async_post_call_success_deployment_hook(
|
||||
request_data: dict, response: Any, call_type: Optional[CallTypes]
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Allow modifying / reviewing the response just after it's received from the deployment.
|
||||
"""
|
||||
try:
|
||||
typed_call_type = CallTypes(call_type)
|
||||
except ValueError:
|
||||
typed_call_type = None # unknown call type
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
result = await callback.async_post_call_success_deployment_hook(
|
||||
request_data, cast(LLMResponseTypes, response), typed_call_type
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
return response
|
||||
|
||||
def post_call_processing(original_response, model, optional_params: Optional[dict]):
|
||||
try:
|
||||
if original_response is None:
|
||||
pass
|
||||
else:
|
||||
call_type = original_function.__name__
|
||||
if (
|
||||
call_type == CallTypes.completion.value
|
||||
or call_type == CallTypes.acompletion.value
|
||||
):
|
||||
is_coroutine = check_coroutine(original_response)
|
||||
if is_coroutine is True:
|
||||
pass
|
||||
else:
|
||||
if (
|
||||
isinstance(original_response, ModelResponse)
|
||||
and len(original_response.choices) > 0
|
||||
):
|
||||
model_response: Optional[str] = original_response.choices[
|
||||
0
|
||||
].message.content # type: ignore
|
||||
if model_response is not None:
|
||||
### POST-CALL RULES ###
|
||||
rules_obj.post_call_rules(
|
||||
input=model_response, model=model
|
||||
)
|
||||
### JSON SCHEMA VALIDATION ###
|
||||
if litellm.enable_json_schema_validation is True:
|
||||
try:
|
||||
if (
|
||||
optional_params is not None
|
||||
and "response_format" in optional_params
|
||||
and optional_params["response_format"]
|
||||
is not None
|
||||
):
|
||||
json_response_format: Optional[dict] = None
|
||||
if (
|
||||
isinstance(
|
||||
optional_params["response_format"],
|
||||
dict,
|
||||
)
|
||||
and optional_params[
|
||||
"response_format"
|
||||
].get("json_schema")
|
||||
is not None
|
||||
):
|
||||
json_response_format = optional_params[
|
||||
"response_format"
|
||||
]
|
||||
elif _parsing._completions.is_basemodel_type(
|
||||
optional_params["response_format"] # type: ignore
|
||||
):
|
||||
json_response_format = (
|
||||
type_to_response_format_param(
|
||||
response_format=optional_params[
|
||||
"response_format"
|
||||
]
|
||||
)
|
||||
)
|
||||
if json_response_format is not None:
|
||||
litellm.litellm_core_utils.json_validation_rule.validate_schema(
|
||||
schema=json_response_format[
|
||||
"json_schema"
|
||||
]["schema"],
|
||||
response=model_response,
|
||||
)
|
||||
except TypeError:
|
||||
pass
|
||||
if (
|
||||
optional_params is not None
|
||||
and "response_format" in optional_params
|
||||
and isinstance(
|
||||
optional_params["response_format"], dict
|
||||
)
|
||||
and "type" in optional_params["response_format"]
|
||||
and optional_params["response_format"]["type"]
|
||||
== "json_object"
|
||||
and "response_schema"
|
||||
in optional_params["response_format"]
|
||||
and isinstance(
|
||||
optional_params["response_format"][
|
||||
"response_schema"
|
||||
],
|
||||
dict,
|
||||
)
|
||||
and "enforce_validation"
|
||||
in optional_params["response_format"]
|
||||
and optional_params["response_format"][
|
||||
"enforce_validation"
|
||||
]
|
||||
is True
|
||||
):
|
||||
# schema given, json response expected, and validation enforced
|
||||
litellm.litellm_core_utils.json_validation_rule.validate_schema(
|
||||
schema=optional_params["response_format"][
|
||||
"response_schema"
|
||||
],
|
||||
response=model_response,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
@wraps(original_function)
|
||||
def wrapper(*args, **kwargs): # noqa: PLR0915
|
||||
|
|
@ -1270,6 +1280,8 @@ def client(original_function): # noqa: PLR0915
|
|||
original_response=result,
|
||||
model=model or None,
|
||||
optional_params=kwargs,
|
||||
original_function=original_function,
|
||||
rules_obj=rules_obj,
|
||||
)
|
||||
|
||||
# [OPTIONAL] ADD TO CACHE
|
||||
|
|
@ -1487,7 +1499,11 @@ def client(original_function): # noqa: PLR0915
|
|||
return result
|
||||
### POST-CALL RULES ###
|
||||
post_call_processing(
|
||||
original_response=result, model=model, optional_params=kwargs
|
||||
original_response=result,
|
||||
model=model,
|
||||
optional_params=kwargs,
|
||||
original_function=original_function,
|
||||
rules_obj=rules_obj,
|
||||
)
|
||||
# Only run if call_type is a valid value in CallTypes
|
||||
if call_type in [ct.value for ct in CallTypes]:
|
||||
|
|
|
|||
|
|
@ -12872,6 +12872,39 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"gpt-5-pro": {
|
||||
"input_cost_per_token": 1.5e-05,
|
||||
"input_cost_per_token_batches": 7.5e-06,
|
||||
"litellm_provider": "openai",
|
||||
"max_input_tokens": 400000,
|
||||
"max_output_tokens": 272000,
|
||||
"max_tokens": 272000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 1.2e-04,
|
||||
"output_cost_per_token_batches": 6e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": false,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"gpt-5-codex": {
|
||||
"cache_read_input_token_cost": 1.25e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
|
|
@ -13189,6 +13222,20 @@
|
|||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"gpt-image-1-mini": {
|
||||
"cache_read_input_image_token_cost": 2.5e-07,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"input_cost_per_image_token": 2.5e-06,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
"output_cost_per_image_token": 8e-06,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supports_vision": true
|
||||
},
|
||||
"gpt-realtime": {
|
||||
"cache_creation_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
|
|
@ -14623,6 +14670,54 @@
|
|||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"low/1024-x-1024/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.005,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"low/1024-x-1536/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.006,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"low/1536-x-1024/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.006,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"medium/1024-x-1024/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.011,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"medium/1024-x-1536/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.015,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"medium/1536-x-1024/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.015,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"medlm-large": {
|
||||
"input_cost_per_character": 5e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
|
|
|
|||
|
|
@ -156,6 +156,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
object_permission_id String @id @default(uuid())
|
||||
mcp_servers String[] @default([])
|
||||
mcp_access_groups String[] @default([])
|
||||
mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]}
|
||||
vector_stores String[] @default([])
|
||||
teams LiteLLM_TeamTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from litellm.types.llms.openai import (
|
|||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from base_responses_api import BaseResponsesAPITest
|
||||
|
||||
|
||||
class TestAzureResponsesAPITest(BaseResponsesAPITest):
|
||||
def get_base_completion_call_args(self):
|
||||
return {
|
||||
|
|
@ -43,4 +44,55 @@ async def test_azure_responses_api_preview_api_version():
|
|||
api_base=os.getenv("AZURE_RESPONSES_OPENAI_ENDPOINT"),
|
||||
api_key=os.getenv("AZURE_RESPONSES_OPENAI_API_KEY"),
|
||||
input="Hello, can you tell me a short joke?",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_responses_api_status_error():
|
||||
"""
|
||||
Ensure new azure preview api version is working
|
||||
"""
|
||||
litellm._turn_on_debug()
|
||||
|
||||
request_data = {
|
||||
"model": "gpt-5-mini",
|
||||
"input": [
|
||||
{"content": "tell me an interesting fact", "role": "user"},
|
||||
{
|
||||
"id": "rs_0ab687487834d9df0068e462a1b2d88197aabbc832c9ba5316",
|
||||
"summary": [],
|
||||
"type": "reasoning",
|
||||
"content": None,
|
||||
"encrypted_content": None,
|
||||
"status": "completed",
|
||||
},
|
||||
{
|
||||
"id": "msg_0ab687487834d9df0068e462a1df188197b74b1eef05102c18",
|
||||
"content": [
|
||||
{
|
||||
"annotations": [],
|
||||
"text": "Octopuses have three hearts: two pump blood to the gills, while the third pumps it to the rest of the body. Even more unusual, their blood is blue because it uses the copper-containing protein hemocyanin to carry oxygen, which is more efficient than hemoglobin in cold, low-oxygen environments.",
|
||||
"type": "output_text",
|
||||
"logprobs": [],
|
||||
}
|
||||
],
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"type": "message",
|
||||
},
|
||||
{"role": "user", "content": "tell me another"},
|
||||
],
|
||||
"include": [],
|
||||
"instructions": "You are a helpful assistant.",
|
||||
"reasoning": {"effort": "minimal"},
|
||||
"stream": False,
|
||||
"tools": [],
|
||||
}
|
||||
response = await litellm.aresponses(
|
||||
model="azure/gpt-5-mini-2",
|
||||
truncation="auto",
|
||||
api_version="preview",
|
||||
api_base=os.getenv("AZURE_GPT5_MINI_API_BASE"),
|
||||
api_key=os.getenv("AZURE_GPT5_MINI_API_KEY"),
|
||||
input=request_data["input"],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -61,6 +61,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
|
|||
list_keys,
|
||||
regenerate_key_fn,
|
||||
update_key_fn,
|
||||
key_aliases,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
new_team,
|
||||
|
|
@ -151,7 +152,6 @@ def prisma_client():
|
|||
@pytest.mark.flaky(retries=6, delay=1)
|
||||
async def test_new_user_response(prisma_client):
|
||||
try:
|
||||
|
||||
print("prisma client=", prisma_client)
|
||||
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
|
|
@ -424,7 +424,6 @@ async def test_call_with_valid_model_using_all_models(prisma_client):
|
|||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
try:
|
||||
|
||||
await litellm.proxy.proxy_server.prisma_client.connect()
|
||||
|
||||
team_request = NewTeamRequest(
|
||||
|
|
@ -1789,7 +1788,6 @@ async def test_call_with_key_over_model_budget(
|
|||
litellm.callbacks.append(model_budget_limiter)
|
||||
|
||||
try:
|
||||
|
||||
# set budget for chatgpt-v-3 to 0.000001, expect the next request to fail
|
||||
model_max_budget = {
|
||||
"gpt-4o-mini": {
|
||||
|
|
@ -3531,6 +3529,58 @@ async def test_list_keys(prisma_client):
|
|||
assert _key in response["keys"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_aliases(prisma_client):
|
||||
"""
|
||||
Test the key_aliases function:
|
||||
- Returns a list
|
||||
- Includes alias from a newly created key
|
||||
- Aliases are unique and sorted
|
||||
"""
|
||||
import asyncio
|
||||
import uuid
|
||||
import litellm
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
# Wire up test prisma client
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
await litellm.proxy.proxy_server.prisma_client.connect()
|
||||
|
||||
# Basic call
|
||||
response = await key_aliases()
|
||||
assert "aliases" in response
|
||||
assert isinstance(response["aliases"], list)
|
||||
|
||||
# Create a new user (and key) with a unique alias
|
||||
unique_id = str(uuid.uuid4())
|
||||
test_alias = f"key-aliases-test-{unique_id}"
|
||||
test_user_id = f"key-aliases-user-{unique_id}"
|
||||
|
||||
await new_user(
|
||||
data=NewUserRequest(
|
||||
user_id=test_user_id,
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
key_alias=test_alias,
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
# Allow async DB writes to settle
|
||||
await asyncio.sleep(2)
|
||||
|
||||
# Call again and validate
|
||||
response_after = await key_aliases()
|
||||
aliases = response_after["aliases"]
|
||||
|
||||
# Contains the new alias
|
||||
assert test_alias in aliases
|
||||
|
||||
# Unique & sorted (endpoint dedupes and orders ascending)
|
||||
assert len(aliases) == len(set(aliases))
|
||||
assert aliases == sorted(aliases)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_vertex_ai_route(prisma_client):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,10 +1,13 @@
|
|||
import os
|
||||
import ssl
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
# Add the parent directory to the path so we can import litellm
|
||||
sys.path.insert(0, '../../../')
|
||||
sys.path.insert(0, "../../../")
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.types.mcp import MCPStdioConfig, MCPTransport
|
||||
|
|
@ -16,57 +19,54 @@ class TestMCPClient:
|
|||
def test_mcp_client_stdio_init(self):
|
||||
"""Test MCPClient initialization with stdio config"""
|
||||
stdio_config = MCPStdioConfig(
|
||||
command="python",
|
||||
args=["-m", "my_mcp_server"],
|
||||
env={"DEBUG": "1"}
|
||||
command="python", args=["-m", "my_mcp_server"], env={"DEBUG": "1"}
|
||||
)
|
||||
|
||||
client = MCPClient(
|
||||
transport_type=MCPTransport.stdio,
|
||||
stdio_config=stdio_config
|
||||
)
|
||||
|
||||
|
||||
client = MCPClient(transport_type=MCPTransport.stdio, stdio_config=stdio_config)
|
||||
|
||||
assert client.transport_type == MCPTransport.stdio
|
||||
assert client.stdio_config == stdio_config
|
||||
assert client.stdio_config["command"] == "python"
|
||||
assert client.stdio_config["args"] == ["-m", "my_mcp_server"]
|
||||
assert client.stdio_config is not None
|
||||
assert client.stdio_config.get("command") == "python"
|
||||
assert client.stdio_config.get("args") == ["-m", "my_mcp_server"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_client_stdio_connect_error(self):
|
||||
"""Test MCP client stdio connection error handling"""
|
||||
# Test missing stdio_config
|
||||
client = MCPClient(transport_type=MCPTransport.stdio)
|
||||
|
||||
with pytest.raises(ValueError, match="stdio_config is required for stdio transport"):
|
||||
|
||||
with pytest.raises(
|
||||
ValueError, match="stdio_config is required for stdio transport"
|
||||
):
|
||||
await client.connect()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('litellm.experimental_mcp_client.client.stdio_client')
|
||||
@patch('litellm.experimental_mcp_client.client.ClientSession')
|
||||
async def test_mcp_client_stdio_connect_success(self, mock_session, mock_stdio_client):
|
||||
@patch("litellm.experimental_mcp_client.client.stdio_client")
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_mcp_client_stdio_connect_success(
|
||||
self, mock_session, mock_stdio_client
|
||||
):
|
||||
"""Test successful stdio connection"""
|
||||
# Setup mocks
|
||||
mock_transport = (MagicMock(), MagicMock())
|
||||
mock_stdio_client.return_value.__aenter__ = AsyncMock(return_value=mock_transport)
|
||||
|
||||
mock_stdio_client.return_value.__aenter__ = AsyncMock(
|
||||
return_value=mock_transport
|
||||
)
|
||||
|
||||
mock_session_instance = MagicMock()
|
||||
mock_session_instance.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session.return_value = mock_session_instance
|
||||
|
||||
|
||||
stdio_config = MCPStdioConfig(
|
||||
command="python",
|
||||
args=["-m", "my_mcp_server"],
|
||||
env={"DEBUG": "1"}
|
||||
command="python", args=["-m", "my_mcp_server"], env={"DEBUG": "1"}
|
||||
)
|
||||
|
||||
client = MCPClient(
|
||||
transport_type=MCPTransport.stdio,
|
||||
stdio_config=stdio_config
|
||||
)
|
||||
|
||||
|
||||
client = MCPClient(transport_type=MCPTransport.stdio, stdio_config=stdio_config)
|
||||
|
||||
await client.connect()
|
||||
|
||||
|
||||
# Verify stdio_client was called with correct parameters
|
||||
mock_stdio_client.assert_called_once()
|
||||
call_args = mock_stdio_client.call_args[0][0]
|
||||
|
|
@ -74,6 +74,162 @@ class TestMCPClient:
|
|||
assert call_args.args == ["-m", "my_mcp_server"]
|
||||
assert call_args.env == {"DEBUG": "1"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.streamablehttp_client")
|
||||
@patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"SSL_CERT_FILE": "/path/to/custom/ca-bundle.pem",
|
||||
"SSL_CERTIFICATE": "/path/to/client-cert.pem",
|
||||
},
|
||||
)
|
||||
async def test_mcp_client_ssl_configuration_from_env(
|
||||
self, mock_streamablehttp_client
|
||||
):
|
||||
"""Test that MCP client uses SSL configuration from environment variables"""
|
||||
# Setup mocks
|
||||
mock_transport = (MagicMock(), MagicMock())
|
||||
mock_streamablehttp_client.return_value.__aenter__ = AsyncMock(
|
||||
return_value=mock_transport
|
||||
)
|
||||
|
||||
# Mock the session
|
||||
with patch(
|
||||
"litellm.experimental_mcp_client.client.ClientSession"
|
||||
) as mock_session:
|
||||
mock_session_instance = MagicMock()
|
||||
mock_session_instance.__aenter__ = AsyncMock(
|
||||
return_value=mock_session_instance
|
||||
)
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session.return_value = mock_session_instance
|
||||
|
||||
client = MCPClient(
|
||||
server_url="https://mcp-server.example.com",
|
||||
transport_type=MCPTransport.http,
|
||||
)
|
||||
|
||||
await client.connect()
|
||||
|
||||
# Verify streamablehttp_client was called
|
||||
mock_streamablehttp_client.assert_called_once()
|
||||
call_kwargs = mock_streamablehttp_client.call_args[1]
|
||||
|
||||
# Verify httpx_client_factory was passed
|
||||
assert "httpx_client_factory" in call_kwargs
|
||||
httpx_factory = call_kwargs["httpx_client_factory"]
|
||||
|
||||
# Test the factory creates a client with proper SSL config
|
||||
# When SSL_CERT_FILE is set, the factory should use get_ssl_configuration
|
||||
test_client = httpx_factory(headers={"test": "header"})
|
||||
|
||||
# Verify the client was created successfully with SSL configuration
|
||||
assert test_client is not None
|
||||
assert isinstance(test_client, httpx.AsyncClient)
|
||||
# Verify it has the expected properties
|
||||
assert test_client.headers is not None
|
||||
# Clean up
|
||||
await test_client.aclose()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.sse_client")
|
||||
async def test_mcp_client_ssl_verify_parameter(self, mock_sse_client):
|
||||
"""Test that MCP client uses ssl_verify parameter when provided"""
|
||||
# Setup mocks
|
||||
mock_transport = (MagicMock(), MagicMock())
|
||||
mock_sse_client.return_value.__aenter__ = AsyncMock(return_value=mock_transport)
|
||||
|
||||
# Mock the session
|
||||
with patch(
|
||||
"litellm.experimental_mcp_client.client.ClientSession"
|
||||
) as mock_session:
|
||||
mock_session_instance = MagicMock()
|
||||
mock_session_instance.__aenter__ = AsyncMock(
|
||||
return_value=mock_session_instance
|
||||
)
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session.return_value = mock_session_instance
|
||||
|
||||
# Test with ssl_verify=False
|
||||
client = MCPClient(
|
||||
server_url="https://mcp-server.example.com",
|
||||
transport_type=MCPTransport.sse,
|
||||
ssl_verify=False,
|
||||
)
|
||||
|
||||
await client.connect()
|
||||
|
||||
# Verify sse_client was called
|
||||
mock_sse_client.assert_called_once()
|
||||
call_kwargs = mock_sse_client.call_args[1]
|
||||
|
||||
# Verify httpx_client_factory was passed
|
||||
assert "httpx_client_factory" in call_kwargs
|
||||
httpx_factory = call_kwargs["httpx_client_factory"]
|
||||
|
||||
# Test the factory creates a client with SSL verification disabled
|
||||
# When ssl_verify=False, the factory should disable SSL verification
|
||||
test_client = httpx_factory(headers={"test": "header"})
|
||||
|
||||
# Verify the client was created successfully
|
||||
assert test_client is not None
|
||||
assert isinstance(test_client, httpx.AsyncClient)
|
||||
# Verify it has the expected properties
|
||||
assert test_client.headers is not None
|
||||
# Clean up
|
||||
await test_client.aclose()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.streamablehttp_client")
|
||||
async def test_mcp_client_ssl_verify_custom_path(self, mock_streamablehttp_client):
|
||||
"""Test that MCP client uses custom CA bundle path from ssl_verify parameter"""
|
||||
# Setup mocks
|
||||
mock_transport = (MagicMock(), MagicMock())
|
||||
mock_streamablehttp_client.return_value.__aenter__ = AsyncMock(
|
||||
return_value=mock_transport
|
||||
)
|
||||
|
||||
# Mock the session
|
||||
with patch(
|
||||
"litellm.experimental_mcp_client.client.ClientSession"
|
||||
) as mock_session:
|
||||
mock_session_instance = MagicMock()
|
||||
mock_session_instance.__aenter__ = AsyncMock(
|
||||
return_value=mock_session_instance
|
||||
)
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session.return_value = mock_session_instance
|
||||
|
||||
# Test with custom CA bundle path
|
||||
custom_ca_path = "/custom/path/to/ca-bundle.pem"
|
||||
client = MCPClient(
|
||||
server_url="https://mcp-server.example.com",
|
||||
transport_type=MCPTransport.http,
|
||||
ssl_verify=custom_ca_path,
|
||||
)
|
||||
|
||||
await client.connect()
|
||||
|
||||
# Verify streamablehttp_client was called
|
||||
mock_streamablehttp_client.assert_called_once()
|
||||
call_kwargs = mock_streamablehttp_client.call_args[1]
|
||||
|
||||
# Verify httpx_client_factory was passed
|
||||
assert "httpx_client_factory" in call_kwargs
|
||||
httpx_factory = call_kwargs["httpx_client_factory"]
|
||||
|
||||
# Test the factory creates a client with custom CA bundle path
|
||||
# When ssl_verify is a path, the factory should use that path for SSL verification
|
||||
test_client = httpx_factory(headers={"test": "header"})
|
||||
|
||||
# Verify the client was created successfully
|
||||
assert test_client is not None
|
||||
assert isinstance(test_client, httpx.AsyncClient)
|
||||
# Verify it has the expected properties
|
||||
assert test_client.headers is not None
|
||||
# Clean up
|
||||
await test_client.aclose()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
pytest.main([__file__])
|
||||
|
|
|
|||
|
|
@ -43,3 +43,25 @@ def test_azure_ai_validate_environment():
|
|||
litellm_params={},
|
||||
)
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
|
||||
|
||||
def test_azure_ai_grok_stop_parameter_handling():
|
||||
"""
|
||||
Test that Grok models properly handle stop parameter filtering in Azure AI Studio.
|
||||
"""
|
||||
config = AzureAIStudioConfig()
|
||||
|
||||
# Test Grok model detection
|
||||
assert config._supports_stop_reason("grok-4-fast") == False
|
||||
assert config._supports_stop_reason("grok-4") == False
|
||||
assert config._supports_stop_reason("grok-3-mini") == False
|
||||
assert config._supports_stop_reason("grok-code-fast") == False
|
||||
assert config._supports_stop_reason("gpt-4") == True
|
||||
|
||||
# Test supported parameters for Grok models
|
||||
grok_params = config.get_supported_openai_params("grok-4-fast")
|
||||
assert "stop" not in grok_params, "Grok models should not support stop parameter"
|
||||
|
||||
# Test supported parameters for non-Grok models
|
||||
gpt_params = config.get_supported_openai_params("gpt-4")
|
||||
assert "stop" in gpt_params, "GPT models should support stop parameter"
|
||||
|
|
|
|||
|
|
@ -191,7 +191,8 @@ class TestTTLExtraction:
|
|||
class TestTransformationWithTTL:
|
||||
"""Test the complete transformation with TTL support"""
|
||||
|
||||
def test_transform_with_valid_ttl(self):
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
def test_transform_with_valid_ttl(self, custom_llm_provider):
|
||||
"""Test transformation includes TTL when provided"""
|
||||
messages = [
|
||||
{
|
||||
|
|
@ -205,19 +206,32 @@ class TestTransformationWithTTL:
|
|||
]
|
||||
}
|
||||
]
|
||||
|
||||
vertex_location="test_location"
|
||||
vertex_project="test_project"
|
||||
|
||||
result = transform_openai_messages_to_gemini_context_caching(
|
||||
model="gemini-1.5-pro",
|
||||
messages=messages,
|
||||
cache_key="test-cache-key"
|
||||
cache_key="test-cache-key",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_location="test_location",
|
||||
vertex_project="test_project"
|
||||
)
|
||||
|
||||
assert "ttl" in result
|
||||
assert result["ttl"] == "3600s"
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
|
||||
if custom_llm_provider == "gemini":
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
else:
|
||||
assert result["model"] == f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/gemini-1.5-pro"
|
||||
|
||||
|
||||
assert result["displayName"] == "test-cache-key"
|
||||
|
||||
def test_transform_without_ttl(self):
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
def test_transform_without_ttl(self, custom_llm_provider):
|
||||
"""Test transformation without TTL"""
|
||||
messages = [
|
||||
{
|
||||
|
|
@ -231,18 +245,30 @@ class TestTransformationWithTTL:
|
|||
]
|
||||
}
|
||||
]
|
||||
|
||||
vertex_location="test_location"
|
||||
vertex_project="test_project"
|
||||
|
||||
result = transform_openai_messages_to_gemini_context_caching(
|
||||
model="gemini-1.5-pro",
|
||||
messages=messages,
|
||||
cache_key="test-cache-key"
|
||||
cache_key="test-cache-key",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_location=vertex_location,
|
||||
vertex_project=vertex_project
|
||||
)
|
||||
|
||||
assert "ttl" not in result
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
|
||||
if custom_llm_provider == "gemini":
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
else:
|
||||
assert result["model"] == f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/gemini-1.5-pro"
|
||||
|
||||
assert result["displayName"] == "test-cache-key"
|
||||
|
||||
def test_transform_with_invalid_ttl(self):
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
def test_transform_with_invalid_ttl(self, custom_llm_provider):
|
||||
"""Test transformation with invalid TTL (should be ignored)"""
|
||||
messages = [
|
||||
{
|
||||
|
|
@ -256,18 +282,29 @@ class TestTransformationWithTTL:
|
|||
]
|
||||
}
|
||||
]
|
||||
vertex_location="test_location"
|
||||
vertex_project="test_project"
|
||||
|
||||
result = transform_openai_messages_to_gemini_context_caching(
|
||||
model="gemini-1.5-pro",
|
||||
messages=messages,
|
||||
cache_key="test-cache-key"
|
||||
cache_key="test-cache-key",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_location=vertex_location,
|
||||
vertex_project=vertex_project
|
||||
)
|
||||
|
||||
assert "ttl" not in result
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
|
||||
if custom_llm_provider == "gemini":
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
else:
|
||||
assert result["model"] == f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/gemini-1.5-pro"
|
||||
|
||||
assert result["displayName"] == "test-cache-key"
|
||||
|
||||
def test_transform_with_system_message_and_ttl(self):
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
def test_transform_with_system_message_and_ttl(self, custom_llm_provider):
|
||||
"""Test transformation with system message and TTL"""
|
||||
messages = [
|
||||
{
|
||||
|
|
@ -290,17 +327,28 @@ class TestTransformationWithTTL:
|
|||
]
|
||||
}
|
||||
]
|
||||
|
||||
vertex_location="test_location"
|
||||
vertex_project="test_project"
|
||||
|
||||
result = transform_openai_messages_to_gemini_context_caching(
|
||||
model="gemini-1.5-pro",
|
||||
messages=messages,
|
||||
cache_key="test-cache-key"
|
||||
cache_key="test-cache-key",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_location=vertex_location,
|
||||
vertex_project=vertex_project
|
||||
)
|
||||
|
||||
assert "ttl" in result
|
||||
assert result["ttl"] == "7200s"
|
||||
assert "system_instruction" in result
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
|
||||
if custom_llm_provider == "gemini":
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
else:
|
||||
assert result["model"] == f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/gemini-1.5-pro"
|
||||
|
||||
assert result["displayName"] == "test-cache-key"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -55,6 +55,9 @@ class TestContextCachingEndpoints:
|
|||
|
||||
self.sample_optional_params = {"tools": self.sample_tools.copy()}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -62,12 +65,14 @@ class TestContextCachingEndpoints:
|
|||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj"
|
||||
)
|
||||
def test_check_and_create_cache_with_cached_content(
|
||||
self, mock_cache_obj, mock_separate
|
||||
self, mock_cache_obj, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test check_and_create_cache when cached_content is provided"""
|
||||
# Setup
|
||||
cached_content = "cached_content_123"
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -80,6 +85,10 @@ class TestContextCachingEndpoints:
|
|||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
cached_content=cached_content,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -92,14 +101,21 @@ class TestContextCachingEndpoints:
|
|||
mock_separate.assert_not_called()
|
||||
mock_cache_obj.get_cache_key.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
def test_check_and_create_cache_no_cached_messages(self, mock_separate):
|
||||
def test_check_and_create_cache_no_cached_messages(
|
||||
self, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test check_and_create_cache when no cached messages are found"""
|
||||
# Setup
|
||||
mock_separate.return_value = ([], self.sample_messages) # No cached messages
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -111,6 +127,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -119,6 +139,9 @@ class TestContextCachingEndpoints:
|
|||
assert returned_params == optional_params
|
||||
assert returned_cache is None
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -127,7 +150,7 @@ class TestContextCachingEndpoints:
|
|||
)
|
||||
@patch.object(ContextCachingEndpoints, "check_cache")
|
||||
def test_check_and_create_cache_existing_cache_found(
|
||||
self, mock_check_cache, mock_cache_obj, mock_separate
|
||||
self, mock_check_cache, mock_cache_obj, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test check_and_create_cache when existing cache is found"""
|
||||
# Setup
|
||||
|
|
@ -139,6 +162,8 @@ class TestContextCachingEndpoints:
|
|||
mock_check_cache.return_value = "existing_cache_name"
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -150,6 +175,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -163,6 +192,9 @@ class TestContextCachingEndpoints:
|
|||
messages=cached_messages, tools=self.sample_tools
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -181,6 +213,7 @@ class TestContextCachingEndpoints:
|
|||
mock_transform,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test check_and_create_cache when creating new cache"""
|
||||
# Setup
|
||||
|
|
@ -203,6 +236,8 @@ class TestContextCachingEndpoints:
|
|||
self.mock_client.post.return_value = mock_response
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -214,6 +249,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -228,6 +267,9 @@ class TestContextCachingEndpoints:
|
|||
assert "tools" in call_args.kwargs["json"]
|
||||
assert call_args.kwargs["json"]["tools"] == self.sample_tools
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -237,7 +279,12 @@ class TestContextCachingEndpoints:
|
|||
@patch.object(ContextCachingEndpoints, "check_cache")
|
||||
@patch.object(ContextCachingEndpoints, "_get_token_and_url_context_caching")
|
||||
def test_check_and_create_cache_http_error(
|
||||
self, mock_get_token_url, mock_check_cache, mock_cache_obj, mock_separate
|
||||
self,
|
||||
mock_get_token_url,
|
||||
mock_check_cache,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test check_and_create_cache handles HTTP errors properly"""
|
||||
# Setup
|
||||
|
|
@ -259,6 +306,8 @@ class TestContextCachingEndpoints:
|
|||
self.mock_client.post.side_effect = http_error
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute and Assert
|
||||
with pytest.raises(VertexAIError) as exc_info:
|
||||
|
|
@ -271,12 +320,19 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Bad Request" in str(exc_info.value.message)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -284,12 +340,14 @@ class TestContextCachingEndpoints:
|
|||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj"
|
||||
)
|
||||
async def test_async_check_and_create_cache_with_cached_content(
|
||||
self, mock_cache_obj, mock_separate
|
||||
self, mock_cache_obj, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test async_check_and_create_cache when cached_content is provided"""
|
||||
# Setup
|
||||
cached_content = "cached_content_123"
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -302,6 +360,10 @@ class TestContextCachingEndpoints:
|
|||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
cached_content=cached_content,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -311,14 +373,21 @@ class TestContextCachingEndpoints:
|
|||
assert returned_cache == cached_content
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
async def test_async_check_and_create_cache_no_cached_messages(self, mock_separate):
|
||||
async def test_async_check_and_create_cache_no_cached_messages(
|
||||
self, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test async_check_and_create_cache when no cached messages are found"""
|
||||
# Setup
|
||||
mock_separate.return_value = ([], self.sample_messages)
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -330,6 +399,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -339,6 +412,9 @@ class TestContextCachingEndpoints:
|
|||
assert returned_cache is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -347,7 +423,7 @@ class TestContextCachingEndpoints:
|
|||
)
|
||||
@patch.object(ContextCachingEndpoints, "async_check_cache")
|
||||
async def test_async_check_and_create_cache_existing_cache_found(
|
||||
self, mock_async_check_cache, mock_cache_obj, mock_separate
|
||||
self, mock_async_check_cache, mock_cache_obj, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test async_check_and_create_cache when existing cache is found"""
|
||||
# Setup
|
||||
|
|
@ -359,6 +435,8 @@ class TestContextCachingEndpoints:
|
|||
mock_async_check_cache.return_value = "existing_cache_name"
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -370,6 +448,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -384,6 +466,9 @@ class TestContextCachingEndpoints:
|
|||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -406,6 +491,7 @@ class TestContextCachingEndpoints:
|
|||
mock_transform,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test async_check_and_create_cache when creating new cache"""
|
||||
# Setup
|
||||
|
|
@ -428,6 +514,8 @@ class TestContextCachingEndpoints:
|
|||
self.mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -439,6 +527,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -454,6 +546,9 @@ class TestContextCachingEndpoints:
|
|||
assert call_args.kwargs["json"]["tools"] == self.sample_tools
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -472,6 +567,7 @@ class TestContextCachingEndpoints:
|
|||
mock_async_check_cache,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test async_check_and_create_cache handles timeout errors properly"""
|
||||
# Setup
|
||||
|
|
@ -489,6 +585,8 @@ class TestContextCachingEndpoints:
|
|||
)
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute and Assert
|
||||
with pytest.raises(VertexAIError) as exc_info:
|
||||
|
|
@ -501,12 +599,21 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 408
|
||||
assert "Timeout error occurred" in str(exc_info.value.message)
|
||||
|
||||
def test_check_and_create_cache_tools_popped_from_optional_params(self):
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
def test_check_and_create_cache_tools_popped_from_optional_params(
|
||||
self, custom_llm_provider
|
||||
):
|
||||
"""Test that tools are properly popped from optional_params when there are cached messages"""
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
|
|
@ -520,6 +627,8 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Mock the check_cache to return existing cache so we don't make HTTP calls
|
||||
with patch.object(
|
||||
|
|
@ -535,6 +644,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert tools were popped from optional_params
|
||||
|
|
@ -543,7 +656,12 @@ class TestContextCachingEndpoints:
|
|||
# But original tools should still be available for comparison
|
||||
assert original_tools == self.sample_tools
|
||||
|
||||
def test_check_and_create_cache_tools_not_popped_when_no_cached_messages(self):
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
def test_check_and_create_cache_tools_not_popped_when_no_cached_messages(
|
||||
self, custom_llm_provider
|
||||
):
|
||||
"""Test that tools are NOT popped from optional_params when there are no cached messages"""
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
|
|
@ -555,6 +673,8 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -566,6 +686,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert tools were NOT popped from optional_params (early return)
|
||||
|
|
@ -573,8 +697,11 @@ class TestContextCachingEndpoints:
|
|||
assert optional_params["tools"] == original_tools
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
async def test_async_check_and_create_cache_tools_not_popped_when_no_cached_messages(
|
||||
self,
|
||||
self, custom_llm_provider
|
||||
):
|
||||
"""Test that tools are NOT popped from optional_params in async version when there are no cached messages"""
|
||||
with patch(
|
||||
|
|
@ -587,6 +714,8 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -598,6 +727,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert tools were NOT popped from optional_params (early return)
|
||||
|
|
@ -605,7 +738,12 @@ class TestContextCachingEndpoints:
|
|||
assert optional_params["tools"] == original_tools
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_check_and_create_cache_tools_popped_from_optional_params(self):
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
async def test_async_check_and_create_cache_tools_popped_from_optional_params(
|
||||
self, custom_llm_provider
|
||||
):
|
||||
"""Test that tools are properly popped from optional_params in async version when there are cached messages"""
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
|
|
@ -619,6 +757,8 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Mock the async_check_cache to return existing cache so we don't make HTTP calls
|
||||
with patch.object(
|
||||
|
|
@ -634,6 +774,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert tools were popped from optional_params
|
||||
|
|
|
|||
|
|
@ -0,0 +1,172 @@
|
|||
"""
|
||||
Test to verify that custom headers are correctly forwarded to Gemini/Vertex AI API calls.
|
||||
|
||||
This test verifies the fix for the issue where headers configured via
|
||||
forward_client_headers_to_llm_api were not being passed to Gemini/Vertex AI providers.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import Mock, patch, MagicMock
|
||||
import litellm
|
||||
from litellm import completion
|
||||
|
||||
|
||||
class TestGeminiHeaderForwarding:
|
||||
"""Test cases for verifying header forwarding to Gemini/Vertex AI."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider,model",
|
||||
[
|
||||
("gemini", "gemini/gemini-1.5-pro"),
|
||||
("vertex_ai_beta", "gemini-1.5-pro"),
|
||||
("vertex_ai", "gemini-1.5-pro"),
|
||||
],
|
||||
)
|
||||
def test_headers_forwarded_to_gemini(self, custom_llm_provider, model):
|
||||
"""
|
||||
Test that headers from kwargs are correctly merged and passed to Gemini completion.
|
||||
|
||||
This test verifies that when headers are passed via kwargs (as the proxy does when
|
||||
forward_client_headers_to_llm_api is configured), they are correctly merged with
|
||||
extra_headers and passed to the Vertex AI completion handler.
|
||||
"""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
|
||||
# Headers that would be set by the proxy when forwarding client headers
|
||||
custom_headers = {
|
||||
"X-Custom-Header": "CustomValue",
|
||||
"X-BYOK-Token": "secret-token",
|
||||
}
|
||||
|
||||
# Mock the vertex completion handler
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.VertexLLM.completion"
|
||||
) as mock_vertex_completion:
|
||||
# Configure the mock to return a proper response
|
||||
mock_response = Mock()
|
||||
mock_response.choices = [Mock()]
|
||||
mock_response.choices[0].message.content = "Hello back!"
|
||||
mock_vertex_completion.return_value = mock_response
|
||||
|
||||
try:
|
||||
# Call completion with custom headers via kwargs
|
||||
# This simulates what the proxy does when forward_client_headers_to_llm_api is set
|
||||
completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
headers=custom_headers, # This is how proxy passes forwarded headers
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_key="dummy-key",
|
||||
)
|
||||
|
||||
# Verify that the completion handler was called
|
||||
assert mock_vertex_completion.called, "Vertex completion handler should be called"
|
||||
|
||||
# Get the actual call arguments
|
||||
call_kwargs = mock_vertex_completion.call_args.kwargs
|
||||
|
||||
# Verify that extra_headers parameter contains our custom headers
|
||||
assert "extra_headers" in call_kwargs, "extra_headers should be passed to completion"
|
||||
|
||||
passed_headers = call_kwargs["extra_headers"]
|
||||
assert passed_headers is not None, "extra_headers should not be None"
|
||||
|
||||
# Verify our custom headers are present in the passed headers
|
||||
for header_key, header_value in custom_headers.items():
|
||||
assert (
|
||||
header_key in passed_headers
|
||||
or header_key.lower() in passed_headers
|
||||
), f"Header {header_key} should be in extra_headers"
|
||||
|
||||
print(f"✓ Test passed for {custom_llm_provider}/{model}")
|
||||
print(f" Headers correctly forwarded: {passed_headers}")
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(
|
||||
f"Failed to forward headers to {custom_llm_provider}/{model}: {str(e)}"
|
||||
)
|
||||
|
||||
def test_extra_headers_and_headers_merge(self):
|
||||
"""
|
||||
Test that both extra_headers and headers parameters are correctly merged.
|
||||
|
||||
This ensures that headers from kwargs (forwarded by proxy) and extra_headers
|
||||
(passed explicitly) are both included in the final headers sent to the provider.
|
||||
"""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
|
||||
# Headers from proxy (via kwargs["headers"])
|
||||
proxy_headers = {"X-Forwarded-Header": "ProxyValue"}
|
||||
|
||||
# Explicit extra_headers
|
||||
explicit_headers = {"X-Explicit-Header": "ExplicitValue"}
|
||||
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.VertexLLM.completion"
|
||||
) as mock_vertex_completion:
|
||||
mock_response = Mock()
|
||||
mock_response.choices = [Mock()]
|
||||
mock_response.choices[0].message.content = "Response"
|
||||
mock_vertex_completion.return_value = mock_response
|
||||
|
||||
try:
|
||||
completion(
|
||||
model="gemini/gemini-1.5-pro",
|
||||
messages=messages,
|
||||
headers=proxy_headers, # From proxy forwarding
|
||||
extra_headers=explicit_headers, # Explicitly passed
|
||||
custom_llm_provider="gemini",
|
||||
api_key="dummy-key",
|
||||
)
|
||||
|
||||
call_kwargs = mock_vertex_completion.call_args.kwargs
|
||||
passed_headers = call_kwargs.get("extra_headers", {})
|
||||
|
||||
# Both sets of headers should be present
|
||||
assert (
|
||||
"X-Forwarded-Header" in passed_headers
|
||||
or "x-forwarded-header" in passed_headers
|
||||
), "Proxy forwarded header should be present"
|
||||
|
||||
assert (
|
||||
"X-Explicit-Header" in passed_headers
|
||||
or "x-explicit-header" in passed_headers
|
||||
), "Explicitly passed header should be present"
|
||||
|
||||
print("✓ Both header sources correctly merged and forwarded")
|
||||
print(f" Final headers: {passed_headers}")
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(f"Failed to merge and forward headers: {str(e)}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run the tests
|
||||
test_instance = TestGeminiHeaderForwarding()
|
||||
|
||||
print("\n" + "="*80)
|
||||
print("Testing Gemini/Vertex AI Header Forwarding")
|
||||
print("="*80 + "\n")
|
||||
|
||||
# Test each provider
|
||||
for provider, model in [
|
||||
("gemini", "gemini/gemini-1.5-pro"),
|
||||
("vertex_ai_beta", "gemini-1.5-pro"),
|
||||
("vertex_ai", "gemini-1.5-pro"),
|
||||
]:
|
||||
print(f"\nTesting {provider}/{model}...")
|
||||
try:
|
||||
test_instance.test_headers_forwarded_to_gemini(provider, model)
|
||||
except Exception as e:
|
||||
print(f"✗ Test failed: {e}")
|
||||
|
||||
print("\n\nTesting header merging...")
|
||||
try:
|
||||
test_instance.test_extra_headers_and_headers_merge()
|
||||
except Exception as e:
|
||||
print(f"✗ Test failed: {e}")
|
||||
|
||||
print("\n" + "="*80)
|
||||
print("All tests completed!")
|
||||
print("="*80 + "\n")
|
||||
|
||||
|
|
@ -734,3 +734,264 @@ async def test_call_mcp_tool_user_unauthorized_access():
|
|||
# Verify the exception details
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "User not allowed to call this tool" in exc_info.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_filters_by_key_team_permissions():
|
||||
"""Test that list_tools filters tools based on key/team mcp_tool_permissions"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_get_tools_from_mcp_servers,
|
||||
set_auth_context,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
# Create object permission with tool-level restrictions
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm_123",
|
||||
mcp_tool_permissions={
|
||||
"server1": ["tool1", "tool2"], # Only allow tool1 and tool2
|
||||
},
|
||||
)
|
||||
|
||||
user_api_key_auth = UserAPIKeyAuth(
|
||||
api_key="test_key",
|
||||
user_id="test_user",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
set_auth_context(user_api_key_auth)
|
||||
|
||||
# Mock server
|
||||
server = MagicMock()
|
||||
server.server_id = "server1"
|
||||
server.name = "Test Server"
|
||||
server.alias = "test"
|
||||
server.allowed_tools = None
|
||||
server.disallowed_tools = None
|
||||
|
||||
# Mock manager
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
|
||||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=False
|
||||
):
|
||||
# Return 4 tools, but only 2 should be allowed
|
||||
tool1 = MagicMock()
|
||||
tool1.name = "tool1"
|
||||
tool1.description = "Tool 1"
|
||||
tool1.inputSchema = {}
|
||||
|
||||
tool2 = MagicMock()
|
||||
tool2.name = "tool2"
|
||||
tool2.description = "Tool 2"
|
||||
tool2.inputSchema = {}
|
||||
|
||||
tool3 = MagicMock()
|
||||
tool3.name = "tool3"
|
||||
tool3.description = "Tool 3 - not allowed"
|
||||
tool3.inputSchema = {}
|
||||
|
||||
tool4 = MagicMock()
|
||||
tool4.name = "tool4"
|
||||
tool4.description = "Tool 4 - not allowed"
|
||||
tool4.inputSchema = {}
|
||||
|
||||
return [tool1, tool2, tool3, tool4]
|
||||
|
||||
mock_manager._get_tools_from_server = mock_get_tools_from_server
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
):
|
||||
tools = 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=None,
|
||||
)
|
||||
|
||||
# Should only return tool1 and tool2
|
||||
assert len(tools) == 2
|
||||
tool_names = sorted([t.name for t in tools])
|
||||
assert tool_names == ["tool1", "tool2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_with_team_tool_permissions_inheritance():
|
||||
"""Test that list_tools correctly applies key/team tool permissions inheritance logic"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_get_tools_from_mcp_servers,
|
||||
set_auth_context,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_TeamTable,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
# Team allows tool1, tool2, tool3
|
||||
team_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="team_perm_123",
|
||||
mcp_tool_permissions={
|
||||
"server1": ["tool1", "tool2", "tool3"],
|
||||
},
|
||||
)
|
||||
|
||||
# Key allows tool2, tool3, tool4 - intersection should be tool2, tool3
|
||||
key_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="key_perm_456",
|
||||
mcp_tool_permissions={
|
||||
"server1": ["tool2", "tool3", "tool4"],
|
||||
},
|
||||
)
|
||||
|
||||
user_api_key_auth = UserAPIKeyAuth(
|
||||
api_key="test_key",
|
||||
user_id="test_user",
|
||||
team_id="team_123",
|
||||
object_permission=key_object_permission,
|
||||
)
|
||||
set_auth_context(user_api_key_auth)
|
||||
|
||||
# Mock server
|
||||
server = MagicMock()
|
||||
server.server_id = "server1"
|
||||
server.name = "Test Server"
|
||||
server.alias = "test"
|
||||
server.allowed_tools = None
|
||||
server.disallowed_tools = None
|
||||
|
||||
# Mock manager
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
|
||||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=False
|
||||
):
|
||||
# Return 4 tools
|
||||
tool1 = MagicMock()
|
||||
tool1.name = "tool1"
|
||||
tool1.description = "Tool 1"
|
||||
tool1.inputSchema = {}
|
||||
|
||||
tool2 = MagicMock()
|
||||
tool2.name = "tool2"
|
||||
tool2.description = "Tool 2"
|
||||
tool2.inputSchema = {}
|
||||
|
||||
tool3 = MagicMock()
|
||||
tool3.name = "tool3"
|
||||
tool3.description = "Tool 3"
|
||||
tool3.inputSchema = {}
|
||||
|
||||
tool4 = MagicMock()
|
||||
tool4.name = "tool4"
|
||||
tool4.description = "Tool 4"
|
||||
tool4.inputSchema = {}
|
||||
|
||||
return [tool1, tool2, tool3, tool4]
|
||||
|
||||
mock_manager._get_tools_from_server = mock_get_tools_from_server
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
):
|
||||
# Mock the team object permission retrieval
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_team_object_permission",
|
||||
AsyncMock(return_value=team_object_permission),
|
||||
):
|
||||
tools = 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=None,
|
||||
)
|
||||
|
||||
# Should only return tool2 and tool3 (intersection of key and team permissions)
|
||||
assert len(tools) == 2
|
||||
tool_names = sorted([t.name for t in tools])
|
||||
assert tool_names == ["tool2", "tool3"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_with_no_tool_permissions_shows_all():
|
||||
"""Test that list_tools shows all tools when no mcp_tool_permissions are set"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_get_tools_from_mcp_servers,
|
||||
set_auth_context,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
# No tool-level restrictions
|
||||
user_api_key_auth = UserAPIKeyAuth(
|
||||
api_key="test_key",
|
||||
user_id="test_user",
|
||||
object_permission=None,
|
||||
)
|
||||
set_auth_context(user_api_key_auth)
|
||||
|
||||
# Mock server
|
||||
server = MagicMock()
|
||||
server.server_id = "server1"
|
||||
server.name = "Test Server"
|
||||
server.alias = "test"
|
||||
server.allowed_tools = None
|
||||
server.disallowed_tools = None
|
||||
|
||||
# Mock manager
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
|
||||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=False
|
||||
):
|
||||
# Return 3 tools
|
||||
tool1 = MagicMock()
|
||||
tool1.name = "tool1"
|
||||
tool1.description = "Tool 1"
|
||||
tool1.inputSchema = {}
|
||||
|
||||
tool2 = MagicMock()
|
||||
tool2.name = "tool2"
|
||||
tool2.description = "Tool 2"
|
||||
tool2.inputSchema = {}
|
||||
|
||||
tool3 = MagicMock()
|
||||
tool3.name = "tool3"
|
||||
tool3.description = "Tool 3"
|
||||
tool3.inputSchema = {}
|
||||
|
||||
return [tool1, tool2, tool3]
|
||||
|
||||
mock_manager._get_tools_from_server = mock_get_tools_from_server
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
):
|
||||
tools = 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=None,
|
||||
)
|
||||
|
||||
# Should return all tools when no restrictions
|
||||
assert len(tools) == 3
|
||||
tool_names = sorted([t.name for t in tools])
|
||||
assert tool_names == ["tool1", "tool2", "tool3"]
|
||||
|
|
|
|||
|
|
@ -943,6 +943,201 @@ class TestMCPServerManager:
|
|||
manager.add_update_server(server)
|
||||
assert server.server_id in manager.get_registry()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_tool_permission_allows_permitted_tool(self):
|
||||
"""
|
||||
Test that key can call tool when it's in mcp_tool_permissions allowed list.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="test_server_123",
|
||||
name="Test Server",
|
||||
transport=MCPTransport.http,
|
||||
allowed_tools=None,
|
||||
disallowed_tools=None,
|
||||
)
|
||||
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm_123",
|
||||
mcp_tool_permissions={"test_server_123": ["read_wiki_structure"]},
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
api_key="sk-test",
|
||||
user_id="user-123",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
|
||||
|
||||
# Should succeed
|
||||
await manager.pre_call_tool_check(
|
||||
name="read_wiki_structure",
|
||||
arguments={"repoName": "facebook/react"},
|
||||
server_name_from_prefix="test",
|
||||
user_api_key_auth=user_auth,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
server=server,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_tool_permission_blocks_unpermitted_tool(self):
|
||||
"""
|
||||
Test that key cannot call tool when it's NOT in mcp_tool_permissions allowed list.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="test_server_123",
|
||||
name="Test Server",
|
||||
transport=MCPTransport.http,
|
||||
allowed_tools=None,
|
||||
disallowed_tools=None,
|
||||
)
|
||||
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm_123",
|
||||
mcp_tool_permissions={"test_server_123": ["read_wiki_structure"]},
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
api_key="sk-test",
|
||||
user_id="user-123",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
|
||||
|
||||
# Should fail with 403
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await manager.pre_call_tool_check(
|
||||
name="ask_question",
|
||||
arguments={"question": "test"},
|
||||
server_name_from_prefix="test",
|
||||
user_api_key_auth=user_auth,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
server=server,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_tool_permission_for_key_team_allows_permitted_tool(self):
|
||||
"""
|
||||
Test check_tool_permission_for_key_team directly - should allow permitted tool.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="github_server",
|
||||
name="GitHub Server",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm_456",
|
||||
mcp_tool_permissions={"github_server": ["read_repo", "list_issues"]},
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
api_key="sk-test-key",
|
||||
user_id="user-456",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
# Should not raise exception for allowed tool
|
||||
await manager.check_tool_permission_for_key_team(
|
||||
tool_name="read_repo",
|
||||
server=server,
|
||||
user_api_key_auth=user_auth,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_tool_permission_for_key_team_blocks_unpermitted_tool(self):
|
||||
"""
|
||||
Test check_tool_permission_for_key_team directly - should block unpermitted tool.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="github_server",
|
||||
name="GitHub Server",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm_456",
|
||||
mcp_tool_permissions={"github_server": ["read_repo"]},
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
api_key="sk-test-key",
|
||||
user_id="user-456",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
# Should raise HTTPException for unpermitted tool
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await manager.check_tool_permission_for_key_team(
|
||||
tool_name="delete_repo",
|
||||
server=server,
|
||||
user_api_key_auth=user_auth,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "delete_repo" in exc_info.value.detail["error"]
|
||||
assert "not allowed" in exc_info.value.detail["error"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_tool_permission_for_key_team_allows_all_when_no_restrictions(
|
||||
self,
|
||||
):
|
||||
"""
|
||||
Test check_tool_permission_for_key_team - should allow all tools when no restrictions set.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="github_server",
|
||||
name="GitHub Server",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
|
||||
# No object_permission set on user_auth
|
||||
user_auth = UserAPIKeyAuth(
|
||||
api_key="sk-test-key",
|
||||
user_id="user-456",
|
||||
object_permission=None,
|
||||
)
|
||||
|
||||
# Should allow any tool when no restrictions
|
||||
await manager.check_tool_permission_for_key_team(
|
||||
tool_name="any_tool",
|
||||
server=server,
|
||||
user_api_key_auth=user_auth,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
|
|
|||
|
|
@ -421,6 +421,80 @@ async def test_key_generation_with_object_permission(monkeypatch):
|
|||
assert key_insert_calls[0]["data"].get("object_permission_id") == "objperm123"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_generation_with_mcp_tool_permissions(monkeypatch):
|
||||
"""
|
||||
Test that /key/generate correctly handles mcp_tool_permissions in object_permission.
|
||||
|
||||
This test verifies that:
|
||||
1. mcp_tool_permissions is accepted in the object_permission field
|
||||
2. The field is properly stored in the LiteLLM_ObjectPermissionTable
|
||||
3. The key is correctly linked to the object_permission record
|
||||
"""
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.jsonify_object = lambda data: data
|
||||
|
||||
# Track what data is passed to create
|
||||
created_permission_data = {}
|
||||
|
||||
async def mock_create(**kwargs):
|
||||
created_permission_data.update(kwargs.get("data", {}))
|
||||
return MagicMock(object_permission_id="objperm_mcp_123")
|
||||
|
||||
mock_prisma_client.db = MagicMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable = MagicMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.create = mock_create
|
||||
|
||||
async def _insert_data_side_effect(*args, **kwargs):
|
||||
table_name = kwargs.get("table_name")
|
||||
if table_name == "user":
|
||||
return MagicMock(models=[], spend=0)
|
||||
elif table_name == "key":
|
||||
return MagicMock(
|
||||
token="hashed_token_789",
|
||||
litellm_budget_table=None,
|
||||
object_permission=None,
|
||||
)
|
||||
return MagicMock()
|
||||
|
||||
mock_prisma_client.insert_data = AsyncMock(side_effect=_insert_data_side_effect)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
|
||||
from litellm.proxy._types import (
|
||||
GenerateKeyRequest,
|
||||
LiteLLM_ObjectPermissionBase,
|
||||
LitellmUserRoles,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
generate_key_fn,
|
||||
)
|
||||
|
||||
# Create request with mcp_tool_permissions
|
||||
request_data = GenerateKeyRequest(
|
||||
object_permission=LiteLLM_ObjectPermissionBase(
|
||||
mcp_servers=["server_1"],
|
||||
mcp_tool_permissions={"server_1": ["tool1", "tool2", "tool3"]},
|
||||
)
|
||||
)
|
||||
|
||||
await generate_key_fn(
|
||||
data=request_data,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="user-mcp-1",
|
||||
),
|
||||
)
|
||||
|
||||
# Verify mcp_tool_permissions was stored
|
||||
assert "mcp_tool_permissions" in created_permission_data
|
||||
assert created_permission_data["mcp_tool_permissions"] == {
|
||||
"server_1": ["tool1", "tool2", "tool3"]
|
||||
}
|
||||
assert created_permission_data["mcp_servers"] == ["server_1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_update_object_permissions_existing_permission(monkeypatch):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -336,6 +336,90 @@ async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth):
|
|||
assert created_team_kwargs["data"].get("object_permission_id") == "objperm123"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_team_with_mcp_tool_permissions(mock_db_client, mock_admin_auth):
|
||||
"""
|
||||
Test that /team/new correctly handles mcp_tool_permissions in object_permission.
|
||||
|
||||
This test verifies that:
|
||||
1. mcp_tool_permissions is accepted in the object_permission field
|
||||
2. The field is properly stored in the LiteLLM_ObjectPermissionTable
|
||||
3. The team is correctly linked to the object_permission record
|
||||
"""
|
||||
# Configure mocked prisma client
|
||||
mock_db_client.jsonify_team_object = lambda db_data: db_data
|
||||
mock_db_client.get_data = AsyncMock(return_value=None)
|
||||
mock_db_client.update_data = AsyncMock(return_value=MagicMock())
|
||||
mock_db_client.db = MagicMock()
|
||||
|
||||
# Track what data is passed to object permission create
|
||||
created_permission_data = {}
|
||||
|
||||
async def mock_obj_perm_create(**kwargs):
|
||||
created_permission_data.update(kwargs.get("data", {}))
|
||||
return MagicMock(object_permission_id="objperm_team_mcp_456")
|
||||
|
||||
mock_db_client.db.litellm_objectpermissiontable = MagicMock()
|
||||
mock_db_client.db.litellm_objectpermissiontable.create = mock_obj_perm_create
|
||||
|
||||
# Mock model table
|
||||
mock_db_client.db.litellm_modeltable = MagicMock()
|
||||
mock_db_client.db.litellm_modeltable.create = AsyncMock(
|
||||
return_value=MagicMock(id="model456")
|
||||
)
|
||||
|
||||
# Mock team table
|
||||
team_create_result = MagicMock(
|
||||
team_id="team-mcp-789",
|
||||
object_permission_id="objperm_team_mcp_456",
|
||||
)
|
||||
team_create_result.model_dump.return_value = {
|
||||
"team_id": "team-mcp-789",
|
||||
"object_permission_id": "objperm_team_mcp_456",
|
||||
}
|
||||
mock_db_client.db.litellm_teamtable = MagicMock()
|
||||
mock_db_client.db.litellm_teamtable.create = AsyncMock(return_value=team_create_result)
|
||||
mock_db_client.db.litellm_teamtable.count = AsyncMock(return_value=0)
|
||||
mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=team_create_result)
|
||||
|
||||
# Mock user table
|
||||
mock_db_client.db.litellm_usertable = MagicMock()
|
||||
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionBase, NewTeamRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import new_team
|
||||
|
||||
# Create team with mcp_tool_permissions
|
||||
team_request = NewTeamRequest(
|
||||
team_alias="mcp-team",
|
||||
object_permission=LiteLLM_ObjectPermissionBase(
|
||||
mcp_servers=["server_a", "server_b"],
|
||||
mcp_tool_permissions={
|
||||
"server_a": ["read_wiki_structure", "read_wiki_contents"],
|
||||
"server_b": ["ask_question"],
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
await new_team(
|
||||
data=team_request,
|
||||
http_request=dummy_request,
|
||||
user_api_key_dict=mock_admin_auth,
|
||||
)
|
||||
|
||||
# Verify mcp_tool_permissions was stored
|
||||
assert "mcp_tool_permissions" in created_permission_data
|
||||
assert created_permission_data["mcp_tool_permissions"] == {
|
||||
"server_a": ["read_wiki_structure", "read_wiki_contents"],
|
||||
"server_b": ["ask_question"],
|
||||
}
|
||||
assert created_permission_data["mcp_servers"] == ["server_a", "server_b"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_update_object_permissions_existing_permission(monkeypatch):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -2030,3 +2030,89 @@ def test_add_callback_from_db_to_in_memory_litellm_callbacks():
|
|||
existing_callbacks=existing_callbacks_with_item,
|
||||
)
|
||||
mock_callback_manager.add_litellm_success_callback.assert_not_called()
|
||||
|
||||
|
||||
def test_should_load_db_object_with_supported_db_objects():
|
||||
"""
|
||||
Test _should_load_db_object method with supported_db_objects configuration.
|
||||
|
||||
Verifies that when supported_db_objects is set, only specified object types
|
||||
are loaded from the database.
|
||||
"""
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
proxy_config = ProxyConfig()
|
||||
|
||||
# Test Case 1: supported_db_objects not set - all objects should be loaded
|
||||
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is True
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||||
assert proxy_config._should_load_db_object(object_type="guardrails") is True
|
||||
assert proxy_config._should_load_db_object(object_type="vector_stores") is True
|
||||
|
||||
# Test Case 2: supported_db_objects set to only load MCP
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"supported_db_objects": ["mcp"]},
|
||||
):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is False
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||||
assert proxy_config._should_load_db_object(object_type="guardrails") is False
|
||||
assert proxy_config._should_load_db_object(object_type="vector_stores") is False
|
||||
assert proxy_config._should_load_db_object(object_type="prompts") is False
|
||||
|
||||
# Test Case 3: supported_db_objects set to load multiple types
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"supported_db_objects": ["mcp", "guardrails", "vector_stores"]},
|
||||
):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is False
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||||
assert proxy_config._should_load_db_object(object_type="guardrails") is True
|
||||
assert proxy_config._should_load_db_object(object_type="vector_stores") is True
|
||||
assert proxy_config._should_load_db_object(object_type="prompts") is False
|
||||
|
||||
# Test Case 4: supported_db_objects is not a list (should default to loading all)
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"supported_db_objects": "invalid_type"},
|
||||
):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is True
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||||
|
||||
# Test Case 5: supported_db_objects is an empty list (nothing should be loaded)
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"supported_db_objects": []},
|
||||
):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is False
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is False
|
||||
assert proxy_config._should_load_db_object(object_type="guardrails") is False
|
||||
|
||||
# Test Case 6: Test all available object types
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{
|
||||
"supported_db_objects": [
|
||||
"models",
|
||||
"mcp",
|
||||
"guardrails",
|
||||
"vector_stores",
|
||||
"pass_through_endpoints",
|
||||
"prompts",
|
||||
"model_cost_map",
|
||||
]
|
||||
},
|
||||
):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is True
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||||
assert proxy_config._should_load_db_object(object_type="guardrails") is True
|
||||
assert proxy_config._should_load_db_object(object_type="vector_stores") is True
|
||||
assert (
|
||||
proxy_config._should_load_db_object(object_type="pass_through_endpoints")
|
||||
is True
|
||||
)
|
||||
assert proxy_config._should_load_db_object(object_type="prompts") is True
|
||||
assert proxy_config._should_load_db_object(object_type="model_cost_map") is True
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
"version": "0.1.0",
|
||||
"private": true,
|
||||
"scripts": {
|
||||
"dev": "next dev",
|
||||
"dev": "next dev --turbo",
|
||||
"build": "next build",
|
||||
"start": "next start",
|
||||
"lint": "next lint",
|
||||
|
|
|
|||
|
|
@ -325,14 +325,46 @@ const GeneralSettings: React.FC<GeneralSettingsPageProps> = ({ accessToken, user
|
|||
|
||||
console.log("router_settings", router_settings);
|
||||
|
||||
const numberKeys = new Set(["allowed_fails", "cooldown_time", "num_retries", "timeout", "retry_after"]);
|
||||
const jsonKeys = new Set(["model_group_alias", "retry_policy"]);
|
||||
|
||||
const parseInputValue = (key: string, raw: string | undefined, fallback: unknown) => {
|
||||
if (raw === undefined) return fallback;
|
||||
|
||||
const v = raw.trim();
|
||||
|
||||
if (v.toLowerCase() === "null") return null;
|
||||
|
||||
if (numberKeys.has(key)) {
|
||||
const n = Number(v);
|
||||
return Number.isNaN(n) ? fallback : n;
|
||||
}
|
||||
|
||||
if (jsonKeys.has(key)) {
|
||||
if (v === "") return null;
|
||||
try {
|
||||
return JSON.parse(v);
|
||||
} catch {
|
||||
return fallback;
|
||||
}
|
||||
}
|
||||
|
||||
if (v.toLowerCase() === "true") return true;
|
||||
if (v.toLowerCase() === "false") return false;
|
||||
|
||||
return v;
|
||||
};
|
||||
|
||||
const updatedVariables = Object.fromEntries(
|
||||
Object.entries(router_settings)
|
||||
.map(([key, value]) => {
|
||||
if (key !== "routing_strategy_args" && key !== "routing_strategy") {
|
||||
return [key, (document.querySelector(`input[name="${key}"]`) as HTMLInputElement)?.value || value];
|
||||
} else if (key == "routing_strategy") {
|
||||
const inputEl = document.querySelector(`input[name="${key}"]`) as HTMLInputElement | null;
|
||||
const parsed = parseInputValue(key, inputEl?.value, value);
|
||||
return [key, parsed];
|
||||
} else if (key === "routing_strategy") {
|
||||
return [key, selectedStrategy];
|
||||
} else if (key == "routing_strategy_args" && selectedStrategy == "latency-based-routing") {
|
||||
} else if (key === "routing_strategy_args" && selectedStrategy === "latency-based-routing") {
|
||||
let setRoutingStrategyArgs: routingStrategyArgs = {};
|
||||
|
||||
const lowestLatencyBufferElement = document.querySelector(
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
import { keyListCall, teamListCall, organizationListCall } from "../networking";
|
||||
import { teamListCall, organizationListCall, keyAliasesCall } from "../networking"
|
||||
import { Team } from "./key_list";
|
||||
import { Organization } from "../networking";
|
||||
|
||||
/**
|
||||
* Fetches all key aliases across all pages
|
||||
* Fetches all key aliases via the dedicated /key/aliases endpoint
|
||||
* @param accessToken The access token for API authentication
|
||||
* @returns Array of all unique key aliases
|
||||
*/
|
||||
|
|
@ -11,44 +11,16 @@ export const fetchAllKeyAliases = async (accessToken: string | null): Promise<st
|
|||
if (!accessToken) return [];
|
||||
|
||||
try {
|
||||
// Fetch all pages of keys to extract aliases
|
||||
let allAliases: string[] = [];
|
||||
let currentPage = 1;
|
||||
let hasMorePages = true;
|
||||
|
||||
while (hasMorePages) {
|
||||
const response = await keyListCall(
|
||||
accessToken,
|
||||
null, // organization_id
|
||||
"", // team_id
|
||||
null, // selectedKeyAlias
|
||||
null, // user_id
|
||||
null, // key_hash
|
||||
currentPage,
|
||||
100, // larger page size to reduce number of requests
|
||||
);
|
||||
|
||||
// Extract aliases from this page
|
||||
const pageAliases = response.keys.map((key: any) => key.key_alias).filter(Boolean) as string[];
|
||||
|
||||
allAliases = [...allAliases, ...pageAliases];
|
||||
|
||||
// Check if there are more pages
|
||||
if (currentPage < response.total_pages) {
|
||||
currentPage++;
|
||||
} else {
|
||||
hasMorePages = false;
|
||||
}
|
||||
}
|
||||
|
||||
// Remove duplicates
|
||||
return Array.from(new Set(allAliases));
|
||||
const { aliases } = await keyAliasesCall(accessToken as unknown as String);
|
||||
// Defensive dedupe & null-guard
|
||||
return Array.from(new Set((aliases || []).filter(Boolean)));
|
||||
} catch (error) {
|
||||
console.error("Error fetching all key aliases:", error);
|
||||
return [];
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
/**
|
||||
* Fetches all teams across all pages
|
||||
* @param accessToken The access token for API authentication
|
||||
|
|
|
|||
|
|
@ -80,6 +80,7 @@ export interface KeyResponse {
|
|||
object_permission_id: string;
|
||||
mcp_servers: string[];
|
||||
mcp_access_groups?: string[];
|
||||
mcp_tool_permissions?: Record<string, string[]>;
|
||||
vector_stores: string[];
|
||||
};
|
||||
auto_rotate?: boolean;
|
||||
|
|
|
|||
|
|
@ -4,8 +4,14 @@ import { fetchMCPServers, fetchMCPAccessGroups } from "../networking";
|
|||
import { MCPServer } from "../mcp_tools/types";
|
||||
|
||||
interface MCPServerSelectorProps {
|
||||
onChange: (selected: { servers: string[]; accessGroups: string[] }) => void;
|
||||
value?: { servers: string[]; accessGroups: string[] };
|
||||
onChange: (selected: {
|
||||
servers: string[];
|
||||
accessGroups: string[];
|
||||
}) => void;
|
||||
value?: {
|
||||
servers: string[];
|
||||
accessGroups: string[];
|
||||
};
|
||||
className?: string;
|
||||
accessToken: string;
|
||||
placeholder?: string;
|
||||
|
|
|
|||
|
|
@ -0,0 +1,79 @@
|
|||
import { describe, it, expect, vi, beforeEach } from "vitest";
|
||||
import { render, screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import MCPToolPermissions from "./MCPToolPermissions";
|
||||
import * as networking from "../networking";
|
||||
|
||||
vi.mock("../networking");
|
||||
|
||||
describe("MCPToolPermissions", () => {
|
||||
const mockAccessToken = "test-token";
|
||||
const mockServerId = "server-123";
|
||||
const mockServerName = "Test MCP Server";
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("should update tool permissions when user selects a tool", async () => {
|
||||
/**
|
||||
* Tests that clicking a tool checkbox calls onChange with updated permissions.
|
||||
* This is the core functionality of the component.
|
||||
*/
|
||||
const mockOnChange = vi.fn();
|
||||
const mockTools = [
|
||||
{ name: "read_wiki_structure", description: "Get documentation topics" },
|
||||
{ name: "read_wiki_contents", description: "View documentation" },
|
||||
{ name: "ask_question", description: "Ask questions" },
|
||||
];
|
||||
|
||||
// Mock fetchMCPServers to return server details
|
||||
vi.mocked(networking.fetchMCPServers).mockResolvedValue({
|
||||
data: [
|
||||
{
|
||||
server_id: mockServerId,
|
||||
server_name: mockServerName,
|
||||
alias: mockServerName,
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
// Mock listMCPTools to return tools for the server
|
||||
vi.mocked(networking.listMCPTools).mockResolvedValue({
|
||||
tools: mockTools,
|
||||
error: false,
|
||||
});
|
||||
|
||||
render(
|
||||
<MCPToolPermissions
|
||||
accessToken={mockAccessToken}
|
||||
selectedServers={[mockServerId]}
|
||||
toolPermissions={{}}
|
||||
onChange={mockOnChange}
|
||||
/>
|
||||
);
|
||||
|
||||
// Wait for server and tools to load
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText(mockServerName)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("read_wiki_structure")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
// Get all checkboxes and click the first one (for read_wiki_structure)
|
||||
const checkboxes = screen.getAllByRole("checkbox");
|
||||
await userEvent.click(checkboxes[0]);
|
||||
|
||||
// Verify onChange was called with correct permissions
|
||||
expect(mockOnChange).toHaveBeenCalledWith({
|
||||
[mockServerId]: ["read_wiki_structure"],
|
||||
});
|
||||
|
||||
// Verify API calls
|
||||
expect(networking.fetchMCPServers).toHaveBeenCalledWith(mockAccessToken);
|
||||
expect(networking.listMCPTools).toHaveBeenCalledWith(mockAccessToken, mockServerId);
|
||||
});
|
||||
});
|
||||
|
||||
|
|
@ -0,0 +1,225 @@
|
|||
import React, { useEffect, useState } from "react";
|
||||
import { listMCPTools, fetchMCPServers } from "../networking";
|
||||
import { MCPTool, MCPServer } from "../mcp_tools/types";
|
||||
import { Text } from "@tremor/react";
|
||||
import { Spin, Checkbox } from "antd";
|
||||
import { XIcon } from "lucide-react";
|
||||
|
||||
interface MCPToolPermissionsProps {
|
||||
accessToken: string;
|
||||
selectedServers: string[];
|
||||
toolPermissions: Record<string, string[]>;
|
||||
onChange: (toolPermissions: Record<string, string[]>) => void;
|
||||
disabled?: boolean;
|
||||
}
|
||||
|
||||
const MCPToolPermissions: React.FC<MCPToolPermissionsProps> = ({
|
||||
accessToken,
|
||||
selectedServers,
|
||||
toolPermissions,
|
||||
onChange,
|
||||
disabled = false,
|
||||
}) => {
|
||||
const [servers, setServers] = useState<MCPServer[]>([]);
|
||||
const [serverTools, setServerTools] = useState<Record<string, MCPTool[]>>({});
|
||||
const [loadingTools, setLoadingTools] = useState<Record<string, boolean>>({});
|
||||
const [toolErrors, setToolErrors] = useState<Record<string, string>>({});
|
||||
|
||||
// Fetch server details
|
||||
useEffect(() => {
|
||||
const loadServerDetails = async () => {
|
||||
if (selectedServers.length === 0) {
|
||||
setServers([]);
|
||||
return;
|
||||
}
|
||||
|
||||
try {
|
||||
const response = await fetchMCPServers(accessToken);
|
||||
const allServers = Array.isArray(response) ? response : response.data || [];
|
||||
|
||||
const filteredServers = allServers.filter((server: MCPServer) =>
|
||||
selectedServers.includes(server.server_id)
|
||||
);
|
||||
|
||||
setServers(filteredServers);
|
||||
} catch (error) {
|
||||
console.error("Error fetching MCP servers:", error);
|
||||
setServers([]);
|
||||
}
|
||||
};
|
||||
|
||||
loadServerDetails();
|
||||
}, [selectedServers, accessToken]);
|
||||
|
||||
// Fetch tools for a specific server
|
||||
const fetchToolsForServer = async (serverId: string) => {
|
||||
setLoadingTools(prev => ({ ...prev, [serverId]: true }));
|
||||
setToolErrors(prev => ({ ...prev, [serverId]: "" }));
|
||||
|
||||
try {
|
||||
const response = await listMCPTools(accessToken, serverId);
|
||||
|
||||
if (response.error) {
|
||||
setToolErrors(prev => ({ ...prev, [serverId]: response.message || "Failed to fetch tools" }));
|
||||
setServerTools(prev => ({ ...prev, [serverId]: [] }));
|
||||
} else {
|
||||
setServerTools(prev => ({ ...prev, [serverId]: response.tools || [] }));
|
||||
}
|
||||
} catch (err) {
|
||||
console.error(`Error fetching tools for server ${serverId}:`, err);
|
||||
setToolErrors(prev => ({ ...prev, [serverId]: "Failed to fetch tools" }));
|
||||
setServerTools(prev => ({ ...prev, [serverId]: [] }));
|
||||
} finally {
|
||||
setLoadingTools(prev => ({ ...prev, [serverId]: false }));
|
||||
}
|
||||
};
|
||||
|
||||
// Auto-fetch tools when servers change
|
||||
useEffect(() => {
|
||||
servers.forEach(server => {
|
||||
if (!serverTools[server.server_id] && !loadingTools[server.server_id]) {
|
||||
fetchToolsForServer(server.server_id);
|
||||
}
|
||||
});
|
||||
}, [servers]);
|
||||
|
||||
// Handle tool selection
|
||||
const handleToolToggle = (serverId: string, toolName: string) => {
|
||||
const currentTools = toolPermissions[serverId] || [];
|
||||
const newTools = currentTools.includes(toolName)
|
||||
? currentTools.filter(name => name !== toolName)
|
||||
: [...currentTools, toolName];
|
||||
|
||||
const updatedPermissions = {
|
||||
...toolPermissions,
|
||||
[serverId]: newTools,
|
||||
};
|
||||
onChange(updatedPermissions);
|
||||
};
|
||||
|
||||
const handleSelectAll = (serverId: string) => {
|
||||
const tools = serverTools[serverId] || [];
|
||||
onChange({
|
||||
...toolPermissions,
|
||||
[serverId]: tools.map(t => t.name),
|
||||
});
|
||||
};
|
||||
|
||||
const handleDeselectAll = (serverId: string) => {
|
||||
onChange({
|
||||
...toolPermissions,
|
||||
[serverId]: [],
|
||||
});
|
||||
};
|
||||
|
||||
if (selectedServers.length === 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="space-y-4">
|
||||
{servers.map((server) => {
|
||||
const serverName = server.server_name || server.alias || server.server_id;
|
||||
const tools = serverTools[server.server_id] || [];
|
||||
const selectedTools = toolPermissions[server.server_id] || [];
|
||||
const isLoading = loadingTools[server.server_id];
|
||||
const error = toolErrors[server.server_id];
|
||||
|
||||
return (
|
||||
<div key={server.server_id} className="border rounded-lg bg-gray-50">
|
||||
{/* Header */}
|
||||
<div className="flex items-center justify-between p-4 border-b bg-white rounded-t-lg">
|
||||
<div>
|
||||
<Text className="font-semibold text-gray-900">{serverName}</Text>
|
||||
{server.description && (
|
||||
<Text className="text-sm text-gray-500">{server.description}</Text>
|
||||
)}
|
||||
</div>
|
||||
<div className="flex items-center gap-3">
|
||||
<button
|
||||
className="text-sm text-blue-600 hover:text-blue-700 font-medium"
|
||||
onClick={() => handleSelectAll(server.server_id)}
|
||||
disabled={disabled || isLoading}
|
||||
>
|
||||
Select All
|
||||
</button>
|
||||
<button
|
||||
className="text-sm text-blue-600 hover:text-blue-700 font-medium"
|
||||
onClick={() => handleDeselectAll(server.server_id)}
|
||||
disabled={disabled || isLoading}
|
||||
>
|
||||
Deselect All
|
||||
</button>
|
||||
<button
|
||||
className="text-gray-400 hover:text-gray-600"
|
||||
onClick={() => {
|
||||
// Handle remove server if needed
|
||||
}}
|
||||
>
|
||||
<XIcon className="w-4 h-4" />
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Tools */}
|
||||
<div className="p-4">
|
||||
<Text className="text-sm font-medium text-gray-700 mb-3">Available Tools</Text>
|
||||
|
||||
{/* Loading */}
|
||||
{isLoading && (
|
||||
<div className="flex items-center justify-center py-8">
|
||||
<Spin size="large" />
|
||||
<Text className="ml-3 text-gray-500">Loading tools...</Text>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Error */}
|
||||
{error && !isLoading && (
|
||||
<div className="p-4 bg-red-50 border border-red-200 rounded-lg text-center">
|
||||
<Text className="text-red-600 font-medium">Unable to load tools</Text>
|
||||
<Text className="text-sm text-red-500 mt-1">{error}</Text>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Tool List - Compact */}
|
||||
{!isLoading && !error && tools.length > 0 && (
|
||||
<div className="space-y-2">
|
||||
{tools.map((tool) => {
|
||||
const isSelected = selectedTools.includes(tool.name);
|
||||
|
||||
return (
|
||||
<div key={tool.name} className="flex items-start gap-2">
|
||||
<Checkbox
|
||||
checked={isSelected}
|
||||
onChange={() => handleToolToggle(server.server_id, tool.name)}
|
||||
disabled={disabled}
|
||||
/>
|
||||
<div className="flex-1 min-w-0">
|
||||
<div className="flex items-center gap-2">
|
||||
<Text className="font-medium text-gray-900">{tool.name}</Text>
|
||||
<Text className="text-sm text-gray-500">
|
||||
- {tool.description || "No description"}
|
||||
</Text>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Empty State */}
|
||||
{!isLoading && !error && tools.length === 0 && (
|
||||
<div className="text-center py-6">
|
||||
<Text className="text-gray-500">No tools available</Text>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
export default MCPToolPermissions;
|
||||
|
|
@ -3102,6 +3102,41 @@ export const keyListCall = async (
|
|||
}
|
||||
};
|
||||
|
||||
export const keyAliasesCall = async (
|
||||
accessToken: String
|
||||
): Promise<{ aliases: string[] }> => {
|
||||
/**
|
||||
* Get all key aliases from proxy
|
||||
*/
|
||||
try {
|
||||
let url = proxyBaseUrl ? `${proxyBaseUrl}/key/aliases` : `/key/aliases`;
|
||||
console.log("in keyAliasesCall");
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "GET",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
|
||||
const data = await response.json();
|
||||
console.log("/key/aliases API Response:", data);
|
||||
return data; // { aliases: string[] }
|
||||
} catch (error) {
|
||||
console.error("Failed to fetch key aliases:", error);
|
||||
throw error;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
export const spendUsersCall = async (accessToken: String, userID: String) => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/spend/users` : `/spend/users`;
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ import BudgetDurationDropdown from "../common_components/budget_duration_dropdow
|
|||
import { formatNumberWithCommas } from "@/utils/dataUtils";
|
||||
import { callback_map, mapDisplayToInternalNames } from "../callback_info_helpers";
|
||||
import MCPServerSelector from "../mcp_server_management/MCPServerSelector";
|
||||
import MCPToolPermissions from "../mcp_server_management/MCPToolPermissions";
|
||||
import ModelAliasManager from "../common_components/ModelAliasManager";
|
||||
import NotificationsManager from "../molecules/notifications_manager";
|
||||
import KeyLifecycleSettings from "../common_components/KeyLifecycleSettings";
|
||||
|
|
@ -357,6 +358,16 @@ const CreateKey: React.FC<CreateKeyProps> = ({
|
|||
delete formValues.allowed_mcp_servers_and_groups;
|
||||
}
|
||||
|
||||
// Add MCP tool permissions to object_permission
|
||||
const mcpToolPermissions = formValues.mcp_tool_permissions || {};
|
||||
if (Object.keys(mcpToolPermissions).length > 0) {
|
||||
if (!formValues.object_permission) {
|
||||
formValues.object_permission = {};
|
||||
}
|
||||
formValues.object_permission.mcp_tool_permissions = mcpToolPermissions;
|
||||
}
|
||||
delete formValues.mcp_tool_permissions;
|
||||
|
||||
// Transform allowed_mcp_access_groups into object_permission format
|
||||
if (formValues.allowed_mcp_access_groups && formValues.allowed_mcp_access_groups.length > 0) {
|
||||
if (!formValues.object_permission) {
|
||||
|
|
@ -372,6 +383,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({
|
|||
formValues.aliases = JSON.stringify(modelAliases);
|
||||
}
|
||||
|
||||
|
||||
let response;
|
||||
if (keyOwner === "service_account") {
|
||||
response = await keyCreateServiceAccountCall(accessToken, formValues);
|
||||
|
|
@ -715,15 +727,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({
|
|||
{/* Section 3: Optional Settings */}
|
||||
{!isFormDisabled && (
|
||||
<div className="mb-8">
|
||||
<Accordion
|
||||
className="mt-4 mb-4"
|
||||
onClick={() => {
|
||||
if (!mcpAccessGroupsLoaded) {
|
||||
fetchMcpAccessGroups();
|
||||
setMcpAccessGroupsLoaded(true);
|
||||
}
|
||||
}}
|
||||
>
|
||||
<Accordion className="mt-4 mb-4">
|
||||
<AccordionHeader>
|
||||
<Title className="m-0">Optional Settings</Title>
|
||||
</AccordionHeader>
|
||||
|
|
@ -923,28 +927,6 @@ const CreateKey: React.FC<CreateKeyProps> = ({
|
|||
placeholder="Select vector stores (optional)"
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label={
|
||||
<span>
|
||||
Allowed MCP Servers{" "}
|
||||
<Tooltip title="Select which MCP servers or access groups this key can access. ">
|
||||
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="allowed_mcp_servers_and_groups"
|
||||
className="mt-4"
|
||||
help="Select MCP servers or access groups this key can access. "
|
||||
>
|
||||
<MCPServerSelector
|
||||
onChange={(val: any) => form.setFieldValue("allowed_mcp_servers_and_groups", val)}
|
||||
value={form.getFieldValue("allowed_mcp_servers_and_groups")}
|
||||
accessToken={accessToken}
|
||||
placeholder="Select MCP servers or access groups (optional)"
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label={
|
||||
<span>
|
||||
|
|
@ -980,6 +962,64 @@ const CreateKey: React.FC<CreateKeyProps> = ({
|
|||
options={predefinedTags}
|
||||
/>
|
||||
</Form.Item>
|
||||
<Accordion
|
||||
className="mt-4 mb-4"
|
||||
onClick={() => {
|
||||
if (!mcpAccessGroupsLoaded) {
|
||||
fetchMcpAccessGroups();
|
||||
setMcpAccessGroupsLoaded(true);
|
||||
}
|
||||
}}
|
||||
>
|
||||
<AccordionHeader>
|
||||
<b>MCP Settings</b>
|
||||
</AccordionHeader>
|
||||
<AccordionBody>
|
||||
<Form.Item
|
||||
label={
|
||||
<span>
|
||||
Allowed MCP Servers{" "}
|
||||
<Tooltip title="Select which MCP servers or access groups this key can access">
|
||||
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="allowed_mcp_servers_and_groups"
|
||||
help="Select MCP servers or access groups this key can access"
|
||||
>
|
||||
<MCPServerSelector
|
||||
onChange={(val: any) => form.setFieldValue("allowed_mcp_servers_and_groups", val)}
|
||||
value={form.getFieldValue("allowed_mcp_servers_and_groups")}
|
||||
accessToken={accessToken}
|
||||
placeholder="Select MCP servers or access groups (optional)"
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
{/* Hidden field to register mcp_tool_permissions with the form */}
|
||||
<Form.Item name="mcp_tool_permissions" initialValue={{}} hidden>
|
||||
<Input type="hidden" />
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
noStyle
|
||||
shouldUpdate={(prevValues, currentValues) =>
|
||||
prevValues.allowed_mcp_servers_and_groups !== currentValues.allowed_mcp_servers_and_groups ||
|
||||
prevValues.mcp_tool_permissions !== currentValues.mcp_tool_permissions
|
||||
}
|
||||
>
|
||||
{() => (
|
||||
<div className="mt-6">
|
||||
<MCPToolPermissions
|
||||
accessToken={accessToken}
|
||||
selectedServers={form.getFieldValue("allowed_mcp_servers_and_groups")?.servers || []}
|
||||
toolPermissions={form.getFieldValue("mcp_tool_permissions") || {}}
|
||||
onChange={(toolPerms) => form.setFieldsValue({ mcp_tool_permissions: toolPerms })}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</Form.Item>
|
||||
</AccordionBody>
|
||||
</Accordion>
|
||||
|
||||
{premiumUser ? (
|
||||
<Accordion className="mt-4 mb-4">
|
||||
|
|
|
|||
|
|
@ -43,6 +43,7 @@ import { isAdminRole } from "@/utils/roles";
|
|||
import ObjectPermissionsView from "../object_permissions_view";
|
||||
import VectorStoreSelector from "../vector_store_management/VectorStoreSelector";
|
||||
import MCPServerSelector from "../mcp_server_management/MCPServerSelector";
|
||||
import MCPToolPermissions from "../mcp_server_management/MCPToolPermissions";
|
||||
import { formatNumberWithCommas } from "@/utils/dataUtils";
|
||||
import EditLoggingSettings from "./EditLoggingSettings";
|
||||
import LoggingSettingsView from "../logging_settings_view";
|
||||
|
|
@ -96,6 +97,7 @@ export interface TeamData {
|
|||
object_permission_id: string;
|
||||
mcp_servers: string[];
|
||||
mcp_access_groups?: string[];
|
||||
mcp_tool_permissions?: Record<string, string[]>;
|
||||
vector_stores: string[];
|
||||
};
|
||||
team_member_budget_table: {
|
||||
|
|
@ -348,7 +350,9 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
|
|||
servers: [],
|
||||
accessGroups: [],
|
||||
};
|
||||
if ((servers && servers.length > 0) || (accessGroups && accessGroups.length > 0)) {
|
||||
const mcpToolPermissions = values.mcp_tool_permissions || {};
|
||||
|
||||
if ((servers && servers.length > 0) || (accessGroups && accessGroups.length > 0) || Object.keys(mcpToolPermissions).length > 0) {
|
||||
updateData.object_permission = {};
|
||||
if (servers && servers.length > 0) {
|
||||
updateData.object_permission.mcp_servers = servers;
|
||||
|
|
@ -356,8 +360,12 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
|
|||
if (accessGroups && accessGroups.length > 0) {
|
||||
updateData.object_permission.mcp_access_groups = accessGroups;
|
||||
}
|
||||
if (Object.keys(mcpToolPermissions).length > 0) {
|
||||
updateData.object_permission.mcp_tool_permissions = mcpToolPermissions;
|
||||
}
|
||||
}
|
||||
delete values.mcp_servers_and_groups;
|
||||
delete values.mcp_tool_permissions;
|
||||
|
||||
const response = await teamUpdateCall(accessToken, updateData);
|
||||
|
||||
|
|
@ -552,6 +560,7 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
|
|||
servers: info.object_permission?.mcp_servers || [],
|
||||
accessGroups: info.object_permission?.mcp_access_groups || [],
|
||||
},
|
||||
mcp_tool_permissions: info.object_permission?.mcp_tool_permissions || {},
|
||||
}}
|
||||
layout="vertical"
|
||||
>
|
||||
|
|
@ -671,6 +680,31 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
|
|||
placeholder="Select MCP servers or access groups (optional)"
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
{/* Hidden field to register mcp_tool_permissions with the form */}
|
||||
<Form.Item name="mcp_tool_permissions" initialValue={{}} hidden>
|
||||
<Input type="hidden" />
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
noStyle
|
||||
shouldUpdate={(prevValues, currentValues) =>
|
||||
prevValues.mcp_servers_and_groups !== currentValues.mcp_servers_and_groups ||
|
||||
prevValues.mcp_tool_permissions !== currentValues.mcp_tool_permissions
|
||||
}
|
||||
>
|
||||
{() => (
|
||||
<div className="mb-6">
|
||||
<MCPToolPermissions
|
||||
accessToken={accessToken || ""}
|
||||
selectedServers={form.getFieldValue("mcp_servers_and_groups")?.servers || []}
|
||||
toolPermissions={form.getFieldValue("mcp_tool_permissions") || {}}
|
||||
onChange={(toolPerms) => form.setFieldsValue({ mcp_tool_permissions: toolPerms })}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item label="Organization ID" name="organization_id">
|
||||
<Input type="" />
|
||||
</Form.Item>
|
||||
|
|
|
|||
|
|
@ -66,6 +66,7 @@ import type { KeyResponse, Team } from "./key_team_helpers/key_list";
|
|||
import { formatNumberWithCommas } from "../utils/dataUtils";
|
||||
import { AlertTriangleIcon, XIcon } from "lucide-react";
|
||||
import MCPServerSelector from "./mcp_server_management/MCPServerSelector";
|
||||
import MCPToolPermissions from "./mcp_server_management/MCPToolPermissions";
|
||||
import ModelAliasManager from "./common_components/ModelAliasManager";
|
||||
import NotificationsManager from "./molecules/notifications_manager";
|
||||
|
||||
|
|
@ -378,7 +379,8 @@ const Teams: React.FC<TeamProps> = ({
|
|||
(formValues.allowed_vector_store_ids && formValues.allowed_vector_store_ids.length > 0) ||
|
||||
(formValues.allowed_mcp_servers_and_groups &&
|
||||
(formValues.allowed_mcp_servers_and_groups.servers?.length > 0 ||
|
||||
formValues.allowed_mcp_servers_and_groups.accessGroups?.length > 0))
|
||||
formValues.allowed_mcp_servers_and_groups.accessGroups?.length > 0 ||
|
||||
formValues.allowed_mcp_servers_and_groups.toolPermissions))
|
||||
) {
|
||||
formValues.object_permission = {};
|
||||
if (formValues.allowed_vector_store_ids && formValues.allowed_vector_store_ids.length > 0) {
|
||||
|
|
@ -395,6 +397,15 @@ const Teams: React.FC<TeamProps> = ({
|
|||
}
|
||||
delete formValues.allowed_mcp_servers_and_groups;
|
||||
}
|
||||
|
||||
// Add tool permissions separately
|
||||
if (formValues.mcp_tool_permissions && Object.keys(formValues.mcp_tool_permissions).length > 0) {
|
||||
if (!formValues.object_permission) {
|
||||
formValues.object_permission = {};
|
||||
}
|
||||
formValues.object_permission.mcp_tool_permissions = formValues.mcp_tool_permissions;
|
||||
delete formValues.mcp_tool_permissions;
|
||||
}
|
||||
}
|
||||
|
||||
// Transform allowed_mcp_access_groups into object_permission
|
||||
|
|
@ -1240,18 +1251,26 @@ const Teams: React.FC<TeamProps> = ({
|
|||
placeholder="Select vector stores (optional)"
|
||||
/>
|
||||
</Form.Item>
|
||||
</AccordionBody>
|
||||
</Accordion>
|
||||
|
||||
<Accordion className="mt-8 mb-8">
|
||||
<AccordionHeader>
|
||||
<b>MCP Settings</b>
|
||||
</AccordionHeader>
|
||||
<AccordionBody>
|
||||
<Form.Item
|
||||
label={
|
||||
<span>
|
||||
Allowed MCP Servers{" "}
|
||||
<Tooltip title="Select which MCP servers or access groups this team can access by default. ">
|
||||
<Tooltip title="Select which MCP servers or access groups this team can access">
|
||||
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="allowed_mcp_servers_and_groups"
|
||||
className="mt-8"
|
||||
help="Select MCP servers or access groups this team can access. "
|
||||
className="mt-4"
|
||||
help="Select MCP servers or access groups this team can access"
|
||||
>
|
||||
<MCPServerSelector
|
||||
onChange={(val: any) => form.setFieldValue("allowed_mcp_servers_and_groups", val)}
|
||||
|
|
@ -1260,6 +1279,30 @@ const Teams: React.FC<TeamProps> = ({
|
|||
placeholder="Select MCP servers or access groups (optional)"
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
{/* Hidden field to register mcp_tool_permissions with the form */}
|
||||
<Form.Item name="mcp_tool_permissions" initialValue={{}} hidden>
|
||||
<Input type="hidden" />
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
noStyle
|
||||
shouldUpdate={(prevValues, currentValues) =>
|
||||
prevValues.allowed_mcp_servers_and_groups !== currentValues.allowed_mcp_servers_and_groups ||
|
||||
prevValues.mcp_tool_permissions !== currentValues.mcp_tool_permissions
|
||||
}
|
||||
>
|
||||
{() => (
|
||||
<div className="mt-6">
|
||||
<MCPToolPermissions
|
||||
accessToken={accessToken || ""}
|
||||
selectedServers={form.getFieldValue("allowed_mcp_servers_and_groups")?.servers || []}
|
||||
toolPermissions={form.getFieldValue("mcp_tool_permissions") || {}}
|
||||
onChange={(toolPerms) => form.setFieldsValue({ mcp_tool_permissions: toolPerms })}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</Form.Item>
|
||||
</AccordionBody>
|
||||
</Accordion>
|
||||
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ import { modelAvailableCall, getPromptsList } from "../networking";
|
|||
import NumericalInput from "../shared/numerical_input";
|
||||
import VectorStoreSelector from "../vector_store_management/VectorStoreSelector";
|
||||
import MCPServerSelector from "../mcp_server_management/MCPServerSelector";
|
||||
import MCPToolPermissions from "../mcp_server_management/MCPToolPermissions";
|
||||
import EditLoggingSettings from "../team/EditLoggingSettings";
|
||||
import { extractLoggingSettings, formatMetadataForDisplay } from "../key_info_utils";
|
||||
import { fetchMCPAccessGroups } from "../networking";
|
||||
|
|
@ -145,6 +146,7 @@ export function KeyEditView({
|
|||
servers: keyData.object_permission?.mcp_servers || [],
|
||||
accessGroups: keyData.object_permission?.mcp_access_groups || [],
|
||||
},
|
||||
mcp_tool_permissions: keyData.object_permission?.mcp_tool_permissions || {},
|
||||
logging_settings: extractLoggingSettings(keyData.metadata),
|
||||
disabled_callbacks: Array.isArray(keyData.metadata?.litellm_disabled_callbacks)
|
||||
? mapInternalToDisplayNames(keyData.metadata.litellm_disabled_callbacks)
|
||||
|
|
@ -166,6 +168,7 @@ export function KeyEditView({
|
|||
servers: keyData.object_permission?.mcp_servers || [],
|
||||
accessGroups: keyData.object_permission?.mcp_access_groups || [],
|
||||
},
|
||||
mcp_tool_permissions: keyData.object_permission?.mcp_tool_permissions || {},
|
||||
logging_settings: extractLoggingSettings(keyData.metadata),
|
||||
disabled_callbacks: Array.isArray(keyData.metadata?.litellm_disabled_callbacks)
|
||||
? mapInternalToDisplayNames(keyData.metadata.litellm_disabled_callbacks)
|
||||
|
|
@ -291,6 +294,30 @@ export function KeyEditView({
|
|||
/>
|
||||
</Form.Item>
|
||||
|
||||
{/* Hidden field to register mcp_tool_permissions with the form */}
|
||||
<Form.Item name="mcp_tool_permissions" initialValue={{}} hidden>
|
||||
<Input type="hidden" />
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
noStyle
|
||||
shouldUpdate={(prevValues, currentValues) =>
|
||||
prevValues.mcp_servers_and_groups !== currentValues.mcp_servers_and_groups ||
|
||||
prevValues.mcp_tool_permissions !== currentValues.mcp_tool_permissions
|
||||
}
|
||||
>
|
||||
{() => (
|
||||
<div className="mb-6">
|
||||
<MCPToolPermissions
|
||||
accessToken={accessToken || ""}
|
||||
selectedServers={form.getFieldValue("mcp_servers_and_groups")?.servers || []}
|
||||
toolPermissions={form.getFieldValue("mcp_tool_permissions") || {}}
|
||||
onChange={(toolPerms) => form.setFieldsValue({ mcp_tool_permissions: toolPerms })}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item label="Team ID" name="team_id">
|
||||
<Select placeholder="Select team" style={{ width: "100%" }}>
|
||||
{/* Only show All Team Models if team has models */}
|
||||
|
|
|
|||
|
|
@ -136,6 +136,18 @@ export default function KeyInfoView({
|
|||
delete formValues.mcp_servers_and_groups;
|
||||
}
|
||||
|
||||
// Handle MCP tool permissions
|
||||
if (formValues.mcp_tool_permissions !== undefined) {
|
||||
const mcpToolPermissions = formValues.mcp_tool_permissions || {};
|
||||
if (Object.keys(mcpToolPermissions).length > 0) {
|
||||
formValues.object_permission = {
|
||||
...formValues.object_permission,
|
||||
mcp_tool_permissions: mcpToolPermissions,
|
||||
};
|
||||
}
|
||||
delete formValues.mcp_tool_permissions;
|
||||
}
|
||||
|
||||
// Convert metadata back to an object if it exists and is a string
|
||||
if (formValues.metadata && typeof formValues.metadata === "string") {
|
||||
try {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue