mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge branch 'main' into litellm_dev_10_04_2025_p3
This commit is contained in:
commit
543e00a886
390 changed files with 22677 additions and 20841 deletions
0
MCP_SSL_CHANGES_SUMMARY.md
Normal file
0
MCP_SSL_CHANGES_SUMMARY.md
Normal file
|
|
@ -23,6 +23,43 @@ LiteLLM Proxy provides an MCP Gateway that allows you to use a fixed endpoint fo
|
|||
|
||||
## Adding your MCP
|
||||
|
||||
### Prerequisites
|
||||
|
||||
To store MCP servers in the database, you need to enable database storage:
|
||||
|
||||
**Environment Variable:**
|
||||
```bash
|
||||
export STORE_MODEL_IN_DB=True
|
||||
```
|
||||
|
||||
**OR in config.yaml:**
|
||||
```yaml
|
||||
general_settings:
|
||||
store_model_in_db: true
|
||||
```
|
||||
|
||||
#### Fine-grained Database Storage Control
|
||||
|
||||
By default, when `store_model_in_db` is `true`, all object types (models, MCPs, guardrails, vector stores, etc.) are stored in the database. If you want to store only specific object types, use the `supported_db_objects` setting.
|
||||
|
||||
**Example: Store only MCP servers in the database**
|
||||
|
||||
```yaml title="config.yaml" showLineNumbers
|
||||
general_settings:
|
||||
store_model_in_db: true
|
||||
supported_db_objects: ["mcp"] # Only store MCP servers in DB
|
||||
|
||||
model_list:
|
||||
- model_name: gpt-4o
|
||||
litellm_params:
|
||||
model: openai/gpt-4o
|
||||
api_key: sk-xxxxxxx
|
||||
```
|
||||
|
||||
**See all available object types:** [Config Settings - supported_db_objects](./proxy/config_settings.md#general_settings---reference)
|
||||
|
||||
If `supported_db_objects` is not set, all object types are loaded from the database (default behavior).
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="ui" label="LiteLLM UI">
|
||||
|
||||
|
|
|
|||
|
|
@ -171,6 +171,7 @@ os.environ["OPENAI_BASE_URL"] = "https://your_host/v1" # OPTIONAL
|
|||
| gpt-5-2025-08-07 | `response = completion(model="gpt-5-2025-08-07", messages=messages)` |
|
||||
| gpt-5-mini-2025-08-07 | `response = completion(model="gpt-5-mini-2025-08-07", messages=messages)` |
|
||||
| gpt-5-nano-2025-08-07 | `response = completion(model="gpt-5-nano-2025-08-07", messages=messages)` |
|
||||
| gpt-5-pro | `response = completion(model="gpt-5-pro", messages=messages)` |
|
||||
| gpt-4.1 | `response = completion(model="gpt-4.1", messages=messages)` |
|
||||
| gpt-4.1-mini | `response = completion(model="gpt-4.1-mini", messages=messages)` |
|
||||
| gpt-4.1-nano | `response = completion(model="gpt-4.1-nano", messages=messages)` |
|
||||
|
|
@ -749,4 +750,24 @@ In your logs you should see the forwarded org id
|
|||
```bash
|
||||
LiteLLM:DEBUG: utils.py:255 - Request to litellm:
|
||||
LiteLLM:DEBUG: utils.py:255 - litellm.acompletion(... organization='my-special-org',)
|
||||
```
|
||||
|
||||
## GPT-5 Pro Special Notes
|
||||
|
||||
GPT-5 Pro is OpenAI's most advanced reasoning model with unique characteristics:
|
||||
|
||||
- **Responses API Only**: GPT-5 Pro is only available through the `/v1/responses` endpoint
|
||||
- **No Streaming**: Does not support streaming responses
|
||||
- **High Reasoning**: Designed for complex reasoning tasks with highest effort reasoning
|
||||
- **Context Window**: 400,000 tokens input, 272,000 tokens output
|
||||
- **Pricing**: $15.00 input / $120.00 output per 1M tokens (Standard), $7.50 input / $60.00 output (Batch)
|
||||
- **Tools**: Supports Web Search, File Search, Image Generation, MCP (but not Code Interpreter or Computer Use)
|
||||
- **Modalities**: Text and Image input, Text output only
|
||||
|
||||
```python
|
||||
# GPT-5 Pro usage example
|
||||
response = completion(
|
||||
model="gpt-5-pro",
|
||||
messages=[{"role": "user", "content": "Solve this complex reasoning problem..."}]
|
||||
)
|
||||
```
|
||||
|
|
@ -191,7 +191,7 @@ print(json.loads(completion.choices[0].message.content))
|
|||
model_list:
|
||||
- model_name: gemini-2.5-pro
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-1.5-pro
|
||||
model: vertex_ai/gemini-2.5-pro
|
||||
vertex_project: "project-id"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: "/path/to/service_account.json" # [OPTIONAL] Do this OR `!gcloud auth application-default login` - run this to add vertex credentials to your env
|
||||
|
|
@ -277,7 +277,7 @@ except JSONSchemaValidationError as e:
|
|||
model_list:
|
||||
- model_name: gemini-2.5-pro
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-1.5-pro
|
||||
model: vertex_ai/gemini-2.5-pro
|
||||
vertex_project: "project-id"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: "/path/to/service_account.json" # [OPTIONAL] Do this OR `!gcloud auth application-default login` - run this to add vertex credentials to your env
|
||||
|
|
@ -981,11 +981,155 @@ curl http://0.0.0.0:4000/v1/chat/completions \
|
|||
|
||||
### **Context Caching**
|
||||
|
||||
Use Vertex AI context caching is supported by calling provider api directly. (Unified Endpoint support coming soon.).
|
||||
#### Unified Endpoint
|
||||
|
||||
Use Vertex AI context caching in the same way as [**Google AI Studio - Context Caching**](../providers/gemini.md#context-caching)
|
||||
|
||||
|
||||
##### Example usage
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
for _ in range(2):
|
||||
resp = completion(
|
||||
model="vertex_ai/gemini-2.5-pro",
|
||||
messages=[
|
||||
# System Message
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Here is the full text of a complex legal agreement" * 4000,
|
||||
"cache_control": {"type": "ephemeral"}, # 👈 KEY CHANGE
|
||||
}
|
||||
],
|
||||
},
|
||||
# marked for caching with the cache_control parameter, so that this checkpoint can read from the previous cache.
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "What are the key terms and conditions in this agreement?",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
}]
|
||||
)
|
||||
|
||||
print(resp.usage) # 👈 2nd usage block will be less, since cached tokens used
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="sdk-ttl" label="SDK with Custom TTL">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
# Cache for 2 hours (7200 seconds)
|
||||
resp = completion(
|
||||
model="vertex_ai/gemini-2.5-pro",
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Here is the full text of a complex legal agreement" * 4000,
|
||||
"cache_control": {
|
||||
"type": "ephemeral",
|
||||
"ttl": "7200s" # 👈 Cache for 2 hours
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "What are the key terms and conditions in this agreement?",
|
||||
"cache_control": {
|
||||
"type": "ephemeral",
|
||||
"ttl": "3600s" # 👈 This TTL will be ignored (first one is used)
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
print(resp.usage)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
1. Setup config.yaml
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gemini-2.5-pro
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-2.5-pro
|
||||
vertex_project: "project-id"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: "/path/to/service_account.json"
|
||||
```
|
||||
|
||||
2. Start proxy
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
3. Test it!
|
||||
|
||||
```bash
|
||||
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-d '{
|
||||
"model": "gemini-2.5-flash",
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Long cache message (must be >= 1024 tokens)",
|
||||
"cache_control": {
|
||||
"type": "ephemeral",
|
||||
"ttl": "7200s"
|
||||
}
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "What is the text about?"
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}'
|
||||
|
||||
```
|
||||
|
||||
#### Calling provider api directly
|
||||
|
||||
[**Go straight to provider**](../pass_through/vertex_ai.md#context-caching)
|
||||
|
||||
#### 1. Create the Cache
|
||||
##### 1. Create the Cache
|
||||
|
||||
First, create the cache by sending a `POST` request to the `cachedContents` endpoint via the LiteLLM proxy.
|
||||
|
||||
|
|
@ -1011,7 +1155,7 @@ curl http://0.0.0.0:4000/vertex_ai/v1/projects/{project_id}/locations/{location}
|
|||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
#### 2. Get the Cache Name from the Response
|
||||
##### 2. Get the Cache Name from the Response
|
||||
|
||||
Vertex AI will return a response containing the `name` of the cached content. This name is the identifier for your cached data.
|
||||
|
||||
|
|
@ -1030,7 +1174,7 @@ Vertex AI will return a response containing the `name` of the cached content. Th
|
|||
}
|
||||
```
|
||||
|
||||
#### 3. Use the Cached Content
|
||||
##### 3. Use the Cached Content
|
||||
|
||||
Use the `name` from the response as `cachedContent` or `cached_content` in subsequent API calls to reuse the cached information. This is passed in the body of your request to `/chat/completions`.
|
||||
|
||||
|
|
|
|||
|
|
@ -224,6 +224,7 @@ router_settings:
|
|||
| service_account_settings | List[Dict[str, Any]] | Set `service_account_settings` if you want to create settings that only apply to service account keys (Doc on service accounts)[./service_accounts.md] |
|
||||
| image_generation_model | str | The default model to use for image generation - ignores model set in request |
|
||||
| store_model_in_db | boolean | If true, enables storing model + credential information in the DB. |
|
||||
| supported_db_objects | List[str] | Fine-grained control over which object types to load from the database when `store_model_in_db` is True. Available types: `"models"`, `"mcp"`, `"guardrails"`, `"vector_stores"`, `"pass_through_endpoints"`, `"prompts"`, `"model_cost_map"`. If not set, all object types are loaded (default behavior). Example: `supported_db_objects: ["mcp"]` to only load MCP servers from DB. |
|
||||
| store_prompts_in_spend_logs | boolean | If true, allows prompts and responses to be stored in the spend logs table. |
|
||||
| max_request_size_mb | int | The maximum size for requests in MB. Requests above this size will be rejected. |
|
||||
| max_response_size_mb | int | The maximum size for responses in MB. LLM Responses above this size will not be sent. |
|
||||
|
|
|
|||
354
docs/my-website/docs/proxy/security_encryption_faq.md
Normal file
354
docs/my-website/docs/proxy/security_encryption_faq.md
Normal file
|
|
@ -0,0 +1,354 @@
|
|||
# LiteLLM Self-Hosted Security & Encryption FAQ
|
||||
|
||||
## Data in Transit Encryption
|
||||
|
||||
### Does the product encrypt data in transit?
|
||||
|
||||
**Yes**, LiteLLM encrypts data in transit using TLS/SSL.
|
||||
|
||||
### Available in both OSS and Enterprise?
|
||||
|
||||
**Yes**, TLS encryption is available in both Open Source and Enterprise versions.
|
||||
|
||||
### In transit between the calling client and the product?
|
||||
|
||||
**Yes**, HTTPS/TLS is supported through SSL certificate configuration.
|
||||
|
||||
**Configuration:**
|
||||
```bash
|
||||
# CLI
|
||||
litellm --ssl_keyfile_path /path/to/key.pem --ssl_certfile_path /path/to/cert.pem
|
||||
|
||||
# Environment Variables
|
||||
export SSL_KEYFILE_PATH="/path/to/key.pem"
|
||||
export SSL_CERTFILE_PATH="/path/to/cert.pem"
|
||||
```
|
||||
|
||||
**Documentation Reference:** `docs/my-website/docs/guides/security_settings.md`
|
||||
|
||||
### In transit between the product and the LLM providers?
|
||||
|
||||
**Yes**, all connections to LLM providers use TLS encryption by default.
|
||||
|
||||
**Implementation Details:**
|
||||
- Uses Python's `ssl.create_default_context()`
|
||||
- Leverages HTTPX and aiohttp libraries with SSL/TLS enabled
|
||||
- Uses certifi CA bundle by default for SSL verification
|
||||
|
||||
**Code Reference:** `litellm/llms/custom_httpx/http_handler.py` (lines 43-105)
|
||||
|
||||
### Are TCP sessions to the LLM providers shared?
|
||||
|
||||
**Yes**, TCP connections are pooled and reused.
|
||||
|
||||
**Details:**
|
||||
- Connection pooling is enabled by default
|
||||
- Default: 1000 max concurrent connections with keepalive
|
||||
- Sessions are maintained across requests to the same provider
|
||||
- Reduces overhead of TLS handshakes
|
||||
|
||||
**Code Reference:** `litellm/llms/custom_httpx/http_handler.py` (lines 704-712)
|
||||
|
||||
### Or does the product negotiate a new TLS session with the same LLM provider for every sequential call?
|
||||
|
||||
**No**, TLS sessions are reused through connection pooling. New TLS handshakes are not performed for every request.
|
||||
|
||||
### How is it encrypted?
|
||||
|
||||
**TLS 1.2 and TLS 1.3**
|
||||
|
||||
Uses Python's default SSL context which supports both TLS 1.2 and TLS 1.3. The specific version negotiated depends on:
|
||||
- Python version
|
||||
- System SSL library (typically OpenSSL)
|
||||
- Server capabilities
|
||||
|
||||
**Implementation:** `ssl.create_default_context()` in Python
|
||||
|
||||
### How are these added to the product's configuration?
|
||||
|
||||
#### x.509 Certificate
|
||||
|
||||
**Method 1: CLI Arguments**
|
||||
```bash
|
||||
litellm --ssl_certfile_path /path/to/certificate.pem
|
||||
```
|
||||
|
||||
**Method 2: Environment Variable**
|
||||
```bash
|
||||
export SSL_CERTFILE_PATH="/path/to/certificate.pem"
|
||||
```
|
||||
|
||||
#### Private Key
|
||||
|
||||
**Method 1: CLI Arguments**
|
||||
```bash
|
||||
litellm --ssl_keyfile_path /path/to/private_key.pem
|
||||
```
|
||||
|
||||
**Method 2: Environment Variable**
|
||||
```bash
|
||||
export SSL_KEYFILE_PATH="/path/to/private_key.pem"
|
||||
```
|
||||
|
||||
#### Certificate Bundle/Chain
|
||||
|
||||
**For client-to-proxy connections:**
|
||||
Use standard SSL certificate setup with intermediate certificates bundled in the certfile.
|
||||
|
||||
**For proxy-to-LLM provider connections:**
|
||||
|
||||
**Method 1: Config YAML**
|
||||
```yaml
|
||||
litellm_settings:
|
||||
ssl_verify: "/path/to/ca_bundle.pem"
|
||||
```
|
||||
|
||||
**Method 2: Environment Variable**
|
||||
```bash
|
||||
export SSL_CERT_FILE="/path/to/ca_bundle.pem"
|
||||
```
|
||||
|
||||
**Method 3: Client Certificate Authentication**
|
||||
```yaml
|
||||
litellm_settings:
|
||||
ssl_certificate: "/path/to/client_certificate.pem"
|
||||
```
|
||||
|
||||
or
|
||||
|
||||
```bash
|
||||
export SSL_CERTIFICATE="/path/to/client_certificate.pem"
|
||||
```
|
||||
|
||||
### Documentation Coverage
|
||||
|
||||
**Primary Documentation:**
|
||||
- `docs/my-website/docs/guides/security_settings.md` - SSL/TLS configuration guide
|
||||
|
||||
**Additional References:**
|
||||
- `litellm/proxy/proxy_cli.py` (lines 455-467) - CLI options
|
||||
- `docs/my-website/docs/completion/http_handler_config.md` - Custom HTTP handler configuration
|
||||
|
||||
---
|
||||
|
||||
## Data at Rest Encryption
|
||||
|
||||
### Does the product encrypt data at rest?
|
||||
|
||||
**Partially**. Only specific sensitive data is encrypted at rest.
|
||||
|
||||
### What data is stored in encrypted form?
|
||||
|
||||
#### Encrypted Data:
|
||||
1. **LLM API Keys** - Model credentials in `LiteLLM_ProxyModelTable.litellm_params`
|
||||
2. **Provider Credentials** - Stored in `LiteLLM_CredentialsTable.credential_values`
|
||||
3. **Configuration Secrets** - Sensitive config values in `LiteLLM_Config` table
|
||||
4. **Virtual Keys** - When using secret managers (optional feature)
|
||||
|
||||
#### NOT Encrypted:
|
||||
1. **Spend Logs** - Request/response data in `LiteLLM_SpendLogs`
|
||||
2. **Audit Logs** - Change history in `LiteLLM_AuditLog`
|
||||
3. **User/Team/Organization Data** - Metadata and configuration
|
||||
4. **Cached Prompts and Completions** - Cache data is stored in plaintext
|
||||
|
||||
### Cached prompts and completions?
|
||||
|
||||
**No**, cached prompts and completions are **NOT encrypted**.
|
||||
|
||||
Cache backends (Redis, S3, local disk) store data as plaintext JSON.
|
||||
|
||||
**Code References:**
|
||||
- `litellm/caching/redis_cache.py`
|
||||
- `litellm/caching/s3_cache.py`
|
||||
- `litellm/caching/caching.py`
|
||||
|
||||
### Configuration data?
|
||||
|
||||
**Partially encrypted**.
|
||||
|
||||
#### What IS Encrypted:
|
||||
- LLM API keys and credentials in model configurations
|
||||
- Sensitive values in `LiteLLM_Config` table
|
||||
- Credential values in `LiteLLM_CredentialsTable`
|
||||
|
||||
#### What is NOT Encrypted:
|
||||
- Model names and aliases
|
||||
- Rate limits and budget settings
|
||||
- User/team/organization metadata
|
||||
- Non-sensitive configuration parameters
|
||||
|
||||
**Code Reference:** `litellm/proxy/management_endpoints/model_management_endpoints.py` (lines 275-308)
|
||||
|
||||
### Log data?
|
||||
|
||||
**No**, log data is **NOT encrypted**.
|
||||
|
||||
Log data stored in database tables is in plaintext:
|
||||
- `LiteLLM_SpendLogs` - Contains request/response data, tokens, spend
|
||||
- `LiteLLM_ErrorLogs` - Error information
|
||||
- `LiteLLM_AuditLog` - Audit trail of changes
|
||||
|
||||
**Note:** You can disable logging to avoid storing sensitive data:
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
disable_spend_logs: True # Disable writing spend logs to DB
|
||||
disable_error_logs: True # Disable writing error logs to DB
|
||||
```
|
||||
|
||||
**Documentation:** `docs/my-website/docs/proxy/db_info.md` (lines 52-60)
|
||||
|
||||
### Where is it stored?
|
||||
|
||||
#### In the DB?
|
||||
|
||||
**Yes**, encrypted data is stored in PostgreSQL database.
|
||||
|
||||
**Key Tables with Encrypted Data:**
|
||||
- `LiteLLM_ProxyModelTable` - Model configurations with encrypted API keys
|
||||
- `LiteLLM_CredentialsTable` - Credential values
|
||||
- `LiteLLM_Config` - Configuration secrets
|
||||
|
||||
**Schema Reference:** `schema.prisma`
|
||||
|
||||
#### In the filesystem?
|
||||
|
||||
**No**, encrypted data is not stored in the filesystem by default.
|
||||
|
||||
**Note:** If using disk cache (`disk_cache_dir`), cached data is stored unencrypted.
|
||||
|
||||
#### Somewhere else?
|
||||
|
||||
**Optional:** When using secret managers (AWS Secrets Manager, Azure Key Vault, HashiCorp Vault), encrypted data can be stored externally.
|
||||
|
||||
**Configuration:**
|
||||
```yaml
|
||||
general_settings:
|
||||
key_management_system: "aws_secret_manager" # or "azure_key_vault", "hashicorp_vault"
|
||||
```
|
||||
|
||||
**Documentation:** `docs/my-website/docs/secret.md`
|
||||
|
||||
### How is it encrypted?
|
||||
|
||||
**Algorithm:** NaCl SecretBox (XSalsa20-Poly1305 AEAD)
|
||||
|
||||
**NOT AES-256** - LiteLLM uses NaCl (Networking and Cryptography Library) which provides:
|
||||
- XSalsa20 stream cipher
|
||||
- Poly1305 MAC for authentication
|
||||
- Equivalent security to AES-256
|
||||
|
||||
**Key Derivation:**
|
||||
1. Takes `LITELLM_SALT_KEY` (or `LITELLM_MASTER_KEY` if salt key not set)
|
||||
2. Hashes with SHA-256 to derive 256-bit encryption key
|
||||
3. Uses NaCl SecretBox for authenticated encryption
|
||||
|
||||
**Code Reference:** `litellm/proxy/common_utils/encrypt_decrypt_utils.py` (lines 69-112)
|
||||
|
||||
**Implementation:**
|
||||
```python
|
||||
import hashlib
|
||||
import nacl.secret
|
||||
|
||||
# Derive 256-bit key from salt
|
||||
hash_object = hashlib.sha256(signing_key.encode())
|
||||
hash_bytes = hash_object.digest()
|
||||
|
||||
# Create SecretBox and encrypt
|
||||
box = nacl.secret.SecretBox(hash_bytes)
|
||||
encrypted = box.encrypt(value_bytes)
|
||||
```
|
||||
|
||||
### Setting the Encryption Key
|
||||
|
||||
**Required Environment Variable:**
|
||||
```bash
|
||||
export LITELLM_SALT_KEY="your-strong-random-key-here"
|
||||
```
|
||||
|
||||
**Important Notes:**
|
||||
- ⚠️ **Must be set before adding any models**
|
||||
- ⚠️ **Never change this key** - encrypted data becomes unrecoverable
|
||||
- ⚠️ Use a strong random key (recommended: https://1password.com/password-generator/)
|
||||
- If not set, falls back to `LITELLM_MASTER_KEY`
|
||||
|
||||
**Documentation:** `docs/my-website/docs/proxy/prod.md` (section 8, lines 184-196)
|
||||
|
||||
### Documentation Coverage
|
||||
|
||||
**Primary Documentation:**
|
||||
- `docs/my-website/docs/proxy/prod.md` (section 8) - LITELLM_SALT_KEY setup
|
||||
- `docs/my-website/docs/secret.md` - Secret management systems
|
||||
- `docs/my-website/docs/proxy/db_info.md` - Database information
|
||||
|
||||
**Additional References:**
|
||||
- `security.md` - General security measures
|
||||
- `docs/my-website/docs/data_security.md` - Data privacy overview
|
||||
- `schema.prisma` - Database schema with encrypted fields
|
||||
|
||||
---
|
||||
|
||||
## Summary of Security Features
|
||||
|
||||
### ✅ Provided Out of the Box
|
||||
|
||||
1. **TLS/SSL encryption** for client-to-proxy connections
|
||||
2. **TLS encryption** for proxy-to-LLM provider connections (with connection pooling)
|
||||
3. **Encrypted storage** of LLM API keys and credentials
|
||||
4. **Support for TLS 1.2 and TLS 1.3**
|
||||
5. **Connection pooling** to reduce TLS handshake overhead
|
||||
|
||||
### ⚠️ Important Limitations
|
||||
|
||||
1. **Cached data is NOT encrypted** (Redis, S3, disk cache)
|
||||
2. **Log data is NOT encrypted** (spend logs, audit logs)
|
||||
3. **Request/response payloads in logs are NOT encrypted**
|
||||
4. **Uses NaCl SecretBox, NOT AES-256** (equivalent security)
|
||||
5. **TLS version not explicitly configured** - uses Python/system defaults
|
||||
|
||||
### 🔧 Configuration Requirements
|
||||
|
||||
**For Production Deployments:**
|
||||
|
||||
1. **Set LITELLM_SALT_KEY** before adding any models
|
||||
2. **Configure SSL certificates** for HTTPS client connections
|
||||
3. **Consider disabling logs** if they contain sensitive data
|
||||
4. **Use secret managers** for enhanced security (optional)
|
||||
5. **Configure CA bundles** if using custom certificates
|
||||
|
||||
---
|
||||
|
||||
## Quick Start Security Checklist
|
||||
|
||||
```bash
|
||||
# 1. Generate a strong salt key
|
||||
export LITELLM_SALT_KEY="$(openssl rand -base64 32)"
|
||||
|
||||
# 2. Set up SSL certificates (for HTTPS)
|
||||
export SSL_KEYFILE_PATH="/path/to/private_key.pem"
|
||||
export SSL_CERTFILE_PATH="/path/to/certificate.pem"
|
||||
|
||||
# 3. Configure database
|
||||
export DATABASE_URL="postgresql://user:password@host:port/dbname"
|
||||
|
||||
# 4. (Optional) Disable logs if they contain sensitive data
|
||||
# Add to config.yaml:
|
||||
# general_settings:
|
||||
# disable_spend_logs: True
|
||||
# disable_error_logs: True
|
||||
|
||||
# 5. Start LiteLLM Proxy
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Additional Resources
|
||||
|
||||
- **LiteLLM Documentation:** https://docs.litellm.ai/
|
||||
- **Security Settings Guide:** https://docs.litellm.ai/docs/guides/security_settings
|
||||
- **Production Deployment:** https://docs.litellm.ai/docs/proxy/prod
|
||||
- **Secret Management:** https://docs.litellm.ai/docs/secret
|
||||
|
||||
For security inquiries: support@berri.ai
|
||||
|
||||
BIN
docs/my-website/img/release_notes/perf_77_5.png
Normal file
BIN
docs/my-website/img/release_notes/perf_77_5.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 1.3 MiB |
BIN
docs/my-website/img/release_notes/perf_77_7.png
Normal file
BIN
docs/my-website/img/release_notes/perf_77_7.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 253 KiB |
BIN
docs/my-website/img/release_notes/schedule_key_rotations.png
Normal file
BIN
docs/my-website/img/release_notes/schedule_key_rotations.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 603 KiB |
|
|
@ -50,30 +50,6 @@ pip install litellm==1.75.5.post2
|
|||
- **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.
|
||||
|
||||
|
||||
### 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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
date: 2025-09-29T10:00:00
|
||||
authors:
|
||||
|
|
@ -11,6 +11,10 @@ authors:
|
|||
title: CTO, LiteLLM
|
||||
url: https://www.linkedin.com/in/reffajnaahsi/
|
||||
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
|
||||
---
|
||||
|
|
@ -28,7 +32,7 @@ import TabItem from '@theme/TabItem';
|
|||
docker run \
|
||||
-e STORE_MODEL_IN_DB=True \
|
||||
-p 4000:4000 \
|
||||
ghcr.io/berriai/litellm:v1.77.5.rc.1
|
||||
ghcr.io/berriai/litellm:v1.77.5-stable
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
|
@ -49,7 +53,54 @@ pip install litellm==1.77.5
|
|||
- **MCP OAuth 2.0 Support** - Enhanced authentication for Model Context Protocol integrations
|
||||
- **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
|
||||
- **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
|
||||
|
||||
|
|
|
|||
|
|
@ -59,12 +59,50 @@ pip install litellm==1.77.7.rc.1
|
|||
## 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
|
||||
- **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
|
||||
- **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
|
||||
- **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 Model Support
|
||||
|
|
|
|||
|
|
@ -674,6 +674,7 @@ const sidebars = {
|
|||
items: [
|
||||
"data_security",
|
||||
"data_retention",
|
||||
"proxy/security_encryption_faq",
|
||||
"migration_policy",
|
||||
{
|
||||
type: "category",
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@
|
|||
| gpt-3.5-turbo-16k | `completion('gpt-3.5-turbo-16k', messages)` | `os.environ['OPENAI_API_KEY']` |
|
||||
| gpt-3.5-turbo-16k-0613 | `completion('gpt-3.5-turbo-16k-0613', messages)` | `os.environ['OPENAI_API_KEY']` |
|
||||
| gpt-4 | `completion('gpt-4', messages)` | `os.environ['OPENAI_API_KEY']` |
|
||||
| gpt-5-pro | `completion('gpt-5-pro', messages)` | `os.environ['OPENAI_API_KEY']` |
|
||||
|
||||
## Azure OpenAI Chat Completion Models
|
||||
For Azure calls add the `azure/` prefix to `model`. If your azure deployment name is `gpt-v-2` set `model` = `azure/gpt-v-2`
|
||||
|
|
|
|||
|
|
@ -0,0 +1,3 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "mcp_tool_permissions" JSONB;
|
||||
|
||||
|
|
@ -156,6 +156,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
object_permission_id String @id @default(uuid())
|
||||
mcp_servers String[] @default([])
|
||||
mcp_access_groups String[] @default([])
|
||||
mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]}
|
||||
vector_stores String[] @default([])
|
||||
teams LiteLLM_TeamTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
|
|
|
|||
|
|
@ -5,8 +5,9 @@ LiteLLM Proxy uses this MCP Client to connnect to other MCP servers.
|
|||
import asyncio
|
||||
import base64
|
||||
from datetime import timedelta
|
||||
from typing import Dict, List, Optional, Union
|
||||
from typing import Callable, Dict, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
from mcp import ClientSession, StdioServerParameters
|
||||
from mcp.client.sse import sse_client
|
||||
from mcp.client.stdio import stdio_client
|
||||
|
|
@ -17,6 +18,8 @@ from mcp.types import TextContent
|
|||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
|
||||
from litellm.types.llms.custom_http import VerifyTypes
|
||||
from litellm.types.mcp import (
|
||||
MCPAuth,
|
||||
MCPAuthType,
|
||||
|
|
@ -48,6 +51,7 @@ class MCPClient:
|
|||
timeout: float = 60.0,
|
||||
stdio_config: Optional[MCPStdioConfig] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
ssl_verify: Optional[VerifyTypes] = None,
|
||||
):
|
||||
self.server_url: str = server_url
|
||||
self.transport_type: MCPTransport = transport_type
|
||||
|
|
@ -62,6 +66,7 @@ class MCPClient:
|
|||
self._task: Optional[asyncio.Task] = None
|
||||
self.stdio_config: Optional[MCPStdioConfig] = stdio_config
|
||||
self.extra_headers: Optional[Dict[str, str]] = extra_headers
|
||||
self.ssl_verify: Optional[VerifyTypes] = ssl_verify
|
||||
# handle the basic auth value if provided
|
||||
if auth_value:
|
||||
self.update_auth_value(auth_value)
|
||||
|
|
@ -104,10 +109,12 @@ class MCPClient:
|
|||
await self._session.initialize()
|
||||
elif self.transport_type == MCPTransport.sse:
|
||||
headers = self._get_auth_headers()
|
||||
httpx_client_factory = self._create_httpx_client_factory()
|
||||
self._transport_ctx = sse_client(
|
||||
url=self.server_url,
|
||||
timeout=self.timeout,
|
||||
headers=headers,
|
||||
httpx_client_factory=httpx_client_factory,
|
||||
)
|
||||
self._transport = await self._transport_ctx.__aenter__()
|
||||
self._session_ctx = ClientSession(
|
||||
|
|
@ -117,13 +124,15 @@ class MCPClient:
|
|||
await self._session.initialize()
|
||||
else: # http
|
||||
headers = self._get_auth_headers()
|
||||
httpx_client_factory = self._create_httpx_client_factory()
|
||||
verbose_logger.debug(
|
||||
"litellm headers for streamablehttp_client: ", headers
|
||||
"litellm headers for streamablehttp_client: %s", headers
|
||||
)
|
||||
self._transport_ctx = streamablehttp_client(
|
||||
url=self.server_url,
|
||||
timeout=timedelta(seconds=self.timeout),
|
||||
headers=headers,
|
||||
httpx_client_factory=httpx_client_factory,
|
||||
)
|
||||
self._transport = await self._transport_ctx.__aenter__()
|
||||
self._session_ctx = ClientSession(
|
||||
|
|
@ -215,6 +224,41 @@ class MCPClient:
|
|||
|
||||
return headers
|
||||
|
||||
def _create_httpx_client_factory(self) -> Callable[..., httpx.AsyncClient]:
|
||||
"""
|
||||
Create a custom httpx client factory that uses LiteLLM's SSL configuration.
|
||||
|
||||
This factory follows the same CA bundle path logic as http_handler.py:
|
||||
1. Check ssl_verify parameter (can be SSLContext, bool, or path to CA bundle)
|
||||
2. Check SSL_VERIFY environment variable
|
||||
3. Check SSL_CERT_FILE environment variable
|
||||
4. Fall back to certifi CA bundle
|
||||
"""
|
||||
|
||||
def factory(
|
||||
*,
|
||||
headers: Optional[Dict[str, str]] = None,
|
||||
timeout: Optional[httpx.Timeout] = None,
|
||||
auth: Optional[httpx.Auth] = None,
|
||||
) -> httpx.AsyncClient:
|
||||
"""Create an httpx.AsyncClient with LiteLLM's SSL configuration."""
|
||||
# Get unified SSL configuration using the same logic as http_handler.py
|
||||
ssl_config = get_ssl_configuration(self.ssl_verify)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"MCP client using SSL configuration: {type(ssl_config).__name__}"
|
||||
)
|
||||
|
||||
return httpx.AsyncClient(
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
auth=auth,
|
||||
verify=ssl_config,
|
||||
follow_redirects=True,
|
||||
)
|
||||
|
||||
return factory
|
||||
|
||||
async def list_tools(self) -> List[MCPTool]:
|
||||
"""List available tools from the server."""
|
||||
if not self._session:
|
||||
|
|
|
|||
|
|
@ -665,10 +665,6 @@ class BaseAzureLLM(BaseOpenAILLM):
|
|||
) -> dict:
|
||||
litellm_params = litellm_params or GenericLiteLLMParams()
|
||||
|
||||
# If api-key is already in headers, preserve it
|
||||
if "api-key" in headers:
|
||||
return headers
|
||||
|
||||
api_key = (
|
||||
litellm_params.api_key
|
||||
or litellm.api_key
|
||||
|
|
@ -693,7 +689,7 @@ class BaseAzureLLM(BaseOpenAILLM):
|
|||
def _get_base_azure_url(
|
||||
api_base: Optional[str],
|
||||
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,
|
||||
) -> str:
|
||||
"""
|
||||
|
|
|
|||
85
litellm/llms/azure/passthrough/transformation.py
Normal file
85
litellm/llms/azure/passthrough/transformation.py
Normal 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)
|
||||
|
|
@ -1,6 +1,7 @@
|
|||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
|
||||
|
||||
import httpx
|
||||
from openai.types.responses import ResponseReasoningItem
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.azure.common_utils import BaseAzureLLM
|
||||
|
|
@ -38,6 +39,50 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
model = model.replace("o_series/", "")
|
||||
return model
|
||||
|
||||
def _handle_reasoning_item(self, item: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Handle reasoning items specifically to filter out status=None using OpenAI's model.
|
||||
Issue: https://github.com/BerriAI/litellm/issues/13484
|
||||
OpenAI API does not accept ReasoningItem(status=None), so we need to:
|
||||
1. Check if the item is a reasoning type
|
||||
2. Create a ResponseReasoningItem object with the item data
|
||||
3. Convert it back to dict with exclude_none=True to filter None values
|
||||
"""
|
||||
if item.get("type") == "reasoning":
|
||||
try:
|
||||
# Ensure required fields are present for ResponseReasoningItem
|
||||
item_data = dict(item)
|
||||
if "id" not in item_data:
|
||||
item_data["id"] = f"reasoning_{hash(str(item_data))}"
|
||||
if "summary" not in item_data:
|
||||
item_data["summary"] = (
|
||||
item_data.get("reasoning_content", "")[:100] + "..."
|
||||
if len(item_data.get("reasoning_content", "")) > 100
|
||||
else item_data.get("reasoning_content", "")
|
||||
)
|
||||
|
||||
# Create ResponseReasoningItem object from the item data
|
||||
reasoning_item = ResponseReasoningItem(**item_data)
|
||||
|
||||
# Convert back to dict with exclude_none=True to exclude None fields
|
||||
dict_reasoning_item = reasoning_item.model_dump(exclude_none=True)
|
||||
dict_reasoning_item.pop("status", None)
|
||||
|
||||
return dict_reasoning_item
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
f"Failed to create ResponseReasoningItem, falling back to manual filtering: {e}"
|
||||
)
|
||||
# Fallback: manually filter out known None fields
|
||||
filtered_item = {
|
||||
k: v
|
||||
for k, v in item.items()
|
||||
if v is not None
|
||||
or k not in {"status", "content", "encrypted_content"}
|
||||
}
|
||||
return filtered_item
|
||||
return item
|
||||
|
||||
def transform_responses_api_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -48,12 +93,13 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
) -> Dict:
|
||||
"""No transform applied since inputs are in OpenAI spec already"""
|
||||
stripped_model_name = self.get_stripped_model_name(model)
|
||||
return dict(
|
||||
ResponsesAPIRequestParams(
|
||||
model=stripped_model_name,
|
||||
input=input,
|
||||
**response_api_optional_request_params,
|
||||
)
|
||||
|
||||
return super().transform_responses_api_request(
|
||||
model=stripped_model_name,
|
||||
input=input,
|
||||
response_api_optional_request_params=response_api_optional_request_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
def get_complete_url(
|
||||
|
|
@ -217,15 +263,15 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
at the correct location (before any query parameters).
|
||||
"""
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
|
||||
# Parse the URL to separate its components
|
||||
parsed_url = urlparse(api_base)
|
||||
|
||||
|
||||
# Insert the response_id and /cancel at the end of the path component
|
||||
# Remove trailing slash if present to avoid double slashes
|
||||
path = parsed_url.path.rstrip("/")
|
||||
new_path = f"{path}/{response_id}/cancel"
|
||||
|
||||
|
||||
# Reconstruct the URL with all original components but with the modified path
|
||||
cancel_url = urlunparse(
|
||||
(
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
|
||||
from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error
|
||||
from litellm.llms.openai.openai import OpenAIConfig
|
||||
from litellm.llms.xai.chat.transformation import XAIChatConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse, ProviderField
|
||||
|
|
@ -35,9 +36,24 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
for param in supported_params:
|
||||
if param != "tool_choice":
|
||||
filtered_supported_params.append(param)
|
||||
return filtered_supported_params
|
||||
supported_params = filtered_supported_params
|
||||
|
||||
# Filter out unsupported parameters for specific models
|
||||
if not self._supports_stop_reason(model):
|
||||
supported_params = [param for param in supported_params if param != "stop"]
|
||||
|
||||
return supported_params
|
||||
|
||||
def _supports_stop_reason(self, model: str) -> bool:
|
||||
"""
|
||||
Check if the model supports stop tokens.
|
||||
"""
|
||||
if "grok" in model:
|
||||
# Reuse Xai method for Grok model
|
||||
xai_config = XAIChatConfig()
|
||||
return xai_config._supports_stop_reason(model)
|
||||
return True
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
@ -53,9 +69,7 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
else:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
|
||||
headers["Content-Type"] = (
|
||||
"application/json" # tell Azure AI Studio to expect JSON
|
||||
)
|
||||
headers["Content-Type"] = "application/json" # tell Azure AI Studio to expect JSON
|
||||
|
||||
return headers
|
||||
|
||||
|
|
@ -65,10 +79,7 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
"""
|
||||
parsed_url = urlparse(api_base)
|
||||
host = parsed_url.hostname
|
||||
if host and (
|
||||
host.endswith(".services.ai.azure.com")
|
||||
or host.endswith(".openai.azure.com")
|
||||
):
|
||||
if host and (host.endswith(".services.ai.azure.com") or host.endswith(".openai.azure.com")):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
|
@ -115,13 +126,9 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
|
||||
# Add the path to the base URL
|
||||
if "services.ai.azure.com" in api_base:
|
||||
new_url = _add_path_to_api_base(
|
||||
api_base=api_base, ending_path="/models/chat/completions"
|
||||
)
|
||||
new_url = _add_path_to_api_base(api_base=api_base, ending_path="/models/chat/completions")
|
||||
else:
|
||||
new_url = _add_path_to_api_base(
|
||||
api_base=api_base, ending_path="/chat/completions"
|
||||
)
|
||||
new_url = _add_path_to_api_base(api_base=api_base, ending_path="/chat/completions")
|
||||
|
||||
# Use the new query_params dictionary
|
||||
final_url = httpx.URL(new_url).copy_with(params=query_params)
|
||||
|
|
@ -191,11 +198,7 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
dynamic_api_key = api_key or get_secret_str("AZURE_AI_API_KEY")
|
||||
|
||||
if self._is_azure_openai_model(model=model, api_base=api_base):
|
||||
verbose_logger.debug(
|
||||
"Model={} is Azure OpenAI model. Setting custom_llm_provider='azure'.".format(
|
||||
model
|
||||
)
|
||||
)
|
||||
verbose_logger.debug("Model={} is Azure OpenAI model. Setting custom_llm_provider='azure'.".format(model))
|
||||
custom_llm_provider = "azure"
|
||||
return api_base, dynamic_api_key, custom_llm_provider
|
||||
|
||||
|
|
@ -211,9 +214,7 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
if extra_body and isinstance(extra_body, dict):
|
||||
optional_params.update(extra_body)
|
||||
optional_params.pop("max_retries", None)
|
||||
return super().transform_request(
|
||||
model, messages, optional_params, litellm_params, headers
|
||||
)
|
||||
return super().transform_request(model, messages, optional_params, litellm_params, headers)
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
|
|
@ -252,47 +253,30 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
|
||||
if should_drop_params and "Extra inputs are not permitted" in error_text:
|
||||
return True
|
||||
elif (
|
||||
"unknown field: parameter index is not a valid field" in error_text
|
||||
): # remove index from tool calls
|
||||
elif "unknown field: parameter index is not a valid field" in error_text: # remove index from tool calls
|
||||
return True
|
||||
elif (
|
||||
AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value
|
||||
in error_text
|
||||
AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value in error_text
|
||||
): # remove extra-parameters from tool calls
|
||||
return True
|
||||
return super().should_retry_llm_api_inside_llm_translation_on_http_error(
|
||||
e=e, litellm_params=litellm_params
|
||||
)
|
||||
return super().should_retry_llm_api_inside_llm_translation_on_http_error(e=e, litellm_params=litellm_params)
|
||||
|
||||
@property
|
||||
def max_retry_on_unprocessable_entity_error(self) -> int:
|
||||
return 2
|
||||
|
||||
def transform_request_on_unprocessable_entity_error(
|
||||
self, e: httpx.HTTPStatusError, request_data: dict
|
||||
) -> dict:
|
||||
def transform_request_on_unprocessable_entity_error(self, e: httpx.HTTPStatusError, request_data: dict) -> dict:
|
||||
_messages = cast(Optional[List[AllMessageValues]], request_data.get("messages"))
|
||||
if (
|
||||
"unknown field: parameter index is not a valid field" in e.response.text
|
||||
and _messages is not None
|
||||
):
|
||||
if "unknown field: parameter index is not a valid field" in e.response.text and _messages is not None:
|
||||
litellm.remove_index_from_tool_calls(
|
||||
messages=_messages,
|
||||
)
|
||||
elif (
|
||||
AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value
|
||||
in e.response.text
|
||||
):
|
||||
request_data = self._drop_extra_params_from_request_data(
|
||||
request_data, e.response.text
|
||||
)
|
||||
elif AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value in e.response.text:
|
||||
request_data = self._drop_extra_params_from_request_data(request_data, e.response.text)
|
||||
data = drop_params_from_unprocessable_entity_error(e=e, data=request_data)
|
||||
return data
|
||||
|
||||
def _drop_extra_params_from_request_data(
|
||||
self, request_data: dict, error_text: str
|
||||
) -> dict:
|
||||
def _drop_extra_params_from_request_data(self, request_data: dict, error_text: str) -> dict:
|
||||
params_to_drop = self._extract_params_to_drop_from_error_text(error_text)
|
||||
if params_to_drop:
|
||||
for param in params_to_drop:
|
||||
|
|
@ -300,9 +284,7 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
request_data.pop(param, None)
|
||||
return request_data
|
||||
|
||||
def _extract_params_to_drop_from_error_text(
|
||||
self, error_text: str
|
||||
) -> Optional[List[str]]:
|
||||
def _extract_params_to_drop_from_error_text(self, error_text: str) -> Optional[List[str]]:
|
||||
"""
|
||||
Error text looks like this"
|
||||
"Extra parameters ['stream_options', 'extra-parameters'] are not allowed when extra-parameters is not set or set to be 'error'.
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
|
|||
"presence_penalty",
|
||||
"frequency_penalty",
|
||||
"top_logprobs",
|
||||
"stop",
|
||||
]
|
||||
|
||||
return [
|
||||
|
|
|
|||
|
|
@ -1,12 +1,4 @@
|
|||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Dict,
|
||||
Optional,
|
||||
Union,
|
||||
cast,
|
||||
get_type_hints,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional, Union, cast, get_type_hints
|
||||
|
||||
import httpx
|
||||
from openai.types.responses import ResponseReasoningItem
|
||||
|
|
@ -127,7 +119,6 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
2. Create a ResponseReasoningItem object with the item data
|
||||
3. Convert it back to dict with exclude_none=True to filter None values
|
||||
"""
|
||||
verbose_logger.debug(f"Handling reasoning item: {item}")
|
||||
if item.get("type") == "reasoning":
|
||||
try:
|
||||
# Ensure required fields are present for ResponseReasoningItem
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
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
|
||||
|
||||
|
|
@ -22,15 +23,25 @@ else:
|
|||
|
||||
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
|
||||
def get_config(cls):
|
||||
return super().get_config()
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List:
|
||||
return ["temperature", "max_tokens", "top_p", "response_format"]
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
return [
|
||||
"temperature",
|
||||
"max_tokens",
|
||||
"top_p",
|
||||
"response_format",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
|
|
@ -56,6 +67,57 @@ class SnowflakeConfig(OpenAIGPTConfig):
|
|||
optional_params[param] = value
|
||||
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(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -71,6 +133,7 @@ class SnowflakeConfig(OpenAIGPTConfig):
|
|||
json_mode: Optional[bool] = None,
|
||||
) -> ModelResponse:
|
||||
response_json = raw_response.json()
|
||||
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
|
|
@ -78,6 +141,26 @@ class SnowflakeConfig(OpenAIGPTConfig):
|
|||
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.model = "snowflake/" + (returned_response.model or "")
|
||||
|
|
@ -150,6 +233,95 @@ class SnowflakeConfig(OpenAIGPTConfig):
|
|||
|
||||
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(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -160,6 +332,18 @@ class SnowflakeConfig(OpenAIGPTConfig):
|
|||
) -> dict:
|
||||
stream: bool = optional_params.pop("stream", None) or False
|
||||
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 {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
|
|||
"""
|
||||
|
||||
import re
|
||||
from typing import List, Optional, Tuple
|
||||
from typing import List, Optional, Tuple, Literal
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.llms.vertex_ai import CachedContentRequestBody
|
||||
|
|
@ -155,13 +155,18 @@ def separate_cached_messages(
|
|||
|
||||
|
||||
def transform_openai_messages_to_gemini_context_caching(
|
||||
model: str, messages: List[AllMessageValues], cache_key: str
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
cache_key: str,
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
) -> CachedContentRequestBody:
|
||||
# Extract TTL from cached messages BEFORE system message transformation
|
||||
ttl = extract_ttl_from_cached_messages(messages)
|
||||
|
||||
supports_system_message = get_supports_system_message(
|
||||
model=model, custom_llm_provider="gemini"
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
|
||||
transformed_system_messages, new_messages = _transform_system_message(
|
||||
|
|
@ -170,9 +175,14 @@ def transform_openai_messages_to_gemini_context_caching(
|
|||
|
||||
transformed_messages = _gemini_convert_messages_with_history(messages=new_messages)
|
||||
|
||||
model_name = "models/{}".format(model)
|
||||
|
||||
if custom_llm_provider == "vertex_ai" or custom_llm_provider == "vertex_ai_beta":
|
||||
model_name = f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/{model_name}"
|
||||
|
||||
data = CachedContentRequestBody(
|
||||
contents=transformed_messages,
|
||||
model="models/{}".format(model),
|
||||
model=model_name,
|
||||
displayName=cache_key,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -41,8 +41,11 @@ class ContextCachingEndpoints(VertexBase):
|
|||
def _get_token_and_url_context_caching(
|
||||
self,
|
||||
gemini_api_key: Optional[str],
|
||||
custom_llm_provider: Literal["gemini"],
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
api_base: Optional[str],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
) -> Tuple[Optional[str], str]:
|
||||
"""
|
||||
Internal function. Returns the token and url for the call.
|
||||
|
|
@ -58,9 +61,15 @@ class ContextCachingEndpoints(VertexBase):
|
|||
url = "https://generativelanguage.googleapis.com/v1beta/{}?key={}".format(
|
||||
endpoint, gemini_api_key
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
auth_header = vertex_auth_header
|
||||
endpoint = "cachedContents"
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}"
|
||||
else:
|
||||
raise NotImplementedError
|
||||
auth_header = vertex_auth_header
|
||||
endpoint = "cachedContents"
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1beta1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}"
|
||||
|
||||
|
||||
return self._check_custom_proxy(
|
||||
api_base=api_base,
|
||||
|
|
@ -80,6 +89,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
api_key: str,
|
||||
api_base: Optional[str],
|
||||
logging_obj: Logging,
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Checks if content already cached.
|
||||
|
|
@ -94,8 +107,11 @@ class ContextCachingEndpoints(VertexBase):
|
|||
|
||||
_, url = self._get_token_and_url_context_caching(
|
||||
gemini_api_key=api_key,
|
||||
custom_llm_provider="gemini",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
try:
|
||||
## LOGGING
|
||||
|
|
@ -145,6 +161,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
api_key: str,
|
||||
api_base: Optional[str],
|
||||
logging_obj: Logging,
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str]
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Checks if content already cached.
|
||||
|
|
@ -159,8 +179,11 @@ class ContextCachingEndpoints(VertexBase):
|
|||
|
||||
_, url = self._get_token_and_url_context_caching(
|
||||
gemini_api_key=api_key,
|
||||
custom_llm_provider="gemini",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
try:
|
||||
## LOGGING
|
||||
|
|
@ -212,6 +235,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
client: Optional[HTTPHandler],
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
logging_obj: Logging,
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
extra_headers: Optional[dict] = None,
|
||||
cached_content: Optional[str] = None,
|
||||
) -> Tuple[List[AllMessageValues], dict, Optional[str]]:
|
||||
|
|
@ -240,8 +267,11 @@ class ContextCachingEndpoints(VertexBase):
|
|||
## AUTHORIZATION ##
|
||||
token, url = self._get_token_and_url_context_caching(
|
||||
gemini_api_key=api_key,
|
||||
custom_llm_provider="gemini",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
|
||||
headers = {
|
||||
|
|
@ -273,6 +303,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
if google_cache_name:
|
||||
return non_cached_messages, optional_params, google_cache_name
|
||||
|
|
@ -280,7 +314,12 @@ class ContextCachingEndpoints(VertexBase):
|
|||
## TRANSFORM REQUEST
|
||||
cached_content_request_body = (
|
||||
transform_openai_messages_to_gemini_context_caching(
|
||||
model=model, messages=cached_messages, cache_key=generated_cache_key
|
||||
model=model,
|
||||
messages=cached_messages,
|
||||
cache_key=generated_cache_key,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -328,6 +367,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
client: Optional[AsyncHTTPHandler],
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
logging_obj: Logging,
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
extra_headers: Optional[dict] = None,
|
||||
cached_content: Optional[str] = None,
|
||||
) -> Tuple[List[AllMessageValues], dict, Optional[str]]:
|
||||
|
|
@ -356,8 +399,11 @@ class ContextCachingEndpoints(VertexBase):
|
|||
## AUTHORIZATION ##
|
||||
token, url = self._get_token_and_url_context_caching(
|
||||
gemini_api_key=api_key,
|
||||
custom_llm_provider="gemini",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
|
||||
headers = {
|
||||
|
|
@ -386,6 +432,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header
|
||||
)
|
||||
|
||||
if google_cache_name:
|
||||
|
|
@ -394,7 +444,12 @@ class ContextCachingEndpoints(VertexBase):
|
|||
## TRANSFORM REQUEST
|
||||
cached_content_request_body = (
|
||||
transform_openai_messages_to_gemini_context_caching(
|
||||
model=model, messages=cached_messages, cache_key=generated_cache_key
|
||||
model=model,
|
||||
messages=cached_messages,
|
||||
cache_key=generated_cache_key,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -514,34 +514,35 @@ def sync_transform_request_body(
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
litellm_params: dict,
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
) -> RequestBody:
|
||||
from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints
|
||||
|
||||
context_caching_endpoints = ContextCachingEndpoints()
|
||||
|
||||
if gemini_api_key is not None:
|
||||
(
|
||||
messages,
|
||||
optional_params,
|
||||
cached_content,
|
||||
) = context_caching_endpoints.check_and_create_cache(
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=gemini_api_key,
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
extra_headers=extra_headers,
|
||||
cached_content=optional_params.pop("cached_content", None),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
else: # [TODO] implement context caching for gemini as well
|
||||
cached_content = None
|
||||
if "cached_content" in optional_params:
|
||||
cached_content = optional_params.pop("cached_content")
|
||||
elif "cachedContent" in optional_params:
|
||||
cached_content = optional_params.pop("cachedContent")
|
||||
(
|
||||
messages,
|
||||
optional_params,
|
||||
cached_content,
|
||||
) = context_caching_endpoints.check_and_create_cache(
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=gemini_api_key or "dummy",
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
extra_headers=extra_headers,
|
||||
cached_content=optional_params.pop("cached_content", None),
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header,
|
||||
)
|
||||
|
||||
|
||||
return _transform_request_body(
|
||||
messages=messages,
|
||||
|
|
@ -565,34 +566,34 @@ async def async_transform_request_body(
|
|||
logging_obj: litellm.litellm_core_utils.litellm_logging.Logging, # type: ignore
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
litellm_params: dict,
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
) -> RequestBody:
|
||||
from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints
|
||||
|
||||
context_caching_endpoints = ContextCachingEndpoints()
|
||||
|
||||
if gemini_api_key is not None:
|
||||
(
|
||||
messages,
|
||||
optional_params,
|
||||
cached_content,
|
||||
) = await context_caching_endpoints.async_check_and_create_cache(
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=gemini_api_key,
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
extra_headers=extra_headers,
|
||||
cached_content=optional_params.pop("cached_content", None),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
else: # [TODO] implement context caching for gemini as well
|
||||
cached_content = None
|
||||
if "cached_content" in optional_params:
|
||||
cached_content = optional_params.pop("cached_content")
|
||||
elif "cachedContent" in optional_params:
|
||||
cached_content = optional_params.pop("cachedContent")
|
||||
(
|
||||
messages,
|
||||
optional_params,
|
||||
cached_content,
|
||||
) = await context_caching_endpoints.async_check_and_create_cache(
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=gemini_api_key or "dummy",
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
extra_headers=extra_headers,
|
||||
cached_content=optional_params.pop("cached_content", None),
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=vertex_auth_header,
|
||||
)
|
||||
|
||||
return _transform_request_body(
|
||||
messages=messages,
|
||||
|
|
|
|||
|
|
@ -1792,7 +1792,6 @@ class VertexLLM(VertexBase):
|
|||
gemini_api_key: Optional[str] = None,
|
||||
extra_headers: Optional[dict] = None,
|
||||
) -> CustomStreamWrapper:
|
||||
request_body = await async_transform_request_body(**data) # type: ignore
|
||||
|
||||
should_use_v1beta1_features = self.is_using_v1beta1_features(
|
||||
optional_params=optional_params
|
||||
|
|
@ -1826,6 +1825,13 @@ class VertexLLM(VertexBase):
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
request_body = await async_transform_request_body(
|
||||
**data,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=auth_header) # type: ignore
|
||||
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
|
|
@ -1913,7 +1919,12 @@ class VertexLLM(VertexBase):
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
request_body = await async_transform_request_body(**data) # type: ignore
|
||||
request_body = await async_transform_request_body(
|
||||
**data,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=auth_header) # type: ignore
|
||||
|
||||
_async_client_params = {}
|
||||
if timeout:
|
||||
_async_client_params["timeout"] = timeout
|
||||
|
|
@ -2088,7 +2099,11 @@ class VertexLLM(VertexBase):
|
|||
)
|
||||
|
||||
## TRANSFORMATION ##
|
||||
data = sync_transform_request_body(**transform_request_params)
|
||||
data = sync_transform_request_body(
|
||||
**transform_request_params,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header=auth_header)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
|
|||
|
|
@ -2872,7 +2872,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
custom_llm_provider=custom_llm_provider, # type: ignore
|
||||
client=client,
|
||||
api_base=api_base,
|
||||
extra_headers=extra_headers,
|
||||
extra_headers=headers,
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
|
|
@ -2941,7 +2941,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
custom_llm_provider=custom_llm_provider, # type: ignore
|
||||
client=client,
|
||||
api_base=api_base,
|
||||
extra_headers=extra_headers,
|
||||
extra_headers=headers,
|
||||
)
|
||||
elif "openai" in model:
|
||||
# Vertex Model Garden - OpenAI compatible models
|
||||
|
|
|
|||
|
|
@ -22173,6 +22173,307 @@
|
|||
"supports_tool_choice": true,
|
||||
"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": {
|
||||
"input_cost_per_second": 0.0001,
|
||||
"litellm_provider": "openai",
|
||||
|
|
|
|||
|
|
@ -242,12 +242,14 @@ def llm_passthrough_route(
|
|||
request_query_params=request_query_params,
|
||||
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:
|
||||
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)
|
||||
|
||||
|
||||
# Add or update query parameters
|
||||
provider_api_key = provider_config.get_api_key(api_key)
|
||||
|
||||
|
|
|
|||
|
|
@ -333,6 +333,139 @@ class MCPRequestHandler:
|
|||
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
async def _get_key_object_permission(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
):
|
||||
"""Helper to get key object_permission from cache or DB."""
|
||||
from litellm.proxy.auth.auth_checks import get_object_permission
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if not user_api_key_auth:
|
||||
return None
|
||||
|
||||
# Already loaded
|
||||
if user_api_key_auth.object_permission:
|
||||
return user_api_key_auth.object_permission
|
||||
|
||||
# Need to fetch from DB
|
||||
if user_api_key_auth.object_permission_id and prisma_client:
|
||||
return await get_object_permission(
|
||||
object_permission_id=user_api_key_auth.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def _get_team_object_permission(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
):
|
||||
"""Helper to get team object_permission from cache or DB."""
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client:
|
||||
return None
|
||||
|
||||
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
|
||||
team_id=user_api_key_auth.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
return team_obj.object_permission if team_obj else None
|
||||
|
||||
@staticmethod
|
||||
async def get_allowed_tools_for_server(
|
||||
server_id: str,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> Optional[List[str]]:
|
||||
"""
|
||||
Get list of allowed tool names for a specific server based on key/team permissions.
|
||||
Follows same inheritance logic as get_allowed_mcp_servers.
|
||||
|
||||
Args:
|
||||
server_id: Server ID to check permissions for
|
||||
user_api_key_auth: User auth
|
||||
|
||||
Returns:
|
||||
List[str] if restrictions exist, None if no restrictions (allow all)
|
||||
"""
|
||||
if not user_api_key_auth:
|
||||
return None
|
||||
|
||||
try:
|
||||
# Get key and team object permissions
|
||||
key_obj_perm = await MCPRequestHandler._get_key_object_permission(user_api_key_auth)
|
||||
team_obj_perm = await MCPRequestHandler._get_team_object_permission(user_api_key_auth)
|
||||
|
||||
# Extract tool permissions for this server
|
||||
key_tools = key_obj_perm.mcp_tool_permissions.get(server_id) if key_obj_perm and key_obj_perm.mcp_tool_permissions else None
|
||||
team_tools = team_obj_perm.mcp_tool_permissions.get(server_id) if team_obj_perm and team_obj_perm.mcp_tool_permissions else None
|
||||
|
||||
# Apply same inheritance logic as get_allowed_mcp_servers
|
||||
if team_tools:
|
||||
if key_tools:
|
||||
# Both have restrictions → intersection
|
||||
return list(set(team_tools) & set(key_tools))
|
||||
else:
|
||||
# Only team has restrictions → inherit from team
|
||||
return team_tools
|
||||
else:
|
||||
# No team restrictions → use key restrictions
|
||||
return key_tools
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get allowed tools for server: {str(e)}")
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def is_tool_allowed_for_server(
|
||||
tool_name: str,
|
||||
server_id: str,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if a specific tool is allowed for a server based on key/team permissions.
|
||||
|
||||
Args:
|
||||
tool_name: Name of the tool to check
|
||||
server_id: Server ID
|
||||
user_api_key_auth: User auth
|
||||
|
||||
Returns:
|
||||
True if allowed, False if blocked
|
||||
"""
|
||||
allowed_tools = await MCPRequestHandler.get_allowed_tools_for_server(
|
||||
server_id=server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
# None means no restrictions (allow all)
|
||||
if allowed_tools is None:
|
||||
return True
|
||||
|
||||
# Empty list means no tools allowed
|
||||
if not allowed_tools:
|
||||
return False
|
||||
|
||||
# Check if tool is in allowed list
|
||||
return tool_name in allowed_tools
|
||||
|
||||
@staticmethod
|
||||
def is_tool_allowed(
|
||||
allowed_mcp_servers: List[str],
|
||||
|
|
|
|||
|
|
@ -608,6 +608,45 @@ class MCPServerManager:
|
|||
return tool_name not in server.disallowed_tools
|
||||
return True
|
||||
|
||||
async def check_tool_permission_for_key_team(
|
||||
self,
|
||||
tool_name: str,
|
||||
server: MCPServer,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> None:
|
||||
"""
|
||||
Check if a tool is allowed based on key/team object_permission.mcp_tool_permissions.
|
||||
Uses MCPRequestHandler.is_tool_allowed_for_server for consistent inheritance logic.
|
||||
Raises HTTPException if tool is not allowed.
|
||||
|
||||
Args:
|
||||
tool_name: Name of the tool to check
|
||||
server: MCPServer object
|
||||
user_api_key_auth: User authentication
|
||||
|
||||
Raises:
|
||||
HTTPException: If tool is not allowed for this key/team
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
if not user_api_key_auth:
|
||||
return
|
||||
|
||||
# Check if tool is allowed
|
||||
is_allowed = await MCPRequestHandler.is_tool_allowed_for_server(
|
||||
tool_name=tool_name,
|
||||
server_id=server.server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
if not is_allowed:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": f"Tool '{tool_name}' is not allowed for your key/team on server '{server.name}'. Contact proxy admin for access."
|
||||
},
|
||||
)
|
||||
|
||||
async def pre_call_tool_check(
|
||||
self,
|
||||
name: str,
|
||||
|
|
@ -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 = {
|
||||
"name": name,
|
||||
"arguments": arguments,
|
||||
|
|
|
|||
|
|
@ -499,6 +499,13 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
filtered_tools = await filter_tools_by_key_team_permissions(
|
||||
tools=filtered_tools,
|
||||
server_id=server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
all_tools.extend(filtered_tools)
|
||||
|
||||
verbose_logger.debug(
|
||||
|
|
@ -516,6 +523,25 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
return all_tools
|
||||
|
||||
async def filter_tools_by_key_team_permissions(
|
||||
tools: List[MCPTool],
|
||||
server_id: str,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> List[MCPTool]:
|
||||
"""Filter tools based on key/team mcp_tool_permissions."""
|
||||
# Filter by key/team tool-level permissions
|
||||
allowed_tool_names = await MCPRequestHandler.get_allowed_tools_for_server(
|
||||
server_id=server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
if allowed_tool_names is not None:
|
||||
filtered_tools = [t for t in tools if t.name in allowed_tool_names]
|
||||
else:
|
||||
# No restrictions, return all tools
|
||||
filtered_tools = tools
|
||||
|
||||
return filtered_tools
|
||||
|
||||
async def _list_mcp_tools(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
mcp_auth_header: Optional[str] = None,
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
|
|
@ -1 +1 @@
|
|||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[185],{77401:function(n,e,t){Promise.resolve().then(t.t.bind(t,39974,23)),Promise.resolve().then(t.t.bind(t,2778,23))},2778:function(){},39974:function(n){n.exports={style:{fontFamily:"'__Inter_1c856b', '__Inter_Fallback_1c856b'",fontStyle:"normal"},className:"__className_1c856b"}}},function(n){n.O(0,[919,986,971,117,744],function(){return n(n.s=77401)}),_N_E=n.O()}]);
|
||||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[185],{85210:function(n,e,t){Promise.resolve().then(t.t.bind(t,39974,23)),Promise.resolve().then(t.t.bind(t,2778,23))},2778:function(){},39974:function(n){n.exports={style:{fontFamily:"'__Inter_1c856b', '__Inter_Fallback_1c856b'",fontStyle:"normal"},className:"__className_1c856b"}}},function(n){n.O(0,[919,986,971,117,744],function(){return n(n.s=85210)}),_N_E=n.O()}]);
|
||||
|
|
@ -1 +1 @@
|
|||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[418],{96422:function(e,n,t){Promise.resolve().then(t.bind(t,52829))},52829:function(e,n,t){"use strict";t.r(n),t.d(n,{default:function(){return f}});var u=t(57437),s=t(2265),c=t(99376),r=t(72162);function f(){let e=(0,c.useSearchParams)().get("key"),[n,t]=(0,s.useState)(null);return(0,s.useEffect)(()=>{e&&t(e)},[e]),(0,u.jsx)(r.Z,{accessToken:n})}}},function(e){e.O(0,[50,521,154,162,971,117,744],function(){return e(e.s=96422)}),_N_E=e.O()}]);
|
||||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[418],{67355:function(e,n,t){Promise.resolve().then(t.bind(t,52829))},52829:function(e,n,t){"use strict";t.r(n),t.d(n,{default:function(){return f}});var u=t(57437),s=t(2265),c=t(99376),r=t(72162);function f(){let e=(0,c.useSearchParams)().get("key"),[n,t]=(0,s.useState)(null);return(0,s.useEffect)(()=>{e&&t(e)},[e]),(0,u.jsx)(r.Z,{accessToken:n})}}},function(e){e.O(0,[50,521,49,162,971,117,744],function(){return e(e.s=67355)}),_N_E=e.O()}]);
|
||||
|
|
@ -0,0 +1 @@
|
|||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[25],{38520:function(e,n,u){Promise.resolve().then(u.bind(u,22775))},22775:function(e,n,u){"use strict";u.r(n),u.d(n,{default:function(){return f}});var t=u(57437),s=u(2265),r=u(99376),c=u(97851);function f(){let e=(0,r.useSearchParams)().get("key"),[n,u]=(0,s.useState)(null);return(0,s.useEffect)(()=>{e&&u(e)},[e]),(0,t.jsx)(c.Z,{accessToken:n,publicPage:!0,premiumUser:!1,userRole:null})}}},function(e){e.O(0,[50,521,866,49,162,851,971,117,744],function(){return e(e.s=38520)}),_N_E=e.O()}]);
|
||||
|
|
@ -1 +0,0 @@
|
|||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[25],{9397:function(e,n,u){Promise.resolve().then(u.bind(u,22775))},22775:function(e,n,u){"use strict";u.r(n),u.d(n,{default:function(){return f}});var t=u(57437),s=u(2265),r=u(99376),c=u(97851);function f(){let e=(0,r.useSearchParams)().get("key"),[n,u]=(0,s.useState)(null);return(0,s.useEffect)(()=>{e&&u(e)},[e]),(0,t.jsx)(c.Z,{accessToken:n,publicPage:!0,premiumUser:!1,userRole:null})}}},function(e){e.O(0,[50,521,866,154,162,851,971,117,744],function(){return e(e.s=9397)}),_N_E=e.O()}]);
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
|
|
@ -1 +1 @@
|
|||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[744],{60400:function(e,n,t){Promise.resolve().then(t.t.bind(t,12846,23)),Promise.resolve().then(t.t.bind(t,19107,23)),Promise.resolve().then(t.t.bind(t,61060,23)),Promise.resolve().then(t.t.bind(t,4707,23)),Promise.resolve().then(t.t.bind(t,80,23)),Promise.resolve().then(t.t.bind(t,36423,23))}},function(e){var n=function(n){return e(e.s=n)};e.O(0,[971,117],function(){return n(54278),n(60400)}),_N_E=e.O()}]);
|
||||
(self.webpackChunk_N_E=self.webpackChunk_N_E||[]).push([[744],{78483:function(e,n,t){Promise.resolve().then(t.t.bind(t,12846,23)),Promise.resolve().then(t.t.bind(t,19107,23)),Promise.resolve().then(t.t.bind(t,61060,23)),Promise.resolve().then(t.t.bind(t,4707,23)),Promise.resolve().then(t.t.bind(t,80,23)),Promise.resolve().then(t.t.bind(t,36423,23))}},function(e){var n=function(n){return e(e.s=n)};e.O(0,[971,117],function(){return n(54278),n(78483)}),_N_E=e.O()}]);
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
|
|
@ -0,0 +1,9 @@
|
|||
<svg id="Layer_1" data-name="Layer 1" xmlns="http://www.w3.org/2000/svg" viewBox="0 0 146.36 139.16" xmlns:xlink="http://www.w3.org/1999/xlink">
|
||||
<defs>
|
||||
<style>
|
||||
.cls-1{fill:#29b5e8;fill-rule:evenodd;}
|
||||
</style>
|
||||
</defs>
|
||||
<path class="cls-1" d="M134.81,60.1l-16.47,9.49L134.81,79a8.65,8.65,0,1,1-8.67,15l-29.51-17a8.68,8.68,0,0,1-4.33-7.75,8.48,8.48,0,0,1,.31-2,8.68,8.68,0,0,1,4-5.19l29.51-16.94A8.69,8.69,0,0,1,138,48.31,8.58,8.58,0,0,1,134.81,60.1Zm-15.59,46L89.72,89.13a8.72,8.72,0,0,0-13.06,7.48v33.9a8.69,8.69,0,0,0,17.37,0v-19L110.54,121a8.66,8.66,0,1,0,8.68-15Zm-34-33.16L72.92,85.09a2.44,2.44,0,0,1-1.54.65H67.77a2.51,2.51,0,0,1-1.54-.65L54,72.9a2.45,2.45,0,0,1-.64-1.52v-3.6A2.5,2.5,0,0,1,54,66.25L66.23,54.06a2.5,2.5,0,0,1,1.54-.64h3.61a2.45,2.45,0,0,1,1.54.64L85.18,66.25a2.49,2.49,0,0,1,.63,1.53v3.6A2.44,2.44,0,0,1,85.18,72.9Zm-9.8-3.38A2.59,2.59,0,0,0,74.73,68l-3.55-3.51a2.51,2.51,0,0,0-1.54-.64h-.13a2.46,2.46,0,0,0-1.53.64L64.43,68a2.51,2.51,0,0,0-.63,1.55v.13a2.41,2.41,0,0,0,.63,1.52L68,74.7a2.48,2.48,0,0,0,1.53.64h.13a2.51,2.51,0,0,0,1.54-.64l3.55-3.53a2.49,2.49,0,0,0,.65-1.52ZM19.93,33.08,49.44,50a8.73,8.73,0,0,0,13.07-7.49V8.64a8.69,8.69,0,0,0-17.37,0v19l-16.53-9.5a8.65,8.65,0,1,0-8.68,15ZM84.69,51.16a8.64,8.64,0,0,0,5-1.13l29.5-17a8.65,8.65,0,1,0-8.68-15L94,27.61v-19a8.69,8.69,0,0,0-17.37,0v33.9A8.66,8.66,0,0,0,84.69,51.16ZM54.48,88a8.58,8.58,0,0,0-5,1.13L19.93,106.06a8.66,8.66,0,1,0,8.68,15l16.53-9.49v19a8.69,8.69,0,0,0,17.37,0V96.61A8.65,8.65,0,0,0,54.48,88Zm-8-15.87a8.61,8.61,0,0,0-4-10L13,45.14A8.69,8.69,0,0,0,1.17,48.31,8.59,8.59,0,0,0,4.35,60.1l16.47,9.49L4.35,79A8.65,8.65,0,1,0,13,94l29.48-17A8.59,8.59,0,0,0,46.47,72.13Zm93.15-56.22H138.3v1.63h1.32c.61,0,1-.28,1-.8S140.26,15.91,139.62,15.91Zm-2.94-1.5h3c1.62,0,2.7.89,2.7,2.27a2.16,2.16,0,0,1-1.08,1.9l1.17,1.68v.34h-1.69L139.62,19H138.3V20.6h-1.62Zm8.3,3.22a5.48,5.48,0,0,0-5.58-5.83c-3.31,0-5.51,2.39-5.51,5.83,0,3.28,2.2,5.82,5.51,5.82A5.47,5.47,0,0,0,145,17.63Zm1.38,0c0,3.89-2.6,7.14-7,7.14s-6.89-3.28-6.89-7.14,2.57-7.14,6.89-7.14S146.36,13.73,146.36,17.63Z">
|
||||
</path>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 2 KiB |
File diff suppressed because one or more lines are too long
|
|
@ -1,7 +1,7 @@
|
|||
2:I[19107,[],"ClientPageRoot"]
|
||||
3:I[55139,["665","static/chunks/3014691f-b7b79b78e27792f3.js","990","static/chunks/13b76428-ebdf3012af0e4489.js","50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","866","static/chunks/866-3523e0e07cf314f6.js","313","static/chunks/313-27c820a98e9413e5.js","154","static/chunks/154-f87cf692dcea3018.js","162","static/chunks/162-dd6427ff1a4ad9f4.js","851","static/chunks/851-bbe6d02cf41bb87a.js","931","static/chunks/app/page-f400068ac45ce482.js"],"default",1]
|
||||
3:I[73148,["665","static/chunks/3014691f-b7b79b78e27792f3.js","990","static/chunks/13b76428-ebdf3012af0e4489.js","50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","866","static/chunks/866-9e1803a09e9ae8da.js","313","static/chunks/313-0025fb08e386c4b8.js","49","static/chunks/49-b6f167418ea8dbf4.js","162","static/chunks/162-ffcd7d9fbb033bdf.js","851","static/chunks/851-0a73701a3a0b0187.js","931","static/chunks/app/page-5b9ff2d173a47e2c.js"],"default",1]
|
||||
4:I[4707,[],""]
|
||||
5:I[36423,[],""]
|
||||
0:["WkpkdsewrdPMuTzVGS_5j",[[["",{"children":["__PAGE__",{}]},"$undefined","$undefined",true],["",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/4103fa525703177b.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
0:["_S0i_Y-CCoYQWc9dIuLxF",[[["",{"children":["__PAGE__",{}]},"$undefined","$undefined",true],["",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/c200be8dd8638678.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
6:[["$","meta","0",{"name":"viewport","content":"width=device-width, initial-scale=1"}],["$","meta","1",{"charSet":"utf-8"}],["$","title","2",{"children":"LiteLLM Dashboard"}],["$","meta","3",{"name":"description","content":"LiteLLM Proxy Admin UI"}],["$","link","4",{"rel":"icon","href":"/favicon.ico","type":"image/x-icon","sizes":"16x16"}],["$","link","5",{"rel":"icon","href":"./favicon.ico"}],["$","meta","6",{"name":"next-size-adjust"}]]
|
||||
1:null
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
2:I[19107,[],"ClientPageRoot"]
|
||||
3:I[52829,["50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","154","static/chunks/154-f87cf692dcea3018.js","162","static/chunks/162-dd6427ff1a4ad9f4.js","418","static/chunks/app/model_hub/page-237d2973f13202c4.js"],"default",1]
|
||||
3:I[52829,["50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","49","static/chunks/49-b6f167418ea8dbf4.js","162","static/chunks/162-ffcd7d9fbb033bdf.js","418","static/chunks/app/model_hub/page-d7915f579b770030.js"],"default",1]
|
||||
4:I[4707,[],""]
|
||||
5:I[36423,[],""]
|
||||
0:["WkpkdsewrdPMuTzVGS_5j",[[["",{"children":["model_hub",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["model_hub",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","model_hub","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/4103fa525703177b.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
0:["_S0i_Y-CCoYQWc9dIuLxF",[[["",{"children":["model_hub",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["model_hub",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","model_hub","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/c200be8dd8638678.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
6:[["$","meta","0",{"name":"viewport","content":"width=device-width, initial-scale=1"}],["$","meta","1",{"charSet":"utf-8"}],["$","title","2",{"children":"LiteLLM Dashboard"}],["$","meta","3",{"name":"description","content":"LiteLLM Proxy Admin UI"}],["$","link","4",{"rel":"icon","href":"/favicon.ico","type":"image/x-icon","sizes":"16x16"}],["$","link","5",{"rel":"icon","href":"./favicon.ico"}],["$","meta","6",{"name":"next-size-adjust"}]]
|
||||
1:null
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
2:I[19107,[],"ClientPageRoot"]
|
||||
3:I[22775,["50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","866","static/chunks/866-3523e0e07cf314f6.js","154","static/chunks/154-f87cf692dcea3018.js","162","static/chunks/162-dd6427ff1a4ad9f4.js","851","static/chunks/851-bbe6d02cf41bb87a.js","25","static/chunks/app/model_hub_table/page-5d1aa98a47f9e9fd.js"],"default",1]
|
||||
3:I[22775,["50","static/chunks/50-d0da2dd7acce2eb9.js","521","static/chunks/521-d97d355792d44830.js","866","static/chunks/866-9e1803a09e9ae8da.js","49","static/chunks/49-b6f167418ea8dbf4.js","162","static/chunks/162-ffcd7d9fbb033bdf.js","851","static/chunks/851-0a73701a3a0b0187.js","25","static/chunks/app/model_hub_table/page-2f23c22a47b20607.js"],"default",1]
|
||||
4:I[4707,[],""]
|
||||
5:I[36423,[],""]
|
||||
0:["WkpkdsewrdPMuTzVGS_5j",[[["",{"children":["model_hub_table",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["model_hub_table",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","model_hub_table","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/4103fa525703177b.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
0:["_S0i_Y-CCoYQWc9dIuLxF",[[["",{"children":["model_hub_table",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["model_hub_table",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","model_hub_table","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/c200be8dd8638678.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
6:[["$","meta","0",{"name":"viewport","content":"width=device-width, initial-scale=1"}],["$","meta","1",{"charSet":"utf-8"}],["$","title","2",{"children":"LiteLLM Dashboard"}],["$","meta","3",{"name":"description","content":"LiteLLM Proxy Admin UI"}],["$","link","4",{"rel":"icon","href":"/favicon.ico","type":"image/x-icon","sizes":"16x16"}],["$","link","5",{"rel":"icon","href":"./favicon.ico"}],["$","meta","6",{"name":"next-size-adjust"}]]
|
||||
1:null
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
|
|
@ -1,7 +1,7 @@
|
|||
2:I[19107,[],"ClientPageRoot"]
|
||||
3:I[12011,["665","static/chunks/3014691f-b7b79b78e27792f3.js","50","static/chunks/50-d0da2dd7acce2eb9.js","154","static/chunks/154-f87cf692dcea3018.js","461","static/chunks/app/onboarding/page-099f7aa4c559d470.js"],"default",1]
|
||||
3:I[12011,["665","static/chunks/3014691f-b7b79b78e27792f3.js","50","static/chunks/50-d0da2dd7acce2eb9.js","49","static/chunks/49-b6f167418ea8dbf4.js","461","static/chunks/app/onboarding/page-1fd26064f407ad1d.js"],"default",1]
|
||||
4:I[4707,[],""]
|
||||
5:I[36423,[],""]
|
||||
0:["WkpkdsewrdPMuTzVGS_5j",[[["",{"children":["onboarding",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["onboarding",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","onboarding","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/4103fa525703177b.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
0:["_S0i_Y-CCoYQWc9dIuLxF",[[["",{"children":["onboarding",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],["",{"children":["onboarding",{"children":["__PAGE__",{},[["$L1",["$","$L2",null,{"props":{"params":{},"searchParams":{}},"Component":"$3"}],null],null],null]},[null,["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children","onboarding","children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","notFoundStyles":"$undefined"}]],null]},[[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/349654da14372cd9.css","precedence":"next","crossOrigin":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/css/c200be8dd8638678.css","precedence":"next","crossOrigin":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"__className_1c856b","children":["$","$L4",null,{"parallelRouterKey":"children","segmentPath":["children"],"error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":"404"}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],"notFoundStyles":[]}]}]}]],null],null],["$L6",null]]]]
|
||||
6:[["$","meta","0",{"name":"viewport","content":"width=device-width, initial-scale=1"}],["$","meta","1",{"charSet":"utf-8"}],["$","title","2",{"children":"LiteLLM Dashboard"}],["$","meta","3",{"name":"description","content":"LiteLLM Proxy Admin UI"}],["$","link","4",{"rel":"icon","href":"/favicon.ico","type":"image/x-icon","sizes":"16x16"}],["$","link","5",{"rel":"icon","href":"./favicon.ico"}],["$","meta","6",{"name":"next-size-adjust"}]]
|
||||
1:null
|
||||
|
|
|
|||
|
|
@ -1,26 +1,6 @@
|
|||
model_list:
|
||||
- model_name: openai/gpt-4o
|
||||
- model_name: gpt-5-mini
|
||||
litellm_params:
|
||||
model: openai/gpt-4o-mini
|
||||
api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5"
|
||||
api_key: dummy
|
||||
- model_name: "byok-wildcard/*"
|
||||
litellm_params:
|
||||
model: openai/*
|
||||
- model_name: xai-grok-3
|
||||
litellm_params:
|
||||
model: xai/grok-3
|
||||
- model_name: hosted_vllm/whisper-v3
|
||||
litellm_params:
|
||||
model: hosted_vllm/whisper-v3
|
||||
api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5"
|
||||
api_key: dummy
|
||||
|
||||
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"}
|
||||
model: azure/gpt-5-mini-2
|
||||
api_key: os.environ/AZURE_API_KEY_ALT
|
||||
api_base: os.environ/AZURE_API_BASE_ALT
|
||||
|
|
|
|||
|
|
@ -52,6 +52,24 @@ else:
|
|||
Span = Any
|
||||
|
||||
|
||||
class SupportedDBObjectType(str, enum.Enum):
|
||||
"""
|
||||
Supported database object types for fine-grained DB storage control.
|
||||
Use in general_settings.supported_db_objects to specify which objects to load from DB.
|
||||
"""
|
||||
|
||||
MODELS = "models"
|
||||
MCP = "mcp"
|
||||
GUARDRAILS = "guardrails"
|
||||
VECTOR_STORES = "vector_stores"
|
||||
PASS_THROUGH_ENDPOINTS = "pass_through_endpoints"
|
||||
PROMPTS = "prompts"
|
||||
MODEL_COST_MAP = "model_cost_map"
|
||||
|
||||
def __str__(self):
|
||||
return str(self.value)
|
||||
|
||||
|
||||
class LiteLLMTeamRoles(enum.Enum):
|
||||
# team admin
|
||||
TEAM_ADMIN = "admin"
|
||||
|
|
@ -712,6 +730,7 @@ class ModelParams(LiteLLMPydanticObjectBase):
|
|||
class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase):
|
||||
mcp_servers: Optional[List[str]] = None
|
||||
mcp_access_groups: Optional[List[str]] = None
|
||||
mcp_tool_permissions: Optional[Dict[str, List[str]]] = None
|
||||
vector_stores: Optional[List[str]] = None
|
||||
|
||||
|
||||
|
|
@ -1414,6 +1433,16 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase):
|
|||
object_permission_id: str
|
||||
mcp_servers: Optional[List[str]] = []
|
||||
mcp_access_groups: Optional[List[str]] = []
|
||||
mcp_tool_permissions: Optional[Dict[str, List[str]]] = None
|
||||
"""
|
||||
Mapping - server_id -> list of tools
|
||||
|
||||
Enforces allowed tools for a specific key/team/organization
|
||||
{
|
||||
"1234567890": ["tool_name_1", "tool_name_2"]
|
||||
}
|
||||
"""
|
||||
|
||||
vector_stores: Optional[List[str]] = []
|
||||
|
||||
|
||||
|
|
@ -1749,6 +1778,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
description="[DEPRECATED] Use 'user_header_mappings' instead. When set, the header value is treated as the end user id unless overridden by user_header_mappings.",
|
||||
)
|
||||
user_header_mappings: Optional[List[UserHeaderMapping]] = None
|
||||
supported_db_objects: Optional[List[SupportedDBObjectType]] = Field(
|
||||
None,
|
||||
description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map'. If not set, all objects are loaded (default behavior).",
|
||||
)
|
||||
|
||||
|
||||
class ConfigYAML(LiteLLMPydanticObjectBase):
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ import fastapi
|
|||
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching import DualCache
|
||||
|
|
@ -837,7 +838,7 @@ async def generate_key_fn(
|
|||
- enforced_params: Optional[List[str]] - List of enforced params for the key (Enterprise only). [Docs](https://docs.litellm.ai/docs/proxy/enterprise#enforce-required-params-for-llm-requests)
|
||||
- prompts: Optional[List[str]] - List of prompts that the key is allowed to use.
|
||||
- allowed_routes: Optional[list] - List of allowed routes for the key. Store the actual route or store a wildcard pattern for a set of routes. Example - ["/chat/completions", "/embeddings", "/keys/*"]
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission.
|
||||
- key_type: Optional[str] - Type of key that determines default allowed routes. Options: "llm_api" (can call LLM API routes), "management" (can call management routes), "read_only" (can only call info/read routes), "default" (uses default allowed routes). Defaults to "default".
|
||||
- prompts: Optional[List[str]] - List of allowed prompts for the key. If specified, the key will only be able to use these specific prompts.
|
||||
- auto_rotate: Optional[bool] - Whether this key should be automatically rotated (regenerated)
|
||||
|
|
@ -987,7 +988,7 @@ async def generate_service_account_key_fn(
|
|||
- tags: Optional[List[str]] - Tags for [tracking spend](https://litellm.vercel.app/docs/proxy/enterprise#tracking-spend-for-custom-tags) and/or doing [tag-based routing](https://litellm.vercel.app/docs/proxy/tag_routing).
|
||||
- enforced_params: Optional[List[str]] - List of enforced params for the key (Enterprise only). [Docs](https://docs.litellm.ai/docs/proxy/enterprise#enforce-required-params-for-llm-requests)
|
||||
- allowed_routes: Optional[list] - List of allowed routes for the key. Store the actual route or store a wildcard pattern for a set of routes. Example - ["/chat/completions", "/embeddings", "/keys/*"]
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission.
|
||||
Examples:
|
||||
|
||||
1. Allow users to turn on/off pii masking
|
||||
|
|
@ -1125,6 +1126,13 @@ async def _set_object_permission(
|
|||
return data_json
|
||||
|
||||
if "object_permission" in data_json:
|
||||
# Serialize mcp_tool_permissions JSON field to avoid GraphQL parsing issues
|
||||
# (e.g., server IDs starting with "3e64" being interpreted as floats)
|
||||
if "mcp_tool_permissions" in data_json["object_permission"]:
|
||||
data_json["object_permission"]["mcp_tool_permissions"] = safe_dumps(
|
||||
data_json["object_permission"]["mcp_tool_permissions"]
|
||||
)
|
||||
|
||||
created_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=data_json["object_permission"],
|
||||
|
|
@ -1300,7 +1308,7 @@ async def update_key_fn(
|
|||
- temp_budget_expiry: Optional[str] - Expiry time for the temporary budget increase (Enterprise only).
|
||||
- allowed_routes: Optional[list] - List of allowed routes for the key. Store the actual route or store a wildcard pattern for a set of routes. Example - ["/chat/completions", "/embeddings", "/keys/*"]
|
||||
- prompts: Optional[List[str]] - List of allowed prompts for the key. If specified, the key will only be able to use these specific prompts.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission.
|
||||
- auto_rotate: Optional[bool] - Whether this key should be automatically rotated
|
||||
- rotation_interval: Optional[str] - How often to rotate this key (e.g., '30d', '90d'). Required if auto_rotate=True
|
||||
Example:
|
||||
|
|
@ -2841,6 +2849,79 @@ async def list_keys(
|
|||
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/key/aliases",
|
||||
tags=["key management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def key_aliases() -> Dict[str, List[str]]:
|
||||
"""
|
||||
Lists all key aliases
|
||||
|
||||
Returns:
|
||||
{
|
||||
"aliases": List[str]
|
||||
}
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
verbose_proxy_logger.debug("Entering key_aliases function")
|
||||
|
||||
if prisma_client is None:
|
||||
verbose_proxy_logger.error("Database not connected")
|
||||
raise Exception("Database not connected")
|
||||
|
||||
where: Dict[str, Any] = {}
|
||||
try:
|
||||
where.update(_get_condition_to_filter_out_ui_session_tokens())
|
||||
except NameError:
|
||||
# Helper may not exist in some builds; ignore if missing
|
||||
pass
|
||||
|
||||
rows = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where=where,
|
||||
order=[{"key_alias": "asc"}],
|
||||
)
|
||||
|
||||
seen = set()
|
||||
aliases: List[str] = []
|
||||
for row in rows:
|
||||
alias = getattr(row, "key_alias", None)
|
||||
if alias is None and isinstance(row, dict):
|
||||
alias = row.get("key_alias")
|
||||
|
||||
if not alias:
|
||||
continue
|
||||
|
||||
alias_str = str(alias).strip()
|
||||
if alias_str and alias_str not in seen:
|
||||
seen.add(alias_str)
|
||||
aliases.append(alias_str)
|
||||
|
||||
verbose_proxy_logger.debug(f"Returning {len(aliases)} key aliases")
|
||||
|
||||
return {"aliases": aliases}
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error in key_aliases: {e}")
|
||||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "detail", f"error({str(e)})"),
|
||||
type=ProxyErrorTypes.internal_server_error,
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR),
|
||||
)
|
||||
elif isinstance(e, ProxyException):
|
||||
raise e
|
||||
raise ProxyException(
|
||||
message="Authentication Error, " + str(e),
|
||||
type=ProxyErrorTypes.internal_server_error,
|
||||
param=getattr(e, "param", "None"),
|
||||
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
)
|
||||
|
||||
|
||||
def _validate_sort_params(
|
||||
sort_by: Optional[str], sort_order: str
|
||||
|
|
|
|||
|
|
@ -311,7 +311,7 @@ async def new_team( # noqa: PLR0915
|
|||
- model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias)
|
||||
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
|
||||
- prompts: Optional[List[str]] - List of prompts that the team is allowed to use.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission.
|
||||
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
|
||||
- team_member_rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for individual team members.
|
||||
- team_member_tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for individual team members.
|
||||
|
|
@ -769,7 +769,7 @@ async def update_team(
|
|||
- model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias)
|
||||
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
|
||||
- prompts: Optional[List[str]] - List of prompts that the team is allowed to use.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission.
|
||||
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
|
||||
- team_member_rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for individual team members.
|
||||
- team_member_tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for individual team members.
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ from typing import Dict, Optional, Union
|
|||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
|
||||
|
||||
async def attach_object_permission_to_dict(
|
||||
|
|
@ -114,6 +116,15 @@ async def handle_update_object_permission_common(
|
|||
if isinstance(new_object_permission, dict):
|
||||
existing_object_permissions_dict.update(new_object_permission)
|
||||
|
||||
#########################################################
|
||||
# Serialize mcp_tool_permissions JSON field to avoid GraphQL parsing issues
|
||||
# (e.g., server IDs starting with "3e64" being interpreted as floats)
|
||||
#########################################################
|
||||
if "mcp_tool_permissions" in existing_object_permissions_dict:
|
||||
existing_object_permissions_dict["mcp_tool_permissions"] = safe_dumps(
|
||||
existing_object_permissions_dict["mcp_tool_permissions"]
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Commit the update to the LiteLLM_ObjectPermissionTable
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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.proxy._types import *
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
user_api_key_auth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_read_request_body,
|
||||
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.pass_through_endpoints import (
|
||||
HttpPassThroughEndpointHelpers,
|
||||
create_pass_through_route,
|
||||
create_websocket_passthrough_route,
|
||||
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
|
||||
"""
|
||||
|
|
@ -93,12 +94,16 @@ async def llm_passthrough_factory_proxy_route(
|
|||
model=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()
|
||||
|
||||
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
|
||||
|
||||
|
|
@ -177,11 +182,17 @@ async def gemini_proxy_route(
|
|||
[Docs](https://docs.litellm.ai/docs/pass_through/google_ai_studio)
|
||||
"""
|
||||
## 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
|
||||
|
||||
# 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)
|
||||
"""
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
HttpPassThroughEndpointHelpers,
|
||||
)
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
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)
|
||||
if is_router_model and llm_router:
|
||||
result = cast(
|
||||
|
|
@ -314,7 +324,11 @@ async def vllm_proxy_route(
|
|||
content=None,
|
||||
data=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,
|
||||
headers=None,
|
||||
cookies=None,
|
||||
|
|
@ -492,7 +506,9 @@ async def handle_bedrock_count_tokens(
|
|||
# Extract model from request body
|
||||
model = request_body.get("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
|
||||
litellm_params = {"user_api_key_dict": user_api_key_dict}
|
||||
|
|
@ -531,7 +547,9 @@ async def handle_bedrock_count_tokens(
|
|||
raise
|
||||
except Exception as 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(
|
||||
|
|
@ -583,7 +601,8 @@ async def bedrock_llm_proxy_route(
|
|||
raise HTTPException(
|
||||
status_code=400,
|
||||
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")
|
||||
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:
|
||||
return await bedrock_llm_proxy_route(
|
||||
endpoint=endpoint,
|
||||
|
|
@ -677,7 +698,9 @@ async def bedrock_proxy_route(
|
|||
data = await request.json()
|
||||
except Exception as 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)
|
||||
prepped = _request.prepare()
|
||||
|
||||
|
|
@ -738,8 +761,14 @@ async def assemblyai_proxy_route(
|
|||
[Docs](https://api.assemblyai.com)
|
||||
"""
|
||||
# Set base URL based on the route
|
||||
assembly_region = AssemblyAIPassthroughLoggingHandler._get_assembly_region_from_url(url=str(request.url))
|
||||
base_target_url = AssemblyAIPassthroughLoggingHandler._get_assembly_base_url_from_region(region=assembly_region)
|
||||
assembly_region = AssemblyAIPassthroughLoggingHandler._get_assembly_region_from_url(
|
||||
url=str(request.url)
|
||||
)
|
||||
base_target_url = (
|
||||
AssemblyAIPassthroughLoggingHandler._get_assembly_base_url_from_region(
|
||||
region=assembly_region
|
||||
)
|
||||
)
|
||||
encoded_endpoint = httpx.URL(endpoint).path
|
||||
# Ensure endpoint starts with '/' for proper URL construction
|
||||
if not encoded_endpoint.startswith("/"):
|
||||
|
|
@ -794,17 +823,91 @@ async def azure_proxy_route(
|
|||
Call any azure endpoint using the proxy.
|
||||
|
||||
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")
|
||||
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
|
||||
azure_api_key = passthrough_endpoint_router.get_credentials(
|
||||
custom_llm_provider=litellm.LlmProviders.AZURE.value,
|
||||
region_name=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(
|
||||
endpoint=endpoint,
|
||||
|
|
@ -828,7 +931,9 @@ class BaseVertexAIPassThroughHandler(ABC):
|
|||
|
||||
@staticmethod
|
||||
@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
|
||||
|
||||
|
||||
|
|
@ -838,7 +943,9 @@ class VertexAIDiscoveryPassThroughHandler(BaseVertexAIPassThroughHandler):
|
|||
return "https://discoveryengine.googleapis.com/"
|
||||
|
||||
@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
|
||||
|
||||
|
||||
|
|
@ -848,7 +955,9 @@ class VertexAIPassThroughHandler(BaseVertexAIPassThroughHandler):
|
|||
return get_vertex_base_url(vertex_location)
|
||||
|
||||
@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)
|
||||
|
||||
|
||||
|
|
@ -914,14 +1023,18 @@ async def _base_vertex_proxy_route(
|
|||
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
|
||||
# Use headers from the incoming request if no vertex credentials are found
|
||||
if vertex_credentials is None or vertex_credentials.vertex_project is None:
|
||||
headers = dict(request.headers) or {}
|
||||
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("host", None)
|
||||
else:
|
||||
|
|
@ -1087,7 +1200,9 @@ async def openai_proxy_route(
|
|||
region_name=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(
|
||||
endpoint=endpoint,
|
||||
|
|
@ -1133,7 +1248,9 @@ class BaseOpenAIPassThroughHandler:
|
|||
endpoint_func = create_pass_through_route(
|
||||
endpoint=endpoint,
|
||||
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
|
||||
received_value = await endpoint_func(
|
||||
request,
|
||||
|
|
@ -1150,7 +1267,10 @@ class BaseOpenAIPassThroughHandler:
|
|||
"""
|
||||
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"
|
||||
return headers
|
||||
|
||||
|
|
@ -1166,7 +1286,9 @@ class BaseOpenAIPassThroughHandler:
|
|||
)
|
||||
|
||||
@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.
|
||||
"""
|
||||
|
|
@ -1182,9 +1304,14 @@ class BaseOpenAIPassThroughHandler:
|
|||
joined_path_str = str(base_url.copy_with(path=full_path))
|
||||
|
||||
# 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
|
||||
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
|
||||
|
||||
|
|
@ -1231,9 +1358,7 @@ async def vertex_ai_live_websocket_passthrough(
|
|||
|
||||
if vertex_credentials_config is not None:
|
||||
resolved_project = resolved_project or vertex_credentials_config.vertex_project
|
||||
temp_location = (
|
||||
resolved_location or vertex_credentials_config.vertex_location
|
||||
)
|
||||
temp_location = resolved_location or vertex_credentials_config.vertex_location
|
||||
# Ensure resolved_location is a string
|
||||
if isinstance(temp_location, dict):
|
||||
resolved_location = str(temp_location)
|
||||
|
|
@ -1241,7 +1366,11 @@ async def vertex_ai_live_websocket_passthrough(
|
|||
resolved_location = str(temp_location)
|
||||
else:
|
||||
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:
|
||||
resolved_location = resolved_location or (
|
||||
|
|
@ -1302,7 +1431,7 @@ async def vertex_ai_live_websocket_passthrough(
|
|||
# Use the new WebSocket passthrough pattern
|
||||
if user_api_key_dict is None:
|
||||
raise ValueError("user_api_key_dict is required for WebSocket passthrough")
|
||||
|
||||
|
||||
return await websocket_passthrough_request(
|
||||
websocket=websocket,
|
||||
target=service_url,
|
||||
|
|
|
|||
|
|
@ -2957,6 +2957,40 @@ class ProxyConfig:
|
|||
|
||||
return config
|
||||
|
||||
def _should_load_db_object(
|
||||
self, object_type: Union[str, SupportedDBObjectType]
|
||||
) -> bool:
|
||||
"""
|
||||
Check if an object type should be loaded from the database based on general_settings.supported_db_objects.
|
||||
|
||||
Args:
|
||||
object_type: Type of object to check (e.g., SupportedDBObjectType.MODELS, "models", etc.)
|
||||
|
||||
Returns:
|
||||
True if the object should be loaded, False otherwise
|
||||
"""
|
||||
global general_settings
|
||||
|
||||
# Get the supported_db_objects configuration
|
||||
supported_db_objects = general_settings.get("supported_db_objects", None)
|
||||
|
||||
# If supported_db_objects is not set, load all objects (default behavior)
|
||||
if supported_db_objects is None:
|
||||
return True
|
||||
|
||||
# If supported_db_objects is set, only load specified objects
|
||||
if not isinstance(supported_db_objects, list):
|
||||
verbose_proxy_logger.warning(
|
||||
f"supported_db_objects is not a list, got {type(supported_db_objects)}. Loading all objects."
|
||||
)
|
||||
return True
|
||||
|
||||
# Convert object_type to string for comparison (handles both str and enum)
|
||||
object_type_str = str(object_type)
|
||||
|
||||
# Check if the object type is in the list (supports both str and enum values)
|
||||
return any(str(obj) == object_type_str for obj in supported_db_objects)
|
||||
|
||||
async def _get_models_from_db(self, prisma_client: PrismaClient) -> list:
|
||||
try:
|
||||
new_models = await prisma_client.db.litellm_proxymodeltable.find_many()
|
||||
|
|
@ -2988,12 +3022,14 @@ class ProxyConfig:
|
|||
f"Master key is not initialized or formatted. master_key={master_key}"
|
||||
)
|
||||
|
||||
new_models = await self._get_models_from_db(prisma_client=prisma_client)
|
||||
# Only load models from DB if "models" is in supported_db_objects (or if supported_db_objects is not set)
|
||||
if self._should_load_db_object(object_type="models"):
|
||||
new_models = await self._get_models_from_db(prisma_client=prisma_client)
|
||||
|
||||
# update llm router
|
||||
await self._update_llm_router(
|
||||
new_models=new_models, proxy_logging_obj=proxy_logging_obj
|
||||
)
|
||||
# update llm router
|
||||
await self._update_llm_router(
|
||||
new_models=new_models, proxy_logging_obj=proxy_logging_obj
|
||||
)
|
||||
|
||||
db_general_settings = await prisma_client.db.litellm_config.find_first(
|
||||
where={"param_name": "general_settings"}
|
||||
|
|
@ -3021,12 +3057,23 @@ class ProxyConfig:
|
|||
|
||||
ex. Vector Stores, Guardrails, MCP tools, etc.
|
||||
"""
|
||||
await self._init_guardrails_in_db(prisma_client=prisma_client)
|
||||
await self._init_vector_stores_in_db(prisma_client=prisma_client)
|
||||
await self._init_mcp_servers_in_db()
|
||||
await self._init_pass_through_endpoints_in_db()
|
||||
await self._init_prompts_in_db(prisma_client=prisma_client)
|
||||
await self._check_and_reload_model_cost_map(prisma_client=prisma_client)
|
||||
if self._should_load_db_object(object_type="guardrails"):
|
||||
await self._init_guardrails_in_db(prisma_client=prisma_client)
|
||||
|
||||
if self._should_load_db_object(object_type="vector_stores"):
|
||||
await self._init_vector_stores_in_db(prisma_client=prisma_client)
|
||||
|
||||
if self._should_load_db_object(object_type="mcp"):
|
||||
await self._init_mcp_servers_in_db()
|
||||
|
||||
if self._should_load_db_object(object_type="pass_through_endpoints"):
|
||||
await self._init_pass_through_endpoints_in_db()
|
||||
|
||||
if self._should_load_db_object(object_type="prompts"):
|
||||
await self._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
if self._should_load_db_object(object_type="model_cost_map"):
|
||||
await self._check_and_reload_model_cost_map(prisma_client=prisma_client)
|
||||
|
||||
async def _check_and_reload_model_cost_map(self, prisma_client: PrismaClient):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -156,6 +156,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
object_permission_id String @id @default(uuid())
|
||||
mcp_servers String[] @default([])
|
||||
mcp_access_groups String[] @default([])
|
||||
mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]}
|
||||
vector_stores String[] @default([])
|
||||
teams LiteLLM_TeamTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
|
|
|
|||
|
|
@ -3601,8 +3601,10 @@ def is_known_model(model: Optional[str], llm_router: Optional[Router]) -> bool:
|
|||
return False
|
||||
model_names = llm_router.get_model_names()
|
||||
|
||||
model_names_set = set(model_names)
|
||||
|
||||
is_in_list = False
|
||||
if model in model_names:
|
||||
if model in model_names_set:
|
||||
is_in_list = True
|
||||
|
||||
return is_in_list
|
||||
|
|
|
|||
|
|
@ -416,6 +416,9 @@ class Router:
|
|||
|
||||
# Initialize model ID to deployment index mapping for O(1) lookups
|
||||
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:
|
||||
# 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)
|
||||
self.model_list = []
|
||||
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
|
||||
|
||||
for model in original_model_list:
|
||||
|
|
@ -5138,6 +5142,9 @@ class Router:
|
|||
f"\nInitialized Model List {self.get_model_names()}"
|
||||
)
|
||||
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:
|
||||
import os
|
||||
|
|
@ -5365,20 +5372,27 @@ class Router:
|
|||
self, model: dict, model_id: Optional[str] = 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:
|
||||
- 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"]
|
||||
"""
|
||||
idx = len(self.model_list)
|
||||
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:
|
||||
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:
|
||||
self.model_id_to_deployment_index_map[model["model_info"]["id"]] = (
|
||||
len(self.model_list) - 1
|
||||
)
|
||||
self.model_id_to_deployment_index_map[model["model_info"]["id"]] = idx
|
||||
|
||||
# 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]:
|
||||
"""
|
||||
|
|
@ -6094,6 +6108,22 @@ class Router:
|
|||
additional_headers[header] = value
|
||||
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):
|
||||
"""
|
||||
Build model index from model list to enable O(1) lookups immediately.
|
||||
|
|
@ -6198,18 +6228,41 @@ class Router:
|
|||
Used for accurate 'get_model_list'.
|
||||
|
||||
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] = []
|
||||
for model in 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)
|
||||
|
||||
# O(1) lookup in model_name index
|
||||
if model_name in self.model_name_to_deployment_indices:
|
||||
indices = self.model_name_to_deployment_indices[model_name]
|
||||
|
||||
# O(k) where k = deployments for this model_name (typically 1-10)
|
||||
for idx in indices:
|
||||
model = self.model_list[idx]
|
||||
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)
|
||||
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
|
||||
|
||||
|
|
|
|||
371
litellm/utils.py
371
litellm/utils.py
|
|
@ -532,9 +532,6 @@ def get_dynamic_callbacks(
|
|||
return returned_callbacks
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def function_setup( # noqa: PLR0915
|
||||
original_function: str, rules_obj, start_time, *args, **kwargs
|
||||
): # just run once to check if user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
|
||||
|
|
@ -802,7 +799,7 @@ def function_setup( # noqa: PLR0915
|
|||
call_type=call_type,
|
||||
):
|
||||
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
|
||||
messages=messages,
|
||||
stream=stream,
|
||||
|
|
@ -910,159 +907,169 @@ def _get_wrapper_timeout(
|
|||
|
||||
return timeout
|
||||
|
||||
def check_coroutine(value) -> bool:
|
||||
return get_coroutine_checker().is_async_callable(value)
|
||||
|
||||
|
||||
async def async_pre_call_deployment_hook(kwargs: Dict[str, Any], call_type: str):
|
||||
"""
|
||||
Allow modifying the request just before it's sent to the deployment.
|
||||
|
||||
Use this instead of 'async_pre_call_hook' when you need to modify the request AFTER a deployment is selected, but BEFORE the request is sent.
|
||||
"""
|
||||
try:
|
||||
typed_call_type = CallTypes(call_type)
|
||||
except ValueError:
|
||||
typed_call_type = None # unknown call type
|
||||
|
||||
modified_kwargs = kwargs.copy()
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
result = await callback.async_pre_call_deployment_hook(
|
||||
modified_kwargs, typed_call_type
|
||||
)
|
||||
if result is not None:
|
||||
modified_kwargs = result
|
||||
|
||||
return modified_kwargs
|
||||
|
||||
|
||||
async def async_post_call_success_deployment_hook(
|
||||
request_data: dict, response: Any, call_type: Optional[CallTypes]
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Allow modifying / reviewing the response just after it's received from the deployment.
|
||||
"""
|
||||
try:
|
||||
typed_call_type = CallTypes(call_type)
|
||||
except ValueError:
|
||||
typed_call_type = None # unknown call type
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
result = await callback.async_post_call_success_deployment_hook(
|
||||
request_data, cast(LLMResponseTypes, response), typed_call_type
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
return response
|
||||
|
||||
|
||||
def post_call_processing(
|
||||
original_response,
|
||||
model,
|
||||
optional_params: Optional[dict],
|
||||
original_function,
|
||||
rules_obj,
|
||||
):
|
||||
try:
|
||||
if original_response is None:
|
||||
pass
|
||||
else:
|
||||
call_type = original_function.__name__
|
||||
if (
|
||||
call_type == CallTypes.completion.value
|
||||
or call_type == CallTypes.acompletion.value
|
||||
):
|
||||
is_coroutine = check_coroutine(original_response)
|
||||
if is_coroutine is True:
|
||||
pass
|
||||
else:
|
||||
if (
|
||||
isinstance(original_response, ModelResponse)
|
||||
and len(original_response.choices) > 0
|
||||
):
|
||||
model_response: Optional[str] = original_response.choices[
|
||||
0
|
||||
].message.content # type: ignore
|
||||
if model_response is not None:
|
||||
### POST-CALL RULES ###
|
||||
rules_obj.post_call_rules(
|
||||
input=model_response, model=model
|
||||
)
|
||||
### JSON SCHEMA VALIDATION ###
|
||||
if litellm.enable_json_schema_validation is True:
|
||||
try:
|
||||
if (
|
||||
optional_params is not None
|
||||
and "response_format" in optional_params
|
||||
and optional_params["response_format"]
|
||||
is not None
|
||||
):
|
||||
json_response_format: Optional[dict] = None
|
||||
if (
|
||||
isinstance(
|
||||
optional_params["response_format"],
|
||||
dict,
|
||||
)
|
||||
and optional_params[
|
||||
"response_format"
|
||||
].get("json_schema")
|
||||
is not None
|
||||
):
|
||||
json_response_format = optional_params[
|
||||
"response_format"
|
||||
]
|
||||
elif _parsing._completions.is_basemodel_type(
|
||||
optional_params["response_format"] # type: ignore
|
||||
):
|
||||
json_response_format = (
|
||||
type_to_response_format_param(
|
||||
response_format=optional_params[
|
||||
"response_format"
|
||||
]
|
||||
)
|
||||
)
|
||||
if json_response_format is not None:
|
||||
litellm.litellm_core_utils.json_validation_rule.validate_schema(
|
||||
schema=json_response_format[
|
||||
"json_schema"
|
||||
]["schema"],
|
||||
response=model_response,
|
||||
)
|
||||
except TypeError:
|
||||
pass
|
||||
if (
|
||||
optional_params is not None
|
||||
and "response_format" in optional_params
|
||||
and isinstance(
|
||||
optional_params["response_format"], dict
|
||||
)
|
||||
and "type" in optional_params["response_format"]
|
||||
and optional_params["response_format"]["type"]
|
||||
== "json_object"
|
||||
and "response_schema"
|
||||
in optional_params["response_format"]
|
||||
and isinstance(
|
||||
optional_params["response_format"][
|
||||
"response_schema"
|
||||
],
|
||||
dict,
|
||||
)
|
||||
and "enforce_validation"
|
||||
in optional_params["response_format"]
|
||||
and optional_params["response_format"][
|
||||
"enforce_validation"
|
||||
]
|
||||
is True
|
||||
):
|
||||
# schema given, json response expected, and validation enforced
|
||||
litellm.litellm_core_utils.json_validation_rule.validate_schema(
|
||||
schema=optional_params["response_format"][
|
||||
"response_schema"
|
||||
],
|
||||
response=model_response,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
||||
def client(original_function): # noqa: PLR0915
|
||||
rules_obj = Rules()
|
||||
|
||||
def check_coroutine(value) -> bool:
|
||||
return get_coroutine_checker().is_async_callable(value)
|
||||
|
||||
async def async_pre_call_deployment_hook(kwargs: Dict[str, Any], call_type: str):
|
||||
"""
|
||||
Allow modifying the request just before it's sent to the deployment.
|
||||
|
||||
Use this instead of 'async_pre_call_hook' when you need to modify the request AFTER a deployment is selected, but BEFORE the request is sent.
|
||||
"""
|
||||
try:
|
||||
typed_call_type = CallTypes(call_type)
|
||||
except ValueError:
|
||||
typed_call_type = None # unknown call type
|
||||
|
||||
modified_kwargs = kwargs.copy()
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
result = await callback.async_pre_call_deployment_hook(
|
||||
modified_kwargs, typed_call_type
|
||||
)
|
||||
if result is not None:
|
||||
modified_kwargs = result
|
||||
|
||||
return modified_kwargs
|
||||
|
||||
async def async_post_call_success_deployment_hook(
|
||||
request_data: dict, response: Any, call_type: Optional[CallTypes]
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Allow modifying / reviewing the response just after it's received from the deployment.
|
||||
"""
|
||||
try:
|
||||
typed_call_type = CallTypes(call_type)
|
||||
except ValueError:
|
||||
typed_call_type = None # unknown call type
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
result = await callback.async_post_call_success_deployment_hook(
|
||||
request_data, cast(LLMResponseTypes, response), typed_call_type
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
return response
|
||||
|
||||
def post_call_processing(original_response, model, optional_params: Optional[dict]):
|
||||
try:
|
||||
if original_response is None:
|
||||
pass
|
||||
else:
|
||||
call_type = original_function.__name__
|
||||
if (
|
||||
call_type == CallTypes.completion.value
|
||||
or call_type == CallTypes.acompletion.value
|
||||
):
|
||||
is_coroutine = check_coroutine(original_response)
|
||||
if is_coroutine is True:
|
||||
pass
|
||||
else:
|
||||
if (
|
||||
isinstance(original_response, ModelResponse)
|
||||
and len(original_response.choices) > 0
|
||||
):
|
||||
model_response: Optional[str] = original_response.choices[
|
||||
0
|
||||
].message.content # type: ignore
|
||||
if model_response is not None:
|
||||
### POST-CALL RULES ###
|
||||
rules_obj.post_call_rules(
|
||||
input=model_response, model=model
|
||||
)
|
||||
### JSON SCHEMA VALIDATION ###
|
||||
if litellm.enable_json_schema_validation is True:
|
||||
try:
|
||||
if (
|
||||
optional_params is not None
|
||||
and "response_format" in optional_params
|
||||
and optional_params["response_format"]
|
||||
is not None
|
||||
):
|
||||
json_response_format: Optional[dict] = None
|
||||
if (
|
||||
isinstance(
|
||||
optional_params["response_format"],
|
||||
dict,
|
||||
)
|
||||
and optional_params[
|
||||
"response_format"
|
||||
].get("json_schema")
|
||||
is not None
|
||||
):
|
||||
json_response_format = optional_params[
|
||||
"response_format"
|
||||
]
|
||||
elif _parsing._completions.is_basemodel_type(
|
||||
optional_params["response_format"] # type: ignore
|
||||
):
|
||||
json_response_format = (
|
||||
type_to_response_format_param(
|
||||
response_format=optional_params[
|
||||
"response_format"
|
||||
]
|
||||
)
|
||||
)
|
||||
if json_response_format is not None:
|
||||
litellm.litellm_core_utils.json_validation_rule.validate_schema(
|
||||
schema=json_response_format[
|
||||
"json_schema"
|
||||
]["schema"],
|
||||
response=model_response,
|
||||
)
|
||||
except TypeError:
|
||||
pass
|
||||
if (
|
||||
optional_params is not None
|
||||
and "response_format" in optional_params
|
||||
and isinstance(
|
||||
optional_params["response_format"], dict
|
||||
)
|
||||
and "type" in optional_params["response_format"]
|
||||
and optional_params["response_format"]["type"]
|
||||
== "json_object"
|
||||
and "response_schema"
|
||||
in optional_params["response_format"]
|
||||
and isinstance(
|
||||
optional_params["response_format"][
|
||||
"response_schema"
|
||||
],
|
||||
dict,
|
||||
)
|
||||
and "enforce_validation"
|
||||
in optional_params["response_format"]
|
||||
and optional_params["response_format"][
|
||||
"enforce_validation"
|
||||
]
|
||||
is True
|
||||
):
|
||||
# schema given, json response expected, and validation enforced
|
||||
litellm.litellm_core_utils.json_validation_rule.validate_schema(
|
||||
schema=optional_params["response_format"][
|
||||
"response_schema"
|
||||
],
|
||||
response=model_response,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
@wraps(original_function)
|
||||
def wrapper(*args, **kwargs): # noqa: PLR0915
|
||||
|
|
@ -1273,6 +1280,8 @@ def client(original_function): # noqa: PLR0915
|
|||
original_response=result,
|
||||
model=model or None,
|
||||
optional_params=kwargs,
|
||||
original_function=original_function,
|
||||
rules_obj=rules_obj,
|
||||
)
|
||||
|
||||
# [OPTIONAL] ADD TO CACHE
|
||||
|
|
@ -1417,7 +1426,8 @@ def client(original_function): # noqa: PLR0915
|
|||
if _caching_handler_response is not None:
|
||||
if (
|
||||
_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
|
||||
|
||||
|
|
@ -1489,7 +1499,11 @@ def client(original_function): # noqa: PLR0915
|
|||
return result
|
||||
### POST-CALL RULES ###
|
||||
post_call_processing(
|
||||
original_response=result, model=model, optional_params=kwargs
|
||||
original_response=result,
|
||||
model=model,
|
||||
optional_params=kwargs,
|
||||
original_function=original_function,
|
||||
rules_obj=rules_obj,
|
||||
)
|
||||
# Only run if call_type is a valid value in CallTypes
|
||||
if call_type in [ct.value for ct in CallTypes]:
|
||||
|
|
@ -1683,7 +1697,6 @@ def _is_streaming_request(
|
|||
return False
|
||||
|
||||
|
||||
|
||||
def _select_tokenizer(
|
||||
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_output_tokens=_model_info.get("max_output_tokens", None),
|
||||
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_priority=_model_info.get("input_cost_per_token_priority", None),
|
||||
input_cost_per_token_flex=_model_info.get(
|
||||
"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", None
|
||||
),
|
||||
cache_read_input_token_cost=_model_info.get(
|
||||
"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_priority=_model_info.get("cache_read_input_token_cost_priority", None),
|
||||
cache_read_input_token_cost_flex=_model_info.get(
|
||||
"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", None
|
||||
),
|
||||
|
|
@ -4901,8 +4922,12 @@ def _get_model_info_helper( # noqa: PLR0915
|
|||
"output_cost_per_token_batches"
|
||||
),
|
||||
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_priority=_model_info.get("output_cost_per_token_priority", None),
|
||||
output_cost_per_token_flex=_model_info.get(
|
||||
"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", None
|
||||
),
|
||||
|
|
@ -6434,7 +6459,7 @@ def get_valid_models(
|
|||
|
||||
try:
|
||||
################################
|
||||
# init litellm_params
|
||||
# init litellm_params
|
||||
#################################
|
||||
if litellm_params is None:
|
||||
litellm_params = LiteLLM_Params(model="")
|
||||
|
|
@ -6443,7 +6468,7 @@ def get_valid_models(
|
|||
if api_base is not None:
|
||||
litellm_params.api_base = api_base
|
||||
#################################
|
||||
|
||||
|
||||
check_provider_endpoint = (
|
||||
check_provider_endpoint or litellm.check_provider_endpoint
|
||||
)
|
||||
|
|
@ -6918,7 +6943,10 @@ class ProviderConfigManager:
|
|||
return litellm.LlamaAPIConfig()
|
||||
elif litellm.LlmProviders.TEXT_COMPLETION_OPENAI == provider:
|
||||
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()
|
||||
elif litellm.LlmProviders.SNOWFLAKE == provider:
|
||||
return litellm.SnowflakeConfig()
|
||||
|
|
@ -7345,7 +7373,12 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return VLLMPassthroughConfig()
|
||||
elif LlmProviders.AZURE == provider:
|
||||
from litellm.llms.azure.passthrough.transformation import (
|
||||
AzurePassthroughConfig,
|
||||
)
|
||||
|
||||
return AzurePassthroughConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -7532,9 +7565,7 @@ class ProviderConfigManager:
|
|||
|
||||
return RecraftImageEditConfig()
|
||||
elif LlmProviders.AZURE_AI == provider:
|
||||
from litellm.llms.azure_ai.image_edit import (
|
||||
get_azure_ai_image_edit_config,
|
||||
)
|
||||
from litellm.llms.azure_ai.image_edit import get_azure_ai_image_edit_config
|
||||
|
||||
return get_azure_ai_image_edit_config(model)
|
||||
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.
|
||||
"""
|
||||
_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(
|
||||
Optional[str],
|
||||
|
|
|
|||
|
|
@ -12872,6 +12872,39 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"gpt-5-pro": {
|
||||
"input_cost_per_token": 1.5e-05,
|
||||
"input_cost_per_token_batches": 7.5e-06,
|
||||
"litellm_provider": "openai",
|
||||
"max_input_tokens": 400000,
|
||||
"max_output_tokens": 272000,
|
||||
"max_tokens": 272000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 1.2e-04,
|
||||
"output_cost_per_token_batches": 6e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": false,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"gpt-5-codex": {
|
||||
"cache_read_input_token_cost": 1.25e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
|
|
@ -13189,6 +13222,20 @@
|
|||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"gpt-image-1-mini": {
|
||||
"cache_read_input_image_token_cost": 2.5e-07,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"input_cost_per_image_token": 2.5e-06,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
"output_cost_per_image_token": 8e-06,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supports_vision": true
|
||||
},
|
||||
"gpt-realtime": {
|
||||
"cache_creation_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
|
|
@ -14623,6 +14670,54 @@
|
|||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"low/1024-x-1024/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.005,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"low/1024-x-1536/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.006,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"low/1536-x-1024/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.006,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"medium/1024-x-1024/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.011,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"medium/1024-x-1536/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.015,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"medium/1536-x-1024/gpt-image-1-mini": {
|
||||
"input_cost_per_image": 0.015,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"medlm-large": {
|
||||
"input_cost_per_character": 5e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
|
|
@ -22173,6 +22268,307 @@
|
|||
"supports_tool_choice": true,
|
||||
"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": {
|
||||
"input_cost_per_second": 0.0001,
|
||||
"litellm_provider": "openai",
|
||||
|
|
|
|||
|
|
@ -156,6 +156,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
object_permission_id String @id @default(uuid())
|
||||
mcp_servers String[] @default([])
|
||||
mcp_access_groups String[] @default([])
|
||||
mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]}
|
||||
vector_stores String[] @default([])
|
||||
teams LiteLLM_TeamTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from litellm.types.llms.openai import (
|
|||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from base_responses_api import BaseResponsesAPITest
|
||||
|
||||
|
||||
class TestAzureResponsesAPITest(BaseResponsesAPITest):
|
||||
def get_base_completion_call_args(self):
|
||||
return {
|
||||
|
|
@ -43,4 +44,55 @@ async def test_azure_responses_api_preview_api_version():
|
|||
api_base=os.getenv("AZURE_RESPONSES_OPENAI_ENDPOINT"),
|
||||
api_key=os.getenv("AZURE_RESPONSES_OPENAI_API_KEY"),
|
||||
input="Hello, can you tell me a short joke?",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_responses_api_status_error():
|
||||
"""
|
||||
Ensure new azure preview api version is working
|
||||
"""
|
||||
litellm._turn_on_debug()
|
||||
|
||||
request_data = {
|
||||
"model": "gpt-5-mini",
|
||||
"input": [
|
||||
{"content": "tell me an interesting fact", "role": "user"},
|
||||
{
|
||||
"id": "rs_0ab687487834d9df0068e462a1b2d88197aabbc832c9ba5316",
|
||||
"summary": [],
|
||||
"type": "reasoning",
|
||||
"content": None,
|
||||
"encrypted_content": None,
|
||||
"status": "completed",
|
||||
},
|
||||
{
|
||||
"id": "msg_0ab687487834d9df0068e462a1df188197b74b1eef05102c18",
|
||||
"content": [
|
||||
{
|
||||
"annotations": [],
|
||||
"text": "Octopuses have three hearts: two pump blood to the gills, while the third pumps it to the rest of the body. Even more unusual, their blood is blue because it uses the copper-containing protein hemocyanin to carry oxygen, which is more efficient than hemoglobin in cold, low-oxygen environments.",
|
||||
"type": "output_text",
|
||||
"logprobs": [],
|
||||
}
|
||||
],
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"type": "message",
|
||||
},
|
||||
{"role": "user", "content": "tell me another"},
|
||||
],
|
||||
"include": [],
|
||||
"instructions": "You are a helpful assistant.",
|
||||
"reasoning": {"effort": "minimal"},
|
||||
"stream": False,
|
||||
"tools": [],
|
||||
}
|
||||
response = await litellm.aresponses(
|
||||
model="azure/gpt-5-mini-2",
|
||||
truncation="auto",
|
||||
api_version="preview",
|
||||
api_base=os.getenv("AZURE_GPT5_MINI_API_BASE"),
|
||||
api_key=os.getenv("AZURE_GPT5_MINI_API_KEY"),
|
||||
input=request_data["input"],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from dotenv import load_dotenv
|
|||
load_dotenv()
|
||||
import pytest
|
||||
|
||||
from litellm import completion, acompletion
|
||||
from litellm import completion, acompletion, responses
|
||||
from litellm.exceptions import APIConnectionError
|
||||
|
||||
@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
|
||||
except Exception as 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
|
||||
|
|
|
|||
|
|
@ -61,6 +61,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
|
|||
list_keys,
|
||||
regenerate_key_fn,
|
||||
update_key_fn,
|
||||
key_aliases,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
new_team,
|
||||
|
|
@ -151,7 +152,6 @@ def prisma_client():
|
|||
@pytest.mark.flaky(retries=6, delay=1)
|
||||
async def test_new_user_response(prisma_client):
|
||||
try:
|
||||
|
||||
print("prisma client=", prisma_client)
|
||||
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
|
|
@ -424,7 +424,6 @@ async def test_call_with_valid_model_using_all_models(prisma_client):
|
|||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
try:
|
||||
|
||||
await litellm.proxy.proxy_server.prisma_client.connect()
|
||||
|
||||
team_request = NewTeamRequest(
|
||||
|
|
@ -1789,7 +1788,6 @@ async def test_call_with_key_over_model_budget(
|
|||
litellm.callbacks.append(model_budget_limiter)
|
||||
|
||||
try:
|
||||
|
||||
# set budget for chatgpt-v-3 to 0.000001, expect the next request to fail
|
||||
model_max_budget = {
|
||||
"gpt-4o-mini": {
|
||||
|
|
@ -3531,6 +3529,58 @@ async def test_list_keys(prisma_client):
|
|||
assert _key in response["keys"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_aliases(prisma_client):
|
||||
"""
|
||||
Test the key_aliases function:
|
||||
- Returns a list
|
||||
- Includes alias from a newly created key
|
||||
- Aliases are unique and sorted
|
||||
"""
|
||||
import asyncio
|
||||
import uuid
|
||||
import litellm
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
# Wire up test prisma client
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
await litellm.proxy.proxy_server.prisma_client.connect()
|
||||
|
||||
# Basic call
|
||||
response = await key_aliases()
|
||||
assert "aliases" in response
|
||||
assert isinstance(response["aliases"], list)
|
||||
|
||||
# Create a new user (and key) with a unique alias
|
||||
unique_id = str(uuid.uuid4())
|
||||
test_alias = f"key-aliases-test-{unique_id}"
|
||||
test_user_id = f"key-aliases-user-{unique_id}"
|
||||
|
||||
await new_user(
|
||||
data=NewUserRequest(
|
||||
user_id=test_user_id,
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
key_alias=test_alias,
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
# Allow async DB writes to settle
|
||||
await asyncio.sleep(2)
|
||||
|
||||
# Call again and validate
|
||||
response_after = await key_aliases()
|
||||
aliases = response_after["aliases"]
|
||||
|
||||
# Contains the new alias
|
||||
assert test_alias in aliases
|
||||
|
||||
# Unique & sorted (endpoint dedupes and orders ascending)
|
||||
assert len(aliases) == len(set(aliases))
|
||||
assert aliases == sorted(aliases)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_vertex_ai_route(prisma_client):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -77,7 +77,6 @@ class TestRouterIndexManagement:
|
|||
# Verify: Index map uses model_info.id
|
||||
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):
|
||||
"""Test _add_model_to_list_and_index_map with multiple models to verify indexing"""
|
||||
# Setup: Empty router
|
||||
|
|
@ -127,3 +126,54 @@ class TestRouterIndexManagement:
|
|||
# Test: Empty router
|
||||
empty_router = Router(model_list=[])
|
||||
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]
|
||||
|
|
|
|||
|
|
@ -1,10 +1,13 @@
|
|||
import os
|
||||
import ssl
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
# Add the parent directory to the path so we can import litellm
|
||||
sys.path.insert(0, '../../../')
|
||||
sys.path.insert(0, "../../../")
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.types.mcp import MCPStdioConfig, MCPTransport
|
||||
|
|
@ -16,57 +19,54 @@ class TestMCPClient:
|
|||
def test_mcp_client_stdio_init(self):
|
||||
"""Test MCPClient initialization with stdio config"""
|
||||
stdio_config = MCPStdioConfig(
|
||||
command="python",
|
||||
args=["-m", "my_mcp_server"],
|
||||
env={"DEBUG": "1"}
|
||||
command="python", args=["-m", "my_mcp_server"], env={"DEBUG": "1"}
|
||||
)
|
||||
|
||||
client = MCPClient(
|
||||
transport_type=MCPTransport.stdio,
|
||||
stdio_config=stdio_config
|
||||
)
|
||||
|
||||
|
||||
client = MCPClient(transport_type=MCPTransport.stdio, stdio_config=stdio_config)
|
||||
|
||||
assert client.transport_type == MCPTransport.stdio
|
||||
assert client.stdio_config == stdio_config
|
||||
assert client.stdio_config["command"] == "python"
|
||||
assert client.stdio_config["args"] == ["-m", "my_mcp_server"]
|
||||
assert client.stdio_config is not None
|
||||
assert client.stdio_config.get("command") == "python"
|
||||
assert client.stdio_config.get("args") == ["-m", "my_mcp_server"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_client_stdio_connect_error(self):
|
||||
"""Test MCP client stdio connection error handling"""
|
||||
# Test missing stdio_config
|
||||
client = MCPClient(transport_type=MCPTransport.stdio)
|
||||
|
||||
with pytest.raises(ValueError, match="stdio_config is required for stdio transport"):
|
||||
|
||||
with pytest.raises(
|
||||
ValueError, match="stdio_config is required for stdio transport"
|
||||
):
|
||||
await client.connect()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('litellm.experimental_mcp_client.client.stdio_client')
|
||||
@patch('litellm.experimental_mcp_client.client.ClientSession')
|
||||
async def test_mcp_client_stdio_connect_success(self, mock_session, mock_stdio_client):
|
||||
@patch("litellm.experimental_mcp_client.client.stdio_client")
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_mcp_client_stdio_connect_success(
|
||||
self, mock_session, mock_stdio_client
|
||||
):
|
||||
"""Test successful stdio connection"""
|
||||
# Setup mocks
|
||||
mock_transport = (MagicMock(), MagicMock())
|
||||
mock_stdio_client.return_value.__aenter__ = AsyncMock(return_value=mock_transport)
|
||||
|
||||
mock_stdio_client.return_value.__aenter__ = AsyncMock(
|
||||
return_value=mock_transport
|
||||
)
|
||||
|
||||
mock_session_instance = MagicMock()
|
||||
mock_session_instance.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session.return_value = mock_session_instance
|
||||
|
||||
|
||||
stdio_config = MCPStdioConfig(
|
||||
command="python",
|
||||
args=["-m", "my_mcp_server"],
|
||||
env={"DEBUG": "1"}
|
||||
command="python", args=["-m", "my_mcp_server"], env={"DEBUG": "1"}
|
||||
)
|
||||
|
||||
client = MCPClient(
|
||||
transport_type=MCPTransport.stdio,
|
||||
stdio_config=stdio_config
|
||||
)
|
||||
|
||||
|
||||
client = MCPClient(transport_type=MCPTransport.stdio, stdio_config=stdio_config)
|
||||
|
||||
await client.connect()
|
||||
|
||||
|
||||
# Verify stdio_client was called with correct parameters
|
||||
mock_stdio_client.assert_called_once()
|
||||
call_args = mock_stdio_client.call_args[0][0]
|
||||
|
|
@ -74,6 +74,162 @@ class TestMCPClient:
|
|||
assert call_args.args == ["-m", "my_mcp_server"]
|
||||
assert call_args.env == {"DEBUG": "1"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.streamablehttp_client")
|
||||
@patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"SSL_CERT_FILE": "/path/to/custom/ca-bundle.pem",
|
||||
"SSL_CERTIFICATE": "/path/to/client-cert.pem",
|
||||
},
|
||||
)
|
||||
async def test_mcp_client_ssl_configuration_from_env(
|
||||
self, mock_streamablehttp_client
|
||||
):
|
||||
"""Test that MCP client uses SSL configuration from environment variables"""
|
||||
# Setup mocks
|
||||
mock_transport = (MagicMock(), MagicMock())
|
||||
mock_streamablehttp_client.return_value.__aenter__ = AsyncMock(
|
||||
return_value=mock_transport
|
||||
)
|
||||
|
||||
# Mock the session
|
||||
with patch(
|
||||
"litellm.experimental_mcp_client.client.ClientSession"
|
||||
) as mock_session:
|
||||
mock_session_instance = MagicMock()
|
||||
mock_session_instance.__aenter__ = AsyncMock(
|
||||
return_value=mock_session_instance
|
||||
)
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session.return_value = mock_session_instance
|
||||
|
||||
client = MCPClient(
|
||||
server_url="https://mcp-server.example.com",
|
||||
transport_type=MCPTransport.http,
|
||||
)
|
||||
|
||||
await client.connect()
|
||||
|
||||
# Verify streamablehttp_client was called
|
||||
mock_streamablehttp_client.assert_called_once()
|
||||
call_kwargs = mock_streamablehttp_client.call_args[1]
|
||||
|
||||
# Verify httpx_client_factory was passed
|
||||
assert "httpx_client_factory" in call_kwargs
|
||||
httpx_factory = call_kwargs["httpx_client_factory"]
|
||||
|
||||
# Test the factory creates a client with proper SSL config
|
||||
# When SSL_CERT_FILE is set, the factory should use get_ssl_configuration
|
||||
test_client = httpx_factory(headers={"test": "header"})
|
||||
|
||||
# Verify the client was created successfully with SSL configuration
|
||||
assert test_client is not None
|
||||
assert isinstance(test_client, httpx.AsyncClient)
|
||||
# Verify it has the expected properties
|
||||
assert test_client.headers is not None
|
||||
# Clean up
|
||||
await test_client.aclose()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.sse_client")
|
||||
async def test_mcp_client_ssl_verify_parameter(self, mock_sse_client):
|
||||
"""Test that MCP client uses ssl_verify parameter when provided"""
|
||||
# Setup mocks
|
||||
mock_transport = (MagicMock(), MagicMock())
|
||||
mock_sse_client.return_value.__aenter__ = AsyncMock(return_value=mock_transport)
|
||||
|
||||
# Mock the session
|
||||
with patch(
|
||||
"litellm.experimental_mcp_client.client.ClientSession"
|
||||
) as mock_session:
|
||||
mock_session_instance = MagicMock()
|
||||
mock_session_instance.__aenter__ = AsyncMock(
|
||||
return_value=mock_session_instance
|
||||
)
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session.return_value = mock_session_instance
|
||||
|
||||
# Test with ssl_verify=False
|
||||
client = MCPClient(
|
||||
server_url="https://mcp-server.example.com",
|
||||
transport_type=MCPTransport.sse,
|
||||
ssl_verify=False,
|
||||
)
|
||||
|
||||
await client.connect()
|
||||
|
||||
# Verify sse_client was called
|
||||
mock_sse_client.assert_called_once()
|
||||
call_kwargs = mock_sse_client.call_args[1]
|
||||
|
||||
# Verify httpx_client_factory was passed
|
||||
assert "httpx_client_factory" in call_kwargs
|
||||
httpx_factory = call_kwargs["httpx_client_factory"]
|
||||
|
||||
# Test the factory creates a client with SSL verification disabled
|
||||
# When ssl_verify=False, the factory should disable SSL verification
|
||||
test_client = httpx_factory(headers={"test": "header"})
|
||||
|
||||
# Verify the client was created successfully
|
||||
assert test_client is not None
|
||||
assert isinstance(test_client, httpx.AsyncClient)
|
||||
# Verify it has the expected properties
|
||||
assert test_client.headers is not None
|
||||
# Clean up
|
||||
await test_client.aclose()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.streamablehttp_client")
|
||||
async def test_mcp_client_ssl_verify_custom_path(self, mock_streamablehttp_client):
|
||||
"""Test that MCP client uses custom CA bundle path from ssl_verify parameter"""
|
||||
# Setup mocks
|
||||
mock_transport = (MagicMock(), MagicMock())
|
||||
mock_streamablehttp_client.return_value.__aenter__ = AsyncMock(
|
||||
return_value=mock_transport
|
||||
)
|
||||
|
||||
# Mock the session
|
||||
with patch(
|
||||
"litellm.experimental_mcp_client.client.ClientSession"
|
||||
) as mock_session:
|
||||
mock_session_instance = MagicMock()
|
||||
mock_session_instance.__aenter__ = AsyncMock(
|
||||
return_value=mock_session_instance
|
||||
)
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session.return_value = mock_session_instance
|
||||
|
||||
# Test with custom CA bundle path
|
||||
custom_ca_path = "/custom/path/to/ca-bundle.pem"
|
||||
client = MCPClient(
|
||||
server_url="https://mcp-server.example.com",
|
||||
transport_type=MCPTransport.http,
|
||||
ssl_verify=custom_ca_path,
|
||||
)
|
||||
|
||||
await client.connect()
|
||||
|
||||
# Verify streamablehttp_client was called
|
||||
mock_streamablehttp_client.assert_called_once()
|
||||
call_kwargs = mock_streamablehttp_client.call_args[1]
|
||||
|
||||
# Verify httpx_client_factory was passed
|
||||
assert "httpx_client_factory" in call_kwargs
|
||||
httpx_factory = call_kwargs["httpx_client_factory"]
|
||||
|
||||
# Test the factory creates a client with custom CA bundle path
|
||||
# When ssl_verify is a path, the factory should use that path for SSL verification
|
||||
test_client = httpx_factory(headers={"test": "header"})
|
||||
|
||||
# Verify the client was created successfully
|
||||
assert test_client is not None
|
||||
assert isinstance(test_client, httpx.AsyncClient)
|
||||
# Verify it has the expected properties
|
||||
assert test_client.headers is not None
|
||||
# Clean up
|
||||
await test_client.aclose()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
pytest.main([__file__])
|
||||
|
|
|
|||
|
|
@ -43,3 +43,25 @@ def test_azure_ai_validate_environment():
|
|||
litellm_params={},
|
||||
)
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
|
||||
|
||||
def test_azure_ai_grok_stop_parameter_handling():
|
||||
"""
|
||||
Test that Grok models properly handle stop parameter filtering in Azure AI Studio.
|
||||
"""
|
||||
config = AzureAIStudioConfig()
|
||||
|
||||
# Test Grok model detection
|
||||
assert config._supports_stop_reason("grok-4-fast") == False
|
||||
assert config._supports_stop_reason("grok-4") == False
|
||||
assert config._supports_stop_reason("grok-3-mini") == False
|
||||
assert config._supports_stop_reason("grok-code-fast") == False
|
||||
assert config._supports_stop_reason("gpt-4") == True
|
||||
|
||||
# Test supported parameters for Grok models
|
||||
grok_params = config.get_supported_openai_params("grok-4-fast")
|
||||
assert "stop" not in grok_params, "Grok models should not support stop parameter"
|
||||
|
||||
# Test supported parameters for non-Grok models
|
||||
gpt_params = config.get_supported_openai_params("gpt-4")
|
||||
assert "stop" in gpt_params, "GPT models should support stop parameter"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -191,7 +191,8 @@ class TestTTLExtraction:
|
|||
class TestTransformationWithTTL:
|
||||
"""Test the complete transformation with TTL support"""
|
||||
|
||||
def test_transform_with_valid_ttl(self):
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
def test_transform_with_valid_ttl(self, custom_llm_provider):
|
||||
"""Test transformation includes TTL when provided"""
|
||||
messages = [
|
||||
{
|
||||
|
|
@ -205,19 +206,32 @@ class TestTransformationWithTTL:
|
|||
]
|
||||
}
|
||||
]
|
||||
|
||||
vertex_location="test_location"
|
||||
vertex_project="test_project"
|
||||
|
||||
result = transform_openai_messages_to_gemini_context_caching(
|
||||
model="gemini-1.5-pro",
|
||||
messages=messages,
|
||||
cache_key="test-cache-key"
|
||||
cache_key="test-cache-key",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_location="test_location",
|
||||
vertex_project="test_project"
|
||||
)
|
||||
|
||||
assert "ttl" in result
|
||||
assert result["ttl"] == "3600s"
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
|
||||
if custom_llm_provider == "gemini":
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
else:
|
||||
assert result["model"] == f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/gemini-1.5-pro"
|
||||
|
||||
|
||||
assert result["displayName"] == "test-cache-key"
|
||||
|
||||
def test_transform_without_ttl(self):
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
def test_transform_without_ttl(self, custom_llm_provider):
|
||||
"""Test transformation without TTL"""
|
||||
messages = [
|
||||
{
|
||||
|
|
@ -231,18 +245,30 @@ class TestTransformationWithTTL:
|
|||
]
|
||||
}
|
||||
]
|
||||
|
||||
vertex_location="test_location"
|
||||
vertex_project="test_project"
|
||||
|
||||
result = transform_openai_messages_to_gemini_context_caching(
|
||||
model="gemini-1.5-pro",
|
||||
messages=messages,
|
||||
cache_key="test-cache-key"
|
||||
cache_key="test-cache-key",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_location=vertex_location,
|
||||
vertex_project=vertex_project
|
||||
)
|
||||
|
||||
assert "ttl" not in result
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
|
||||
if custom_llm_provider == "gemini":
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
else:
|
||||
assert result["model"] == f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/gemini-1.5-pro"
|
||||
|
||||
assert result["displayName"] == "test-cache-key"
|
||||
|
||||
def test_transform_with_invalid_ttl(self):
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
def test_transform_with_invalid_ttl(self, custom_llm_provider):
|
||||
"""Test transformation with invalid TTL (should be ignored)"""
|
||||
messages = [
|
||||
{
|
||||
|
|
@ -256,18 +282,29 @@ class TestTransformationWithTTL:
|
|||
]
|
||||
}
|
||||
]
|
||||
vertex_location="test_location"
|
||||
vertex_project="test_project"
|
||||
|
||||
result = transform_openai_messages_to_gemini_context_caching(
|
||||
model="gemini-1.5-pro",
|
||||
messages=messages,
|
||||
cache_key="test-cache-key"
|
||||
cache_key="test-cache-key",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_location=vertex_location,
|
||||
vertex_project=vertex_project
|
||||
)
|
||||
|
||||
assert "ttl" not in result
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
|
||||
if custom_llm_provider == "gemini":
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
else:
|
||||
assert result["model"] == f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/gemini-1.5-pro"
|
||||
|
||||
assert result["displayName"] == "test-cache-key"
|
||||
|
||||
def test_transform_with_system_message_and_ttl(self):
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
def test_transform_with_system_message_and_ttl(self, custom_llm_provider):
|
||||
"""Test transformation with system message and TTL"""
|
||||
messages = [
|
||||
{
|
||||
|
|
@ -290,17 +327,28 @@ class TestTransformationWithTTL:
|
|||
]
|
||||
}
|
||||
]
|
||||
|
||||
vertex_location="test_location"
|
||||
vertex_project="test_project"
|
||||
|
||||
result = transform_openai_messages_to_gemini_context_caching(
|
||||
model="gemini-1.5-pro",
|
||||
messages=messages,
|
||||
cache_key="test-cache-key"
|
||||
cache_key="test-cache-key",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_location=vertex_location,
|
||||
vertex_project=vertex_project
|
||||
)
|
||||
|
||||
assert "ttl" in result
|
||||
assert result["ttl"] == "7200s"
|
||||
assert "system_instruction" in result
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
|
||||
if custom_llm_provider == "gemini":
|
||||
assert result["model"] == "models/gemini-1.5-pro"
|
||||
else:
|
||||
assert result["model"] == f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/gemini-1.5-pro"
|
||||
|
||||
assert result["displayName"] == "test-cache-key"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -55,6 +55,9 @@ class TestContextCachingEndpoints:
|
|||
|
||||
self.sample_optional_params = {"tools": self.sample_tools.copy()}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -62,12 +65,14 @@ class TestContextCachingEndpoints:
|
|||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj"
|
||||
)
|
||||
def test_check_and_create_cache_with_cached_content(
|
||||
self, mock_cache_obj, mock_separate
|
||||
self, mock_cache_obj, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test check_and_create_cache when cached_content is provided"""
|
||||
# Setup
|
||||
cached_content = "cached_content_123"
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -80,6 +85,10 @@ class TestContextCachingEndpoints:
|
|||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
cached_content=cached_content,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -92,14 +101,21 @@ class TestContextCachingEndpoints:
|
|||
mock_separate.assert_not_called()
|
||||
mock_cache_obj.get_cache_key.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
def test_check_and_create_cache_no_cached_messages(self, mock_separate):
|
||||
def test_check_and_create_cache_no_cached_messages(
|
||||
self, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test check_and_create_cache when no cached messages are found"""
|
||||
# Setup
|
||||
mock_separate.return_value = ([], self.sample_messages) # No cached messages
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -111,6 +127,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -119,6 +139,9 @@ class TestContextCachingEndpoints:
|
|||
assert returned_params == optional_params
|
||||
assert returned_cache is None
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -127,7 +150,7 @@ class TestContextCachingEndpoints:
|
|||
)
|
||||
@patch.object(ContextCachingEndpoints, "check_cache")
|
||||
def test_check_and_create_cache_existing_cache_found(
|
||||
self, mock_check_cache, mock_cache_obj, mock_separate
|
||||
self, mock_check_cache, mock_cache_obj, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test check_and_create_cache when existing cache is found"""
|
||||
# Setup
|
||||
|
|
@ -139,6 +162,8 @@ class TestContextCachingEndpoints:
|
|||
mock_check_cache.return_value = "existing_cache_name"
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -150,6 +175,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -163,6 +192,9 @@ class TestContextCachingEndpoints:
|
|||
messages=cached_messages, tools=self.sample_tools
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -181,6 +213,7 @@ class TestContextCachingEndpoints:
|
|||
mock_transform,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test check_and_create_cache when creating new cache"""
|
||||
# Setup
|
||||
|
|
@ -203,6 +236,8 @@ class TestContextCachingEndpoints:
|
|||
self.mock_client.post.return_value = mock_response
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -214,6 +249,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -228,6 +267,9 @@ class TestContextCachingEndpoints:
|
|||
assert "tools" in call_args.kwargs["json"]
|
||||
assert call_args.kwargs["json"]["tools"] == self.sample_tools
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -237,7 +279,12 @@ class TestContextCachingEndpoints:
|
|||
@patch.object(ContextCachingEndpoints, "check_cache")
|
||||
@patch.object(ContextCachingEndpoints, "_get_token_and_url_context_caching")
|
||||
def test_check_and_create_cache_http_error(
|
||||
self, mock_get_token_url, mock_check_cache, mock_cache_obj, mock_separate
|
||||
self,
|
||||
mock_get_token_url,
|
||||
mock_check_cache,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test check_and_create_cache handles HTTP errors properly"""
|
||||
# Setup
|
||||
|
|
@ -259,6 +306,8 @@ class TestContextCachingEndpoints:
|
|||
self.mock_client.post.side_effect = http_error
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute and Assert
|
||||
with pytest.raises(VertexAIError) as exc_info:
|
||||
|
|
@ -271,12 +320,19 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Bad Request" in str(exc_info.value.message)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -284,12 +340,14 @@ class TestContextCachingEndpoints:
|
|||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj"
|
||||
)
|
||||
async def test_async_check_and_create_cache_with_cached_content(
|
||||
self, mock_cache_obj, mock_separate
|
||||
self, mock_cache_obj, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test async_check_and_create_cache when cached_content is provided"""
|
||||
# Setup
|
||||
cached_content = "cached_content_123"
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -302,6 +360,10 @@ class TestContextCachingEndpoints:
|
|||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
cached_content=cached_content,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -311,14 +373,21 @@ class TestContextCachingEndpoints:
|
|||
assert returned_cache == cached_content
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
async def test_async_check_and_create_cache_no_cached_messages(self, mock_separate):
|
||||
async def test_async_check_and_create_cache_no_cached_messages(
|
||||
self, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test async_check_and_create_cache when no cached messages are found"""
|
||||
# Setup
|
||||
mock_separate.return_value = ([], self.sample_messages)
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -330,6 +399,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -339,6 +412,9 @@ class TestContextCachingEndpoints:
|
|||
assert returned_cache is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -347,7 +423,7 @@ class TestContextCachingEndpoints:
|
|||
)
|
||||
@patch.object(ContextCachingEndpoints, "async_check_cache")
|
||||
async def test_async_check_and_create_cache_existing_cache_found(
|
||||
self, mock_async_check_cache, mock_cache_obj, mock_separate
|
||||
self, mock_async_check_cache, mock_cache_obj, mock_separate, custom_llm_provider
|
||||
):
|
||||
"""Test async_check_and_create_cache when existing cache is found"""
|
||||
# Setup
|
||||
|
|
@ -359,6 +435,8 @@ class TestContextCachingEndpoints:
|
|||
mock_async_check_cache.return_value = "existing_cache_name"
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -370,6 +448,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -384,6 +466,9 @@ class TestContextCachingEndpoints:
|
|||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -406,6 +491,7 @@ class TestContextCachingEndpoints:
|
|||
mock_transform,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test async_check_and_create_cache when creating new cache"""
|
||||
# Setup
|
||||
|
|
@ -428,6 +514,8 @@ class TestContextCachingEndpoints:
|
|||
self.mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -439,6 +527,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert
|
||||
|
|
@ -454,6 +546,9 @@ class TestContextCachingEndpoints:
|
|||
assert call_args.kwargs["json"]["tools"] == self.sample_tools
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
|
|
@ -472,6 +567,7 @@ class TestContextCachingEndpoints:
|
|||
mock_async_check_cache,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test async_check_and_create_cache handles timeout errors properly"""
|
||||
# Setup
|
||||
|
|
@ -489,6 +585,8 @@ class TestContextCachingEndpoints:
|
|||
)
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute and Assert
|
||||
with pytest.raises(VertexAIError) as exc_info:
|
||||
|
|
@ -501,12 +599,21 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 408
|
||||
assert "Timeout error occurred" in str(exc_info.value.message)
|
||||
|
||||
def test_check_and_create_cache_tools_popped_from_optional_params(self):
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
def test_check_and_create_cache_tools_popped_from_optional_params(
|
||||
self, custom_llm_provider
|
||||
):
|
||||
"""Test that tools are properly popped from optional_params when there are cached messages"""
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
|
|
@ -520,6 +627,8 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Mock the check_cache to return existing cache so we don't make HTTP calls
|
||||
with patch.object(
|
||||
|
|
@ -535,6 +644,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert tools were popped from optional_params
|
||||
|
|
@ -543,7 +656,12 @@ class TestContextCachingEndpoints:
|
|||
# But original tools should still be available for comparison
|
||||
assert original_tools == self.sample_tools
|
||||
|
||||
def test_check_and_create_cache_tools_not_popped_when_no_cached_messages(self):
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
def test_check_and_create_cache_tools_not_popped_when_no_cached_messages(
|
||||
self, custom_llm_provider
|
||||
):
|
||||
"""Test that tools are NOT popped from optional_params when there are no cached messages"""
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
|
|
@ -555,6 +673,8 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -566,6 +686,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert tools were NOT popped from optional_params (early return)
|
||||
|
|
@ -573,8 +697,11 @@ class TestContextCachingEndpoints:
|
|||
assert optional_params["tools"] == original_tools
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
async def test_async_check_and_create_cache_tools_not_popped_when_no_cached_messages(
|
||||
self,
|
||||
self, custom_llm_provider
|
||||
):
|
||||
"""Test that tools are NOT popped from optional_params in async version when there are no cached messages"""
|
||||
with patch(
|
||||
|
|
@ -587,6 +714,8 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -598,6 +727,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert tools were NOT popped from optional_params (early return)
|
||||
|
|
@ -605,7 +738,12 @@ class TestContextCachingEndpoints:
|
|||
assert optional_params["tools"] == original_tools
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_check_and_create_cache_tools_popped_from_optional_params(self):
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
async def test_async_check_and_create_cache_tools_popped_from_optional_params(
|
||||
self, custom_llm_provider
|
||||
):
|
||||
"""Test that tools are properly popped from optional_params in async version when there are cached messages"""
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
|
|
@ -619,6 +757,8 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
|
||||
# Mock the async_check_cache to return existing cache so we don't make HTTP calls
|
||||
with patch.object(
|
||||
|
|
@ -634,6 +774,10 @@ class TestContextCachingEndpoints:
|
|||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project=test_project,
|
||||
vertex_location=test_location,
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# Assert tools were popped from optional_params
|
||||
|
|
|
|||
|
|
@ -0,0 +1,172 @@
|
|||
"""
|
||||
Test to verify that custom headers are correctly forwarded to Gemini/Vertex AI API calls.
|
||||
|
||||
This test verifies the fix for the issue where headers configured via
|
||||
forward_client_headers_to_llm_api were not being passed to Gemini/Vertex AI providers.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import Mock, patch, MagicMock
|
||||
import litellm
|
||||
from litellm import completion
|
||||
|
||||
|
||||
class TestGeminiHeaderForwarding:
|
||||
"""Test cases for verifying header forwarding to Gemini/Vertex AI."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider,model",
|
||||
[
|
||||
("gemini", "gemini/gemini-1.5-pro"),
|
||||
("vertex_ai_beta", "gemini-1.5-pro"),
|
||||
("vertex_ai", "gemini-1.5-pro"),
|
||||
],
|
||||
)
|
||||
def test_headers_forwarded_to_gemini(self, custom_llm_provider, model):
|
||||
"""
|
||||
Test that headers from kwargs are correctly merged and passed to Gemini completion.
|
||||
|
||||
This test verifies that when headers are passed via kwargs (as the proxy does when
|
||||
forward_client_headers_to_llm_api is configured), they are correctly merged with
|
||||
extra_headers and passed to the Vertex AI completion handler.
|
||||
"""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
|
||||
# Headers that would be set by the proxy when forwarding client headers
|
||||
custom_headers = {
|
||||
"X-Custom-Header": "CustomValue",
|
||||
"X-BYOK-Token": "secret-token",
|
||||
}
|
||||
|
||||
# Mock the vertex completion handler
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.VertexLLM.completion"
|
||||
) as mock_vertex_completion:
|
||||
# Configure the mock to return a proper response
|
||||
mock_response = Mock()
|
||||
mock_response.choices = [Mock()]
|
||||
mock_response.choices[0].message.content = "Hello back!"
|
||||
mock_vertex_completion.return_value = mock_response
|
||||
|
||||
try:
|
||||
# Call completion with custom headers via kwargs
|
||||
# This simulates what the proxy does when forward_client_headers_to_llm_api is set
|
||||
completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
headers=custom_headers, # This is how proxy passes forwarded headers
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_key="dummy-key",
|
||||
)
|
||||
|
||||
# Verify that the completion handler was called
|
||||
assert mock_vertex_completion.called, "Vertex completion handler should be called"
|
||||
|
||||
# Get the actual call arguments
|
||||
call_kwargs = mock_vertex_completion.call_args.kwargs
|
||||
|
||||
# Verify that extra_headers parameter contains our custom headers
|
||||
assert "extra_headers" in call_kwargs, "extra_headers should be passed to completion"
|
||||
|
||||
passed_headers = call_kwargs["extra_headers"]
|
||||
assert passed_headers is not None, "extra_headers should not be None"
|
||||
|
||||
# Verify our custom headers are present in the passed headers
|
||||
for header_key, header_value in custom_headers.items():
|
||||
assert (
|
||||
header_key in passed_headers
|
||||
or header_key.lower() in passed_headers
|
||||
), f"Header {header_key} should be in extra_headers"
|
||||
|
||||
print(f"✓ Test passed for {custom_llm_provider}/{model}")
|
||||
print(f" Headers correctly forwarded: {passed_headers}")
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(
|
||||
f"Failed to forward headers to {custom_llm_provider}/{model}: {str(e)}"
|
||||
)
|
||||
|
||||
def test_extra_headers_and_headers_merge(self):
|
||||
"""
|
||||
Test that both extra_headers and headers parameters are correctly merged.
|
||||
|
||||
This ensures that headers from kwargs (forwarded by proxy) and extra_headers
|
||||
(passed explicitly) are both included in the final headers sent to the provider.
|
||||
"""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
|
||||
# Headers from proxy (via kwargs["headers"])
|
||||
proxy_headers = {"X-Forwarded-Header": "ProxyValue"}
|
||||
|
||||
# Explicit extra_headers
|
||||
explicit_headers = {"X-Explicit-Header": "ExplicitValue"}
|
||||
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.VertexLLM.completion"
|
||||
) as mock_vertex_completion:
|
||||
mock_response = Mock()
|
||||
mock_response.choices = [Mock()]
|
||||
mock_response.choices[0].message.content = "Response"
|
||||
mock_vertex_completion.return_value = mock_response
|
||||
|
||||
try:
|
||||
completion(
|
||||
model="gemini/gemini-1.5-pro",
|
||||
messages=messages,
|
||||
headers=proxy_headers, # From proxy forwarding
|
||||
extra_headers=explicit_headers, # Explicitly passed
|
||||
custom_llm_provider="gemini",
|
||||
api_key="dummy-key",
|
||||
)
|
||||
|
||||
call_kwargs = mock_vertex_completion.call_args.kwargs
|
||||
passed_headers = call_kwargs.get("extra_headers", {})
|
||||
|
||||
# Both sets of headers should be present
|
||||
assert (
|
||||
"X-Forwarded-Header" in passed_headers
|
||||
or "x-forwarded-header" in passed_headers
|
||||
), "Proxy forwarded header should be present"
|
||||
|
||||
assert (
|
||||
"X-Explicit-Header" in passed_headers
|
||||
or "x-explicit-header" in passed_headers
|
||||
), "Explicitly passed header should be present"
|
||||
|
||||
print("✓ Both header sources correctly merged and forwarded")
|
||||
print(f" Final headers: {passed_headers}")
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(f"Failed to merge and forward headers: {str(e)}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run the tests
|
||||
test_instance = TestGeminiHeaderForwarding()
|
||||
|
||||
print("\n" + "="*80)
|
||||
print("Testing Gemini/Vertex AI Header Forwarding")
|
||||
print("="*80 + "\n")
|
||||
|
||||
# Test each provider
|
||||
for provider, model in [
|
||||
("gemini", "gemini/gemini-1.5-pro"),
|
||||
("vertex_ai_beta", "gemini-1.5-pro"),
|
||||
("vertex_ai", "gemini-1.5-pro"),
|
||||
]:
|
||||
print(f"\nTesting {provider}/{model}...")
|
||||
try:
|
||||
test_instance.test_headers_forwarded_to_gemini(provider, model)
|
||||
except Exception as e:
|
||||
print(f"✗ Test failed: {e}")
|
||||
|
||||
print("\n\nTesting header merging...")
|
||||
try:
|
||||
test_instance.test_extra_headers_and_headers_merge()
|
||||
except Exception as e:
|
||||
print(f"✗ Test failed: {e}")
|
||||
|
||||
print("\n" + "="*80)
|
||||
print("All tests completed!")
|
||||
print("="*80 + "\n")
|
||||
|
||||
|
|
@ -53,27 +53,38 @@ def test_llm_passthrough_route():
|
|||
|
||||
def test_bedrock_application_inference_profile_url_encoding():
|
||||
client = HTTPHandler()
|
||||
|
||||
|
||||
mock_provider_config = MagicMock()
|
||||
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"),
|
||||
"https://bedrock-runtime.us-east-1.amazonaws.com"
|
||||
httpx.URL(
|
||||
"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.validate_environment.return_value = {}
|
||||
mock_provider_config.sign_request.return_value = ({}, 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=("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:
|
||||
|
||||
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=("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_obj = MagicMock()
|
||||
mock_logging_obj.update_environment_variables = MagicMock()
|
||||
|
||||
|
||||
response = llm_passthrough_route(
|
||||
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",
|
||||
|
|
@ -86,7 +97,7 @@ def test_bedrock_application_inference_profile_url_encoding():
|
|||
# Verify that build_request was called with the encoded URL
|
||||
mock_build_request.assert_called_once()
|
||||
call_args = mock_build_request.call_args
|
||||
|
||||
|
||||
# The URL should have the application-inference-profile ID encoded
|
||||
actual_url = str(call_args.kwargs["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():
|
||||
client = HTTPHandler()
|
||||
|
||||
|
||||
# Mock the provider config and its methods
|
||||
mock_provider_config = MagicMock()
|
||||
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"),
|
||||
"https://bedrock-runtime.us-east-1.amazonaws.com"
|
||||
httpx.URL(
|
||||
"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.validate_environment.return_value = {}
|
||||
mock_provider_config.sign_request.return_value = ({}, 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=("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:
|
||||
|
||||
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=("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_obj = MagicMock()
|
||||
mock_logging_obj.update_environment_variables = MagicMock()
|
||||
|
||||
|
||||
response = llm_passthrough_route(
|
||||
model="anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
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)
|
||||
mock_build_request.assert_called_once()
|
||||
call_args = mock_build_request.call_args
|
||||
|
||||
|
||||
# The URL should NOT have application-inference-profile encoding
|
||||
actual_url = str(call_args.kwargs["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
|
||||
)
|
||||
assert result is True
|
||||
|
||||
|
||||
# Test 2: no stream in request body should return original stream param
|
||||
parsed_body = {"model": "test-model"}
|
||||
result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
|
||||
parsed_body=parsed_body, stream=False
|
||||
)
|
||||
assert result is False
|
||||
|
||||
|
||||
# Test 3: stream=False in request body should return False
|
||||
parsed_body = {"stream": False, "model": "test-model"}
|
||||
result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
|
||||
parsed_body=parsed_body, stream=True
|
||||
)
|
||||
assert result is False
|
||||
|
||||
|
||||
# Test 4: no stream param provided, no stream in body
|
||||
parsed_body = {"model": "test-model"}
|
||||
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():
|
||||
"""Create a mock request with headers"""
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class QueryParams:
|
||||
def __init__(self):
|
||||
self._dict = {}
|
||||
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self._dict)
|
||||
|
||||
|
||||
def items(self):
|
||||
return self._dict.items()
|
||||
|
||||
|
|
@ -210,6 +232,7 @@ def mock_request():
|
|||
def mock_user_api_key_dict():
|
||||
"""Create a mock user API key dictionary"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
return UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
|
|
@ -223,8 +246,8 @@ async def test_pass_through_request_stream_param_override(
|
|||
mock_request, mock_user_api_key_dict
|
||||
):
|
||||
"""
|
||||
Test that when stream=None is passed as parameter but stream=True
|
||||
is in request body, the request body value takes precedence and
|
||||
Test that when stream=None is passed as parameter but stream=True
|
||||
is in request body, the request body value takes precedence and
|
||||
the eventual POST request uses streaming.
|
||||
"""
|
||||
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",
|
||||
"max_tokens": 256,
|
||||
"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
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "text/event-stream"}
|
||||
|
||||
|
||||
# Mock the streaming response behavior
|
||||
async def mock_aiter_bytes():
|
||||
yield b'data: {"content": "Hello"}\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
|
||||
|
||||
# Create mocks for the async client
|
||||
mock_async_client = AsyncMock()
|
||||
mock_request_obj = AsyncMock()
|
||||
|
||||
|
||||
# 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 send to return the streaming 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
|
||||
|
||||
# Create the request
|
||||
request = mock_request(
|
||||
headers={}, method="POST", request_body=request_body
|
||||
)
|
||||
request = mock_request(headers={}, method="POST", request_body=request_body)
|
||||
|
||||
with patch(
|
||||
"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"),
|
||||
json=request_body,
|
||||
params={},
|
||||
headers={
|
||||
"Authorization": "Bearer test-key"
|
||||
},
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
|
||||
|
||||
# Verify that send was called with stream=True
|
||||
mock_async_client.send.assert_called_once_with(
|
||||
mock_request_obj,
|
||||
stream=True # This proves that stream=True from request body was used
|
||||
mock_request_obj,
|
||||
stream=True, # This proves that stream=True from request body was used
|
||||
)
|
||||
|
||||
|
||||
# Verify that the non-streaming request method was NOT called
|
||||
mock_async_client.request.assert_not_called()
|
||||
|
||||
|
||||
# Verify response is a StreamingResponse
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_stream_param_no_override(
|
||||
mock_request, mock_user_api_key_dict
|
||||
):
|
||||
"""
|
||||
Test that when stream=False is passed as parameter and no stream
|
||||
is in request body, the function parameter is used and
|
||||
Test that when stream=False is passed as parameter and no stream
|
||||
is in request body, the function parameter is used and
|
||||
the eventual request uses non-streaming.
|
||||
"""
|
||||
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
|
||||
request_body = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"max_tokens": 256,
|
||||
"messages": [{"role": "user", "content": "Hello, world"}],
|
||||
# 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.headers = {"content-type": "application/json"}
|
||||
mock_response._content = b'{"response": "Hello world"}'
|
||||
|
||||
|
||||
async def mock_aread():
|
||||
return mock_response._content
|
||||
|
||||
|
||||
mock_response.aread = mock_aread
|
||||
|
||||
# Create mocks for the async client
|
||||
mock_async_client = AsyncMock()
|
||||
|
||||
|
||||
# Mock request to return the non-streaming 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
|
||||
|
||||
# Create the request
|
||||
request = mock_request(
|
||||
headers={}, method="POST", request_body=request_body
|
||||
)
|
||||
request = mock_request(headers={}, method="POST", request_body=request_body)
|
||||
|
||||
with patch(
|
||||
"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)
|
||||
mock_async_client.build_request.assert_not_called()
|
||||
|
||||
|
||||
# Verify that send was NOT called (no streaming path)
|
||||
mock_async_client.send.assert_not_called()
|
||||
|
||||
|
||||
# Verify that the non-streaming request method WAS called
|
||||
mock_async_client.request.assert_called_once_with(
|
||||
method="POST",
|
||||
url=httpx.URL("https://api.anthropic.com/v1/messages"),
|
||||
headers={
|
||||
"Authorization": "Bearer test-key"
|
||||
},
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
params={},
|
||||
json=request_body,
|
||||
)
|
||||
|
||||
|
||||
# Verify response is a regular Response (not StreamingResponse)
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
|
||||
assert not isinstance(response, StreamingResponse)
|
||||
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
|
||||
|
|
|
|||
|
|
@ -734,3 +734,264 @@ async def test_call_mcp_tool_user_unauthorized_access():
|
|||
# Verify the exception details
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "User not allowed to call this tool" in exc_info.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_filters_by_key_team_permissions():
|
||||
"""Test that list_tools filters tools based on key/team mcp_tool_permissions"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_get_tools_from_mcp_servers,
|
||||
set_auth_context,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
# Create object permission with tool-level restrictions
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm_123",
|
||||
mcp_tool_permissions={
|
||||
"server1": ["tool1", "tool2"], # Only allow tool1 and tool2
|
||||
},
|
||||
)
|
||||
|
||||
user_api_key_auth = UserAPIKeyAuth(
|
||||
api_key="test_key",
|
||||
user_id="test_user",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
set_auth_context(user_api_key_auth)
|
||||
|
||||
# Mock server
|
||||
server = MagicMock()
|
||||
server.server_id = "server1"
|
||||
server.name = "Test Server"
|
||||
server.alias = "test"
|
||||
server.allowed_tools = None
|
||||
server.disallowed_tools = None
|
||||
|
||||
# Mock manager
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
|
||||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=False
|
||||
):
|
||||
# Return 4 tools, but only 2 should be allowed
|
||||
tool1 = MagicMock()
|
||||
tool1.name = "tool1"
|
||||
tool1.description = "Tool 1"
|
||||
tool1.inputSchema = {}
|
||||
|
||||
tool2 = MagicMock()
|
||||
tool2.name = "tool2"
|
||||
tool2.description = "Tool 2"
|
||||
tool2.inputSchema = {}
|
||||
|
||||
tool3 = MagicMock()
|
||||
tool3.name = "tool3"
|
||||
tool3.description = "Tool 3 - not allowed"
|
||||
tool3.inputSchema = {}
|
||||
|
||||
tool4 = MagicMock()
|
||||
tool4.name = "tool4"
|
||||
tool4.description = "Tool 4 - not allowed"
|
||||
tool4.inputSchema = {}
|
||||
|
||||
return [tool1, tool2, tool3, tool4]
|
||||
|
||||
mock_manager._get_tools_from_server = mock_get_tools_from_server
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
):
|
||||
tools = await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=None,
|
||||
mcp_servers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
)
|
||||
|
||||
# Should only return tool1 and tool2
|
||||
assert len(tools) == 2
|
||||
tool_names = sorted([t.name for t in tools])
|
||||
assert tool_names == ["tool1", "tool2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_with_team_tool_permissions_inheritance():
|
||||
"""Test that list_tools correctly applies key/team tool permissions inheritance logic"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_get_tools_from_mcp_servers,
|
||||
set_auth_context,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_TeamTable,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
# Team allows tool1, tool2, tool3
|
||||
team_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="team_perm_123",
|
||||
mcp_tool_permissions={
|
||||
"server1": ["tool1", "tool2", "tool3"],
|
||||
},
|
||||
)
|
||||
|
||||
# Key allows tool2, tool3, tool4 - intersection should be tool2, tool3
|
||||
key_object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="key_perm_456",
|
||||
mcp_tool_permissions={
|
||||
"server1": ["tool2", "tool3", "tool4"],
|
||||
},
|
||||
)
|
||||
|
||||
user_api_key_auth = UserAPIKeyAuth(
|
||||
api_key="test_key",
|
||||
user_id="test_user",
|
||||
team_id="team_123",
|
||||
object_permission=key_object_permission,
|
||||
)
|
||||
set_auth_context(user_api_key_auth)
|
||||
|
||||
# Mock server
|
||||
server = MagicMock()
|
||||
server.server_id = "server1"
|
||||
server.name = "Test Server"
|
||||
server.alias = "test"
|
||||
server.allowed_tools = None
|
||||
server.disallowed_tools = None
|
||||
|
||||
# Mock manager
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
|
||||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=False
|
||||
):
|
||||
# Return 4 tools
|
||||
tool1 = MagicMock()
|
||||
tool1.name = "tool1"
|
||||
tool1.description = "Tool 1"
|
||||
tool1.inputSchema = {}
|
||||
|
||||
tool2 = MagicMock()
|
||||
tool2.name = "tool2"
|
||||
tool2.description = "Tool 2"
|
||||
tool2.inputSchema = {}
|
||||
|
||||
tool3 = MagicMock()
|
||||
tool3.name = "tool3"
|
||||
tool3.description = "Tool 3"
|
||||
tool3.inputSchema = {}
|
||||
|
||||
tool4 = MagicMock()
|
||||
tool4.name = "tool4"
|
||||
tool4.description = "Tool 4"
|
||||
tool4.inputSchema = {}
|
||||
|
||||
return [tool1, tool2, tool3, tool4]
|
||||
|
||||
mock_manager._get_tools_from_server = mock_get_tools_from_server
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
):
|
||||
# Mock the team object permission retrieval
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_team_object_permission",
|
||||
AsyncMock(return_value=team_object_permission),
|
||||
):
|
||||
tools = await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=None,
|
||||
mcp_servers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
)
|
||||
|
||||
# Should only return tool2 and tool3 (intersection of key and team permissions)
|
||||
assert len(tools) == 2
|
||||
tool_names = sorted([t.name for t in tools])
|
||||
assert tool_names == ["tool2", "tool3"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_with_no_tool_permissions_shows_all():
|
||||
"""Test that list_tools shows all tools when no mcp_tool_permissions are set"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_get_tools_from_mcp_servers,
|
||||
set_auth_context,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
# No tool-level restrictions
|
||||
user_api_key_auth = UserAPIKeyAuth(
|
||||
api_key="test_key",
|
||||
user_id="test_user",
|
||||
object_permission=None,
|
||||
)
|
||||
set_auth_context(user_api_key_auth)
|
||||
|
||||
# Mock server
|
||||
server = MagicMock()
|
||||
server.server_id = "server1"
|
||||
server.name = "Test Server"
|
||||
server.alias = "test"
|
||||
server.allowed_tools = None
|
||||
server.disallowed_tools = None
|
||||
|
||||
# Mock manager
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
|
||||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=False
|
||||
):
|
||||
# Return 3 tools
|
||||
tool1 = MagicMock()
|
||||
tool1.name = "tool1"
|
||||
tool1.description = "Tool 1"
|
||||
tool1.inputSchema = {}
|
||||
|
||||
tool2 = MagicMock()
|
||||
tool2.name = "tool2"
|
||||
tool2.description = "Tool 2"
|
||||
tool2.inputSchema = {}
|
||||
|
||||
tool3 = MagicMock()
|
||||
tool3.name = "tool3"
|
||||
tool3.description = "Tool 3"
|
||||
tool3.inputSchema = {}
|
||||
|
||||
return [tool1, tool2, tool3]
|
||||
|
||||
mock_manager._get_tools_from_server = mock_get_tools_from_server
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
):
|
||||
tools = await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=None,
|
||||
mcp_servers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
)
|
||||
|
||||
# Should return all tools when no restrictions
|
||||
assert len(tools) == 3
|
||||
tool_names = sorted([t.name for t in tools])
|
||||
assert tool_names == ["tool1", "tool2", "tool3"]
|
||||
|
|
|
|||
|
|
@ -943,6 +943,201 @@ class TestMCPServerManager:
|
|||
manager.add_update_server(server)
|
||||
assert server.server_id in manager.get_registry()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_tool_permission_allows_permitted_tool(self):
|
||||
"""
|
||||
Test that key can call tool when it's in mcp_tool_permissions allowed list.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="test_server_123",
|
||||
name="Test Server",
|
||||
transport=MCPTransport.http,
|
||||
allowed_tools=None,
|
||||
disallowed_tools=None,
|
||||
)
|
||||
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm_123",
|
||||
mcp_tool_permissions={"test_server_123": ["read_wiki_structure"]},
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
api_key="sk-test",
|
||||
user_id="user-123",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
|
||||
|
||||
# Should succeed
|
||||
await manager.pre_call_tool_check(
|
||||
name="read_wiki_structure",
|
||||
arguments={"repoName": "facebook/react"},
|
||||
server_name_from_prefix="test",
|
||||
user_api_key_auth=user_auth,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
server=server,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_tool_permission_blocks_unpermitted_tool(self):
|
||||
"""
|
||||
Test that key cannot call tool when it's NOT in mcp_tool_permissions allowed list.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="test_server_123",
|
||||
name="Test Server",
|
||||
transport=MCPTransport.http,
|
||||
allowed_tools=None,
|
||||
disallowed_tools=None,
|
||||
)
|
||||
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm_123",
|
||||
mcp_tool_permissions={"test_server_123": ["read_wiki_structure"]},
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
api_key="sk-test",
|
||||
user_id="user-123",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
|
||||
|
||||
# Should fail with 403
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await manager.pre_call_tool_check(
|
||||
name="ask_question",
|
||||
arguments={"question": "test"},
|
||||
server_name_from_prefix="test",
|
||||
user_api_key_auth=user_auth,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
server=server,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_tool_permission_for_key_team_allows_permitted_tool(self):
|
||||
"""
|
||||
Test check_tool_permission_for_key_team directly - should allow permitted tool.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="github_server",
|
||||
name="GitHub Server",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm_456",
|
||||
mcp_tool_permissions={"github_server": ["read_repo", "list_issues"]},
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
api_key="sk-test-key",
|
||||
user_id="user-456",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
# Should not raise exception for allowed tool
|
||||
await manager.check_tool_permission_for_key_team(
|
||||
tool_name="read_repo",
|
||||
server=server,
|
||||
user_api_key_auth=user_auth,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_tool_permission_for_key_team_blocks_unpermitted_tool(self):
|
||||
"""
|
||||
Test check_tool_permission_for_key_team directly - should block unpermitted tool.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="github_server",
|
||||
name="GitHub Server",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm_456",
|
||||
mcp_tool_permissions={"github_server": ["read_repo"]},
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
api_key="sk-test-key",
|
||||
user_id="user-456",
|
||||
object_permission=object_permission,
|
||||
)
|
||||
|
||||
# Should raise HTTPException for unpermitted tool
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await manager.check_tool_permission_for_key_team(
|
||||
tool_name="delete_repo",
|
||||
server=server,
|
||||
user_api_key_auth=user_auth,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "delete_repo" in exc_info.value.detail["error"]
|
||||
assert "not allowed" in exc_info.value.detail["error"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_tool_permission_for_key_team_allows_all_when_no_restrictions(
|
||||
self,
|
||||
):
|
||||
"""
|
||||
Test check_tool_permission_for_key_team - should allow all tools when no restrictions set.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="github_server",
|
||||
name="GitHub Server",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
|
||||
# No object_permission set on user_auth
|
||||
user_auth = UserAPIKeyAuth(
|
||||
api_key="sk-test-key",
|
||||
user_id="user-456",
|
||||
object_permission=None,
|
||||
)
|
||||
|
||||
# Should allow any tool when no restrictions
|
||||
await manager.check_tool_permission_for_key_team(
|
||||
tool_name="any_tool",
|
||||
server=server,
|
||||
user_api_key_auth=user_auth,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
|
|
|||
|
|
@ -421,6 +421,80 @@ async def test_key_generation_with_object_permission(monkeypatch):
|
|||
assert key_insert_calls[0]["data"].get("object_permission_id") == "objperm123"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_generation_with_mcp_tool_permissions(monkeypatch):
|
||||
"""
|
||||
Test that /key/generate correctly handles mcp_tool_permissions in object_permission.
|
||||
|
||||
This test verifies that:
|
||||
1. mcp_tool_permissions is accepted in the object_permission field
|
||||
2. The field is properly stored in the LiteLLM_ObjectPermissionTable
|
||||
3. The key is correctly linked to the object_permission record
|
||||
"""
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.jsonify_object = lambda data: data
|
||||
|
||||
# Track what data is passed to create
|
||||
created_permission_data = {}
|
||||
|
||||
async def mock_create(**kwargs):
|
||||
created_permission_data.update(kwargs.get("data", {}))
|
||||
return MagicMock(object_permission_id="objperm_mcp_123")
|
||||
|
||||
mock_prisma_client.db = MagicMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable = MagicMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.create = mock_create
|
||||
|
||||
async def _insert_data_side_effect(*args, **kwargs):
|
||||
table_name = kwargs.get("table_name")
|
||||
if table_name == "user":
|
||||
return MagicMock(models=[], spend=0)
|
||||
elif table_name == "key":
|
||||
return MagicMock(
|
||||
token="hashed_token_789",
|
||||
litellm_budget_table=None,
|
||||
object_permission=None,
|
||||
)
|
||||
return MagicMock()
|
||||
|
||||
mock_prisma_client.insert_data = AsyncMock(side_effect=_insert_data_side_effect)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
|
||||
from litellm.proxy._types import (
|
||||
GenerateKeyRequest,
|
||||
LiteLLM_ObjectPermissionBase,
|
||||
LitellmUserRoles,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
generate_key_fn,
|
||||
)
|
||||
|
||||
# Create request with mcp_tool_permissions
|
||||
request_data = GenerateKeyRequest(
|
||||
object_permission=LiteLLM_ObjectPermissionBase(
|
||||
mcp_servers=["server_1"],
|
||||
mcp_tool_permissions={"server_1": ["tool1", "tool2", "tool3"]},
|
||||
)
|
||||
)
|
||||
|
||||
await generate_key_fn(
|
||||
data=request_data,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="user-mcp-1",
|
||||
),
|
||||
)
|
||||
|
||||
# Verify mcp_tool_permissions was stored
|
||||
assert "mcp_tool_permissions" in created_permission_data
|
||||
assert created_permission_data["mcp_tool_permissions"] == {
|
||||
"server_1": ["tool1", "tool2", "tool3"]
|
||||
}
|
||||
assert created_permission_data["mcp_servers"] == ["server_1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_update_object_permissions_existing_permission(monkeypatch):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -336,6 +336,90 @@ async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth):
|
|||
assert created_team_kwargs["data"].get("object_permission_id") == "objperm123"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_team_with_mcp_tool_permissions(mock_db_client, mock_admin_auth):
|
||||
"""
|
||||
Test that /team/new correctly handles mcp_tool_permissions in object_permission.
|
||||
|
||||
This test verifies that:
|
||||
1. mcp_tool_permissions is accepted in the object_permission field
|
||||
2. The field is properly stored in the LiteLLM_ObjectPermissionTable
|
||||
3. The team is correctly linked to the object_permission record
|
||||
"""
|
||||
# Configure mocked prisma client
|
||||
mock_db_client.jsonify_team_object = lambda db_data: db_data
|
||||
mock_db_client.get_data = AsyncMock(return_value=None)
|
||||
mock_db_client.update_data = AsyncMock(return_value=MagicMock())
|
||||
mock_db_client.db = MagicMock()
|
||||
|
||||
# Track what data is passed to object permission create
|
||||
created_permission_data = {}
|
||||
|
||||
async def mock_obj_perm_create(**kwargs):
|
||||
created_permission_data.update(kwargs.get("data", {}))
|
||||
return MagicMock(object_permission_id="objperm_team_mcp_456")
|
||||
|
||||
mock_db_client.db.litellm_objectpermissiontable = MagicMock()
|
||||
mock_db_client.db.litellm_objectpermissiontable.create = mock_obj_perm_create
|
||||
|
||||
# Mock model table
|
||||
mock_db_client.db.litellm_modeltable = MagicMock()
|
||||
mock_db_client.db.litellm_modeltable.create = AsyncMock(
|
||||
return_value=MagicMock(id="model456")
|
||||
)
|
||||
|
||||
# Mock team table
|
||||
team_create_result = MagicMock(
|
||||
team_id="team-mcp-789",
|
||||
object_permission_id="objperm_team_mcp_456",
|
||||
)
|
||||
team_create_result.model_dump.return_value = {
|
||||
"team_id": "team-mcp-789",
|
||||
"object_permission_id": "objperm_team_mcp_456",
|
||||
}
|
||||
mock_db_client.db.litellm_teamtable = MagicMock()
|
||||
mock_db_client.db.litellm_teamtable.create = AsyncMock(return_value=team_create_result)
|
||||
mock_db_client.db.litellm_teamtable.count = AsyncMock(return_value=0)
|
||||
mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=team_create_result)
|
||||
|
||||
# Mock user table
|
||||
mock_db_client.db.litellm_usertable = MagicMock()
|
||||
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionBase, NewTeamRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import new_team
|
||||
|
||||
# Create team with mcp_tool_permissions
|
||||
team_request = NewTeamRequest(
|
||||
team_alias="mcp-team",
|
||||
object_permission=LiteLLM_ObjectPermissionBase(
|
||||
mcp_servers=["server_a", "server_b"],
|
||||
mcp_tool_permissions={
|
||||
"server_a": ["read_wiki_structure", "read_wiki_contents"],
|
||||
"server_b": ["ask_question"],
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
await new_team(
|
||||
data=team_request,
|
||||
http_request=dummy_request,
|
||||
user_api_key_dict=mock_admin_auth,
|
||||
)
|
||||
|
||||
# Verify mcp_tool_permissions was stored
|
||||
assert "mcp_tool_permissions" in created_permission_data
|
||||
assert created_permission_data["mcp_tool_permissions"] == {
|
||||
"server_a": ["read_wiki_structure", "read_wiki_contents"],
|
||||
"server_b": ["ask_question"],
|
||||
}
|
||||
assert created_permission_data["mcp_servers"] == ["server_a", "server_b"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_update_object_permissions_existing_permission(monkeypatch):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -2030,3 +2030,89 @@ def test_add_callback_from_db_to_in_memory_litellm_callbacks():
|
|||
existing_callbacks=existing_callbacks_with_item,
|
||||
)
|
||||
mock_callback_manager.add_litellm_success_callback.assert_not_called()
|
||||
|
||||
|
||||
def test_should_load_db_object_with_supported_db_objects():
|
||||
"""
|
||||
Test _should_load_db_object method with supported_db_objects configuration.
|
||||
|
||||
Verifies that when supported_db_objects is set, only specified object types
|
||||
are loaded from the database.
|
||||
"""
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
proxy_config = ProxyConfig()
|
||||
|
||||
# Test Case 1: supported_db_objects not set - all objects should be loaded
|
||||
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is True
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||||
assert proxy_config._should_load_db_object(object_type="guardrails") is True
|
||||
assert proxy_config._should_load_db_object(object_type="vector_stores") is True
|
||||
|
||||
# Test Case 2: supported_db_objects set to only load MCP
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"supported_db_objects": ["mcp"]},
|
||||
):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is False
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||||
assert proxy_config._should_load_db_object(object_type="guardrails") is False
|
||||
assert proxy_config._should_load_db_object(object_type="vector_stores") is False
|
||||
assert proxy_config._should_load_db_object(object_type="prompts") is False
|
||||
|
||||
# Test Case 3: supported_db_objects set to load multiple types
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"supported_db_objects": ["mcp", "guardrails", "vector_stores"]},
|
||||
):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is False
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||||
assert proxy_config._should_load_db_object(object_type="guardrails") is True
|
||||
assert proxy_config._should_load_db_object(object_type="vector_stores") is True
|
||||
assert proxy_config._should_load_db_object(object_type="prompts") is False
|
||||
|
||||
# Test Case 4: supported_db_objects is not a list (should default to loading all)
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"supported_db_objects": "invalid_type"},
|
||||
):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is True
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||||
|
||||
# Test Case 5: supported_db_objects is an empty list (nothing should be loaded)
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"supported_db_objects": []},
|
||||
):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is False
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is False
|
||||
assert proxy_config._should_load_db_object(object_type="guardrails") is False
|
||||
|
||||
# Test Case 6: Test all available object types
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{
|
||||
"supported_db_objects": [
|
||||
"models",
|
||||
"mcp",
|
||||
"guardrails",
|
||||
"vector_stores",
|
||||
"pass_through_endpoints",
|
||||
"prompts",
|
||||
"model_cost_map",
|
||||
]
|
||||
},
|
||||
):
|
||||
assert proxy_config._should_load_db_object(object_type="models") is True
|
||||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||||
assert proxy_config._should_load_db_object(object_type="guardrails") is True
|
||||
assert proxy_config._should_load_db_object(object_type="vector_stores") is True
|
||||
assert (
|
||||
proxy_config._should_load_db_object(object_type="pass_through_endpoints")
|
||||
is True
|
||||
)
|
||||
assert proxy_config._should_load_db_object(object_type="prompts") is True
|
||||
assert proxy_config._should_load_db_object(object_type="model_cost_map") is True
|
||||
|
|
|
|||
11
ui/litellm-dashboard/.prettierignore
Normal file
11
ui/litellm-dashboard/.prettierignore
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
node_modules
|
||||
.next
|
||||
.out
|
||||
dist
|
||||
build
|
||||
.coverage
|
||||
.vercel
|
||||
.turbo
|
||||
.next-static
|
||||
*.min.js
|
||||
coverage/
|
||||
7
ui/litellm-dashboard/.prettierrc
Normal file
7
ui/litellm-dashboard/.prettierrc
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
{
|
||||
"semi": true,
|
||||
"singleQuote": false,
|
||||
"tabWidth": 2,
|
||||
"printWidth": 120,
|
||||
"trailingComma": "all"
|
||||
}
|
||||
|
|
@ -1,7 +0,0 @@
|
|||
{
|
||||
"semi": false,
|
||||
"tabWidth": 2,
|
||||
"printWidth": 120,
|
||||
"trailingComma": "all",
|
||||
"jsxBracketSameLine": false
|
||||
}
|
||||
|
|
@ -1,12 +1,12 @@
|
|||
/** @type {import('next').NextConfig} */
|
||||
const nextConfig = {
|
||||
output: 'export',
|
||||
basePath: '',
|
||||
assetPrefix: '/litellm-asset-prefix', // If a server_root_path is set, this will be overridden by runtime injection
|
||||
output: "export",
|
||||
basePath: "",
|
||||
assetPrefix: "/litellm-asset-prefix", // If a server_root_path is set, this will be overridden by runtime injection
|
||||
};
|
||||
|
||||
nextConfig.experimental = {
|
||||
missingSuspenseWithCSRBailout: false
|
||||
}
|
||||
missingSuspenseWithCSRBailout: false,
|
||||
};
|
||||
|
||||
export default nextConfig;
|
||||
|
|
|
|||
1
ui/litellm-dashboard/package-lock.json
generated
1
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -17973,6 +17973,7 @@
|
|||
"resolved": "https://registry.npmjs.org/prettier/-/prettier-3.2.5.tgz",
|
||||
"integrity": "sha512-3/GWa9aOC0YeD7LUfvOG2NiDyhOWRvt1k+rcKhOuYnMY24iiCphgneUfJDyFXd6rZCAnuLBv6UeAULtrhT/F4A==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"bin": {
|
||||
"prettier": "bin/prettier.cjs"
|
||||
},
|
||||
|
|
|
|||
|
|
@ -3,12 +3,14 @@
|
|||
"version": "0.1.0",
|
||||
"private": true,
|
||||
"scripts": {
|
||||
"dev": "next dev",
|
||||
"dev": "next dev --turbo",
|
||||
"build": "next build",
|
||||
"start": "next start",
|
||||
"lint": "next lint",
|
||||
"test": "vitest",
|
||||
"test:watch": "vitest -w"
|
||||
"test:watch": "vitest -w",
|
||||
"format": "prettier --write .",
|
||||
"format:check": "prettier --check ."
|
||||
},
|
||||
"dependencies": {
|
||||
"@anthropic-ai/sdk": "^0.54.0",
|
||||
|
|
|
|||
|
|
@ -19,12 +19,7 @@
|
|||
|
||||
body {
|
||||
color: rgb(var(--foreground-rgb));
|
||||
background: linear-gradient(
|
||||
to bottom,
|
||||
transparent,
|
||||
rgb(var(--background-end-rgb))
|
||||
)
|
||||
rgb(var(--background-start-rgb));
|
||||
background: linear-gradient(to bottom, transparent, rgb(var(--background-end-rgb))) rgb(var(--background-start-rgb));
|
||||
}
|
||||
|
||||
@layer utilities {
|
||||
|
|
|
|||
|
|
@ -19,7 +19,5 @@ export default function PublicModelHub() {
|
|||
* populate navbar
|
||||
*
|
||||
*/
|
||||
return (
|
||||
<PublicModelHubPage accessToken={accessToken} />
|
||||
);
|
||||
return <PublicModelHubPage accessToken={accessToken} />;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -19,7 +19,5 @@ export default function PublicModelHubTable() {
|
|||
* populate navbar
|
||||
*
|
||||
*/
|
||||
return (
|
||||
<ModelHubTable accessToken={accessToken} publicPage={true} premiumUser={false} userRole={null}/>
|
||||
);
|
||||
}
|
||||
return <ModelHubTable accessToken={accessToken} publicPage={true} premiumUser={false} userRole={null} />;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,16 +1,7 @@
|
|||
"use client";
|
||||
import React, { Suspense, useEffect, useState } from "react";
|
||||
import { useSearchParams } from "next/navigation";
|
||||
import {
|
||||
Card,
|
||||
Title,
|
||||
Text,
|
||||
TextInput,
|
||||
Callout,
|
||||
Button,
|
||||
Grid,
|
||||
Col,
|
||||
} from "@tremor/react";
|
||||
import { Card, Title, Text, TextInput, Callout, Button, Grid, Col } from "@tremor/react";
|
||||
import { RiAlarmWarningLine, RiCheckboxCircleLine } from "@remixicon/react";
|
||||
import {
|
||||
invitationClaimCall,
|
||||
|
|
@ -18,7 +9,7 @@ import {
|
|||
getOnboardingCredentials,
|
||||
claimOnboardingToken,
|
||||
getUiConfig,
|
||||
getProxyBaseUrl
|
||||
getProxyBaseUrl,
|
||||
} from "@/components/networking";
|
||||
import { jwtDecode } from "jwt-decode";
|
||||
import { Form, Button as Button2, message } from "antd";
|
||||
|
|
@ -27,7 +18,7 @@ import { getCookie } from "@/utils/cookieUtils";
|
|||
export default function Onboarding() {
|
||||
const [form] = Form.useForm();
|
||||
const searchParams = useSearchParams()!;
|
||||
const token = getCookie('token');
|
||||
const token = getCookie("token");
|
||||
const inviteID = searchParams.get("invitation_id");
|
||||
const action = searchParams.get("action");
|
||||
const [accessToken, setAccessToken] = useState<string | null>(null);
|
||||
|
|
@ -39,14 +30,16 @@ export default function Onboarding() {
|
|||
const [getUiConfigLoading, setGetUiConfigLoading] = useState<boolean>(true);
|
||||
|
||||
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);
|
||||
setGetUiConfigLoading(false);
|
||||
});
|
||||
}, []);
|
||||
|
||||
useEffect(() => {
|
||||
if (!inviteID || getUiConfigLoading) { // wait for the ui config to be loaded
|
||||
if (!inviteID || getUiConfigLoading) {
|
||||
// wait for the ui config to be loaded
|
||||
return;
|
||||
}
|
||||
|
||||
|
|
@ -72,14 +65,7 @@ export default function Onboarding() {
|
|||
}, [inviteID, getUiConfigLoading]);
|
||||
|
||||
const handleSubmit = (formValues: Record<string, any>) => {
|
||||
console.log(
|
||||
"in handle submit. accessToken:",
|
||||
accessToken,
|
||||
"token:",
|
||||
jwtToken,
|
||||
"formValues:",
|
||||
formValues
|
||||
);
|
||||
console.log("in handle submit. accessToken:", accessToken, "token:", jwtToken, "formValues:", formValues);
|
||||
if (!accessToken || !jwtToken) {
|
||||
return;
|
||||
}
|
||||
|
|
@ -89,12 +75,7 @@ export default function Onboarding() {
|
|||
if (!userID || !inviteID) {
|
||||
return;
|
||||
}
|
||||
claimOnboardingToken(
|
||||
accessToken,
|
||||
inviteID,
|
||||
userID,
|
||||
formValues.password
|
||||
).then((data) => {
|
||||
claimOnboardingToken(accessToken, inviteID, userID, formValues.password).then((data) => {
|
||||
let litellm_dashboard_ui = "/ui/";
|
||||
litellm_dashboard_ui += "?login=success";
|
||||
|
||||
|
|
@ -119,15 +100,14 @@ export default function Onboarding() {
|
|||
<Card>
|
||||
<Title className="text-sm mb-5 text-center">🚅 LiteLLM</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" && (
|
||||
<Callout
|
||||
className="mt-4"
|
||||
title="SSO"
|
||||
icon={RiCheckboxCircleLine}
|
||||
color="sky"
|
||||
>
|
||||
<Callout className="mt-4" title="SSO" icon={RiCheckboxCircleLine} color="sky">
|
||||
<Grid numItems={2} className="flex justify-between items-center">
|
||||
<Col>SSO is under the Enterprise Tier.</Col>
|
||||
|
||||
|
|
@ -142,28 +122,16 @@ export default function Onboarding() {
|
|||
</Callout>
|
||||
)}
|
||||
|
||||
<Form
|
||||
className="mt-10 mb-5 mx-auto"
|
||||
layout="vertical"
|
||||
onFinish={handleSubmit}
|
||||
>
|
||||
<Form className="mt-10 mb-5 mx-auto" layout="vertical" onFinish={handleSubmit}>
|
||||
<>
|
||||
<Form.Item label="Email Address" name="user_email">
|
||||
<TextInput
|
||||
type="email"
|
||||
disabled={true}
|
||||
value={userEmail}
|
||||
defaultValue={userEmail}
|
||||
className="max-w-md"
|
||||
/>
|
||||
<TextInput type="email" disabled={true} value={userEmail} defaultValue={userEmail} className="max-w-md" />
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label="Password"
|
||||
name="password"
|
||||
rules={[
|
||||
{ required: true, message: "password required to sign up" },
|
||||
]}
|
||||
rules={[{ required: true, message: "password required to sign up" }]}
|
||||
help={action === "reset_password" ? "Enter your new password" : "Create a password for your account"}
|
||||
>
|
||||
<TextInput placeholder="" type="password" className="max-w-md" />
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue