Merge branch 'main' into litellm_dev_10_04_2025_p3

This commit is contained in:
Krish Dholakia 2025-10-06 20:11:52 -07:00 • committed by GitHub
commit 543e00a886
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
390 changed files with 22677 additions and 20841 deletions

View file

View file

@ -23,6 +23,43 @@ LiteLLM Proxy provides an MCP Gateway that allows you to use a fixed endpoint fo
## Adding your MCP ## 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> <Tabs>
<TabItem value="ui" label="LiteLLM UI"> <TabItem value="ui" label="LiteLLM UI">

View file

@ -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-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-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-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 | `response = completion(model="gpt-4.1", messages=messages)` |
| gpt-4.1-mini | `response = completion(model="gpt-4.1-mini", 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)` | | 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 ```bash
LiteLLM:DEBUG: utils.py:255 - Request to litellm: LiteLLM:DEBUG: utils.py:255 - Request to litellm:
LiteLLM:DEBUG: utils.py:255 - litellm.acompletion(... organization='my-special-org',) 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..."}]
)
``` ```

View file

@ -191,7 +191,7 @@ print(json.loads(completion.choices[0].message.content))
model_list: model_list:
- model_name: gemini-2.5-pro - model_name: gemini-2.5-pro
litellm_params: litellm_params:
model: vertex_ai/gemini-1.5-pro model: vertex_ai/gemini-2.5-pro
vertex_project: "project-id" vertex_project: "project-id"
vertex_location: "us-central1" 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 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_list:
- model_name: gemini-2.5-pro - model_name: gemini-2.5-pro
litellm_params: litellm_params:
model: vertex_ai/gemini-1.5-pro model: vertex_ai/gemini-2.5-pro
vertex_project: "project-id" vertex_project: "project-id"
vertex_location: "us-central1" 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 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** ### **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) [**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. 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> </TabItem>
</Tabs> </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. 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`. 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`.

View file

@ -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] | | 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 | | 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. | | 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. | | 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_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. | | max_response_size_mb | int | The maximum size for responses in MB. LLM Responses above this size will not be sent. |

View 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

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 253 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 603 KiB

View file

@ -50,30 +50,6 @@ pip install litellm==1.75.5.post2
- **Oracle Cloud Infrastructure** - New LLM provider for calling models on Oracle Cloud Infrastructure. - **Oracle Cloud Infrastructure** - New LLM provider for calling models on Oracle Cloud Infrastructure.
- **Digital Ocean's Gradient AI** - New LLM provider for calling models on Digital Ocean's Gradient AI platform. - **Digital Ocean's Gradient AI** - New LLM provider for calling models on Digital Ocean's Gradient AI platform.
### 54% RPS Improvement
Throughput increased by 54% (1,040 → 1,602 RPS, aggregated) per instance while maintaining a 40 ms median overhead. The improvement comes from fixing major O(n²) inefficiencies in the router, primarily caused by repeated use of in statements inside loops over large arrays. Tests were run with a database-only setup (no cache hits). As a result, p95 latency improved by 30% (2,700 → 1,900 ms), enhancing overall stability and scalability under heavy load.
---
### Test Setup
All benchmarks were executed using Locust with 1,000 concurrent users and a ramp-up of 500. The environment was configured to stress the routing layer and eliminate caching as a variable.
**System Specs**
- **CPU:** 8 vCPUs
- **Memory:** 32 GB RAM
**Configuration (config.yaml)**
View the complete configuration: [gist.github.com/AlexsanderHamir/config.yaml](https://gist.github.com/AlexsanderHamir/53f7d554a5d2afcf2c4edb5b6be68ff4)
**Load Script (no_cache_hits.py)**
View the complete load testing script: [gist.github.com/AlexsanderHamir/no_cache_hits.py](https://gist.github.com/AlexsanderHamir/42c33d7a4dc7a57f56a78b560dee3a42)
--- ---
### Risk of Upgrade ### Risk of Upgrade

View file

@ -1,5 +1,5 @@
--- ---
title: "[Preview] v1.77.5-stable - MCP OAuth 2.0 Support" title: "v1.77.5-stable - MCP OAuth 2.0 Support"
slug: "v1-77-5" slug: "v1-77-5"
date: 2025-09-29T10:00:00 date: 2025-09-29T10:00:00
authors: authors:
@ -11,6 +11,10 @@ authors:
title: CTO, LiteLLM title: CTO, LiteLLM
url: https://www.linkedin.com/in/reffajnaahsi/ url: https://www.linkedin.com/in/reffajnaahsi/
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
- name: Alexsander Hamir
title: Backend Performance Engineer
url: https://www.linkedin.com/in/alexsander-baptista/
image_url: https://media.licdn.com/dms/image/v2/D5603AQGXnziu4kqNCQ/profile-displayphoto-crop_800_800/B56ZkxEcuOKEAI-/0/1757464874550?e=1762387200&v=beta&t=9SNXLsWhx8OnYPAMQ9fqAr02oevDYEAL2vMYg2f9ieg
hide_table_of_contents: false hide_table_of_contents: false
--- ---
@ -28,7 +32,7 @@ import TabItem from '@theme/TabItem';
docker run \ docker run \
-e STORE_MODEL_IN_DB=True \ -e STORE_MODEL_IN_DB=True \
-p 4000:4000 \ -p 4000:4000 \
ghcr.io/berriai/litellm:v1.77.5.rc.1 ghcr.io/berriai/litellm:v1.77.5-stable
``` ```
</TabItem> </TabItem>
@ -49,7 +53,54 @@ pip install litellm==1.77.5
- **MCP OAuth 2.0 Support** - Enhanced authentication for Model Context Protocol integrations - **MCP OAuth 2.0 Support** - Enhanced authentication for Model Context Protocol integrations
- **Scheduled Key Rotations** - Automated key rotation capabilities for enhanced security - **Scheduled Key Rotations** - Automated key rotation capabilities for enhanced security
- **New Gemini 2.5 Flash & Flash-lite Models** - Latest September 2025 preview models with improved pricing and features - **New Gemini 2.5 Flash & Flash-lite Models** - Latest September 2025 preview models with improved pricing and features
- **Performance Improvements** - Critical InMemoryCache unbounded growth resolution - **Performance Improvements** - 54% RPS improvement
---
### Scheduled Key Rotations
<Image img={require('../../img/release_notes/schedule_key_rotations.png')} style={{ width: '800px', height: 'auto' }} />
<br/>
This release brings support for scheduling virtual key rotations on LiteLLM AI Gateway.
This is great for Proxy Admins looking to enforce Enterprise Grade security for use cases going through LiteLLM AI Gateway.
From this release you can enforce Virtual Keys to rotate on a schedule of your choice e.g every 15 days/30 days/60 days etc.
---
### Performance Improvements - 54% RPS Improvement
<Image img={require('../../img/release_notes/perf_77_5.png')} style={{ width: '800px', height: 'auto' }} />
<br/>
This release brings a 54% RPS improvement (1,040 → 1,602 RPS, aggregated) per instance.
The improvement comes from fixing O(n²) inefficiencies in the LiteLLM Router, primarily caused by repeated use of `in` statements inside loops over large arrays.
Tests were run with a database-only setup (no cache hits).
#### Test Setup
All benchmarks were executed using Locust with 1,000 concurrent users and a ramp-up of 500. The environment was configured to stress the routing layer and eliminate caching as a variable.
**System Specs**
- **CPU:** 8 vCPUs
- **Memory:** 32 GB RAM
**Configuration (config.yaml)**
View the complete configuration: [gist.github.com/AlexsanderHamir/config.yaml](https://gist.github.com/AlexsanderHamir/53f7d554a5d2afcf2c4edb5b6be68ff4)
**Load Script (no_cache_hits.py)**
View the complete load testing script: [gist.github.com/AlexsanderHamir/no_cache_hits.py](https://gist.github.com/AlexsanderHamir/42c33d7a4dc7a57f56a78b560dee3a42)
---
## New Models / Updated Models ## New Models / Updated Models

View file

@ -59,12 +59,50 @@ pip install litellm==1.77.7.rc.1
## Key Highlights ## Key Highlights
- **Dynamic Rate Limiter v3** - Automatically maximizes throughput when capacity is available (< 80% saturation) by allowing lower-priority requests to use unused capacity, then switches to fair priority-based allocation under high load (≥ 80%) to prevent blocking - **Dynamic Rate Limiter v3** - Automatically maximizes throughput when capacity is available (< 80% saturation) by allowing lower-priority requests to use unused capacity, then switches to fair priority-based allocation under high load (≥ 80%) to prevent blocking
- **Major Performance Improvements** - Router optimization reducing P99 latency by 62.5%, cache improvements from O(n*log(n)) to O(log(n)) - **Major Performance Improvements** - 2.9x lower median latency at 1,000 concurrent users.
- **Claude Sonnet 4.5** - Support for Anthropic's new Claude Sonnet 4.5 model family with 200K+ context and tiered pricing - **Claude Sonnet 4.5** - Support for Anthropic's new Claude Sonnet 4.5 model family with 200K+ context and tiered pricing
- **MCP Gateway Enhancements** - Fine-grained tool control, server permissions, and forwardable headers - **MCP Gateway Enhancements** - Fine-grained tool control, server permissions, and forwardable headers
- **AMD Lemonade & Nvidia NIM** - New provider support for AMD Lemonade and Nvidia NIM Rerank - **AMD Lemonade & Nvidia NIM** - New provider support for AMD Lemonade and Nvidia NIM Rerank
- **GitLab Prompt Management** - GitLab-based prompt management integration - **GitLab Prompt Management** - GitLab-based prompt management integration
### Performance - 2.9x Lower Median Latency
<Image img={require('../../img/release_notes/perf_77_7.png')} style={{ width: '800px', height: 'auto' }} />
<br/>
This update removes LiteLLM router inefficiencies, reducing complexity from O(M×N) to O(1). Previously, it built a new array and ran repeated checks like data["model"] in llm_router.get_model_ids(). Now, a direct ID-to-deployment map eliminates redundant allocations and scans.
As a result, performance improved across all latency percentiles:
- **Median latency:** 320 ms → **110 ms** (−65.6%)
- **p95 latency:** 850 ms → **440 ms** (−48.2%)
- **p99 latency:** 1,400 ms → **810 ms** (−42.1%)
- **Average latency:** 864 ms → **310 ms** (−64%)
#### Test Setup
**Locust**
- **Concurrent users:** 1,000
- **Ramp-up:** 500
**System Specs**
- **CPU:** 4 vCPUs
- **Memory:** 8 GB RAM
- **LiteLLM Workers:** 4
- **Instances**: 4
**Configuration (config.yaml)**
View the complete configuration: [gist.github.com/AlexsanderHamir/config.yaml](https://gist.github.com/AlexsanderHamir/53f7d554a5d2afcf2c4edb5b6be68ff4)
**Load Script (no_cache_hits.py)**
View the complete load testing script: [gist.github.com/AlexsanderHamir/no_cache_hits.py](https://gist.github.com/AlexsanderHamir/42c33d7a4dc7a57f56a78b560dee3a42)
## New Models / Updated Models ## New Models / Updated Models
#### New Model Support #### New Model Support

View file

@ -674,6 +674,7 @@ const sidebars = {
items: [ items: [
"data_security", "data_security",
"data_retention", "data_retention",
"proxy/security_encryption_faq",
"migration_policy", "migration_policy",
{ {
type: "category", type: "category",

View file

@ -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 | `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-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-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 ## 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` For Azure calls add the `azure/` prefix to `model`. If your azure deployment name is `gpt-v-2` set `model` = `azure/gpt-v-2`

View file

@ -0,0 +1,3 @@
-- AlterTable
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "mcp_tool_permissions" JSONB;

View file

@ -156,6 +156,7 @@ model LiteLLM_ObjectPermissionTable {
object_permission_id String @id @default(uuid()) object_permission_id String @id @default(uuid())
mcp_servers String[] @default([]) mcp_servers String[] @default([])
mcp_access_groups 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([]) vector_stores String[] @default([])
teams LiteLLM_TeamTable[] teams LiteLLM_TeamTable[]
verification_tokens LiteLLM_VerificationToken[] verification_tokens LiteLLM_VerificationToken[]

View file

@ -5,8 +5,9 @@ LiteLLM Proxy uses this MCP Client to connnect to other MCP servers.
import asyncio import asyncio
import base64 import base64
from datetime import timedelta 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 import ClientSession, StdioServerParameters
from mcp.client.sse import sse_client from mcp.client.sse import sse_client
from mcp.client.stdio import stdio_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 mcp.types import Tool as MCPTool
from litellm._logging import verbose_logger 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 ( from litellm.types.mcp import (
MCPAuth, MCPAuth,
MCPAuthType, MCPAuthType,
@ -48,6 +51,7 @@ class MCPClient:
timeout: float = 60.0, timeout: float = 60.0,
stdio_config: Optional[MCPStdioConfig] = None, stdio_config: Optional[MCPStdioConfig] = None,
extra_headers: Optional[Dict[str, str]] = None, extra_headers: Optional[Dict[str, str]] = None,
ssl_verify: Optional[VerifyTypes] = None,
): ):
self.server_url: str = server_url self.server_url: str = server_url
self.transport_type: MCPTransport = transport_type self.transport_type: MCPTransport = transport_type
@ -62,6 +66,7 @@ class MCPClient:
self._task: Optional[asyncio.Task] = None self._task: Optional[asyncio.Task] = None
self.stdio_config: Optional[MCPStdioConfig] = stdio_config self.stdio_config: Optional[MCPStdioConfig] = stdio_config
self.extra_headers: Optional[Dict[str, str]] = extra_headers self.extra_headers: Optional[Dict[str, str]] = extra_headers
self.ssl_verify: Optional[VerifyTypes] = ssl_verify
# handle the basic auth value if provided # handle the basic auth value if provided
if auth_value: if auth_value:
self.update_auth_value(auth_value) self.update_auth_value(auth_value)
@ -104,10 +109,12 @@ class MCPClient:
await self._session.initialize() await self._session.initialize()
elif self.transport_type == MCPTransport.sse: elif self.transport_type == MCPTransport.sse:
headers = self._get_auth_headers() headers = self._get_auth_headers()
httpx_client_factory = self._create_httpx_client_factory()
self._transport_ctx = sse_client( self._transport_ctx = sse_client(
url=self.server_url, url=self.server_url,
timeout=self.timeout, timeout=self.timeout,
headers=headers, headers=headers,
httpx_client_factory=httpx_client_factory,
) )
self._transport = await self._transport_ctx.__aenter__() self._transport = await self._transport_ctx.__aenter__()
self._session_ctx = ClientSession( self._session_ctx = ClientSession(
@ -117,13 +124,15 @@ class MCPClient:
await self._session.initialize() await self._session.initialize()
else: # http else: # http
headers = self._get_auth_headers() headers = self._get_auth_headers()
httpx_client_factory = self._create_httpx_client_factory()
verbose_logger.debug( verbose_logger.debug(
"litellm headers for streamablehttp_client: ", headers "litellm headers for streamablehttp_client: %s", headers
) )
self._transport_ctx = streamablehttp_client( self._transport_ctx = streamablehttp_client(
url=self.server_url, url=self.server_url,
timeout=timedelta(seconds=self.timeout), timeout=timedelta(seconds=self.timeout),
headers=headers, headers=headers,
httpx_client_factory=httpx_client_factory,
) )
self._transport = await self._transport_ctx.__aenter__() self._transport = await self._transport_ctx.__aenter__()
self._session_ctx = ClientSession( self._session_ctx = ClientSession(
@ -215,6 +224,41 @@ class MCPClient:
return headers 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]: async def list_tools(self) -> List[MCPTool]:
"""List available tools from the server.""" """List available tools from the server."""
if not self._session: if not self._session:

View file

@ -665,10 +665,6 @@ class BaseAzureLLM(BaseOpenAILLM):
) -> dict: ) -> dict:
litellm_params = litellm_params or GenericLiteLLMParams() litellm_params = litellm_params or GenericLiteLLMParams()
# If api-key is already in headers, preserve it
if "api-key" in headers:
return headers
api_key = ( api_key = (
litellm_params.api_key litellm_params.api_key
or litellm.api_key or litellm.api_key
@ -693,7 +689,7 @@ class BaseAzureLLM(BaseOpenAILLM):
def _get_base_azure_url( def _get_base_azure_url(
api_base: Optional[str], api_base: Optional[str],
litellm_params: Optional[Union[GenericLiteLLMParams, Dict[str, Any]]], litellm_params: Optional[Union[GenericLiteLLMParams, Dict[str, Any]]],
route: Literal["/openai/responses", "/openai/vector_stores"], route: Union[Literal["/openai/responses", "/openai/vector_stores"], str],
default_api_version: Optional[Union[str, Literal["latest", "preview"]]] = None, default_api_version: Optional[Union[str, Literal["latest", "preview"]]] = None,
) -> str: ) -> str:
""" """

View file

@ -0,0 +1,85 @@
from typing import TYPE_CHECKING, List, Optional, Tuple
import httpx
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.router import GenericLiteLLMParams
if TYPE_CHECKING:
from httpx import URL
class AzurePassthroughConfig(BasePassthroughConfig):
def is_streaming_request(self, endpoint: str, request_data: dict) -> bool:
return "stream" in request_data
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
endpoint: str,
request_query_params: Optional[dict],
litellm_params: dict,
) -> Tuple["URL", str]:
base_target_url = self.get_api_base(api_base)
if base_target_url is None:
raise Exception("Azure api base not found")
litellm_metadata = litellm_params.get("litellm_metadata") or {}
model_group = litellm_metadata.get("model_group")
if model_group and model_group in endpoint:
endpoint = endpoint.replace(model_group, model)
complete_url = BaseAzureLLM._get_base_azure_url(
api_base=base_target_url,
litellm_params=litellm_params,
route=endpoint,
default_api_version=litellm_params.get("api_version"),
)
return (
httpx.URL(complete_url),
base_target_url,
)
def validate_environment(
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
return BaseAzureLLM._base_validate_azure_environment(
headers=headers,
litellm_params=GenericLiteLLMParams(
**{**litellm_params, "api_key": api_key}
),
)
@staticmethod
def get_api_base(
api_base: Optional[str] = None,
) -> Optional[str]:
return api_base or get_secret_str("AZURE_API_BASE")
@staticmethod
def get_api_key(
api_key: Optional[str] = None,
) -> Optional[str]:
return api_key or get_secret_str("AZURE_API_KEY")
@staticmethod
def get_base_model(model: str) -> Optional[str]:
return model
def get_models(
self, api_key: Optional[str] = None, api_base: Optional[str] = None
) -> List[str]:
return super().get_models(api_key, api_base)

View file

@ -1,6 +1,7 @@
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
import httpx import httpx
from openai.types.responses import ResponseReasoningItem
from litellm._logging import verbose_logger from litellm._logging import verbose_logger
from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.azure.common_utils import BaseAzureLLM
@ -38,6 +39,50 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
model = model.replace("o_series/", "") model = model.replace("o_series/", "")
return model 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( def transform_responses_api_request(
self, self,
model: str, model: str,
@ -48,12 +93,13 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
) -> Dict: ) -> Dict:
"""No transform applied since inputs are in OpenAI spec already""" """No transform applied since inputs are in OpenAI spec already"""
stripped_model_name = self.get_stripped_model_name(model) stripped_model_name = self.get_stripped_model_name(model)
return dict(
ResponsesAPIRequestParams( return super().transform_responses_api_request(
model=stripped_model_name, model=stripped_model_name,
input=input, input=input,
**response_api_optional_request_params, response_api_optional_request_params=response_api_optional_request_params,
) litellm_params=litellm_params,
headers=headers,
) )
def get_complete_url( def get_complete_url(
@ -217,15 +263,15 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
at the correct location (before any query parameters). at the correct location (before any query parameters).
""" """
from urllib.parse import urlparse, urlunparse from urllib.parse import urlparse, urlunparse
# Parse the URL to separate its components # Parse the URL to separate its components
parsed_url = urlparse(api_base) parsed_url = urlparse(api_base)
# Insert the response_id and /cancel at the end of the path component # Insert the response_id and /cancel at the end of the path component
# Remove trailing slash if present to avoid double slashes # Remove trailing slash if present to avoid double slashes
path = parsed_url.path.rstrip("/") path = parsed_url.path.rstrip("/")
new_path = f"{path}/{response_id}/cancel" new_path = f"{path}/{response_id}/cancel"
# Reconstruct the URL with all original components but with the modified path # Reconstruct the URL with all original components but with the modified path
cancel_url = urlunparse( cancel_url = urlunparse(
( (

View file

@ -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.base_llm.chat.transformation import LiteLLMLoggingObj
from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error
from litellm.llms.openai.openai import OpenAIConfig 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.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse, ProviderField from litellm.types.utils import ModelResponse, ProviderField
@ -35,9 +36,24 @@ class AzureAIStudioConfig(OpenAIConfig):
for param in supported_params: for param in supported_params:
if param != "tool_choice": if param != "tool_choice":
filtered_supported_params.append(param) 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 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( def validate_environment(
self, self,
headers: dict, headers: dict,
@ -53,9 +69,7 @@ class AzureAIStudioConfig(OpenAIConfig):
else: else:
headers["Authorization"] = f"Bearer {api_key}" headers["Authorization"] = f"Bearer {api_key}"
headers["Content-Type"] = ( headers["Content-Type"] = "application/json" # tell Azure AI Studio to expect JSON
"application/json" # tell Azure AI Studio to expect JSON
)
return headers return headers
@ -65,10 +79,7 @@ class AzureAIStudioConfig(OpenAIConfig):
""" """
parsed_url = urlparse(api_base) parsed_url = urlparse(api_base)
host = parsed_url.hostname host = parsed_url.hostname
if host and ( if host and (host.endswith(".services.ai.azure.com") or host.endswith(".openai.azure.com")):
host.endswith(".services.ai.azure.com")
or host.endswith(".openai.azure.com")
):
return True return True
return False return False
@ -115,13 +126,9 @@ class AzureAIStudioConfig(OpenAIConfig):
# Add the path to the base URL # Add the path to the base URL
if "services.ai.azure.com" in api_base: if "services.ai.azure.com" in api_base:
new_url = _add_path_to_api_base( new_url = _add_path_to_api_base(api_base=api_base, ending_path="/models/chat/completions")
api_base=api_base, ending_path="/models/chat/completions"
)
else: else:
new_url = _add_path_to_api_base( new_url = _add_path_to_api_base(api_base=api_base, ending_path="/chat/completions")
api_base=api_base, ending_path="/chat/completions"
)
# Use the new query_params dictionary # Use the new query_params dictionary
final_url = httpx.URL(new_url).copy_with(params=query_params) 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") 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): if self._is_azure_openai_model(model=model, api_base=api_base):
verbose_logger.debug( verbose_logger.debug("Model={} is Azure OpenAI model. Setting custom_llm_provider='azure'.".format(model))
"Model={} is Azure OpenAI model. Setting custom_llm_provider='azure'.".format(
model
)
)
custom_llm_provider = "azure" custom_llm_provider = "azure"
return api_base, dynamic_api_key, custom_llm_provider return api_base, dynamic_api_key, custom_llm_provider
@ -211,9 +214,7 @@ class AzureAIStudioConfig(OpenAIConfig):
if extra_body and isinstance(extra_body, dict): if extra_body and isinstance(extra_body, dict):
optional_params.update(extra_body) optional_params.update(extra_body)
optional_params.pop("max_retries", None) optional_params.pop("max_retries", None)
return super().transform_request( return super().transform_request(model, messages, optional_params, litellm_params, headers)
model, messages, optional_params, litellm_params, headers
)
def transform_response( def transform_response(
self, self,
@ -252,47 +253,30 @@ class AzureAIStudioConfig(OpenAIConfig):
if should_drop_params and "Extra inputs are not permitted" in error_text: if should_drop_params and "Extra inputs are not permitted" in error_text:
return True return True
elif ( elif "unknown field: parameter index is not a valid field" in error_text: # remove index from tool calls
"unknown field: parameter index is not a valid field" in error_text
): # remove index from tool calls
return True return True
elif ( elif (
AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value in error_text
in error_text
): # remove extra-parameters from tool calls ): # remove extra-parameters from tool calls
return True return True
return super().should_retry_llm_api_inside_llm_translation_on_http_error( return super().should_retry_llm_api_inside_llm_translation_on_http_error(e=e, litellm_params=litellm_params)
e=e, litellm_params=litellm_params
)
@property @property
def max_retry_on_unprocessable_entity_error(self) -> int: def max_retry_on_unprocessable_entity_error(self) -> int:
return 2 return 2
def transform_request_on_unprocessable_entity_error( def transform_request_on_unprocessable_entity_error(self, e: httpx.HTTPStatusError, request_data: dict) -> dict:
self, e: httpx.HTTPStatusError, request_data: dict
) -> dict:
_messages = cast(Optional[List[AllMessageValues]], request_data.get("messages")) _messages = cast(Optional[List[AllMessageValues]], request_data.get("messages"))
if ( if "unknown field: parameter index is not a valid field" in e.response.text and _messages is not None:
"unknown field: parameter index is not a valid field" in e.response.text
and _messages is not None
):
litellm.remove_index_from_tool_calls( litellm.remove_index_from_tool_calls(
messages=_messages, messages=_messages,
) )
elif ( elif AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value in e.response.text:
AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value request_data = self._drop_extra_params_from_request_data(request_data, e.response.text)
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) data = drop_params_from_unprocessable_entity_error(e=e, data=request_data)
return data return data
def _drop_extra_params_from_request_data( def _drop_extra_params_from_request_data(self, request_data: dict, error_text: str) -> dict:
self, request_data: dict, error_text: str
) -> dict:
params_to_drop = self._extract_params_to_drop_from_error_text(error_text) params_to_drop = self._extract_params_to_drop_from_error_text(error_text)
if params_to_drop: if params_to_drop:
for param in params_to_drop: for param in params_to_drop:
@ -300,9 +284,7 @@ class AzureAIStudioConfig(OpenAIConfig):
request_data.pop(param, None) request_data.pop(param, None)
return request_data return request_data
def _extract_params_to_drop_from_error_text( def _extract_params_to_drop_from_error_text(self, error_text: str) -> Optional[List[str]]:
self, error_text: str
) -> Optional[List[str]]:
""" """
Error text looks like this" 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'. "Extra parameters ['stream_options', 'extra-parameters'] are not allowed when extra-parameters is not set or set to be 'error'.

View file

@ -41,6 +41,7 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
"presence_penalty", "presence_penalty",
"frequency_penalty", "frequency_penalty",
"top_logprobs", "top_logprobs",
"stop",
] ]
return [ return [

View file

@ -1,12 +1,4 @@
from typing import ( from typing import TYPE_CHECKING, Any, Dict, Optional, Union, cast, get_type_hints
TYPE_CHECKING,
Any,
Dict,
Optional,
Union,
cast,
get_type_hints,
)
import httpx import httpx
from openai.types.responses import ResponseReasoningItem from openai.types.responses import ResponseReasoningItem
@ -127,7 +119,6 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
2. Create a ResponseReasoningItem object with the item data 2. Create a ResponseReasoningItem object with the item data
3. Convert it back to dict with exclude_none=True to filter None values 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": if item.get("type") == "reasoning":
try: try:
# Ensure required fields are present for ResponseReasoningItem # Ensure required fields are present for ResponseReasoningItem

View file

@ -1,14 +1,15 @@
""" """
Support for Snowflake REST API Support for Snowflake REST API
""" """
from typing import TYPE_CHECKING, Any, List, Optional, Tuple import json
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
import httpx import httpx
from litellm.secret_managers.main import get_secret_str from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse from litellm.types.utils import ChatCompletionMessageToolCall, Function, ModelResponse
from ...openai_like.chat.transformation import OpenAIGPTConfig from ...openai_like.chat.transformation import OpenAIGPTConfig
@ -22,15 +23,25 @@ else:
class SnowflakeConfig(OpenAIGPTConfig): class SnowflakeConfig(OpenAIGPTConfig):
""" """
source: https://docs.snowflake.com/en/sql-reference/functions/complete-snowflake-cortex Reference: https://docs.snowflake.com/en/user-guide/snowflake-cortex/cortex-llm-rest-api
Snowflake Cortex LLM REST API supports function calling with specific models (e.g., Claude 3.5 Sonnet).
This config handles transformation between OpenAI format and Snowflake's tool_spec format.
""" """
@classmethod @classmethod
def get_config(cls): def get_config(cls):
return super().get_config() return super().get_config()
def get_supported_openai_params(self, model: str) -> List: def get_supported_openai_params(self, model: str) -> List[str]:
return ["temperature", "max_tokens", "top_p", "response_format"] return [
"temperature",
"max_tokens",
"top_p",
"response_format",
"tools",
"tool_choice",
]
def map_openai_params( def map_openai_params(
self, self,
@ -56,6 +67,57 @@ class SnowflakeConfig(OpenAIGPTConfig):
optional_params[param] = value optional_params[param] = value
return optional_params return optional_params
def _transform_tool_calls_from_snowflake_to_openai(
self, content_list: List[Dict[str, Any]]
) -> Tuple[str, Optional[List[ChatCompletionMessageToolCall]]]:
"""
Transform Snowflake tool calls to OpenAI format.
Args:
content_list: Snowflake's content_list array containing text and tool_use items
Returns:
Tuple of (text_content, tool_calls)
Snowflake format in content_list:
{
"type": "tool_use",
"tool_use": {
"tool_use_id": "tooluse_...",
"name": "get_weather",
"input": {"location": "Paris"}
}
}
OpenAI format (returned tool_calls):
ChatCompletionMessageToolCall(
id="tooluse_...",
type="function",
function=Function(name="get_weather", arguments='{"location": "Paris"}')
)
"""
text_content = ""
tool_calls: List[ChatCompletionMessageToolCall] = []
for idx, content_item in enumerate(content_list):
if content_item.get("type") == "text":
text_content += content_item.get("text", "")
## TOOL CALLING
elif content_item.get("type") == "tool_use":
tool_use_data = content_item.get("tool_use", {})
tool_call = ChatCompletionMessageToolCall(
id=tool_use_data.get("tool_use_id", ""),
type="function",
function=Function(
name=tool_use_data.get("name", ""),
arguments=json.dumps(tool_use_data.get("input", {})),
),
)
tool_calls.append(tool_call)
return text_content, tool_calls if tool_calls else None
def transform_response( def transform_response(
self, self,
model: str, model: str,
@ -71,6 +133,7 @@ class SnowflakeConfig(OpenAIGPTConfig):
json_mode: Optional[bool] = None, json_mode: Optional[bool] = None,
) -> ModelResponse: ) -> ModelResponse:
response_json = raw_response.json() response_json = raw_response.json()
logging_obj.post_call( logging_obj.post_call(
input=messages, input=messages,
api_key="", api_key="",
@ -78,6 +141,26 @@ class SnowflakeConfig(OpenAIGPTConfig):
additional_args={"complete_input_dict": request_data}, additional_args={"complete_input_dict": request_data},
) )
## RESPONSE TRANSFORMATION
# Snowflake returns content_list (not content) with tool_use objects
# We need to transform this to OpenAI's format with content + tool_calls
if "choices" in response_json and len(response_json["choices"]) > 0:
choice = response_json["choices"][0]
if "message" in choice and "content_list" in choice["message"]:
content_list = choice["message"]["content_list"]
(
text_content,
tool_calls,
) = self._transform_tool_calls_from_snowflake_to_openai(content_list)
# Update the choice message with OpenAI format
choice["message"]["content"] = text_content
if tool_calls:
choice["message"]["tool_calls"] = tool_calls
# Remove Snowflake-specific content_list
del choice["message"]["content_list"]
returned_response = ModelResponse(**response_json) returned_response = ModelResponse(**response_json)
returned_response.model = "snowflake/" + (returned_response.model or "") returned_response.model = "snowflake/" + (returned_response.model or "")
@ -150,6 +233,95 @@ class SnowflakeConfig(OpenAIGPTConfig):
return api_base return api_base
def _transform_tools(self, tools: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""
Transform OpenAI tool format to Snowflake tool format.
Args:
tools: List of tools in OpenAI format
Returns:
List of tools in Snowflake format
OpenAI format:
{
"type": "function",
"function": {
"name": "get_weather",
"description": "...",
"parameters": {...}
}
}
Snowflake format:
{
"tool_spec": {
"type": "generic",
"name": "get_weather",
"description": "...",
"input_schema": {...}
}
}
"""
snowflake_tools: List[Dict[str, Any]] = []
for tool in tools:
if tool.get("type") == "function":
function = tool.get("function", {})
snowflake_tool: Dict[str, Any] = {
"tool_spec": {
"type": "generic",
"name": function.get("name"),
"input_schema": function.get(
"parameters",
{"type": "object", "properties": {}},
),
}
}
# Add description if present
if "description" in function:
snowflake_tool["tool_spec"]["description"] = function[
"description"
]
snowflake_tools.append(snowflake_tool)
return snowflake_tools
def _transform_tool_choice(
self, tool_choice: Union[str, Dict[str, Any]]
) -> Union[str, Dict[str, Any]]:
"""
Transform OpenAI tool_choice format to Snowflake format.
Args:
tool_choice: Tool choice in OpenAI format (str or dict)
Returns:
Tool choice in Snowflake format
OpenAI format:
{"type": "function", "function": {"name": "get_weather"}}
Snowflake format:
{"type": "tool", "name": ["get_weather"]}
Note: String values ("auto", "required", "none") pass through unchanged.
"""
if isinstance(tool_choice, str):
# "auto", "required", "none" pass through as-is
return tool_choice
if isinstance(tool_choice, dict):
if tool_choice.get("type") == "function":
function_name = tool_choice.get("function", {}).get("name")
if function_name:
return {
"type": "tool",
"name": [function_name], # Snowflake expects array
}
return tool_choice
def transform_request( def transform_request(
self, self,
model: str, model: str,
@ -160,6 +332,18 @@ class SnowflakeConfig(OpenAIGPTConfig):
) -> dict: ) -> dict:
stream: bool = optional_params.pop("stream", None) or False stream: bool = optional_params.pop("stream", None) or False
extra_body = optional_params.pop("extra_body", {}) extra_body = optional_params.pop("extra_body", {})
## TOOL CALLING
# Transform tools from OpenAI format to Snowflake's tool_spec format
tools = optional_params.pop("tools", None)
if tools:
optional_params["tools"] = self._transform_tools(tools)
# Transform tool_choice from OpenAI format to Snowflake's tool name array format
tool_choice = optional_params.pop("tool_choice", None)
if tool_choice:
optional_params["tool_choice"] = self._transform_tool_choice(tool_choice)
return { return {
"model": model, "model": model,
"messages": messages, "messages": messages,

View file

@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
""" """
import re 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.openai import AllMessageValues
from litellm.types.llms.vertex_ai import CachedContentRequestBody from litellm.types.llms.vertex_ai import CachedContentRequestBody
@ -155,13 +155,18 @@ def separate_cached_messages(
def transform_openai_messages_to_gemini_context_caching( 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: ) -> CachedContentRequestBody:
# Extract TTL from cached messages BEFORE system message transformation # Extract TTL from cached messages BEFORE system message transformation
ttl = extract_ttl_from_cached_messages(messages) ttl = extract_ttl_from_cached_messages(messages)
supports_system_message = get_supports_system_message( 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( 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) 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( data = CachedContentRequestBody(
contents=transformed_messages, contents=transformed_messages,
model="models/{}".format(model), model=model_name,
displayName=cache_key, displayName=cache_key,
) )

View file

@ -41,8 +41,11 @@ class ContextCachingEndpoints(VertexBase):
def _get_token_and_url_context_caching( def _get_token_and_url_context_caching(
self, self,
gemini_api_key: Optional[str], gemini_api_key: Optional[str],
custom_llm_provider: Literal["gemini"], custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
api_base: Optional[str], api_base: Optional[str],
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
) -> Tuple[Optional[str], str]: ) -> Tuple[Optional[str], str]:
""" """
Internal function. Returns the token and url for the call. 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( url = "https://generativelanguage.googleapis.com/v1beta/{}?key={}".format(
endpoint, gemini_api_key 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: 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( return self._check_custom_proxy(
api_base=api_base, api_base=api_base,
@ -80,6 +89,10 @@ class ContextCachingEndpoints(VertexBase):
api_key: str, api_key: str,
api_base: Optional[str], api_base: Optional[str],
logging_obj: Logging, 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]: ) -> Optional[str]:
""" """
Checks if content already cached. Checks if content already cached.
@ -94,8 +107,11 @@ class ContextCachingEndpoints(VertexBase):
_, url = self._get_token_and_url_context_caching( _, url = self._get_token_and_url_context_caching(
gemini_api_key=api_key, gemini_api_key=api_key,
custom_llm_provider="gemini", custom_llm_provider=custom_llm_provider,
api_base=api_base, api_base=api_base,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
) )
try: try:
## LOGGING ## LOGGING
@ -145,6 +161,10 @@ class ContextCachingEndpoints(VertexBase):
api_key: str, api_key: str,
api_base: Optional[str], api_base: Optional[str],
logging_obj: Logging, 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]: ) -> Optional[str]:
""" """
Checks if content already cached. Checks if content already cached.
@ -159,8 +179,11 @@ class ContextCachingEndpoints(VertexBase):
_, url = self._get_token_and_url_context_caching( _, url = self._get_token_and_url_context_caching(
gemini_api_key=api_key, gemini_api_key=api_key,
custom_llm_provider="gemini", custom_llm_provider=custom_llm_provider,
api_base=api_base, api_base=api_base,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
) )
try: try:
## LOGGING ## LOGGING
@ -212,6 +235,10 @@ class ContextCachingEndpoints(VertexBase):
client: Optional[HTTPHandler], client: Optional[HTTPHandler],
timeout: Optional[Union[float, httpx.Timeout]], timeout: Optional[Union[float, httpx.Timeout]],
logging_obj: Logging, 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, extra_headers: Optional[dict] = None,
cached_content: Optional[str] = None, cached_content: Optional[str] = None,
) -> Tuple[List[AllMessageValues], dict, Optional[str]]: ) -> Tuple[List[AllMessageValues], dict, Optional[str]]:
@ -240,8 +267,11 @@ class ContextCachingEndpoints(VertexBase):
## AUTHORIZATION ## ## AUTHORIZATION ##
token, url = self._get_token_and_url_context_caching( token, url = self._get_token_and_url_context_caching(
gemini_api_key=api_key, gemini_api_key=api_key,
custom_llm_provider="gemini", custom_llm_provider=custom_llm_provider,
api_base=api_base, api_base=api_base,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
) )
headers = { headers = {
@ -273,6 +303,10 @@ class ContextCachingEndpoints(VertexBase):
api_key=api_key, api_key=api_key,
api_base=api_base, api_base=api_base,
logging_obj=logging_obj, 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: if google_cache_name:
return non_cached_messages, optional_params, google_cache_name return non_cached_messages, optional_params, google_cache_name
@ -280,7 +314,12 @@ class ContextCachingEndpoints(VertexBase):
## TRANSFORM REQUEST ## TRANSFORM REQUEST
cached_content_request_body = ( cached_content_request_body = (
transform_openai_messages_to_gemini_context_caching( 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], client: Optional[AsyncHTTPHandler],
timeout: Optional[Union[float, httpx.Timeout]], timeout: Optional[Union[float, httpx.Timeout]],
logging_obj: Logging, 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, extra_headers: Optional[dict] = None,
cached_content: Optional[str] = None, cached_content: Optional[str] = None,
) -> Tuple[List[AllMessageValues], dict, Optional[str]]: ) -> Tuple[List[AllMessageValues], dict, Optional[str]]:
@ -356,8 +399,11 @@ class ContextCachingEndpoints(VertexBase):
## AUTHORIZATION ## ## AUTHORIZATION ##
token, url = self._get_token_and_url_context_caching( token, url = self._get_token_and_url_context_caching(
gemini_api_key=api_key, gemini_api_key=api_key,
custom_llm_provider="gemini", custom_llm_provider=custom_llm_provider,
api_base=api_base, api_base=api_base,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
) )
headers = { headers = {
@ -386,6 +432,10 @@ class ContextCachingEndpoints(VertexBase):
api_key=api_key, api_key=api_key,
api_base=api_base, api_base=api_base,
logging_obj=logging_obj, 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: if google_cache_name:
@ -394,7 +444,12 @@ class ContextCachingEndpoints(VertexBase):
## TRANSFORM REQUEST ## TRANSFORM REQUEST
cached_content_request_body = ( cached_content_request_body = (
transform_openai_messages_to_gemini_context_caching( 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,
) )
) )

View file

@ -514,34 +514,35 @@ def sync_transform_request_body(
logging_obj: LiteLLMLoggingObj, logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
litellm_params: dict, litellm_params: dict,
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
) -> RequestBody: ) -> RequestBody:
from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints
context_caching_endpoints = ContextCachingEndpoints() context_caching_endpoints = ContextCachingEndpoints()
if gemini_api_key is not None: (
( messages,
messages, optional_params,
optional_params, cached_content,
cached_content, ) = context_caching_endpoints.check_and_create_cache(
) = context_caching_endpoints.check_and_create_cache( messages=messages,
messages=messages, optional_params=optional_params,
optional_params=optional_params, api_key=gemini_api_key or "dummy",
api_key=gemini_api_key, api_base=api_base,
api_base=api_base, model=model,
model=model, client=client,
client=client, timeout=timeout,
timeout=timeout, extra_headers=extra_headers,
extra_headers=extra_headers, cached_content=optional_params.pop("cached_content", None),
cached_content=optional_params.pop("cached_content", None), logging_obj=logging_obj,
logging_obj=logging_obj, custom_llm_provider=custom_llm_provider,
) vertex_project=vertex_project,
else: # [TODO] implement context caching for gemini as well vertex_location=vertex_location,
cached_content = None vertex_auth_header=vertex_auth_header,
if "cached_content" in optional_params: )
cached_content = optional_params.pop("cached_content")
elif "cachedContent" in optional_params:
cached_content = optional_params.pop("cachedContent")
return _transform_request_body( return _transform_request_body(
messages=messages, messages=messages,
@ -565,34 +566,34 @@ async def async_transform_request_body(
logging_obj: litellm.litellm_core_utils.litellm_logging.Logging, # type: ignore logging_obj: litellm.litellm_core_utils.litellm_logging.Logging, # type: ignore
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
litellm_params: dict, litellm_params: dict,
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
) -> RequestBody: ) -> RequestBody:
from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints
context_caching_endpoints = ContextCachingEndpoints() context_caching_endpoints = ContextCachingEndpoints()
if gemini_api_key is not None: (
( messages,
messages, optional_params,
optional_params, cached_content,
cached_content, ) = await context_caching_endpoints.async_check_and_create_cache(
) = await context_caching_endpoints.async_check_and_create_cache( messages=messages,
messages=messages, optional_params=optional_params,
optional_params=optional_params, api_key=gemini_api_key or "dummy",
api_key=gemini_api_key, api_base=api_base,
api_base=api_base, model=model,
model=model, client=client,
client=client, timeout=timeout,
timeout=timeout, extra_headers=extra_headers,
extra_headers=extra_headers, cached_content=optional_params.pop("cached_content", None),
cached_content=optional_params.pop("cached_content", None), logging_obj=logging_obj,
logging_obj=logging_obj, custom_llm_provider=custom_llm_provider,
) vertex_project=vertex_project,
else: # [TODO] implement context caching for gemini as well vertex_location=vertex_location,
cached_content = None vertex_auth_header=vertex_auth_header,
if "cached_content" in optional_params: )
cached_content = optional_params.pop("cached_content")
elif "cachedContent" in optional_params:
cached_content = optional_params.pop("cachedContent")
return _transform_request_body( return _transform_request_body(
messages=messages, messages=messages,

View file

@ -1792,7 +1792,6 @@ class VertexLLM(VertexBase):
gemini_api_key: Optional[str] = None, gemini_api_key: Optional[str] = None,
extra_headers: Optional[dict] = None, extra_headers: Optional[dict] = None,
) -> CustomStreamWrapper: ) -> CustomStreamWrapper:
request_body = await async_transform_request_body(**data) # type: ignore
should_use_v1beta1_features = self.is_using_v1beta1_features( should_use_v1beta1_features = self.is_using_v1beta1_features(
optional_params=optional_params optional_params=optional_params
@ -1826,6 +1825,13 @@ class VertexLLM(VertexBase):
litellm_params=litellm_params, 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
logging_obj.pre_call( logging_obj.pre_call(
input=messages, input=messages,
@ -1913,7 +1919,12 @@ class VertexLLM(VertexBase):
litellm_params=litellm_params, 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 = {} _async_client_params = {}
if timeout: if timeout:
_async_client_params["timeout"] = timeout _async_client_params["timeout"] = timeout
@ -2088,7 +2099,11 @@ class VertexLLM(VertexBase):
) )
## TRANSFORMATION ## ## 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
logging_obj.pre_call( logging_obj.pre_call(

View file

@ -2872,7 +2872,7 @@ def completion( # type: ignore # noqa: PLR0915
custom_llm_provider=custom_llm_provider, # type: ignore custom_llm_provider=custom_llm_provider, # type: ignore
client=client, client=client,
api_base=api_base, api_base=api_base,
extra_headers=extra_headers, extra_headers=headers,
) )
elif custom_llm_provider == "vertex_ai": elif custom_llm_provider == "vertex_ai":
@ -2941,7 +2941,7 @@ def completion( # type: ignore # noqa: PLR0915
custom_llm_provider=custom_llm_provider, # type: ignore custom_llm_provider=custom_llm_provider, # type: ignore
client=client, client=client,
api_base=api_base, api_base=api_base,
extra_headers=extra_headers, extra_headers=headers,
) )
elif "openai" in model: elif "openai" in model:
# Vertex Model Garden - OpenAI compatible models # Vertex Model Garden - OpenAI compatible models

View file

@ -22173,6 +22173,307 @@
"supports_tool_choice": true, "supports_tool_choice": true,
"supports_vision": false "supports_vision": false
}, },
"watsonx/bigscience/mt0-xxl-13b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/core42/jais-13b-chat": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/google/flan-t5-xl-3b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0001,
"output_cost_per_token": 0.00025,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-13b-chat-v2": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-13b-instruct-v2": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-3-3-8b-instruct": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.00025,
"output_cost_per_token": 0.001,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/ibm/granite-4-h-small": {
"max_tokens": 20480,
"max_input_tokens": 20480,
"max_output_tokens": 20480,
"input_cost_per_token": 0.000625,
"output_cost_per_token": 0.0025,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/ibm/granite-guardian-3-2-2b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.00015,
"output_cost_per_token": 0.0006,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-guardian-3-3-8b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.00025,
"output_cost_per_token": 0.001,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-ttm-1024-96-r2": {
"max_tokens": 512,
"max_input_tokens": 512,
"max_output_tokens": 512,
"input_cost_per_token": 0.000625,
"output_cost_per_token": 0.000625,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-ttm-1536-96-r2": {
"max_tokens": 512,
"max_input_tokens": 512,
"max_output_tokens": 512,
"input_cost_per_token": 0.000625,
"output_cost_per_token": 0.000625,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-ttm-512-96-r2": {
"max_tokens": 512,
"max_input_tokens": 512,
"max_output_tokens": 512,
"input_cost_per_token": 0.000625,
"output_cost_per_token": 0.000625,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-vision-3-2-2b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.00015,
"output_cost_per_token": 0.0006,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": true
},
"watsonx/meta-llama/llama-3-2-11b-vision-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.00025,
"output_cost_per_token": 0.001,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": true
},
"watsonx/meta-llama/llama-3-2-1b-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.0001,
"output_cost_per_token": 0.0002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/meta-llama/llama-3-2-3b-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.00015,
"output_cost_per_token": 0.0006,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/meta-llama/llama-3-2-90b-vision-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.002,
"output_cost_per_token": 0.008,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": true
},
"watsonx/meta-llama/llama-3-3-70b-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.002,
"output_cost_per_token": 0.006,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/meta-llama/llama-4-maverick-17b": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/meta-llama/llama-guard-3-11b-vision": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.00025,
"output_cost_per_token": 0.001,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": true
},
"watsonx/mistralai/mistral-medium-2505": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.00225,
"output_cost_per_token": 0.00675,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/mistralai/mistral-small-2503": {
"max_tokens": 32000,
"max_input_tokens": 32000,
"max_output_tokens": 32000,
"input_cost_per_token": 0.0002,
"output_cost_per_token": 0.0006,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/mistralai/pixtral-12b-2409": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.00015,
"output_cost_per_token": 0.00015,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": true
},
"watsonx/openai/gpt-oss-120b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.004,
"output_cost_per_token": 0.016,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/sdaia/allam-1-13b-instruct": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"whisper-1": { "whisper-1": {
"input_cost_per_second": 0.0001, "input_cost_per_second": 0.0001,
"litellm_provider": "openai", "litellm_provider": "openai",

View file

@ -242,12 +242,14 @@ def llm_passthrough_route(
request_query_params=request_query_params, request_query_params=request_query_params,
litellm_params=litellm_params_dict, litellm_params=litellm_params_dict,
) )
# need to encode the id of application-inference-profile for bedrock # [TODO: Refactor to bedrockpassthroughconfig] need to encode the id of application-inference-profile for bedrock
if custom_llm_provider == "bedrock" and "application-inference-profile" in endpoint: if custom_llm_provider == "bedrock" and "application-inference-profile" in endpoint:
encoded_url_str = CommonUtils.encode_bedrock_runtime_modelid_arn(str(updated_url)) encoded_url_str = CommonUtils.encode_bedrock_runtime_modelid_arn(
str(updated_url)
)
updated_url = httpx.URL(encoded_url_str) updated_url = httpx.URL(encoded_url_str)
# Add or update query parameters # Add or update query parameters
provider_api_key = provider_config.get_api_key(api_key) provider_api_key = provider_config.get_api_key(api_key)

View file

@ -333,6 +333,139 @@ class MCPRequestHandler:
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}") verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}")
return [] 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 @staticmethod
def is_tool_allowed( def is_tool_allowed(
allowed_mcp_servers: List[str], allowed_mcp_servers: List[str],

View file

@ -608,6 +608,45 @@ class MCPServerManager:
return tool_name not in server.disallowed_tools return tool_name not in server.disallowed_tools
return True 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( async def pre_call_tool_check(
self, self,
name: str, name: str,
@ -627,6 +666,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 = { pre_hook_kwargs = {
"name": name, "name": name,
"arguments": arguments, "arguments": arguments,

View file

@ -499,6 +499,13 @@ if MCP_AVAILABLE:
) )
filtered_tools = filter_tools_by_allowed_tools(tools, server) 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) all_tools.extend(filtered_tools)
verbose_logger.debug( verbose_logger.debug(
@ -516,6 +523,25 @@ if MCP_AVAILABLE:
) )
return all_tools 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( async def _list_mcp_tools(
user_api_key_auth: Optional[UserAPIKeyAuth] = None, user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = 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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -1,7 +1,7 @@
2:I[19107,[],"ClientPageRoot"] 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,[],""] 4:I[4707,[],""]
5:I[36423,[],""] 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"}]] 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:null

View file

@ -1,7 +1,7 @@
2:I[19107,[],"ClientPageRoot"] 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,[],""] 4:I[4707,[],""]
5:I[36423,[],""] 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"}]] 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:null

View file

@ -1,7 +1,7 @@
2:I[19107,[],"ClientPageRoot"] 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,[],""] 4:I[4707,[],""]
5:I[36423,[],""] 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"}]] 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:null

File diff suppressed because one or more lines are too long

View file

@ -1,7 +1,7 @@
2:I[19107,[],"ClientPageRoot"] 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,[],""] 4:I[4707,[],""]
5:I[36423,[],""] 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"}]] 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:null

View file

@ -1,26 +1,6 @@
model_list: model_list:
- model_name: openai/gpt-4o - model_name: gpt-5-mini
litellm_params: litellm_params:
model: openai/gpt-4o-mini model: azure/gpt-5-mini-2
api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5" api_key: os.environ/AZURE_API_KEY_ALT
api_key: dummy api_base: os.environ/AZURE_API_BASE_ALT
- 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
mcp_servers:
local_fake_mcp:
url: "http://127.0.0.1:8001/mcp"
transport: "http"
description: "My custom MCP server"
auth_type: "api_key"
auth_value: "abc123"
static_headers: {"X-API-Key": "abc123"}

View file

@ -52,6 +52,24 @@ else:
Span = Any 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): class LiteLLMTeamRoles(enum.Enum):
# team admin # team admin
TEAM_ADMIN = "admin" TEAM_ADMIN = "admin"
@ -712,6 +730,7 @@ class ModelParams(LiteLLMPydanticObjectBase):
class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase): class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase):
mcp_servers: Optional[List[str]] = None mcp_servers: Optional[List[str]] = None
mcp_access_groups: 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 vector_stores: Optional[List[str]] = None
@ -1414,6 +1433,16 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase):
object_permission_id: str object_permission_id: str
mcp_servers: Optional[List[str]] = [] mcp_servers: Optional[List[str]] = []
mcp_access_groups: 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]] = [] 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.", 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 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): class ConfigYAML(LiteLLMPydanticObjectBase):

View file

@ -21,6 +21,7 @@ import fastapi
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status
import litellm import litellm
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm._logging import verbose_proxy_logger from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid from litellm._uuid import uuid
from litellm.caching import DualCache 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) - 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. - 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/*"] - 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". - 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. - 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) - 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). - 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) - 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/*"] - 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: Examples:
1. Allow users to turn on/off pii masking 1. Allow users to turn on/off pii masking
@ -1125,6 +1126,13 @@ async def _set_object_permission(
return data_json return data_json
if "object_permission" in 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 = ( created_object_permission = (
await prisma_client.db.litellm_objectpermissiontable.create( await prisma_client.db.litellm_objectpermissiontable.create(
data=data_json["object_permission"], 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). - 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/*"] - 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. - 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 - 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 - rotation_interval: Optional[str] - How often to rotate this key (e.g., '30d', '90d'). Required if auto_rotate=True
Example: Example:
@ -2841,6 +2849,79 @@ async def list_keys(
code=status.HTTP_500_INTERNAL_SERVER_ERROR, 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( def _validate_sort_params(
sort_by: Optional[str], sort_order: str sort_by: Optional[str], sort_order: str

View file

@ -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) - 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) - 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. - 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_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_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. - 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) - 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) - 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. - 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_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_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. - team_member_tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for individual team members.

View file

@ -9,6 +9,8 @@ from typing import Dict, Optional, Union
from litellm._logging import verbose_proxy_logger from litellm._logging import verbose_proxy_logger
from litellm.proxy.utils import PrismaClient from litellm.proxy.utils import PrismaClient
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
async def attach_object_permission_to_dict( async def attach_object_permission_to_dict(
@ -114,6 +116,15 @@ async def handle_update_object_permission_common(
if isinstance(new_object_permission, dict): if isinstance(new_object_permission, dict):
existing_object_permissions_dict.update(new_object_permission) 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 # Commit the update to the LiteLLM_ObjectPermissionTable
######################################################### #########################################################

View file

@ -21,9 +21,7 @@ from litellm.constants import BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import * from litellm.proxy._types import *
from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import ( from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
user_api_key_auth,
)
from litellm.proxy.common_utils.http_parsing_utils import ( from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body, _read_request_body,
get_form_data, get_form_data,
@ -31,6 +29,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
) )
from litellm.proxy.pass_through_endpoints.common_utils import get_litellm_virtual_key from litellm.proxy.pass_through_endpoints.common_utils import get_litellm_virtual_key
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
HttpPassThroughEndpointHelpers,
create_pass_through_route, create_pass_through_route,
create_websocket_passthrough_route, create_websocket_passthrough_route,
websocket_passthrough_request, websocket_passthrough_request,
@ -57,7 +56,9 @@ def create_request_copy(request: Request):
} }
def is_passthrough_request_using_router_model(request_body: dict, llm_router: Optional[litellm.Router]) -> bool: def is_passthrough_request_using_router_model(
request_body: dict, llm_router: Optional[litellm.Router]
) -> bool:
""" """
Returns True if the model is in the llm_router model names Returns True if the model is in the llm_router model names
""" """
@ -93,12 +94,16 @@ async def llm_passthrough_factory_proxy_route(
model=None, model=None,
) )
if provider_config is None: if provider_config is None:
raise HTTPException(status_code=404, detail=f"Provider {custom_llm_provider} not found") raise HTTPException(
status_code=404, detail=f"Provider {custom_llm_provider} not found"
)
base_target_url = provider_config.get_api_base() base_target_url = provider_config.get_api_base()
if base_target_url is None: if base_target_url is None:
raise HTTPException(status_code=404, detail=f"Provider {custom_llm_provider} api base not found") raise HTTPException(
status_code=404, detail=f"Provider {custom_llm_provider} api base not found"
)
encoded_endpoint = httpx.URL(endpoint).path encoded_endpoint = httpx.URL(endpoint).path
@ -177,11 +182,17 @@ async def gemini_proxy_route(
[Docs](https://docs.litellm.ai/docs/pass_through/google_ai_studio) [Docs](https://docs.litellm.ai/docs/pass_through/google_ai_studio)
""" """
## CHECK FOR LITELLM API KEY IN THE QUERY PARAMS - ?..key=LITELLM_API_KEY ## CHECK FOR LITELLM API KEY IN THE QUERY PARAMS - ?..key=LITELLM_API_KEY
google_ai_studio_api_key = request.query_params.get("key") or request.headers.get("x-goog-api-key") google_ai_studio_api_key = request.query_params.get("key") or request.headers.get(
"x-goog-api-key"
)
user_api_key_dict = await user_api_key_auth(request=request, api_key=f"Bearer {google_ai_studio_api_key}") user_api_key_dict = await user_api_key_auth(
request=request, api_key=f"Bearer {google_ai_studio_api_key}"
)
base_target_url = os.getenv("GEMINI_API_BASE") or "https://generativelanguage.googleapis.com" base_target_url = (
os.getenv("GEMINI_API_BASE") or "https://generativelanguage.googleapis.com"
)
encoded_endpoint = httpx.URL(endpoint).path encoded_endpoint = httpx.URL(endpoint).path
# Ensure endpoint starts with '/' for proper URL construction # Ensure endpoint starts with '/' for proper URL construction
@ -293,13 +304,12 @@ async def vllm_proxy_route(
""" """
[Docs](https://docs.litellm.ai/docs/pass_through/vllm) [Docs](https://docs.litellm.ai/docs/pass_through/vllm)
""" """
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
HttpPassThroughEndpointHelpers,
)
from litellm.proxy.proxy_server import llm_router from litellm.proxy.proxy_server import llm_router
request_body = await get_request_body(request) request_body = await get_request_body(request)
is_router_model = is_passthrough_request_using_router_model(request_body, llm_router) is_router_model = is_passthrough_request_using_router_model(
request_body, llm_router
)
is_streaming_request = is_passthrough_request_streaming(request_body) is_streaming_request = is_passthrough_request_streaming(request_body)
if is_router_model and llm_router: if is_router_model and llm_router:
result = cast( result = cast(
@ -314,7 +324,11 @@ async def vllm_proxy_route(
content=None, content=None,
data=None, data=None,
files=None, files=None,
json=(request_body if request.headers.get("content-type") == "application/json" else None), json=(
request_body
if request.headers.get("content-type") == "application/json"
else None
),
params=None, params=None,
headers=None, headers=None,
cookies=None, cookies=None,
@ -492,7 +506,9 @@ async def handle_bedrock_count_tokens(
# Extract model from request body # Extract model from request body
model = request_body.get("model") model = request_body.get("model")
if not model: if not model:
raise HTTPException(status_code=400, detail={"error": "Model is required in request body"}) raise HTTPException(
status_code=400, detail={"error": "Model is required in request body"}
)
# Get model parameters from router # Get model parameters from router
litellm_params = {"user_api_key_dict": user_api_key_dict} litellm_params = {"user_api_key_dict": user_api_key_dict}
@ -531,7 +547,9 @@ async def handle_bedrock_count_tokens(
raise raise
except Exception as e: except Exception as e:
verbose_proxy_logger.error(f"Error in handle_bedrock_count_tokens: {str(e)}") verbose_proxy_logger.error(f"Error in handle_bedrock_count_tokens: {str(e)}")
raise HTTPException(status_code=500, detail={"error": f"CountTokens processing error: {str(e)}"}) raise HTTPException(
status_code=500, detail={"error": f"CountTokens processing error: {str(e)}"}
)
async def bedrock_llm_proxy_route( async def bedrock_llm_proxy_route(
@ -583,7 +601,8 @@ async def bedrock_llm_proxy_route(
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={
"error": "Model missing from endpoint. Expected format: /model/<Model>/<endpoint>. Got: " + endpoint, "error": "Model missing from endpoint. Expected format: /model/<Model>/<endpoint>. Got: "
+ endpoint,
}, },
) )
@ -647,7 +666,9 @@ async def bedrock_proxy_route(
aws_region_name = litellm.utils.get_secret(secret_name="AWS_REGION_NAME") aws_region_name = litellm.utils.get_secret(secret_name="AWS_REGION_NAME")
if _is_bedrock_agent_runtime_route(endpoint=endpoint): # handle bedrock agents if _is_bedrock_agent_runtime_route(endpoint=endpoint): # handle bedrock agents
base_target_url = f"https://bedrock-agent-runtime.{aws_region_name}.amazonaws.com" base_target_url = (
f"https://bedrock-agent-runtime.{aws_region_name}.amazonaws.com"
)
else: else:
return await bedrock_llm_proxy_route( return await bedrock_llm_proxy_route(
endpoint=endpoint, endpoint=endpoint,
@ -677,7 +698,9 @@ async def bedrock_proxy_route(
data = await request.json() data = await request.json()
except Exception as e: except Exception as e:
raise HTTPException(status_code=400, detail={"error": e}) raise HTTPException(status_code=400, detail={"error": e})
_request = AWSRequest(method="POST", url=str(updated_url), data=json.dumps(data), headers=headers) _request = AWSRequest(
method="POST", url=str(updated_url), data=json.dumps(data), headers=headers
)
sigv4.add_auth(_request) sigv4.add_auth(_request)
prepped = _request.prepare() prepped = _request.prepare()
@ -738,8 +761,14 @@ async def assemblyai_proxy_route(
[Docs](https://api.assemblyai.com) [Docs](https://api.assemblyai.com)
""" """
# Set base URL based on the route # Set base URL based on the route
assembly_region = AssemblyAIPassthroughLoggingHandler._get_assembly_region_from_url(url=str(request.url)) assembly_region = AssemblyAIPassthroughLoggingHandler._get_assembly_region_from_url(
base_target_url = AssemblyAIPassthroughLoggingHandler._get_assembly_base_url_from_region(region=assembly_region) url=str(request.url)
)
base_target_url = (
AssemblyAIPassthroughLoggingHandler._get_assembly_base_url_from_region(
region=assembly_region
)
)
encoded_endpoint = httpx.URL(endpoint).path encoded_endpoint = httpx.URL(endpoint).path
# Ensure endpoint starts with '/' for proper URL construction # Ensure endpoint starts with '/' for proper URL construction
if not encoded_endpoint.startswith("/"): if not encoded_endpoint.startswith("/"):
@ -794,17 +823,91 @@ async def azure_proxy_route(
Call any azure endpoint using the proxy. Call any azure endpoint using the proxy.
Just use `{PROXY_BASE_URL}/azure/{endpoint:path}` Just use `{PROXY_BASE_URL}/azure/{endpoint:path}`
Checks if the deployment id in the url is a litellm model name. If so, it will route using the llm_router.allm_passthrough_route.
""" """
from litellm.proxy.proxy_server import llm_router
parts = endpoint.split(
"/"
) # azure model is in the url - e.g. https://{endpoint}/openai/deployments/{deployment-id}/completions?api-version=2024-10-21
if len(parts) > 1 and llm_router:
for part in parts:
is_router_model = is_passthrough_request_using_router_model(
request_body={"model": part}, llm_router=llm_router
)
if is_router_model:
request_body = await get_request_body(request)
is_streaming_request = is_passthrough_request_streaming(request_body)
result = await llm_router.allm_passthrough_route(
model=part,
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=dict(request.headers),
stream=request_body.get("stream", False),
content=None,
data=None,
files=None,
json=(
request_body
if request.headers.get("content-type") == "application/json"
else None
),
params=None,
headers=None,
cookies=None,
)
if is_streaming_request:
# Check if result is an async generator (from _async_streaming)
import inspect
if inspect.isasyncgen(result):
# Result is already an async generator, use it directly
return StreamingResponse(
content=result,
status_code=200,
headers={"content-type": "text/event-stream"},
)
else:
# Result is an httpx.Response, use aiter_bytes()
result = cast(httpx.Response, result)
return StreamingResponse(
content=result.aiter_bytes(),
status_code=result.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=result.headers,
custom_headers=None,
),
)
# Non-streaming response
result = cast(httpx.Response, result)
content = await result.aread()
return Response(
content=content,
status_code=result.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=result.headers,
custom_headers=None,
),
)
base_target_url = get_secret_str(secret_name="AZURE_API_BASE") base_target_url = get_secret_str(secret_name="AZURE_API_BASE")
if base_target_url is None: if base_target_url is None:
raise Exception("Required 'AZURE_API_BASE' in environment to make pass-through calls to Azure.") raise Exception(
"Required 'AZURE_API_BASE' in environment to make pass-through calls to Azure."
)
# Add or update query parameters # Add or update query parameters
azure_api_key = passthrough_endpoint_router.get_credentials( azure_api_key = passthrough_endpoint_router.get_credentials(
custom_llm_provider=litellm.LlmProviders.AZURE.value, custom_llm_provider=litellm.LlmProviders.AZURE.value,
region_name=None, region_name=None,
) )
if azure_api_key is None: if azure_api_key is None:
raise Exception("Required 'AZURE_API_KEY' in environment to make pass-through calls to Azure.") raise Exception(
"Required 'AZURE_API_KEY' in environment to make pass-through calls to Azure."
)
return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler( return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler(
endpoint=endpoint, endpoint=endpoint,
@ -828,7 +931,9 @@ class BaseVertexAIPassThroughHandler(ABC):
@staticmethod @staticmethod
@abstractmethod @abstractmethod
def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: Optional[str]) -> str: def update_base_target_url_with_credential_location(
base_target_url: str, vertex_location: Optional[str]
) -> str:
pass pass
@ -838,7 +943,9 @@ class VertexAIDiscoveryPassThroughHandler(BaseVertexAIPassThroughHandler):
return "https://discoveryengine.googleapis.com/" return "https://discoveryengine.googleapis.com/"
@staticmethod @staticmethod
def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: Optional[str]) -> str: def update_base_target_url_with_credential_location(
base_target_url: str, vertex_location: Optional[str]
) -> str:
return base_target_url return base_target_url
@ -848,7 +955,9 @@ class VertexAIPassThroughHandler(BaseVertexAIPassThroughHandler):
return get_vertex_base_url(vertex_location) return get_vertex_base_url(vertex_location)
@staticmethod @staticmethod
def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: Optional[str]) -> str: def update_base_target_url_with_credential_location(
base_target_url: str, vertex_location: Optional[str]
) -> str:
return get_vertex_base_url(vertex_location) return get_vertex_base_url(vertex_location)
@ -914,14 +1023,18 @@ async def _base_vertex_proxy_route(
location=vertex_location, location=vertex_location,
) )
base_target_url = get_vertex_pass_through_handler.get_default_base_target_url(vertex_location) base_target_url = get_vertex_pass_through_handler.get_default_base_target_url(
vertex_location
)
headers_passed_through = False headers_passed_through = False
# Use headers from the incoming request if no vertex credentials are found # Use headers from the incoming request if no vertex credentials are found
if vertex_credentials is None or vertex_credentials.vertex_project is None: if vertex_credentials is None or vertex_credentials.vertex_project is None:
headers = dict(request.headers) or {} headers = dict(request.headers) or {}
headers_passed_through = True headers_passed_through = True
verbose_proxy_logger.debug("default_vertex_config not set, incoming request headers %s", headers) verbose_proxy_logger.debug(
"default_vertex_config not set, incoming request headers %s", headers
)
headers.pop("content-length", None) headers.pop("content-length", None)
headers.pop("host", None) headers.pop("host", None)
else: else:
@ -1087,7 +1200,9 @@ async def openai_proxy_route(
region_name=None, region_name=None,
) )
if openai_api_key is None: if openai_api_key is None:
raise Exception("Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI.") raise Exception(
"Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI."
)
return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler( return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler(
endpoint=endpoint, endpoint=endpoint,
@ -1133,7 +1248,9 @@ class BaseOpenAIPassThroughHandler:
endpoint_func = create_pass_through_route( endpoint_func = create_pass_through_route(
endpoint=endpoint, endpoint=endpoint,
target=str(updated_url), target=str(updated_url),
custom_headers=BaseOpenAIPassThroughHandler._assemble_headers(api_key=api_key, request=request), custom_headers=BaseOpenAIPassThroughHandler._assemble_headers(
api_key=api_key, request=request
),
) # dynamically construct pass-through endpoint based on incoming path ) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func( received_value = await endpoint_func(
request, request,
@ -1150,7 +1267,10 @@ class BaseOpenAIPassThroughHandler:
""" """
Appends the OpenAI-Beta header to the headers if the request is an OpenAI Assistants API request Appends the OpenAI-Beta header to the headers if the request is an OpenAI Assistants API request
""" """
if RouteChecks._is_assistants_api_request(request) is True and "OpenAI-Beta" not in headers: if (
RouteChecks._is_assistants_api_request(request) is True
and "OpenAI-Beta" not in headers
):
headers["OpenAI-Beta"] = "assistants=v2" headers["OpenAI-Beta"] = "assistants=v2"
return headers return headers
@ -1166,7 +1286,9 @@ class BaseOpenAIPassThroughHandler:
) )
@staticmethod @staticmethod
def _join_url_paths(base_url: httpx.URL, path: str, custom_llm_provider: litellm.LlmProviders) -> str: def _join_url_paths(
base_url: httpx.URL, path: str, custom_llm_provider: litellm.LlmProviders
) -> str:
""" """
Properly joins a base URL with a path, preserving any existing path in the base URL. Properly joins a base URL with a path, preserving any existing path in the base URL.
""" """
@ -1182,9 +1304,14 @@ class BaseOpenAIPassThroughHandler:
joined_path_str = str(base_url.copy_with(path=full_path)) joined_path_str = str(base_url.copy_with(path=full_path))
# Apply OpenAI-specific path handling for both branches # Apply OpenAI-specific path handling for both branches
if custom_llm_provider == litellm.LlmProviders.OPENAI and "/v1/" not in joined_path_str: if (
custom_llm_provider == litellm.LlmProviders.OPENAI
and "/v1/" not in joined_path_str
):
# Insert v1 after api.openai.com for OpenAI requests # Insert v1 after api.openai.com for OpenAI requests
joined_path_str = joined_path_str.replace("api.openai.com/", "api.openai.com/v1/") joined_path_str = joined_path_str.replace(
"api.openai.com/", "api.openai.com/v1/"
)
return joined_path_str return joined_path_str
@ -1231,9 +1358,7 @@ async def vertex_ai_live_websocket_passthrough(
if vertex_credentials_config is not None: if vertex_credentials_config is not None:
resolved_project = resolved_project or vertex_credentials_config.vertex_project resolved_project = resolved_project or vertex_credentials_config.vertex_project
temp_location = ( temp_location = resolved_location or vertex_credentials_config.vertex_location
resolved_location or vertex_credentials_config.vertex_location
)
# Ensure resolved_location is a string # Ensure resolved_location is a string
if isinstance(temp_location, dict): if isinstance(temp_location, dict):
resolved_location = str(temp_location) resolved_location = str(temp_location)
@ -1241,7 +1366,11 @@ async def vertex_ai_live_websocket_passthrough(
resolved_location = str(temp_location) resolved_location = str(temp_location)
else: else:
resolved_location = None resolved_location = None
credentials_value = str(vertex_credentials_config.vertex_credentials) if vertex_credentials_config.vertex_credentials is not None else None credentials_value = (
str(vertex_credentials_config.vertex_credentials)
if vertex_credentials_config.vertex_credentials is not None
else None
)
try: try:
resolved_location = resolved_location or ( resolved_location = resolved_location or (
@ -1302,7 +1431,7 @@ async def vertex_ai_live_websocket_passthrough(
# Use the new WebSocket passthrough pattern # Use the new WebSocket passthrough pattern
if user_api_key_dict is None: if user_api_key_dict is None:
raise ValueError("user_api_key_dict is required for WebSocket passthrough") raise ValueError("user_api_key_dict is required for WebSocket passthrough")
return await websocket_passthrough_request( return await websocket_passthrough_request(
websocket=websocket, websocket=websocket,
target=service_url, target=service_url,

View file

@ -2957,6 +2957,40 @@ class ProxyConfig:
return config 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: async def _get_models_from_db(self, prisma_client: PrismaClient) -> list:
try: try:
new_models = await prisma_client.db.litellm_proxymodeltable.find_many() 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}" 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 # update llm router
await self._update_llm_router( await self._update_llm_router(
new_models=new_models, proxy_logging_obj=proxy_logging_obj new_models=new_models, proxy_logging_obj=proxy_logging_obj
) )
db_general_settings = await prisma_client.db.litellm_config.find_first( db_general_settings = await prisma_client.db.litellm_config.find_first(
where={"param_name": "general_settings"} where={"param_name": "general_settings"}
@ -3021,12 +3057,23 @@ class ProxyConfig:
ex. Vector Stores, Guardrails, MCP tools, etc. ex. Vector Stores, Guardrails, MCP tools, etc.
""" """
await self._init_guardrails_in_db(prisma_client=prisma_client) if self._should_load_db_object(object_type="guardrails"):
await self._init_vector_stores_in_db(prisma_client=prisma_client) await self._init_guardrails_in_db(prisma_client=prisma_client)
await self._init_mcp_servers_in_db()
await self._init_pass_through_endpoints_in_db() if self._should_load_db_object(object_type="vector_stores"):
await self._init_prompts_in_db(prisma_client=prisma_client) await self._init_vector_stores_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="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): async def _check_and_reload_model_cost_map(self, prisma_client: PrismaClient):
""" """

View file

@ -156,6 +156,7 @@ model LiteLLM_ObjectPermissionTable {
object_permission_id String @id @default(uuid()) object_permission_id String @id @default(uuid())
mcp_servers String[] @default([]) mcp_servers String[] @default([])
mcp_access_groups 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([]) vector_stores String[] @default([])
teams LiteLLM_TeamTable[] teams LiteLLM_TeamTable[]
verification_tokens LiteLLM_VerificationToken[] verification_tokens LiteLLM_VerificationToken[]

View file

@ -3601,8 +3601,10 @@ def is_known_model(model: Optional[str], llm_router: Optional[Router]) -> bool:
return False return False
model_names = llm_router.get_model_names() model_names = llm_router.get_model_names()
model_names_set = set(model_names)
is_in_list = False is_in_list = False
if model in model_names: if model in model_names_set:
is_in_list = True is_in_list = True
return is_in_list return is_in_list

View file

@ -416,6 +416,9 @@ class Router:
# Initialize model ID to deployment index mapping for O(1) lookups # Initialize model ID to deployment index mapping for O(1) lookups
self.model_id_to_deployment_index_map: Dict[str, int] = {} self.model_id_to_deployment_index_map: Dict[str, int] = {}
# Initialize model name to deployment indices mapping for O(1) lookups
# Maps model_name -> list of indices in model_list
self.model_name_to_deployment_indices: Dict[str, List[int]] = {}
if model_list is not None: if model_list is not None:
# Build model index immediately to enable O(1) lookups from the start # Build model index immediately to enable O(1) lookups from the start
@ -5097,6 +5100,7 @@ class Router:
original_model_list = copy.deepcopy(model_list) original_model_list = copy.deepcopy(model_list)
self.model_list = [] self.model_list = []
self.model_id_to_deployment_index_map = {} # Reset the index self.model_id_to_deployment_index_map = {} # Reset the index
self.model_name_to_deployment_indices = {} # Reset the model_name index
# we add api_base/api_key each model so load balancing between azure/gpt on api_base1 and api_base2 works # we add api_base/api_key each model so load balancing between azure/gpt on api_base1 and api_base2 works
for model in original_model_list: for model in original_model_list:
@ -5138,6 +5142,9 @@ class Router:
f"\nInitialized Model List {self.get_model_names()}" f"\nInitialized Model List {self.get_model_names()}"
) )
self.model_names = [m["model_name"] for m in model_list] self.model_names = [m["model_name"] for m in model_list]
# Build model_name index for O(1) lookups
self._build_model_name_index(self.model_list)
def _add_deployment(self, deployment: Deployment) -> Deployment: def _add_deployment(self, deployment: Deployment) -> Deployment:
import os import os
@ -5365,20 +5372,27 @@ class Router:
self, model: dict, model_id: Optional[str] = None self, model: dict, model_id: Optional[str] = None
) -> None: ) -> None:
""" """
Helper method to add a model to the model_list and update the model_id_to_deployment_index_map. Helper method to add a model to the model_list and update both indices.
Parameters: Parameters:
- model: dict - the model to add to the list - model: dict - the model to add to the list
- model_id: Optional[str] - the model ID to use for indexing. If None, will try to get from model["model_info"]["id"] - model_id: Optional[str] - the model ID to use for indexing. If None, will try to get from model["model_info"]["id"]
""" """
idx = len(self.model_list)
self.model_list.append(model) self.model_list.append(model)
# Update model index for O(1) lookup
# Update model_id index for O(1) lookup
if model_id is not None: if model_id is not None:
self.model_id_to_deployment_index_map[model_id] = len(self.model_list) - 1 self.model_id_to_deployment_index_map[model_id] = idx
elif model.get("model_info", {}).get("id") is not None: elif model.get("model_info", {}).get("id") is not None:
self.model_id_to_deployment_index_map[model["model_info"]["id"]] = ( self.model_id_to_deployment_index_map[model["model_info"]["id"]] = idx
len(self.model_list) - 1
) # Update model_name index for O(1) lookup
model_name = model.get("model_name")
if model_name:
if model_name not in self.model_name_to_deployment_indices:
self.model_name_to_deployment_indices[model_name] = []
self.model_name_to_deployment_indices[model_name].append(idx)
def upsert_deployment(self, deployment: Deployment) -> Optional[Deployment]: def upsert_deployment(self, deployment: Deployment) -> Optional[Deployment]:
""" """
@ -6094,6 +6108,22 @@ class Router:
additional_headers[header] = value additional_headers[header] = value
return response return response
def _build_model_name_index(self, model_list: list) -> None:
"""
Build model_name -> deployment indices mapping for O(1) lookups.
This index allows us to find all deployments for a given model_name in O(1) time
instead of O(n) linear scan through the entire model_list.
"""
self.model_name_to_deployment_indices.clear()
for idx, model in enumerate(model_list):
model_name = model.get("model_name")
if model_name:
if model_name not in self.model_name_to_deployment_indices:
self.model_name_to_deployment_indices[model_name] = []
self.model_name_to_deployment_indices[model_name].append(idx)
def _build_model_id_to_deployment_index_map(self, model_list: list): def _build_model_id_to_deployment_index_map(self, model_list: list):
""" """
Build model index from model list to enable O(1) lookups immediately. Build model index from model list to enable O(1) lookups immediately.
@ -6198,18 +6228,41 @@ class Router:
Used for accurate 'get_model_list'. Used for accurate 'get_model_list'.
if team_id specified, only return team-specific models if team_id specified, only return team-specific models
Optimized with O(1) index lookup instead of O(n) linear scan.
""" """
returned_models: List[DeploymentTypedDict] = [] returned_models: List[DeploymentTypedDict] = []
for model in self.model_list:
if self.should_include_deployment( # O(1) lookup in model_name index
model_name=model_name, model=model, team_id=team_id if model_name in self.model_name_to_deployment_indices:
): indices = self.model_name_to_deployment_indices[model_name]
if model_alias is not None:
alias_model = copy.deepcopy(model) # O(k) where k = deployments for this model_name (typically 1-10)
alias_model["model_name"] = model_alias for idx in indices:
returned_models.append(alias_model) model = self.model_list[idx]
else: if self.should_include_deployment(
returned_models.append(model) 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)
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 return returned_models

View file

@ -532,9 +532,6 @@ def get_dynamic_callbacks(
return returned_callbacks return returned_callbacks
def function_setup( # noqa: PLR0915 def function_setup( # noqa: PLR0915
original_function: str, rules_obj, start_time, *args, **kwargs original_function: str, rules_obj, start_time, *args, **kwargs
): # just run once to check if user wants to send their data anywhere - PostHog/Sentry/Slack/etc. ): # just run once to check if user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
@ -802,7 +799,7 @@ def function_setup( # noqa: PLR0915
call_type=call_type, call_type=call_type,
): ):
stream = True stream = True
logging_obj = get_litellm_logging_class()( # Victim for object pool logging_obj = get_litellm_logging_class()( # Victim for object pool
model=model, # type: ignore model=model, # type: ignore
messages=messages, messages=messages,
stream=stream, stream=stream,
@ -910,159 +907,169 @@ def _get_wrapper_timeout(
return 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 def client(original_function): # noqa: PLR0915
rules_obj = Rules() 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) @wraps(original_function)
def wrapper(*args, **kwargs): # noqa: PLR0915 def wrapper(*args, **kwargs): # noqa: PLR0915
@ -1273,6 +1280,8 @@ def client(original_function): # noqa: PLR0915
original_response=result, original_response=result,
model=model or None, model=model or None,
optional_params=kwargs, optional_params=kwargs,
original_function=original_function,
rules_obj=rules_obj,
) )
# [OPTIONAL] ADD TO CACHE # [OPTIONAL] ADD TO CACHE
@ -1417,7 +1426,8 @@ def client(original_function): # noqa: PLR0915
if _caching_handler_response is not None: if _caching_handler_response is not None:
if ( if (
_caching_handler_response.cached_result is not None _caching_handler_response.cached_result is not None
and _caching_handler_response.final_embedding_cached_response is None and _caching_handler_response.final_embedding_cached_response
is None
): ):
return _caching_handler_response.cached_result return _caching_handler_response.cached_result
@ -1489,7 +1499,11 @@ def client(original_function): # noqa: PLR0915
return result return result
### POST-CALL RULES ### ### POST-CALL RULES ###
post_call_processing( 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 # Only run if call_type is a valid value in CallTypes
if call_type in [ct.value for ct in CallTypes]: if call_type in [ct.value for ct in CallTypes]:
@ -1683,7 +1697,6 @@ def _is_streaming_request(
return False return False
def _select_tokenizer( def _select_tokenizer(
model: str, custom_tokenizer: Optional[CustomHuggingfaceTokenizer] = None model: str, custom_tokenizer: Optional[CustomHuggingfaceTokenizer] = None
): ):
@ -4867,16 +4880,24 @@ def _get_model_info_helper( # noqa: PLR0915
max_input_tokens=_model_info.get("max_input_tokens", None), max_input_tokens=_model_info.get("max_input_tokens", None),
max_output_tokens=_model_info.get("max_output_tokens", None), max_output_tokens=_model_info.get("max_output_tokens", None),
input_cost_per_token=_input_cost_per_token, input_cost_per_token=_input_cost_per_token,
input_cost_per_token_flex=_model_info.get("input_cost_per_token_flex", None), input_cost_per_token_flex=_model_info.get(
input_cost_per_token_priority=_model_info.get("input_cost_per_token_priority", None), "input_cost_per_token_flex", None
),
input_cost_per_token_priority=_model_info.get(
"input_cost_per_token_priority", None
),
cache_creation_input_token_cost=_model_info.get( cache_creation_input_token_cost=_model_info.get(
"cache_creation_input_token_cost", None "cache_creation_input_token_cost", None
), ),
cache_read_input_token_cost=_model_info.get( cache_read_input_token_cost=_model_info.get(
"cache_read_input_token_cost", None "cache_read_input_token_cost", None
), ),
cache_read_input_token_cost_flex=_model_info.get("cache_read_input_token_cost_flex", None), cache_read_input_token_cost_flex=_model_info.get(
cache_read_input_token_cost_priority=_model_info.get("cache_read_input_token_cost_priority", None), "cache_read_input_token_cost_flex", None
),
cache_read_input_token_cost_priority=_model_info.get(
"cache_read_input_token_cost_priority", None
),
cache_creation_input_token_cost_above_1hr=_model_info.get( cache_creation_input_token_cost_above_1hr=_model_info.get(
"cache_creation_input_token_cost_above_1hr", None "cache_creation_input_token_cost_above_1hr", None
), ),
@ -4901,8 +4922,12 @@ def _get_model_info_helper( # noqa: PLR0915
"output_cost_per_token_batches" "output_cost_per_token_batches"
), ),
output_cost_per_token=_output_cost_per_token, output_cost_per_token=_output_cost_per_token,
output_cost_per_token_flex=_model_info.get("output_cost_per_token_flex", None), output_cost_per_token_flex=_model_info.get(
output_cost_per_token_priority=_model_info.get("output_cost_per_token_priority", None), "output_cost_per_token_flex", None
),
output_cost_per_token_priority=_model_info.get(
"output_cost_per_token_priority", None
),
output_cost_per_audio_token=_model_info.get( output_cost_per_audio_token=_model_info.get(
"output_cost_per_audio_token", None "output_cost_per_audio_token", None
), ),
@ -6434,7 +6459,7 @@ def get_valid_models(
try: try:
################################ ################################
# init litellm_params # init litellm_params
################################# #################################
if litellm_params is None: if litellm_params is None:
litellm_params = LiteLLM_Params(model="") litellm_params = LiteLLM_Params(model="")
@ -6443,7 +6468,7 @@ def get_valid_models(
if api_base is not None: if api_base is not None:
litellm_params.api_base = api_base litellm_params.api_base = api_base
################################# #################################
check_provider_endpoint = ( check_provider_endpoint = (
check_provider_endpoint or litellm.check_provider_endpoint check_provider_endpoint or litellm.check_provider_endpoint
) )
@ -6918,7 +6943,10 @@ class ProviderConfigManager:
return litellm.LlamaAPIConfig() return litellm.LlamaAPIConfig()
elif litellm.LlmProviders.TEXT_COMPLETION_OPENAI == provider: elif litellm.LlmProviders.TEXT_COMPLETION_OPENAI == provider:
return litellm.OpenAITextCompletionConfig() return litellm.OpenAITextCompletionConfig()
elif litellm.LlmProviders.COHERE_CHAT == provider or litellm.LlmProviders.COHERE == provider: elif (
litellm.LlmProviders.COHERE_CHAT == provider
or litellm.LlmProviders.COHERE == provider
):
return litellm.CohereChatConfig() return litellm.CohereChatConfig()
elif litellm.LlmProviders.SNOWFLAKE == provider: elif litellm.LlmProviders.SNOWFLAKE == provider:
return litellm.SnowflakeConfig() return litellm.SnowflakeConfig()
@ -7345,7 +7373,12 @@ class ProviderConfigManager:
) )
return VLLMPassthroughConfig() return VLLMPassthroughConfig()
elif LlmProviders.AZURE == provider:
from litellm.llms.azure.passthrough.transformation import (
AzurePassthroughConfig,
)
return AzurePassthroughConfig()
return None return None
@staticmethod @staticmethod
@ -7532,9 +7565,7 @@ class ProviderConfigManager:
return RecraftImageEditConfig() return RecraftImageEditConfig()
elif LlmProviders.AZURE_AI == provider: elif LlmProviders.AZURE_AI == provider:
from litellm.llms.azure_ai.image_edit import ( from litellm.llms.azure_ai.image_edit import get_azure_ai_image_edit_config
get_azure_ai_image_edit_config,
)
return get_azure_ai_image_edit_config(model) return get_azure_ai_image_edit_config(model)
elif LlmProviders.LITELLM_PROXY == provider: elif LlmProviders.LITELLM_PROXY == provider:
@ -7589,7 +7620,9 @@ def get_end_user_id_for_cost_tracking(
service_type: "litellm_logging" or "prometheus" - used to allow prometheus only disable cost tracking. service_type: "litellm_logging" or "prometheus" - used to allow prometheus only disable cost tracking.
""" """
_metadata = cast(dict, get_litellm_metadata_from_kwargs(dict(litellm_params=litellm_params))) _metadata = cast(
dict, get_litellm_metadata_from_kwargs(dict(litellm_params=litellm_params))
)
end_user_id = cast( end_user_id = cast(
Optional[str], Optional[str],

View file

@ -12872,6 +12872,39 @@
"supports_tool_choice": true, "supports_tool_choice": true,
"supports_vision": 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": { "gpt-5-codex": {
"cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost": 1.25e-07,
"input_cost_per_token": 1.25e-06, "input_cost_per_token": 1.25e-06,
@ -13189,6 +13222,20 @@
"/v1/images/generations" "/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": { "gpt-realtime": {
"cache_creation_input_audio_token_cost": 4e-07, "cache_creation_input_audio_token_cost": 4e-07,
"cache_read_input_token_cost": 4e-07, "cache_read_input_token_cost": 4e-07,
@ -14623,6 +14670,54 @@
"/v1/images/generations" "/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": { "medlm-large": {
"input_cost_per_character": 5e-06, "input_cost_per_character": 5e-06,
"litellm_provider": "vertex_ai-language-models", "litellm_provider": "vertex_ai-language-models",
@ -22173,6 +22268,307 @@
"supports_tool_choice": true, "supports_tool_choice": true,
"supports_vision": false "supports_vision": false
}, },
"watsonx/bigscience/mt0-xxl-13b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/core42/jais-13b-chat": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/google/flan-t5-xl-3b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0001,
"output_cost_per_token": 0.00025,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-13b-chat-v2": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-13b-instruct-v2": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-3-3-8b-instruct": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.00025,
"output_cost_per_token": 0.001,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/ibm/granite-4-h-small": {
"max_tokens": 20480,
"max_input_tokens": 20480,
"max_output_tokens": 20480,
"input_cost_per_token": 0.000625,
"output_cost_per_token": 0.0025,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/ibm/granite-guardian-3-2-2b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.00015,
"output_cost_per_token": 0.0006,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-guardian-3-3-8b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.00025,
"output_cost_per_token": 0.001,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-ttm-1024-96-r2": {
"max_tokens": 512,
"max_input_tokens": 512,
"max_output_tokens": 512,
"input_cost_per_token": 0.000625,
"output_cost_per_token": 0.000625,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-ttm-1536-96-r2": {
"max_tokens": 512,
"max_input_tokens": 512,
"max_output_tokens": 512,
"input_cost_per_token": 0.000625,
"output_cost_per_token": 0.000625,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-ttm-512-96-r2": {
"max_tokens": 512,
"max_input_tokens": 512,
"max_output_tokens": 512,
"input_cost_per_token": 0.000625,
"output_cost_per_token": 0.000625,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/ibm/granite-vision-3-2-2b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.00015,
"output_cost_per_token": 0.0006,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": true
},
"watsonx/meta-llama/llama-3-2-11b-vision-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.00025,
"output_cost_per_token": 0.001,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": true
},
"watsonx/meta-llama/llama-3-2-1b-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.0001,
"output_cost_per_token": 0.0002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/meta-llama/llama-3-2-3b-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.00015,
"output_cost_per_token": 0.0006,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/meta-llama/llama-3-2-90b-vision-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.002,
"output_cost_per_token": 0.008,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": true
},
"watsonx/meta-llama/llama-3-3-70b-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.002,
"output_cost_per_token": 0.006,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/meta-llama/llama-4-maverick-17b": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/meta-llama/llama-guard-3-11b-vision": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.00025,
"output_cost_per_token": 0.001,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": true
},
"watsonx/mistralai/mistral-medium-2505": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.00225,
"output_cost_per_token": 0.00675,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/mistralai/mistral-small-2503": {
"max_tokens": 32000,
"max_input_tokens": 32000,
"max_output_tokens": 32000,
"input_cost_per_token": 0.0002,
"output_cost_per_token": 0.0006,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": false
},
"watsonx/mistralai/pixtral-12b-2409": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.00015,
"output_cost_per_token": 0.00015,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": true
},
"watsonx/openai/gpt-oss-120b": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.004,
"output_cost_per_token": 0.016,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"watsonx/sdaia/allam-1-13b-instruct": {
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"input_cost_per_token": 0.0005,
"output_cost_per_token": 0.002,
"litellm_provider": "watsonx",
"mode": "chat",
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": false
},
"whisper-1": { "whisper-1": {
"input_cost_per_second": 0.0001, "input_cost_per_second": 0.0001,
"litellm_provider": "openai", "litellm_provider": "openai",

View file

@ -156,6 +156,7 @@ model LiteLLM_ObjectPermissionTable {
object_permission_id String @id @default(uuid()) object_permission_id String @id @default(uuid())
mcp_servers String[] @default([]) mcp_servers String[] @default([])
mcp_access_groups 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([]) vector_stores String[] @default([])
teams LiteLLM_TeamTable[] teams LiteLLM_TeamTable[]
verification_tokens LiteLLM_VerificationToken[] verification_tokens LiteLLM_VerificationToken[]

View file

@ -19,6 +19,7 @@ from litellm.types.llms.openai import (
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from base_responses_api import BaseResponsesAPITest from base_responses_api import BaseResponsesAPITest
class TestAzureResponsesAPITest(BaseResponsesAPITest): class TestAzureResponsesAPITest(BaseResponsesAPITest):
def get_base_completion_call_args(self): def get_base_completion_call_args(self):
return { return {
@ -43,4 +44,55 @@ async def test_azure_responses_api_preview_api_version():
api_base=os.getenv("AZURE_RESPONSES_OPENAI_ENDPOINT"), api_base=os.getenv("AZURE_RESPONSES_OPENAI_ENDPOINT"),
api_key=os.getenv("AZURE_RESPONSES_OPENAI_API_KEY"), api_key=os.getenv("AZURE_RESPONSES_OPENAI_API_KEY"),
input="Hello, can you tell me a short joke?", 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"],
)

View file

@ -6,7 +6,7 @@ from dotenv import load_dotenv
load_dotenv() load_dotenv()
import pytest import pytest
from litellm import completion, acompletion from litellm import completion, acompletion, responses
from litellm.exceptions import APIConnectionError from litellm.exceptions import APIConnectionError
@pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.parametrize("sync_mode", [True, False])
@ -87,3 +87,70 @@ async def test_chat_completion_snowflake_stream(sync_mode):
raise # Re-raise if it's a different APIConnectionError raise # Re-raise if it's a different APIConnectionError
except Exception as e: except Exception as e:
pytest.fail(f"Error occurred: {e}") pytest.fail(f"Error occurred: {e}")
@pytest.mark.skip(reason="Requires Snowflake credentials - run manually when needed")
def test_snowflake_tool_calling_responses_api():
"""
Test Snowflake tool calling with Responses API.
Requires SNOWFLAKE_JWT and SNOWFLAKE_ACCOUNT_ID environment variables.
"""
import litellm
# Skip if credentials not available
if not os.getenv("SNOWFLAKE_JWT") or not os.getenv("SNOWFLAKE_ACCOUNT_ID"):
pytest.skip("Snowflake credentials not available")
litellm.drop_params = False # We now support tools!
tools = [
{
"type": "function",
"name": "get_weather",
"description": "Get the current weather in a given location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state, e.g. San Francisco, CA",
}
},
"required": ["location"],
},
}
]
try:
# Test with tool_choice to force tool use
response = responses(
model="snowflake/claude-3-5-sonnet",
input="What's the weather in Paris?",
tools=tools,
tool_choice={"type": "function", "function": {"name": "get_weather"}},
max_output_tokens=200,
)
assert response is not None
assert hasattr(response, "output")
assert len(response.output) > 0
# Verify tool call was made
tool_call_found = False
for item in response.output:
if hasattr(item, "type") and item.type == "function_call":
tool_call_found = True
assert item.name == "get_weather"
assert hasattr(item, "arguments")
print(f"✅ Tool call detected: {item.name}({item.arguments})")
break
assert tool_call_found, "Expected tool call but none was found"
except APIConnectionError as e:
if "JWT token is invalid" in str(e):
pytest.skip("Invalid Snowflake JWT token")
elif "Application failed to respond" in str(e) or "502" in str(e):
pytest.skip(f"Snowflake API unavailable: {e}")
else:
raise

View file

@ -61,6 +61,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
list_keys, list_keys,
regenerate_key_fn, regenerate_key_fn,
update_key_fn, update_key_fn,
key_aliases,
) )
from litellm.proxy.management_endpoints.team_endpoints import ( from litellm.proxy.management_endpoints.team_endpoints import (
new_team, new_team,
@ -151,7 +152,6 @@ def prisma_client():
@pytest.mark.flaky(retries=6, delay=1) @pytest.mark.flaky(retries=6, delay=1)
async def test_new_user_response(prisma_client): async def test_new_user_response(prisma_client):
try: try:
print("prisma client=", prisma_client) print("prisma client=", prisma_client)
setattr(litellm.proxy.proxy_server, "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, "prisma_client", prisma_client)
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
try: try:
await litellm.proxy.proxy_server.prisma_client.connect() await litellm.proxy.proxy_server.prisma_client.connect()
team_request = NewTeamRequest( team_request = NewTeamRequest(
@ -1789,7 +1788,6 @@ async def test_call_with_key_over_model_budget(
litellm.callbacks.append(model_budget_limiter) litellm.callbacks.append(model_budget_limiter)
try: try:
# set budget for chatgpt-v-3 to 0.000001, expect the next request to fail # set budget for chatgpt-v-3 to 0.000001, expect the next request to fail
model_max_budget = { model_max_budget = {
"gpt-4o-mini": { "gpt-4o-mini": {
@ -3531,6 +3529,58 @@ async def test_list_keys(prisma_client):
assert _key in response["keys"] 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 @pytest.mark.asyncio
async def test_auth_vertex_ai_route(prisma_client): async def test_auth_vertex_ai_route(prisma_client):
""" """

View file

@ -77,7 +77,6 @@ class TestRouterIndexManagement:
# Verify: Index map uses model_info.id # Verify: Index map uses model_info.id
assert router.model_id_to_deployment_index_map["model-info-id"] == 0 assert router.model_id_to_deployment_index_map["model-info-id"] == 0
def test_add_model_to_list_and_index_map_multiple_models(self, router): def test_add_model_to_list_and_index_map_multiple_models(self, router):
"""Test _add_model_to_list_and_index_map with multiple models to verify indexing""" """Test _add_model_to_list_and_index_map with multiple models to verify indexing"""
# Setup: Empty router # Setup: Empty router
@ -127,3 +126,54 @@ class TestRouterIndexManagement:
# Test: Empty router # Test: Empty router
empty_router = Router(model_list=[]) empty_router = Router(model_list=[])
assert empty_router.has_model_id("any-id") == False assert empty_router.has_model_id("any-id") == False
def test_build_model_name_index(self, router):
"""Test _build_model_name_index function"""
model_list = [
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {"model": "gpt-3.5-turbo"},
"model_info": {"id": "model-1"},
},
{
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4"},
"model_info": {"id": "model-2"},
},
{
"model_name": "gpt-4", # Duplicate model_name, different deployment
"litellm_params": {"model": "gpt-4"},
"model_info": {"id": "model-3"},
},
]
# Test: Build index from model list
router._build_model_name_index(model_list)
# Verify: model_name_to_deployment_indices is correctly built
assert "gpt-3.5-turbo" in router.model_name_to_deployment_indices
assert "gpt-4" in router.model_name_to_deployment_indices
# Verify: gpt-3.5-turbo has single deployment
assert router.model_name_to_deployment_indices["gpt-3.5-turbo"] == [0]
# Verify: gpt-4 has multiple deployments
assert router.model_name_to_deployment_indices["gpt-4"] == [1, 2]
# Test: Rebuild index (should clear and rebuild)
new_model_list = [
{
"model_name": "claude-3",
"litellm_params": {"model": "claude-3"},
"model_info": {"id": "model-4"},
},
]
router._build_model_name_index(new_model_list)
# Verify: Old entries are cleared
assert "gpt-3.5-turbo" not in router.model_name_to_deployment_indices
assert "gpt-4" not in router.model_name_to_deployment_indices
# Verify: New entry is added
assert "claude-3" in router.model_name_to_deployment_indices
assert router.model_name_to_deployment_indices["claude-3"] == [0]

View file

@ -1,10 +1,13 @@
import os
import ssl
import sys import sys
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest import pytest
# Add the parent directory to the path so we can import litellm # 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.experimental_mcp_client.client import MCPClient
from litellm.types.mcp import MCPStdioConfig, MCPTransport from litellm.types.mcp import MCPStdioConfig, MCPTransport
@ -16,57 +19,54 @@ class TestMCPClient:
def test_mcp_client_stdio_init(self): def test_mcp_client_stdio_init(self):
"""Test MCPClient initialization with stdio config""" """Test MCPClient initialization with stdio config"""
stdio_config = MCPStdioConfig( stdio_config = MCPStdioConfig(
command="python", command="python", args=["-m", "my_mcp_server"], env={"DEBUG": "1"}
args=["-m", "my_mcp_server"],
env={"DEBUG": "1"}
) )
client = MCPClient( client = MCPClient(transport_type=MCPTransport.stdio, stdio_config=stdio_config)
transport_type=MCPTransport.stdio,
stdio_config=stdio_config
)
assert client.transport_type == MCPTransport.stdio assert client.transport_type == MCPTransport.stdio
assert client.stdio_config == stdio_config assert client.stdio_config == stdio_config
assert client.stdio_config["command"] == "python" assert client.stdio_config is not None
assert client.stdio_config["args"] == ["-m", "my_mcp_server"] assert client.stdio_config.get("command") == "python"
assert client.stdio_config.get("args") == ["-m", "my_mcp_server"]
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_mcp_client_stdio_connect_error(self): async def test_mcp_client_stdio_connect_error(self):
"""Test MCP client stdio connection error handling""" """Test MCP client stdio connection error handling"""
# Test missing stdio_config # Test missing stdio_config
client = MCPClient(transport_type=MCPTransport.stdio) 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() await client.connect()
@pytest.mark.asyncio @pytest.mark.asyncio
@patch('litellm.experimental_mcp_client.client.stdio_client') @patch("litellm.experimental_mcp_client.client.stdio_client")
@patch('litellm.experimental_mcp_client.client.ClientSession') @patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_mcp_client_stdio_connect_success(self, mock_session, mock_stdio_client): async def test_mcp_client_stdio_connect_success(
self, mock_session, mock_stdio_client
):
"""Test successful stdio connection""" """Test successful stdio connection"""
# Setup mocks # Setup mocks
mock_transport = (MagicMock(), MagicMock()) 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 = MagicMock()
mock_session_instance.__aenter__ = AsyncMock(return_value=mock_session_instance) mock_session_instance.__aenter__ = AsyncMock(return_value=mock_session_instance)
mock_session_instance.initialize = AsyncMock() mock_session_instance.initialize = AsyncMock()
mock_session.return_value = mock_session_instance mock_session.return_value = mock_session_instance
stdio_config = MCPStdioConfig( stdio_config = MCPStdioConfig(
command="python", command="python", args=["-m", "my_mcp_server"], env={"DEBUG": "1"}
args=["-m", "my_mcp_server"],
env={"DEBUG": "1"}
) )
client = MCPClient( client = MCPClient(transport_type=MCPTransport.stdio, stdio_config=stdio_config)
transport_type=MCPTransport.stdio,
stdio_config=stdio_config
)
await client.connect() await client.connect()
# Verify stdio_client was called with correct parameters # Verify stdio_client was called with correct parameters
mock_stdio_client.assert_called_once() mock_stdio_client.assert_called_once()
call_args = mock_stdio_client.call_args[0][0] 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.args == ["-m", "my_mcp_server"]
assert call_args.env == {"DEBUG": "1"} 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__": if __name__ == "__main__":
pytest.main([__file__]) pytest.main([__file__])

View file

@ -43,3 +43,25 @@ def test_azure_ai_validate_environment():
litellm_params={}, litellm_params={},
) )
assert headers["Content-Type"] == "application/json" 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"

View file

@ -0,0 +1,315 @@
"""
Unit tests for Snowflake chat transformation
Tests tool calling request/response transformations
"""
import json
from unittest.mock import MagicMock
import httpx
import pytest
import litellm
from litellm.llms.snowflake.chat.transformation import SnowflakeConfig
from litellm.types.utils import ModelResponse
class TestSnowflakeToolTransformation:
"""Test suite for Snowflake tool calling transformations"""
def test_transform_request_with_tools(self):
"""
Test that OpenAI tool format is correctly transformed to Snowflake's tool_spec format.
"""
config = SnowflakeConfig()
# OpenAI format tools
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather in a given location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state, e.g. San Francisco, CA",
},
"unit": {
"type": "string",
"enum": ["celsius", "fahrenheit"],
},
},
"required": ["location"],
},
},
}
]
optional_params = {"tools": tools}
transformed_request = config.transform_request(
model="claude-3-5-sonnet",
messages=[{"role": "user", "content": "What's the weather?"}],
optional_params=optional_params,
litellm_params={},
headers={},
)
# Verify tools were transformed to Snowflake format
assert "tools" in transformed_request
assert len(transformed_request["tools"]) == 1
snowflake_tool = transformed_request["tools"][0]
assert "tool_spec" in snowflake_tool
assert snowflake_tool["tool_spec"]["type"] == "generic"
assert snowflake_tool["tool_spec"]["name"] == "get_weather"
assert snowflake_tool["tool_spec"]["description"] == "Get the current weather in a given location"
assert "input_schema" in snowflake_tool["tool_spec"]
assert snowflake_tool["tool_spec"]["input_schema"]["type"] == "object"
assert "location" in snowflake_tool["tool_spec"]["input_schema"]["properties"]
def test_transform_request_with_tool_choice(self):
"""
Test that OpenAI tool_choice format is correctly transformed to Snowflake format.
"""
config = SnowflakeConfig()
# OpenAI format tool_choice
tool_choice = {"type": "function", "function": {"name": "get_weather"}}
optional_params = {"tool_choice": tool_choice}
transformed_request = config.transform_request(
model="claude-3-5-sonnet",
messages=[{"role": "user", "content": "What's the weather?"}],
optional_params=optional_params,
litellm_params={},
headers={},
)
# Verify tool_choice was transformed to Snowflake format
assert "tool_choice" in transformed_request
assert transformed_request["tool_choice"]["type"] == "tool"
assert transformed_request["tool_choice"]["name"] == ["get_weather"] # Array format
def test_transform_request_with_string_tool_choice(self):
"""
Test that string tool_choice values pass through unchanged.
"""
config = SnowflakeConfig()
for value in ["auto", "required", "none"]:
optional_params = {"tool_choice": value}
transformed_request = config.transform_request(
model="claude-3-5-sonnet",
messages=[{"role": "user", "content": "Test"}],
optional_params=optional_params,
litellm_params={},
headers={},
)
assert transformed_request["tool_choice"] == value
def test_transform_response_with_tool_calls(self):
"""
Test that Snowflake's content_list with tool_use is transformed to OpenAI format.
"""
config = SnowflakeConfig()
# Mock Snowflake response with tool call
mock_snowflake_response = {
"choices": [
{
"message": {
"content_list": [
{"type": "text", "text": ""},
{
"type": "tool_use",
"tool_use": {
"tool_use_id": "tooluse_abc123",
"name": "get_weather",
"input": {"location": "Paris, France", "unit": "celsius"},
},
},
]
}
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
}
response = httpx.Response(
status_code=200,
json=mock_snowflake_response,
headers={"Content-Type": "application/json"},
)
model_response = ModelResponse(
choices=[litellm.Choices(index=0, message=litellm.Message())]
)
logging_obj = MagicMock()
result = config.transform_response(
model="claude-3-5-sonnet",
raw_response=response,
model_response=model_response,
logging_obj=logging_obj,
request_data={},
messages=[],
optional_params={},
litellm_params={},
encoding={},
)
# General assertions
assert isinstance(result, ModelResponse)
assert len(result.choices) == 1
choice = result.choices[0]
assert isinstance(choice, litellm.Choices)
# Message and tool_calls assertions
message = choice.message
assert isinstance(message, litellm.Message)
assert hasattr(message, "tool_calls")
assert isinstance(message.tool_calls, list)
assert len(message.tool_calls) == 1
# Specific tool_call assertions
tool_call = message.tool_calls[0]
assert isinstance(tool_call, litellm.utils.ChatCompletionMessageToolCall)
assert tool_call.id == "tooluse_abc123"
assert tool_call.type == "function"
assert tool_call.function.name == "get_weather"
# Verify arguments are properly JSON serialized
arguments = json.loads(tool_call.function.arguments)
assert arguments["location"] == "Paris, France"
assert arguments["unit"] == "celsius"
# Verify content_list was removed and content was set
assert message.content == ""
def test_transform_response_with_mixed_content(self):
"""
Test that responses with both text and tool calls are handled correctly.
"""
config = SnowflakeConfig()
# Mock Snowflake response with text and tool call
mock_snowflake_response = {
"choices": [
{
"message": {
"content_list": [
{"type": "text", "text": "Let me check the weather for you. "},
{
"type": "tool_use",
"tool_use": {
"tool_use_id": "tooluse_xyz789",
"name": "get_weather",
"input": {"location": "Tokyo, Japan"},
},
},
]
}
}
],
"usage": {"prompt_tokens": 15, "completion_tokens": 25, "total_tokens": 40},
}
response = httpx.Response(
status_code=200,
json=mock_snowflake_response,
headers={"Content-Type": "application/json"},
)
model_response = ModelResponse(
choices=[litellm.Choices(index=0, message=litellm.Message())]
)
logging_obj = MagicMock()
result = config.transform_response(
model="claude-3-5-sonnet",
raw_response=response,
model_response=model_response,
logging_obj=logging_obj,
request_data={},
messages=[],
optional_params={},
litellm_params={},
encoding={},
)
# Verify text content was extracted
message = result.choices[0].message
assert message.content == "Let me check the weather for you. "
# Verify tool call was also extracted
assert len(message.tool_calls) == 1
assert message.tool_calls[0].function.name == "get_weather"
def test_transform_response_without_tool_calls(self):
"""
Test that regular text responses (without tools) work correctly.
"""
config = SnowflakeConfig()
# Mock Snowflake response without tool calls (standard response)
mock_snowflake_response = {
"choices": [
{
"message": {
"content": "Hello! I'm doing well, thank you for asking.",
"role": "assistant",
}
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 15, "total_tokens": 25},
}
response = httpx.Response(
status_code=200,
json=mock_snowflake_response,
headers={"Content-Type": "application/json"},
)
model_response = ModelResponse(
choices=[litellm.Choices(index=0, message=litellm.Message())]
)
logging_obj = MagicMock()
result = config.transform_response(
model="mistral-7b",
raw_response=response,
model_response=model_response,
logging_obj=logging_obj,
request_data={},
messages=[],
optional_params={},
litellm_params={},
encoding={},
)
# Verify standard response works
assert isinstance(result, ModelResponse)
assert result.choices[0].message.content == "Hello! I'm doing well, thank you for asking."
def test_get_supported_openai_params_includes_tools(self):
"""
Test that tools and tool_choice are in supported params.
"""
config = SnowflakeConfig()
supported_params = config.get_supported_openai_params("claude-3-5-sonnet")
assert "tools" in supported_params
assert "tool_choice" in supported_params
assert "temperature" in supported_params
assert "max_tokens" in supported_params

View file

@ -191,7 +191,8 @@ class TestTTLExtraction:
class TestTransformationWithTTL: class TestTransformationWithTTL:
"""Test the complete transformation with TTL support""" """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""" """Test transformation includes TTL when provided"""
messages = [ messages = [
{ {
@ -205,19 +206,32 @@ class TestTransformationWithTTL:
] ]
} }
] ]
vertex_location="test_location"
vertex_project="test_project"
result = transform_openai_messages_to_gemini_context_caching( result = transform_openai_messages_to_gemini_context_caching(
model="gemini-1.5-pro", model="gemini-1.5-pro",
messages=messages, 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 "ttl" in result
assert result["ttl"] == "3600s" 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" 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""" """Test transformation without TTL"""
messages = [ messages = [
{ {
@ -231,18 +245,30 @@ class TestTransformationWithTTL:
] ]
} }
] ]
vertex_location="test_location"
vertex_project="test_project"
result = transform_openai_messages_to_gemini_context_caching( result = transform_openai_messages_to_gemini_context_caching(
model="gemini-1.5-pro", model="gemini-1.5-pro",
messages=messages, 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 "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" 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)""" """Test transformation with invalid TTL (should be ignored)"""
messages = [ messages = [
{ {
@ -256,18 +282,29 @@ class TestTransformationWithTTL:
] ]
} }
] ]
vertex_location="test_location"
vertex_project="test_project"
result = transform_openai_messages_to_gemini_context_caching( result = transform_openai_messages_to_gemini_context_caching(
model="gemini-1.5-pro", model="gemini-1.5-pro",
messages=messages, 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 "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" 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""" """Test transformation with system message and TTL"""
messages = [ messages = [
{ {
@ -290,17 +327,28 @@ class TestTransformationWithTTL:
] ]
} }
] ]
vertex_location="test_location"
vertex_project="test_project"
result = transform_openai_messages_to_gemini_context_caching( result = transform_openai_messages_to_gemini_context_caching(
model="gemini-1.5-pro", model="gemini-1.5-pro",
messages=messages, 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 "ttl" in result
assert result["ttl"] == "7200s" assert result["ttl"] == "7200s"
assert "system_instruction" in result 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" assert result["displayName"] == "test-cache-key"

View file

@ -55,6 +55,9 @@ class TestContextCachingEndpoints:
self.sample_optional_params = {"tools": self.sample_tools.copy()} self.sample_optional_params = {"tools": self.sample_tools.copy()}
@pytest.mark.parametrize(
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
)
@patch( @patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "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" "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj"
) )
def test_check_and_create_cache_with_cached_content( 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""" """Test check_and_create_cache when cached_content is provided"""
# Setup # Setup
cached_content = "cached_content_123" cached_content = "cached_content_123"
optional_params = self.sample_optional_params.copy() optional_params = self.sample_optional_params.copy()
test_project = "test_project"
test_location = "test_location"
# Execute # Execute
result = self.context_caching.check_and_create_cache( result = self.context_caching.check_and_create_cache(
@ -80,6 +85,10 @@ class TestContextCachingEndpoints:
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, logging_obj=self.mock_logging,
cached_content=cached_content, 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 # Assert
@ -92,14 +101,21 @@ class TestContextCachingEndpoints:
mock_separate.assert_not_called() mock_separate.assert_not_called()
mock_cache_obj.get_cache_key.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( @patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "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""" """Test check_and_create_cache when no cached messages are found"""
# Setup # Setup
mock_separate.return_value = ([], self.sample_messages) # No cached messages mock_separate.return_value = ([], self.sample_messages) # No cached messages
optional_params = self.sample_optional_params.copy() optional_params = self.sample_optional_params.copy()
test_project = "test_project"
test_location = "test_location"
# Execute # Execute
result = self.context_caching.check_and_create_cache( result = self.context_caching.check_and_create_cache(
@ -111,6 +127,10 @@ class TestContextCachingEndpoints:
client=self.mock_client, client=self.mock_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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 # Assert
@ -119,6 +139,9 @@ class TestContextCachingEndpoints:
assert returned_params == optional_params assert returned_params == optional_params
assert returned_cache is None assert returned_cache is None
@pytest.mark.parametrize(
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
)
@patch( @patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
) )
@ -127,7 +150,7 @@ class TestContextCachingEndpoints:
) )
@patch.object(ContextCachingEndpoints, "check_cache") @patch.object(ContextCachingEndpoints, "check_cache")
def test_check_and_create_cache_existing_cache_found( 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""" """Test check_and_create_cache when existing cache is found"""
# Setup # Setup
@ -139,6 +162,8 @@ class TestContextCachingEndpoints:
mock_check_cache.return_value = "existing_cache_name" mock_check_cache.return_value = "existing_cache_name"
optional_params = self.sample_optional_params.copy() optional_params = self.sample_optional_params.copy()
test_project = "test_project"
test_location = "test_location"
# Execute # Execute
result = self.context_caching.check_and_create_cache( result = self.context_caching.check_and_create_cache(
@ -150,6 +175,10 @@ class TestContextCachingEndpoints:
client=self.mock_client, client=self.mock_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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 # Assert
@ -163,6 +192,9 @@ class TestContextCachingEndpoints:
messages=cached_messages, tools=self.sample_tools messages=cached_messages, tools=self.sample_tools
) )
@pytest.mark.parametrize(
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
)
@patch( @patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
) )
@ -181,6 +213,7 @@ class TestContextCachingEndpoints:
mock_transform, mock_transform,
mock_cache_obj, mock_cache_obj,
mock_separate, mock_separate,
custom_llm_provider,
): ):
"""Test check_and_create_cache when creating new cache""" """Test check_and_create_cache when creating new cache"""
# Setup # Setup
@ -203,6 +236,8 @@ class TestContextCachingEndpoints:
self.mock_client.post.return_value = mock_response self.mock_client.post.return_value = mock_response
optional_params = self.sample_optional_params.copy() optional_params = self.sample_optional_params.copy()
test_project = "test_project"
test_location = "test_location"
# Execute # Execute
result = self.context_caching.check_and_create_cache( result = self.context_caching.check_and_create_cache(
@ -214,6 +249,10 @@ class TestContextCachingEndpoints:
client=self.mock_client, client=self.mock_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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 # Assert
@ -228,6 +267,9 @@ class TestContextCachingEndpoints:
assert "tools" in call_args.kwargs["json"] assert "tools" in call_args.kwargs["json"]
assert call_args.kwargs["json"]["tools"] == self.sample_tools assert call_args.kwargs["json"]["tools"] == self.sample_tools
@pytest.mark.parametrize(
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
)
@patch( @patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "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, "check_cache")
@patch.object(ContextCachingEndpoints, "_get_token_and_url_context_caching") @patch.object(ContextCachingEndpoints, "_get_token_and_url_context_caching")
def test_check_and_create_cache_http_error( 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""" """Test check_and_create_cache handles HTTP errors properly"""
# Setup # Setup
@ -259,6 +306,8 @@ class TestContextCachingEndpoints:
self.mock_client.post.side_effect = http_error self.mock_client.post.side_effect = http_error
optional_params = self.sample_optional_params.copy() optional_params = self.sample_optional_params.copy()
test_project = "test_project"
test_location = "test_location"
# Execute and Assert # Execute and Assert
with pytest.raises(VertexAIError) as exc_info: with pytest.raises(VertexAIError) as exc_info:
@ -271,12 +320,19 @@ class TestContextCachingEndpoints:
client=self.mock_client, client=self.mock_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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 exc_info.value.status_code == 400
assert "Bad Request" in str(exc_info.value.message) assert "Bad Request" in str(exc_info.value.message)
@pytest.mark.asyncio @pytest.mark.asyncio
@pytest.mark.parametrize(
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
)
@patch( @patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "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" "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj"
) )
async def test_async_check_and_create_cache_with_cached_content( 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""" """Test async_check_and_create_cache when cached_content is provided"""
# Setup # Setup
cached_content = "cached_content_123" cached_content = "cached_content_123"
optional_params = self.sample_optional_params.copy() optional_params = self.sample_optional_params.copy()
test_project = "test_project"
test_location = "test_location"
# Execute # Execute
result = await self.context_caching.async_check_and_create_cache( result = await self.context_caching.async_check_and_create_cache(
@ -302,6 +360,10 @@ class TestContextCachingEndpoints:
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, logging_obj=self.mock_logging,
cached_content=cached_content, 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 # Assert
@ -311,14 +373,21 @@ class TestContextCachingEndpoints:
assert returned_cache == cached_content assert returned_cache == cached_content
@pytest.mark.asyncio @pytest.mark.asyncio
@pytest.mark.parametrize(
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
)
@patch( @patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "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""" """Test async_check_and_create_cache when no cached messages are found"""
# Setup # Setup
mock_separate.return_value = ([], self.sample_messages) mock_separate.return_value = ([], self.sample_messages)
optional_params = self.sample_optional_params.copy() optional_params = self.sample_optional_params.copy()
test_project = "test_project"
test_location = "test_location"
# Execute # Execute
result = await self.context_caching.async_check_and_create_cache( result = await self.context_caching.async_check_and_create_cache(
@ -330,6 +399,10 @@ class TestContextCachingEndpoints:
client=self.mock_async_client, client=self.mock_async_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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 # Assert
@ -339,6 +412,9 @@ class TestContextCachingEndpoints:
assert returned_cache is None assert returned_cache is None
@pytest.mark.asyncio @pytest.mark.asyncio
@pytest.mark.parametrize(
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
)
@patch( @patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "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") @patch.object(ContextCachingEndpoints, "async_check_cache")
async def test_async_check_and_create_cache_existing_cache_found( 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""" """Test async_check_and_create_cache when existing cache is found"""
# Setup # Setup
@ -359,6 +435,8 @@ class TestContextCachingEndpoints:
mock_async_check_cache.return_value = "existing_cache_name" mock_async_check_cache.return_value = "existing_cache_name"
optional_params = self.sample_optional_params.copy() optional_params = self.sample_optional_params.copy()
test_project = "test_project"
test_location = "test_location"
# Execute # Execute
result = await self.context_caching.async_check_and_create_cache( result = await self.context_caching.async_check_and_create_cache(
@ -370,6 +448,10 @@ class TestContextCachingEndpoints:
client=self.mock_async_client, client=self.mock_async_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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 # Assert
@ -384,6 +466,9 @@ class TestContextCachingEndpoints:
) )
@pytest.mark.asyncio @pytest.mark.asyncio
@pytest.mark.parametrize(
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
)
@patch( @patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
) )
@ -406,6 +491,7 @@ class TestContextCachingEndpoints:
mock_transform, mock_transform,
mock_cache_obj, mock_cache_obj,
mock_separate, mock_separate,
custom_llm_provider,
): ):
"""Test async_check_and_create_cache when creating new cache""" """Test async_check_and_create_cache when creating new cache"""
# Setup # Setup
@ -428,6 +514,8 @@ class TestContextCachingEndpoints:
self.mock_async_client.post = AsyncMock(return_value=mock_response) self.mock_async_client.post = AsyncMock(return_value=mock_response)
optional_params = self.sample_optional_params.copy() optional_params = self.sample_optional_params.copy()
test_project = "test_project"
test_location = "test_location"
# Execute # Execute
result = await self.context_caching.async_check_and_create_cache( result = await self.context_caching.async_check_and_create_cache(
@ -439,6 +527,10 @@ class TestContextCachingEndpoints:
client=self.mock_async_client, client=self.mock_async_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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 # Assert
@ -454,6 +546,9 @@ class TestContextCachingEndpoints:
assert call_args.kwargs["json"]["tools"] == self.sample_tools assert call_args.kwargs["json"]["tools"] == self.sample_tools
@pytest.mark.asyncio @pytest.mark.asyncio
@pytest.mark.parametrize(
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
)
@patch( @patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
) )
@ -472,6 +567,7 @@ class TestContextCachingEndpoints:
mock_async_check_cache, mock_async_check_cache,
mock_cache_obj, mock_cache_obj,
mock_separate, mock_separate,
custom_llm_provider,
): ):
"""Test async_check_and_create_cache handles timeout errors properly""" """Test async_check_and_create_cache handles timeout errors properly"""
# Setup # Setup
@ -489,6 +585,8 @@ class TestContextCachingEndpoints:
) )
optional_params = self.sample_optional_params.copy() optional_params = self.sample_optional_params.copy()
test_project = "test_project"
test_location = "test_location"
# Execute and Assert # Execute and Assert
with pytest.raises(VertexAIError) as exc_info: with pytest.raises(VertexAIError) as exc_info:
@ -501,12 +599,21 @@ class TestContextCachingEndpoints:
client=self.mock_async_client, client=self.mock_async_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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 exc_info.value.status_code == 408
assert "Timeout error occurred" in str(exc_info.value.message) 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""" """Test that tools are properly popped from optional_params when there are cached messages"""
with patch( with patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "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() optional_params = self.sample_optional_params.copy()
original_tools = optional_params["tools"].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 # Mock the check_cache to return existing cache so we don't make HTTP calls
with patch.object( with patch.object(
@ -535,6 +644,10 @@ class TestContextCachingEndpoints:
client=self.mock_client, client=self.mock_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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 # Assert tools were popped from optional_params
@ -543,7 +656,12 @@ class TestContextCachingEndpoints:
# But original tools should still be available for comparison # But original tools should still be available for comparison
assert original_tools == self.sample_tools 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""" """Test that tools are NOT popped from optional_params when there are no cached messages"""
with patch( with patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "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() optional_params = self.sample_optional_params.copy()
original_tools = optional_params["tools"].copy() original_tools = optional_params["tools"].copy()
test_project = "test_project"
test_location = "test_location"
# Execute # Execute
result = self.context_caching.check_and_create_cache( result = self.context_caching.check_and_create_cache(
@ -566,6 +686,10 @@ class TestContextCachingEndpoints:
client=self.mock_client, client=self.mock_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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) # Assert tools were NOT popped from optional_params (early return)
@ -573,8 +697,11 @@ class TestContextCachingEndpoints:
assert optional_params["tools"] == original_tools assert optional_params["tools"] == original_tools
@pytest.mark.asyncio @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( 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""" """Test that tools are NOT popped from optional_params in async version when there are no cached messages"""
with patch( with patch(
@ -587,6 +714,8 @@ class TestContextCachingEndpoints:
optional_params = self.sample_optional_params.copy() optional_params = self.sample_optional_params.copy()
original_tools = optional_params["tools"].copy() original_tools = optional_params["tools"].copy()
test_project = "test_project"
test_location = "test_location"
# Execute # Execute
result = await self.context_caching.async_check_and_create_cache( result = await self.context_caching.async_check_and_create_cache(
@ -598,6 +727,10 @@ class TestContextCachingEndpoints:
client=self.mock_async_client, client=self.mock_async_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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) # Assert tools were NOT popped from optional_params (early return)
@ -605,7 +738,12 @@ class TestContextCachingEndpoints:
assert optional_params["tools"] == original_tools assert optional_params["tools"] == original_tools
@pytest.mark.asyncio @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""" """Test that tools are properly popped from optional_params in async version when there are cached messages"""
with patch( with patch(
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" "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() optional_params = self.sample_optional_params.copy()
original_tools = optional_params["tools"].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 # Mock the async_check_cache to return existing cache so we don't make HTTP calls
with patch.object( with patch.object(
@ -634,6 +774,10 @@ class TestContextCachingEndpoints:
client=self.mock_async_client, client=self.mock_async_client,
timeout=30.0, timeout=30.0,
logging_obj=self.mock_logging, 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 # Assert tools were popped from optional_params

View file

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

View file

@ -53,27 +53,38 @@ def test_llm_passthrough_route():
def test_bedrock_application_inference_profile_url_encoding(): def test_bedrock_application_inference_profile_url_encoding():
client = HTTPHandler() client = HTTPHandler()
mock_provider_config = MagicMock() mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = ( mock_provider_config.get_complete_url.return_value = (
httpx.URL("https://bedrock-runtime.us-east-1.amazonaws.com/model/arn:aws:bedrock:us-east-1:123456789123:application-inference-profile/r742sbn2zckd/converse"), httpx.URL(
"https://bedrock-runtime.us-east-1.amazonaws.com" "https://bedrock-runtime.us-east-1.amazonaws.com/model/arn:aws:bedrock:us-east-1:123456789123:application-inference-profile/r742sbn2zckd/converse"
),
"https://bedrock-runtime.us-east-1.amazonaws.com",
) )
mock_provider_config.get_api_key.return_value = "test-key" mock_provider_config.get_api_key.return_value = "test-key"
mock_provider_config.validate_environment.return_value = {} mock_provider_config.validate_environment.return_value = {}
mock_provider_config.sign_request.return_value = ({}, None) mock_provider_config.sign_request.return_value = ({}, None)
mock_provider_config.is_streaming_request.return_value = False mock_provider_config.is_streaming_request.return_value = False
with patch("litellm.utils.ProviderConfigManager.get_provider_passthrough_config", return_value=mock_provider_config), \ with patch(
patch("litellm.litellm_core_utils.get_litellm_params.get_litellm_params", return_value={}), \ "litellm.utils.ProviderConfigManager.get_provider_passthrough_config",
patch("litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider", return_value=("test-model", "bedrock", "test-key", "test-base")), \ return_value=mock_provider_config,
patch.object(client.client, "send", return_value=MagicMock(status_code=200)) as mock_send, \ ), patch(
patch.object(client.client, "build_request") as mock_build_request: "litellm.litellm_core_utils.get_litellm_params.get_litellm_params",
return_value={},
), patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("test-model", "bedrock", "test-key", "test-base"),
), patch.object(
client.client, "send", return_value=MagicMock(status_code=200)
) as mock_send, patch.object(
client.client, "build_request"
) as mock_build_request:
# Mock logging object # Mock logging object
mock_logging_obj = MagicMock() mock_logging_obj = MagicMock()
mock_logging_obj.update_environment_variables = MagicMock() mock_logging_obj.update_environment_variables = MagicMock()
response = llm_passthrough_route( response = llm_passthrough_route(
model="arn:aws:bedrock:us-east-1:123456789123:application-inference-profile/r742sbn2zckd", model="arn:aws:bedrock:us-east-1:123456789123:application-inference-profile/r742sbn2zckd",
endpoint="model/arn:aws:bedrock:us-east-1:123456789123:application-inference-profile/r742sbn2zckd/converse", endpoint="model/arn:aws:bedrock:us-east-1:123456789123:application-inference-profile/r742sbn2zckd/converse",
@ -86,7 +97,7 @@ def test_bedrock_application_inference_profile_url_encoding():
# Verify that build_request was called with the encoded URL # Verify that build_request was called with the encoded URL
mock_build_request.assert_called_once() mock_build_request.assert_called_once()
call_args = mock_build_request.call_args call_args = mock_build_request.call_args
# The URL should have the application-inference-profile ID encoded # The URL should have the application-inference-profile ID encoded
actual_url = str(call_args.kwargs["url"]) actual_url = str(call_args.kwargs["url"])
assert "application-inference-profile%2Fr742sbn2zckd" in actual_url assert "application-inference-profile%2Fr742sbn2zckd" in actual_url
@ -95,28 +106,39 @@ def test_bedrock_application_inference_profile_url_encoding():
def test_bedrock_non_application_inference_profile_no_encoding(): def test_bedrock_non_application_inference_profile_no_encoding():
client = HTTPHandler() client = HTTPHandler()
# Mock the provider config and its methods # Mock the provider config and its methods
mock_provider_config = MagicMock() mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = ( mock_provider_config.get_complete_url.return_value = (
httpx.URL("https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-3-sonnet-20240229-v1:0/converse"), httpx.URL(
"https://bedrock-runtime.us-east-1.amazonaws.com" "https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-3-sonnet-20240229-v1:0/converse"
),
"https://bedrock-runtime.us-east-1.amazonaws.com",
) )
mock_provider_config.get_api_key.return_value = "test-key" mock_provider_config.get_api_key.return_value = "test-key"
mock_provider_config.validate_environment.return_value = {} mock_provider_config.validate_environment.return_value = {}
mock_provider_config.sign_request.return_value = ({}, None) mock_provider_config.sign_request.return_value = ({}, None)
mock_provider_config.is_streaming_request.return_value = False mock_provider_config.is_streaming_request.return_value = False
with patch("litellm.utils.ProviderConfigManager.get_provider_passthrough_config", return_value=mock_provider_config), \ with patch(
patch("litellm.litellm_core_utils.get_litellm_params.get_litellm_params", return_value={}), \ "litellm.utils.ProviderConfigManager.get_provider_passthrough_config",
patch("litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider", return_value=("test-model", "bedrock", "test-key", "test-base")), \ return_value=mock_provider_config,
patch.object(client.client, "send", return_value=MagicMock(status_code=200)) as mock_send, \ ), patch(
patch.object(client.client, "build_request") as mock_build_request: "litellm.litellm_core_utils.get_litellm_params.get_litellm_params",
return_value={},
), patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("test-model", "bedrock", "test-key", "test-base"),
), patch.object(
client.client, "send", return_value=MagicMock(status_code=200)
) as mock_send, patch.object(
client.client, "build_request"
) as mock_build_request:
# Mock logging object # Mock logging object
mock_logging_obj = MagicMock() mock_logging_obj = MagicMock()
mock_logging_obj.update_environment_variables = MagicMock() mock_logging_obj.update_environment_variables = MagicMock()
response = llm_passthrough_route( response = llm_passthrough_route(
model="anthropic.claude-3-sonnet-20240229-v1:0", model="anthropic.claude-3-sonnet-20240229-v1:0",
endpoint="model/anthropic.claude-3-sonnet-20240229-v1:0/converse", endpoint="model/anthropic.claude-3-sonnet-20240229-v1:0/converse",
@ -129,7 +151,7 @@ def test_bedrock_non_application_inference_profile_no_encoding():
# Verify that build_request was called with the original URL (no encoding) # Verify that build_request was called with the original URL (no encoding)
mock_build_request.assert_called_once() mock_build_request.assert_called_once()
call_args = mock_build_request.call_args call_args = mock_build_request.call_args
# The URL should NOT have application-inference-profile encoding # The URL should NOT have application-inference-profile encoding
actual_url = str(call_args.kwargs["url"]) actual_url = str(call_args.kwargs["url"])
assert "application-inference-profile%2F" not in actual_url assert "application-inference-profile%2F" not in actual_url
@ -151,21 +173,21 @@ def test_update_stream_param_based_on_request_body():
parsed_body=parsed_body, stream=False parsed_body=parsed_body, stream=False
) )
assert result is True assert result is True
# Test 2: no stream in request body should return original stream param # Test 2: no stream in request body should return original stream param
parsed_body = {"model": "test-model"} parsed_body = {"model": "test-model"}
result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body( result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
parsed_body=parsed_body, stream=False parsed_body=parsed_body, stream=False
) )
assert result is False assert result is False
# Test 3: stream=False in request body should return False # Test 3: stream=False in request body should return False
parsed_body = {"stream": False, "model": "test-model"} parsed_body = {"stream": False, "model": "test-model"}
result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body( result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
parsed_body=parsed_body, stream=True parsed_body=parsed_body, stream=True
) )
assert result is False assert result is False
# Test 4: no stream param provided, no stream in body # Test 4: no stream param provided, no stream in body
parsed_body = {"model": "test-model"} parsed_body = {"model": "test-model"}
result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body( result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
@ -178,14 +200,14 @@ def test_update_stream_param_based_on_request_body():
def mock_request(): def mock_request():
"""Create a mock request with headers""" """Create a mock request with headers"""
from typing import Optional from typing import Optional
class QueryParams: class QueryParams:
def __init__(self): def __init__(self):
self._dict = {} self._dict = {}
def __iter__(self): def __iter__(self):
return iter(self._dict) return iter(self._dict)
def items(self): def items(self):
return self._dict.items() return self._dict.items()
@ -210,6 +232,7 @@ def mock_request():
def mock_user_api_key_dict(): def mock_user_api_key_dict():
"""Create a mock user API key dictionary""" """Create a mock user API key dictionary"""
from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy._types import UserAPIKeyAuth
return UserAPIKeyAuth( return UserAPIKeyAuth(
api_key="test-key", api_key="test-key",
user_id="test-user", user_id="test-user",
@ -223,8 +246,8 @@ async def test_pass_through_request_stream_param_override(
mock_request, mock_user_api_key_dict mock_request, mock_user_api_key_dict
): ):
""" """
Test that when stream=None is passed as parameter but stream=True Test that when stream=None is passed as parameter but stream=True
is in request body, the request body value takes precedence and is in request body, the request body value takes precedence and
the eventual POST request uses streaming. the eventual POST request uses streaming.
""" """
from unittest.mock import AsyncMock, Mock, patch from unittest.mock import AsyncMock, Mock, patch
@ -238,29 +261,29 @@ async def test_pass_through_request_stream_param_override(
"model": "claude-3-5-sonnet-20241022", "model": "claude-3-5-sonnet-20241022",
"max_tokens": 256, "max_tokens": 256,
"messages": [{"role": "user", "content": "Hello, world"}], "messages": [{"role": "user", "content": "Hello, world"}],
"stream": True # This should override the function parameter "stream": True, # This should override the function parameter
} }
# Create a mock streaming response # Create a mock streaming response
mock_response = AsyncMock() mock_response = AsyncMock()
mock_response.status_code = 200 mock_response.status_code = 200
mock_response.headers = {"content-type": "text/event-stream"} mock_response.headers = {"content-type": "text/event-stream"}
# Mock the streaming response behavior # Mock the streaming response behavior
async def mock_aiter_bytes(): async def mock_aiter_bytes():
yield b'data: {"content": "Hello"}\n\n' yield b'data: {"content": "Hello"}\n\n'
yield b'data: {"content": "World"}\n\n' yield b'data: {"content": "World"}\n\n'
yield b'data: [DONE]\n\n' yield b"data: [DONE]\n\n"
mock_response.aiter_bytes = mock_aiter_bytes mock_response.aiter_bytes = mock_aiter_bytes
# Create mocks for the async client # Create mocks for the async client
mock_async_client = AsyncMock() mock_async_client = AsyncMock()
mock_request_obj = AsyncMock() mock_request_obj = AsyncMock()
# Mock build_request to return a request object (it's a sync method) # Mock build_request to return a request object (it's a sync method)
mock_async_client.build_request = Mock(return_value=mock_request_obj) mock_async_client.build_request = Mock(return_value=mock_request_obj)
# Mock send to return the streaming response # Mock send to return the streaming response
mock_async_client.send.return_value = mock_response mock_async_client.send.return_value = mock_response
@ -269,9 +292,7 @@ async def test_pass_through_request_stream_param_override(
mock_client_obj.client = mock_async_client mock_client_obj.client = mock_async_client
# Create the request # Create the request
request = mock_request( request = mock_request(headers={}, method="POST", request_body=request_body)
headers={}, method="POST", request_body=request_body
)
with patch( with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client", "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client",
@ -298,33 +319,32 @@ async def test_pass_through_request_stream_param_override(
httpx.URL("https://api.anthropic.com/v1/messages"), httpx.URL("https://api.anthropic.com/v1/messages"),
json=request_body, json=request_body,
params={}, params={},
headers={ headers={"Authorization": "Bearer test-key"},
"Authorization": "Bearer test-key"
},
) )
# Verify that send was called with stream=True # Verify that send was called with stream=True
mock_async_client.send.assert_called_once_with( mock_async_client.send.assert_called_once_with(
mock_request_obj, mock_request_obj,
stream=True # This proves that stream=True from request body was used stream=True, # This proves that stream=True from request body was used
) )
# Verify that the non-streaming request method was NOT called # Verify that the non-streaming request method was NOT called
mock_async_client.request.assert_not_called() mock_async_client.request.assert_not_called()
# Verify response is a StreamingResponse # Verify response is a StreamingResponse
from fastapi.responses import StreamingResponse from fastapi.responses import StreamingResponse
assert isinstance(response, StreamingResponse) assert isinstance(response, StreamingResponse)
assert response.status_code == 200 assert response.status_code == 200
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_pass_through_request_stream_param_no_override( async def test_pass_through_request_stream_param_no_override(
mock_request, mock_user_api_key_dict mock_request, mock_user_api_key_dict
): ):
""" """
Test that when stream=False is passed as parameter and no stream Test that when stream=False is passed as parameter and no stream
is in request body, the function parameter is used and is in request body, the function parameter is used and
the eventual request uses non-streaming. the eventual request uses non-streaming.
""" """
from unittest.mock import AsyncMock, Mock, patch from unittest.mock import AsyncMock, Mock, patch
@ -335,7 +355,7 @@ async def test_pass_through_request_stream_param_no_override(
# Create request body without stream parameter # Create request body without stream parameter
request_body = { request_body = {
"model": "claude-3-5-sonnet-20241022", "model": "claude-3-5-sonnet-20241022",
"max_tokens": 256, "max_tokens": 256,
"messages": [{"role": "user", "content": "Hello, world"}], "messages": [{"role": "user", "content": "Hello, world"}],
# No stream parameter - should use function parameter stream=False # No stream parameter - should use function parameter stream=False
@ -346,15 +366,15 @@ async def test_pass_through_request_stream_param_no_override(
mock_response.status_code = 200 mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"} mock_response.headers = {"content-type": "application/json"}
mock_response._content = b'{"response": "Hello world"}' mock_response._content = b'{"response": "Hello world"}'
async def mock_aread(): async def mock_aread():
return mock_response._content return mock_response._content
mock_response.aread = mock_aread mock_response.aread = mock_aread
# Create mocks for the async client # Create mocks for the async client
mock_async_client = AsyncMock() mock_async_client = AsyncMock()
# Mock request to return the non-streaming response # Mock request to return the non-streaming response
mock_async_client.request.return_value = mock_response mock_async_client.request.return_value = mock_response
@ -363,9 +383,7 @@ async def test_pass_through_request_stream_param_no_override(
mock_client_obj.client = mock_async_client mock_client_obj.client = mock_async_client
# Create the request # Create the request
request = mock_request( request = mock_request(headers={}, method="POST", request_body=request_body)
headers={}, method="POST", request_body=request_body
)
with patch( with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client", "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client",
@ -388,23 +406,105 @@ async def test_pass_through_request_stream_param_no_override(
# Verify that build_request was NOT called (no streaming path) # Verify that build_request was NOT called (no streaming path)
mock_async_client.build_request.assert_not_called() mock_async_client.build_request.assert_not_called()
# Verify that send was NOT called (no streaming path) # Verify that send was NOT called (no streaming path)
mock_async_client.send.assert_not_called() mock_async_client.send.assert_not_called()
# Verify that the non-streaming request method WAS called # Verify that the non-streaming request method WAS called
mock_async_client.request.assert_called_once_with( mock_async_client.request.assert_called_once_with(
method="POST", method="POST",
url=httpx.URL("https://api.anthropic.com/v1/messages"), url=httpx.URL("https://api.anthropic.com/v1/messages"),
headers={ headers={"Authorization": "Bearer test-key"},
"Authorization": "Bearer test-key"
},
params={}, params={},
json=request_body, json=request_body,
) )
# Verify response is a regular Response (not StreamingResponse) # Verify response is a regular Response (not StreamingResponse)
from fastapi.responses import Response, StreamingResponse from fastapi.responses import Response, StreamingResponse
assert not isinstance(response, StreamingResponse) assert not isinstance(response, StreamingResponse)
assert isinstance(response, Response) assert isinstance(response, Response)
assert response.status_code == 200 assert response.status_code == 200
def test_azure_with_custom_api_base_and_key():
"""
Test that llm_passthrough_route correctly handles Azure OpenAI
with custom api_base and api_key.
"""
client = HTTPHandler()
# Mock the provider config and its methods
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
httpx.URL(
"https://my-custom-base/openai/deployments/gpt-4.1/chat/completions?api-version=2024-02-01"
),
"https://my-custom-base",
)
mock_provider_config.get_api_key.return_value = "my-custom-key"
mock_provider_config.validate_environment.return_value = {
"api-key": "my-custom-key"
}
mock_provider_config.sign_request.return_value = (
{"api-key": "my-custom-key"},
None,
)
mock_provider_config.is_streaming_request.return_value = False
with patch(
"litellm.utils.ProviderConfigManager.get_provider_passthrough_config",
return_value=mock_provider_config,
), patch(
"litellm.litellm_core_utils.get_litellm_params.get_litellm_params",
return_value={},
), patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("gpt-4.1", "azure", "my-custom-key", "https://my-custom-base"),
), patch.object(
client.client,
"send",
return_value=MagicMock(
status_code=200, json=lambda: {"id": "chatcmpl-123", "choices": []}
),
) as mock_send, patch.object(
client.client, "build_request"
) as mock_build_request:
# Mock logging object
mock_logging_obj = MagicMock()
mock_logging_obj.update_environment_variables = MagicMock()
response = llm_passthrough_route(
model="azure/gpt-4.1",
endpoint="openai/deployments/gpt-4.1/chat/completions",
method="POST",
custom_llm_provider="azure",
api_base="https://my-custom-base",
api_key="my-custom-key",
json={
"model": "gpt-4.1",
"messages": [{"role": "user", "content": "Hello!"}],
},
client=client,
litellm_logging_obj=mock_logging_obj,
)
# Verify that build_request was called with the correct parameters
mock_build_request.assert_called_once()
call_args = mock_build_request.call_args
# Verify the URL contains the custom base
actual_url = str(call_args.kwargs["url"])
assert "my-custom-base" in actual_url
assert "gpt-4.1" in actual_url
# Verify the headers contain the custom API key
headers = call_args.kwargs["headers"]
assert headers["api-key"] == "my-custom-key"
# Verify the model in JSON body is updated
json_body = call_args.kwargs["json"]
assert json_body["model"] == "gpt-4.1"
assert response.status_code == 200

View file

@ -734,3 +734,264 @@ async def test_call_mcp_tool_user_unauthorized_access():
# Verify the exception details # Verify the exception details
assert exc_info.value.status_code == 403 assert exc_info.value.status_code == 403
assert "User not allowed to call this tool" in exc_info.value.detail 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"]

View file

@ -943,6 +943,201 @@ class TestMCPServerManager:
manager.add_update_server(server) manager.add_update_server(server)
assert server.server_id in manager.get_registry() 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__": if __name__ == "__main__":
pytest.main([__file__]) pytest.main([__file__])

View 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" 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 @pytest.mark.asyncio
async def test_key_update_object_permissions_existing_permission(monkeypatch): async def test_key_update_object_permissions_existing_permission(monkeypatch):
""" """

View file

@ -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" 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 @pytest.mark.asyncio
async def test_team_update_object_permissions_existing_permission(monkeypatch): async def test_team_update_object_permissions_existing_permission(monkeypatch):
""" """

View file

@ -2030,3 +2030,89 @@ def test_add_callback_from_db_to_in_memory_litellm_callbacks():
existing_callbacks=existing_callbacks_with_item, existing_callbacks=existing_callbacks_with_item,
) )
mock_callback_manager.add_litellm_success_callback.assert_not_called() 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

View file

@ -0,0 +1,11 @@
node_modules
.next
.out
dist
build
.coverage
.vercel
.turbo
.next-static
*.min.js
coverage/

View file

@ -0,0 +1,7 @@
{
"semi": true,
"singleQuote": false,
"tabWidth": 2,
"printWidth": 120,
"trailingComma": "all"
}

View file

@ -1,7 +0,0 @@
{
"semi": false,
"tabWidth": 2,
"printWidth": 120,
"trailingComma": "all",
"jsxBracketSameLine": false
}

View file

@ -1,12 +1,12 @@
/** @type {import('next').NextConfig} */ /** @type {import('next').NextConfig} */
const nextConfig = { const nextConfig = {
output: 'export', output: "export",
basePath: '', basePath: "",
assetPrefix: '/litellm-asset-prefix', // If a server_root_path is set, this will be overridden by runtime injection assetPrefix: "/litellm-asset-prefix", // If a server_root_path is set, this will be overridden by runtime injection
}; };
nextConfig.experimental = { nextConfig.experimental = {
missingSuspenseWithCSRBailout: false missingSuspenseWithCSRBailout: false,
} };
export default nextConfig; export default nextConfig;

View file

@ -17973,6 +17973,7 @@
"resolved": "https://registry.npmjs.org/prettier/-/prettier-3.2.5.tgz", "resolved": "https://registry.npmjs.org/prettier/-/prettier-3.2.5.tgz",
"integrity": "sha512-3/GWa9aOC0YeD7LUfvOG2NiDyhOWRvt1k+rcKhOuYnMY24iiCphgneUfJDyFXd6rZCAnuLBv6UeAULtrhT/F4A==", "integrity": "sha512-3/GWa9aOC0YeD7LUfvOG2NiDyhOWRvt1k+rcKhOuYnMY24iiCphgneUfJDyFXd6rZCAnuLBv6UeAULtrhT/F4A==",
"dev": true, "dev": true,
"license": "MIT",
"bin": { "bin": {
"prettier": "bin/prettier.cjs" "prettier": "bin/prettier.cjs"
}, },

View file

@ -3,12 +3,14 @@
"version": "0.1.0", "version": "0.1.0",
"private": true, "private": true,
"scripts": { "scripts": {
"dev": "next dev", "dev": "next dev --turbo",
"build": "next build", "build": "next build",
"start": "next start", "start": "next start",
"lint": "next lint", "lint": "next lint",
"test": "vitest", "test": "vitest",
"test:watch": "vitest -w" "test:watch": "vitest -w",
"format": "prettier --write .",
"format:check": "prettier --check ."
}, },
"dependencies": { "dependencies": {
"@anthropic-ai/sdk": "^0.54.0", "@anthropic-ai/sdk": "^0.54.0",

View file

@ -19,12 +19,7 @@
body { body {
color: rgb(var(--foreground-rgb)); color: rgb(var(--foreground-rgb));
background: linear-gradient( background: linear-gradient(to bottom, transparent, rgb(var(--background-end-rgb))) rgb(var(--background-start-rgb));
to bottom,
transparent,
rgb(var(--background-end-rgb))
)
rgb(var(--background-start-rgb));
} }
@layer utilities { @layer utilities {

View file

@ -19,7 +19,5 @@ export default function PublicModelHub() {
* populate navbar * populate navbar
* *
*/ */
return ( return <PublicModelHubPage accessToken={accessToken} />;
<PublicModelHubPage accessToken={accessToken} />
);
} }

View file

@ -19,7 +19,5 @@ export default function PublicModelHubTable() {
* populate navbar * populate navbar
* *
*/ */
return ( return <ModelHubTable accessToken={accessToken} publicPage={true} premiumUser={false} userRole={null} />;
<ModelHubTable accessToken={accessToken} publicPage={true} premiumUser={false} userRole={null}/> }
);
}

View file

@ -1,16 +1,7 @@
"use client"; "use client";
import React, { Suspense, useEffect, useState } from "react"; import React, { Suspense, useEffect, useState } from "react";
import { useSearchParams } from "next/navigation"; import { useSearchParams } from "next/navigation";
import { import { Card, Title, Text, TextInput, Callout, Button, Grid, Col } from "@tremor/react";
Card,
Title,
Text,
TextInput,
Callout,
Button,
Grid,
Col,
} from "@tremor/react";
import { RiAlarmWarningLine, RiCheckboxCircleLine } from "@remixicon/react"; import { RiAlarmWarningLine, RiCheckboxCircleLine } from "@remixicon/react";
import { import {
invitationClaimCall, invitationClaimCall,
@ -18,7 +9,7 @@ import {
getOnboardingCredentials, getOnboardingCredentials,
claimOnboardingToken, claimOnboardingToken,
getUiConfig, getUiConfig,
getProxyBaseUrl getProxyBaseUrl,
} from "@/components/networking"; } from "@/components/networking";
import { jwtDecode } from "jwt-decode"; import { jwtDecode } from "jwt-decode";
import { Form, Button as Button2, message } from "antd"; import { Form, Button as Button2, message } from "antd";
@ -27,7 +18,7 @@ import { getCookie } from "@/utils/cookieUtils";
export default function Onboarding() { export default function Onboarding() {
const [form] = Form.useForm(); const [form] = Form.useForm();
const searchParams = useSearchParams()!; const searchParams = useSearchParams()!;
const token = getCookie('token'); const token = getCookie("token");
const inviteID = searchParams.get("invitation_id"); const inviteID = searchParams.get("invitation_id");
const action = searchParams.get("action"); const action = searchParams.get("action");
const [accessToken, setAccessToken] = useState<string | null>(null); const [accessToken, setAccessToken] = useState<string | null>(null);
@ -39,14 +30,16 @@ export default function Onboarding() {
const [getUiConfigLoading, setGetUiConfigLoading] = useState<boolean>(true); const [getUiConfigLoading, setGetUiConfigLoading] = useState<boolean>(true);
useEffect(() => { useEffect(() => {
getUiConfig().then((data) => { // get the information for constructing the proxy base url, and then set the token and auth loading getUiConfig().then((data) => {
// get the information for constructing the proxy base url, and then set the token and auth loading
console.log("ui config in onboarding.tsx:", data); console.log("ui config in onboarding.tsx:", data);
setGetUiConfigLoading(false); setGetUiConfigLoading(false);
}); });
}, []); }, []);
useEffect(() => { useEffect(() => {
if (!inviteID || getUiConfigLoading) { // wait for the ui config to be loaded if (!inviteID || getUiConfigLoading) {
// wait for the ui config to be loaded
return; return;
} }
@ -72,14 +65,7 @@ export default function Onboarding() {
}, [inviteID, getUiConfigLoading]); }, [inviteID, getUiConfigLoading]);
const handleSubmit = (formValues: Record<string, any>) => { const handleSubmit = (formValues: Record<string, any>) => {
console.log( console.log("in handle submit. accessToken:", accessToken, "token:", jwtToken, "formValues:", formValues);
"in handle submit. accessToken:",
accessToken,
"token:",
jwtToken,
"formValues:",
formValues
);
if (!accessToken || !jwtToken) { if (!accessToken || !jwtToken) {
return; return;
} }
@ -89,12 +75,7 @@ export default function Onboarding() {
if (!userID || !inviteID) { if (!userID || !inviteID) {
return; return;
} }
claimOnboardingToken( claimOnboardingToken(accessToken, inviteID, userID, formValues.password).then((data) => {
accessToken,
inviteID,
userID,
formValues.password
).then((data) => {
let litellm_dashboard_ui = "/ui/"; let litellm_dashboard_ui = "/ui/";
litellm_dashboard_ui += "?login=success"; litellm_dashboard_ui += "?login=success";
@ -119,15 +100,14 @@ export default function Onboarding() {
<Card> <Card>
<Title className="text-sm mb-5 text-center">🚅 LiteLLM</Title> <Title className="text-sm mb-5 text-center">🚅 LiteLLM</Title>
<Title className="text-xl">{action === "reset_password" ? "Reset Password" : "Sign up"}</Title> <Title className="text-xl">{action === "reset_password" ? "Reset Password" : "Sign up"}</Title>
<Text>{action === "reset_password" ? "Reset your password to access Admin UI." : "Claim your user account to login to Admin UI."}</Text> <Text>
{action === "reset_password"
? "Reset your password to access Admin UI."
: "Claim your user account to login to Admin UI."}
</Text>
{action !== "reset_password" && ( {action !== "reset_password" && (
<Callout <Callout className="mt-4" title="SSO" icon={RiCheckboxCircleLine} color="sky">
className="mt-4"
title="SSO"
icon={RiCheckboxCircleLine}
color="sky"
>
<Grid numItems={2} className="flex justify-between items-center"> <Grid numItems={2} className="flex justify-between items-center">
<Col>SSO is under the Enterprise Tier.</Col> <Col>SSO is under the Enterprise Tier.</Col>
@ -142,28 +122,16 @@ export default function Onboarding() {
</Callout> </Callout>
)} )}
<Form <Form className="mt-10 mb-5 mx-auto" layout="vertical" onFinish={handleSubmit}>
className="mt-10 mb-5 mx-auto"
layout="vertical"
onFinish={handleSubmit}
>
<> <>
<Form.Item label="Email Address" name="user_email"> <Form.Item label="Email Address" name="user_email">
<TextInput <TextInput type="email" disabled={true} value={userEmail} defaultValue={userEmail} className="max-w-md" />
type="email"
disabled={true}
value={userEmail}
defaultValue={userEmail}
className="max-w-md"
/>
</Form.Item> </Form.Item>
<Form.Item <Form.Item
label="Password" label="Password"
name="password" name="password"
rules={[ rules={[{ required: true, message: "password required to sign up" }]}
{ required: true, message: "password required to sign up" },
]}
help={action === "reset_password" ? "Enter your new password" : "Create a password for your account"} help={action === "reset_password" ? "Enter your new password" : "Create a password for your account"}
> >
<TextInput placeholder="" type="password" className="max-w-md" /> <TextInput placeholder="" type="password" className="max-w-md" />

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