mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge branch 'main' into litellm_oct_staging2
This commit is contained in:
commit
ea69f4547d
71 changed files with 6415 additions and 469 deletions
|
|
@ -1041,6 +1041,49 @@ jobs:
|
|||
paths:
|
||||
- llm_responses_api_coverage.xml
|
||||
- llm_responses_api_coverage
|
||||
ocr_testing:
|
||||
docker:
|
||||
- image: cimg/python:3.11
|
||||
auth:
|
||||
username: ${DOCKERHUB_USERNAME}
|
||||
password: ${DOCKERHUB_PASSWORD}
|
||||
working_directory: ~/project
|
||||
|
||||
steps:
|
||||
- checkout
|
||||
- setup_google_dns
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
command: |
|
||||
python -m pip install --upgrade pip
|
||||
python -m pip install -r requirements.txt
|
||||
pip install "pytest==7.3.1"
|
||||
pip install "pytest-retry==1.6.3"
|
||||
pip install "pytest-cov==5.0.0"
|
||||
pip install "pytest-asyncio==0.21.1"
|
||||
pip install "respx==0.22.0"
|
||||
# Run pytest and generate JUnit XML report
|
||||
- run:
|
||||
name: Run tests
|
||||
command: |
|
||||
pwd
|
||||
ls
|
||||
python -m pytest -vv tests/ocr_tests --cov=litellm --cov-report=xml -x -s -v --junitxml=test-results/junit.xml --durations=5
|
||||
no_output_timeout: 120m
|
||||
- run:
|
||||
name: Rename the coverage files
|
||||
command: |
|
||||
mv coverage.xml ocr_coverage.xml
|
||||
mv .coverage ocr_coverage
|
||||
|
||||
# Store test results
|
||||
- store_test_results:
|
||||
path: test-results
|
||||
- persist_to_workspace:
|
||||
root: .
|
||||
paths:
|
||||
- ocr_coverage.xml
|
||||
- ocr_coverage
|
||||
litellm_mapped_tests:
|
||||
docker:
|
||||
- image: cimg/python:3.11
|
||||
|
|
@ -2741,7 +2784,7 @@ jobs:
|
|||
python -m venv venv
|
||||
. venv/bin/activate
|
||||
pip install coverage
|
||||
coverage combine llm_translation_coverage llm_responses_api_coverage mcp_coverage logging_coverage litellm_router_coverage local_testing_coverage litellm_assistants_api_coverage auth_ui_unit_tests_coverage langfuse_coverage caching_coverage litellm_proxy_unit_tests_coverage image_gen_coverage pass_through_unit_tests_coverage batches_coverage litellm_security_tests_coverage guardrails_coverage
|
||||
coverage combine llm_translation_coverage llm_responses_api_coverage ocr_coverage mcp_coverage logging_coverage litellm_router_coverage local_testing_coverage litellm_assistants_api_coverage auth_ui_unit_tests_coverage langfuse_coverage caching_coverage litellm_proxy_unit_tests_coverage image_gen_coverage pass_through_unit_tests_coverage batches_coverage litellm_security_tests_coverage guardrails_coverage
|
||||
coverage xml
|
||||
- codecov/upload:
|
||||
file: ./coverage.xml
|
||||
|
|
@ -3289,6 +3332,12 @@ workflows:
|
|||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- ocr_testing:
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- litellm_mapped_enterprise_tests:
|
||||
filters:
|
||||
branches:
|
||||
|
|
@ -3338,6 +3387,7 @@ workflows:
|
|||
- google_generate_content_endpoint_testing
|
||||
- guardrails_testing
|
||||
- llm_responses_api_testing
|
||||
- ocr_testing
|
||||
- litellm_mapped_tests
|
||||
- litellm_mapped_enterprise_tests
|
||||
- batches_testing
|
||||
|
|
@ -3400,6 +3450,7 @@ workflows:
|
|||
- mcp_testing
|
||||
- google_generate_content_endpoint_testing
|
||||
- llm_responses_api_testing
|
||||
- ocr_testing
|
||||
- litellm_mapped_tests
|
||||
- litellm_mapped_enterprise_tests
|
||||
- batches_testing
|
||||
|
|
|
|||
|
|
@ -347,6 +347,7 @@ curl 'http://0.0.0.0:4000/key/generate' \
|
|||
| [Nebius AI Studio](https://docs.litellm.ai/docs/providers/nebius) | ✅ | ✅ | ✅ | ✅ | ✅ | |
|
||||
| [Heroku](https://docs.litellm.ai/docs/providers/heroku) | ✅ | ✅ | | | | |
|
||||
| [OVHCloud AI Endpoints](https://docs.litellm.ai/docs/providers/ovhcloud) | ✅ | ✅ | | | | |
|
||||
| [CometAPI](https://docs.litellm.ai/docs/providers/cometapi) | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
|
||||
|
||||
[**Read the Docs**](https://docs.litellm.ai/docs/)
|
||||
|
||||
|
|
|
|||
474
cookbook/LiteLLM_CometAPI.ipynb
vendored
Normal file
474
cookbook/LiteLLM_CometAPI.ipynb
vendored
Normal file
File diff suppressed because one or more lines are too long
151
docs/my-website/docs/bedrock_converse.md
Normal file
151
docs/my-website/docs/bedrock_converse.md
Normal file
|
|
@ -0,0 +1,151 @@
|
|||
# /converse
|
||||
|
||||
Call Bedrock's `/converse` endpoint through LiteLLM Proxy.
|
||||
|
||||
| Feature | Supported |
|
||||
|---------|-----------|
|
||||
| Cost Tracking | ✅ |
|
||||
| Logging | ✅ |
|
||||
| Streaming | ✅ via `/converse-stream` |
|
||||
| Load Balancing | ✅ |
|
||||
|
||||
## Quick Start
|
||||
|
||||
### 1. Setup config.yaml
|
||||
|
||||
```yaml showLineNumbers
|
||||
model_list:
|
||||
- model_name: my-bedrock-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-west-2
|
||||
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID # reads from environment
|
||||
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
|
||||
custom_llm_provider: bedrock
|
||||
```
|
||||
|
||||
Set AWS credentials in your environment:
|
||||
|
||||
```bash showLineNumbers
|
||||
export AWS_ACCESS_KEY_ID="your-access-key"
|
||||
export AWS_SECRET_ACCESS_KEY="your-secret-key"
|
||||
```
|
||||
|
||||
### 2. Start Proxy
|
||||
|
||||
```bash showLineNumbers
|
||||
litellm --config config.yaml
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
### 3. Call /converse endpoint
|
||||
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/converse' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"text": "Hello, how are you?"}]
|
||||
}
|
||||
],
|
||||
"inferenceConfig": {
|
||||
"temperature": 0.5,
|
||||
"maxTokens": 100
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
## Streaming
|
||||
|
||||
For streaming responses, use `/converse-stream`:
|
||||
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/converse-stream' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"text": "Tell me a short story"}]
|
||||
}
|
||||
],
|
||||
"inferenceConfig": {
|
||||
"temperature": 0.7,
|
||||
"maxTokens": 200
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
## Load Balancing
|
||||
|
||||
Define multiple deployments with the same `model_name` for automatic load balancing:
|
||||
|
||||
```yaml showLineNumbers
|
||||
model_list:
|
||||
# Deployment 1 - us-west-2
|
||||
- model_name: my-bedrock-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-west-2
|
||||
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
|
||||
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
|
||||
custom_llm_provider: bedrock
|
||||
|
||||
# Deployment 2 - us-east-1
|
||||
- model_name: my-bedrock-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-east-1
|
||||
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
|
||||
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
|
||||
custom_llm_provider: bedrock
|
||||
```
|
||||
|
||||
The proxy automatically distributes requests across both regions.
|
||||
|
||||
## Using boto3 SDK
|
||||
|
||||
```python showLineNumbers
|
||||
import boto3
|
||||
import json
|
||||
import os
|
||||
|
||||
# Set dummy AWS credentials (required by boto3, but not used by LiteLLM proxy)
|
||||
os.environ['AWS_ACCESS_KEY_ID'] = 'dummy'
|
||||
os.environ['AWS_SECRET_ACCESS_KEY'] = 'dummy'
|
||||
os.environ['AWS_BEARER_TOKEN_BEDROCK'] = "sk-1234" # your litellm proxy api key
|
||||
|
||||
# Point boto3 to the LiteLLM proxy
|
||||
bedrock_runtime = boto3.client(
|
||||
service_name='bedrock-runtime',
|
||||
region_name='us-west-2',
|
||||
endpoint_url='http://0.0.0.0:4000/bedrock'
|
||||
)
|
||||
|
||||
response = bedrock_runtime.converse(
|
||||
modelId='my-bedrock-model', # Your model_name from config.yaml
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"text": "Hello, how are you?"}]
|
||||
}
|
||||
],
|
||||
inferenceConfig={
|
||||
"temperature": 0.5,
|
||||
"maxTokens": 100
|
||||
}
|
||||
)
|
||||
|
||||
print(response['output']['message']['content'][0]['text'])
|
||||
```
|
||||
|
||||
## More Info
|
||||
|
||||
For complete documentation including Guardrails, Knowledge Bases, and Agents, see:
|
||||
- [Full Bedrock Passthrough Docs](./pass_through/bedrock)
|
||||
|
||||
145
docs/my-website/docs/bedrock_invoke.md
Normal file
145
docs/my-website/docs/bedrock_invoke.md
Normal file
|
|
@ -0,0 +1,145 @@
|
|||
# /invoke
|
||||
|
||||
Call Bedrock's `/invoke` endpoint through LiteLLM Proxy.
|
||||
|
||||
| Feature | Supported |
|
||||
|---------|-----------|
|
||||
| Cost Tracking | ✅ |
|
||||
| Logging | ✅ |
|
||||
| Streaming | ✅ via `/invoke-with-response-stream` |
|
||||
| Load Balancing | ✅ |
|
||||
|
||||
## Quick Start
|
||||
|
||||
### 1. Setup config.yaml
|
||||
|
||||
```yaml showLineNumbers
|
||||
model_list:
|
||||
- model_name: my-bedrock-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-west-2
|
||||
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID # reads from environment
|
||||
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
|
||||
custom_llm_provider: bedrock
|
||||
```
|
||||
|
||||
Set AWS credentials in your environment:
|
||||
|
||||
```bash showLineNumbers
|
||||
export AWS_ACCESS_KEY_ID="your-access-key"
|
||||
export AWS_SECRET_ACCESS_KEY="your-secret-key"
|
||||
```
|
||||
|
||||
### 2. Start Proxy
|
||||
|
||||
```bash showLineNumbers
|
||||
litellm --config config.yaml
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
### 3. Call /invoke endpoint
|
||||
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/invoke' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"max_tokens": 100,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hello, how are you?"
|
||||
}
|
||||
],
|
||||
"anthropic_version": "bedrock-2023-05-31"
|
||||
}'
|
||||
```
|
||||
|
||||
## Streaming
|
||||
|
||||
For streaming responses, use `/invoke-with-response-stream`:
|
||||
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/invoke-with-response-stream' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"max_tokens": 100,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Tell me a short story"
|
||||
}
|
||||
],
|
||||
"anthropic_version": "bedrock-2023-05-31"
|
||||
}'
|
||||
```
|
||||
|
||||
## Load Balancing
|
||||
|
||||
Define multiple deployments with the same `model_name` for automatic load balancing:
|
||||
|
||||
```yaml showLineNumbers
|
||||
model_list:
|
||||
# Deployment 1 - us-west-2
|
||||
- model_name: my-bedrock-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-west-2
|
||||
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
|
||||
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
|
||||
custom_llm_provider: bedrock
|
||||
|
||||
# Deployment 2 - us-east-1
|
||||
- model_name: my-bedrock-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-east-1
|
||||
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
|
||||
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
|
||||
custom_llm_provider: bedrock
|
||||
```
|
||||
|
||||
The proxy automatically distributes requests across both regions.
|
||||
|
||||
## Using boto3 SDK
|
||||
|
||||
```python showLineNumbers
|
||||
import boto3
|
||||
import json
|
||||
import os
|
||||
|
||||
# Set dummy AWS credentials (required by boto3, but not used by LiteLLM proxy)
|
||||
os.environ['AWS_ACCESS_KEY_ID'] = 'dummy'
|
||||
os.environ['AWS_SECRET_ACCESS_KEY'] = 'dummy'
|
||||
os.environ['AWS_BEARER_TOKEN_BEDROCK'] = "sk-1234" # your litellm proxy api key
|
||||
|
||||
# Point boto3 to the LiteLLM proxy
|
||||
bedrock_runtime = boto3.client(
|
||||
service_name='bedrock-runtime',
|
||||
region_name='us-west-2',
|
||||
endpoint_url='http://0.0.0.0:4000/bedrock'
|
||||
)
|
||||
|
||||
response = bedrock_runtime.invoke_model(
|
||||
modelId='my-bedrock-model', # Your model_name from config.yaml
|
||||
contentType='application/json',
|
||||
accept='application/json',
|
||||
body=json.dumps({
|
||||
"max_tokens": 100,
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"anthropic_version": "bedrock-2023-05-31"
|
||||
})
|
||||
)
|
||||
|
||||
response_body = json.loads(response['body'].read())
|
||||
print(response_body['content'][0]['text'])
|
||||
```
|
||||
|
||||
## More Info
|
||||
|
||||
For complete documentation including Guardrails, Knowledge Bases, and Agents, see:
|
||||
- [Full Bedrock Passthrough Docs](./pass_through/bedrock)
|
||||
|
||||
|
|
@ -117,10 +117,52 @@ litellm_settings:
|
|||
```bash
|
||||
export SSL_CERTIFICATE="/path/to/certificate.pem"
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## 5. Use HTTP_PROXY environment variable
|
||||
## 5. Configure ECDH Curve for SSL/TLS Performance
|
||||
|
||||
The `ssl_ecdh_curve` setting allows you to configure the Elliptic Curve Diffie-Hellman (ECDH) curve used for SSL/TLS key exchange. This is particularly useful for disabling Post-Quantum Cryptography (PQC) to improve performance in environments where PQC is not required.
|
||||
|
||||
**Use Case:** Some OpenSSL 3.x systems enable PQC by default, which can slow down TLS handshakes. Setting the ECDH curve to `X25519` disables PQC and can significantly improve connection performance.
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
import litellm
|
||||
litellm.ssl_ecdh_curve = "X25519" # Disables PQC for better performance
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
```yaml
|
||||
litellm_settings:
|
||||
ssl_ecdh_curve: "X25519"
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="env_var" label="Environment Variables">
|
||||
|
||||
```bash
|
||||
export SSL_ECDH_CURVE="X25519"
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
**Common Valid Curves:**
|
||||
|
||||
- `X25519` - Modern, fast curve (recommended for disabling PQC)
|
||||
- `prime256v1` - NIST P-256 curve
|
||||
- `secp384r1` - NIST P-384 curve
|
||||
- `secp521r1` - NIST P-521 curve
|
||||
|
||||
**Note:** If an invalid curve name is provided or if your Python/OpenSSL version doesn't support this feature, LiteLLM will log a warning and continue with default curves.
|
||||
|
||||
## 6. Use HTTP_PROXY environment variable
|
||||
|
||||
Both httpx and aiohttp libraries use `urllib.request.getproxies` from environment variables. Before client initialization, you may set proxy (and optional SSL_CERT_FILE) by setting the environment variables:
|
||||
|
||||
|
|
|
|||
257
docs/my-website/docs/ocr.md
Normal file
257
docs/my-website/docs/ocr.md
Normal file
|
|
@ -0,0 +1,257 @@
|
|||
# /ocr
|
||||
|
||||
:::tip
|
||||
|
||||
LiteLLM follows the [Mistral API request/response for the OCR API](https://docs.mistral.ai/capabilities/vision/#optical-character-recognition-ocr)
|
||||
|
||||
:::
|
||||
|
||||
## **LiteLLM Python SDK Usage**
|
||||
### Quick Start
|
||||
|
||||
```python
|
||||
from litellm import ocr
|
||||
import os
|
||||
|
||||
os.environ["MISTRAL_API_KEY"] = "sk-.."
|
||||
|
||||
response = ocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": "https://arxiv.org/pdf/2201.04234"
|
||||
}
|
||||
)
|
||||
|
||||
# Access extracted text
|
||||
for page in response.pages:
|
||||
print(f"Page {page.index}:")
|
||||
print(page.markdown)
|
||||
```
|
||||
|
||||
### Async Usage
|
||||
|
||||
```python
|
||||
from litellm import aocr
|
||||
import os, asyncio
|
||||
|
||||
os.environ["MISTRAL_API_KEY"] = "sk-.."
|
||||
|
||||
async def test_async_ocr():
|
||||
response = await aocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": "https://arxiv.org/pdf/2201.04234"
|
||||
}
|
||||
)
|
||||
|
||||
# Access extracted text
|
||||
for page in response.pages:
|
||||
print(f"Page {page.index}:")
|
||||
print(page.markdown)
|
||||
|
||||
asyncio.run(test_async_ocr())
|
||||
```
|
||||
|
||||
### Using Base64 Encoded Documents
|
||||
|
||||
```python
|
||||
import base64
|
||||
from litellm import ocr
|
||||
|
||||
# Encode PDF to base64
|
||||
with open("document.pdf", "rb") as f:
|
||||
base64_pdf = base64.b64encode(f.read()).decode('utf-8')
|
||||
|
||||
response = ocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": f"data:application/pdf;base64,{base64_pdf}"
|
||||
}
|
||||
)
|
||||
```
|
||||
|
||||
### Optional Parameters
|
||||
|
||||
```python
|
||||
response = ocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/doc.pdf"
|
||||
},
|
||||
# Optional Mistral parameters
|
||||
pages=[0, 1, 2], # Only process specific pages
|
||||
include_image_base64=True, # Include extracted images
|
||||
image_limit=10, # Max images to return
|
||||
image_min_size=100 # Min image size to include
|
||||
)
|
||||
```
|
||||
|
||||
## **LiteLLM Proxy Usage**
|
||||
|
||||
LiteLLM provides a Mistral API compatible `/ocr` endpoint for OCR calls.
|
||||
|
||||
**Setup**
|
||||
|
||||
Add this to your litellm proxy config.yaml
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: mistral-ocr
|
||||
litellm_params:
|
||||
model: mistral/mistral-ocr-latest
|
||||
api_key: os.environ/MISTRAL_API_KEY
|
||||
```
|
||||
|
||||
Start litellm
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
Test request
|
||||
|
||||
```bash
|
||||
curl http://0.0.0.0:4000/v1/ocr \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "mistral-ocr",
|
||||
"document": {
|
||||
"type": "document_url",
|
||||
"document_url": "https://arxiv.org/pdf/2201.04234"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
|
||||
## **Request/Response Format**
|
||||
|
||||
:::info
|
||||
|
||||
LiteLLM follows the **Mistral OCR API specification**.
|
||||
|
||||
See the [official Mistral OCR documentation](https://docs.mistral.ai/capabilities/vision/#optical-character-recognition-ocr) for complete details.
|
||||
|
||||
:::
|
||||
|
||||
### Example Request
|
||||
|
||||
```python
|
||||
{
|
||||
"model": "mistral/mistral-ocr-latest",
|
||||
"document": {
|
||||
"type": "document_url",
|
||||
"document_url": "https://arxiv.org/pdf/2201.04234"
|
||||
},
|
||||
"pages": [0, 1, 2], # Optional: specific pages to process
|
||||
"include_image_base64": True, # Optional: include extracted images
|
||||
"image_limit": 10, # Optional: max images to return
|
||||
"image_min_size": 100 # Optional: min image size in pixels
|
||||
}
|
||||
```
|
||||
|
||||
### Request Parameters
|
||||
|
||||
| Parameter | Type | Required | Description |
|
||||
|-----------|------|----------|-------------|
|
||||
| `model` | string | Yes | The OCR model to use (e.g., `"mistral/mistral-ocr-latest"`) |
|
||||
| `document` | object | Yes | Document to process. Must contain `type` and URL field |
|
||||
| `document.type` | string | Yes | Either `"document_url"` for PDFs/docs or `"image_url"` for images |
|
||||
| `document.document_url` | string | Conditional | URL to the document (required if `type` is `"document_url"`) |
|
||||
| `document.image_url` | string | Conditional | URL to the image (required if `type` is `"image_url"`) |
|
||||
| `pages` | array | No | List of specific page indices to process (0-indexed) |
|
||||
| `include_image_base64` | boolean | No | Whether to include extracted images as base64 strings |
|
||||
| `image_limit` | integer | No | Maximum number of images to return |
|
||||
| `image_min_size` | integer | No | Minimum size (in pixels) for images to include |
|
||||
|
||||
#### Document Format Examples
|
||||
|
||||
**For PDFs and documents:**
|
||||
```json
|
||||
{
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/document.pdf"
|
||||
}
|
||||
```
|
||||
|
||||
**For images:**
|
||||
```json
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": "https://example.com/image.png"
|
||||
}
|
||||
```
|
||||
|
||||
**For base64-encoded content:**
|
||||
```json
|
||||
{
|
||||
"type": "document_url",
|
||||
"document_url": "data:application/pdf;base64,JVBERi0xLjQKJ..."
|
||||
}
|
||||
```
|
||||
|
||||
### Response Format
|
||||
|
||||
The response follows Mistral's OCR format with the following structure:
|
||||
|
||||
```json
|
||||
{
|
||||
"pages": [
|
||||
{
|
||||
"index": 0,
|
||||
"markdown": "# Document Title\n\nExtracted text content...",
|
||||
"dimensions": {
|
||||
"dpi": 200,
|
||||
"height": 2200,
|
||||
"width": 1700
|
||||
},
|
||||
"images": [
|
||||
{
|
||||
"image_base64": "base64string...",
|
||||
"bbox": {
|
||||
"x": 100,
|
||||
"y": 200,
|
||||
"width": 300,
|
||||
"height": 400
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"model": "mistral-ocr-2505-completion",
|
||||
"usage_info": {
|
||||
"pages_processed": 29,
|
||||
"doc_size_bytes": 3002783
|
||||
},
|
||||
"document_annotation": null,
|
||||
"object": "ocr"
|
||||
}
|
||||
```
|
||||
|
||||
#### Response Fields
|
||||
|
||||
| Field | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| `pages` | array | List of processed pages with extracted content |
|
||||
| `pages[].index` | integer | Page number (0-indexed) |
|
||||
| `pages[].markdown` | string | Extracted text in Markdown format |
|
||||
| `pages[].dimensions` | object | Page dimensions (dpi, height, width in pixels) |
|
||||
| `pages[].images` | array | Extracted images from the page (if `include_image_base64=true`) |
|
||||
| `model` | string | The model used for OCR processing |
|
||||
| `usage_info` | object | Processing statistics (pages processed, document size) |
|
||||
| `document_annotation` | object | Optional document-level annotations |
|
||||
| `object` | string | Always `"ocr"` for OCR responses |
|
||||
|
||||
|
||||
## **Supported Providers**
|
||||
|
||||
| Provider | Link to Usage |
|
||||
|-------------|--------------------|
|
||||
| Mistral AI | [Usage](#quick-start) |
|
||||
|
||||
|
|
@ -5,24 +5,55 @@ Pass-through endpoints for Bedrock - call provider-specific endpoint, in native
|
|||
| Feature | Supported | Notes |
|
||||
|-------|-------|-------|
|
||||
| Cost Tracking | ✅ | For `/invoke` and `/converse` endpoints |
|
||||
| Logging | ✅ | works across all integrations |
|
||||
| Load Balancing | ✅ | You can load balance `/invoke`, `/converse` routes across multiple deployments| Logging | ✅ | works across all integrations |
|
||||
| End-user Tracking | ❌ | [Tell us if you need this](https://github.com/BerriAI/litellm/issues/new) |
|
||||
| Streaming | ✅ | |
|
||||
|
||||
Just replace `https://bedrock-runtime.{aws_region_name}.amazonaws.com` with `LITELLM_PROXY_BASE_URL/bedrock` 🚀
|
||||
|
||||
#### **Example Usage**
|
||||
```bash
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse' \
|
||||
-H 'Authorization: Bearer anything' \
|
||||
## Overview
|
||||
|
||||
LiteLLM supports two ways to call Bedrock endpoints:
|
||||
|
||||
### 1. **Using config.yaml** (Recommended for model endpoints)
|
||||
|
||||
Define your Bedrock models in `config.yaml` and reference them by name. The proxy handles authentication and routing.
|
||||
|
||||
**Use for**: `/converse`, `/converse-stream`, `/invoke`, `/invoke-with-response-stream`
|
||||
|
||||
```yaml showLineNumbers
|
||||
model_list:
|
||||
- model_name: my-bedrock-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-west-2
|
||||
custom_llm_provider: bedrock
|
||||
```
|
||||
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/converse' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"messages": [
|
||||
{"role": "user",
|
||||
"content": [{"text": "Hello"}]
|
||||
}
|
||||
]
|
||||
}'
|
||||
-d '{"messages": [{"role": "user", "content": [{"text": "Hello"}]}]}'
|
||||
```
|
||||
|
||||
### 2. **Direct passthrough** (For non-model endpoints)
|
||||
|
||||
Set AWS credentials via environment variables and call Bedrock endpoints directly.
|
||||
|
||||
**Use for**: Guardrails, Knowledge Bases, Agents, and other non-model endpoints
|
||||
|
||||
```bash showLineNumbers
|
||||
export AWS_ACCESS_KEY_ID=""
|
||||
export AWS_SECRET_ACCESS_KEY=""
|
||||
export AWS_REGION_NAME="us-west-2"
|
||||
```
|
||||
|
||||
```bash showLineNumbers
|
||||
curl "http://0.0.0.0:4000/bedrock/guardrail/my-guardrail-id/version/1/apply" \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{"contents": [{"text": {"text": "Hello"}}], "source": "INPUT"}'
|
||||
```
|
||||
|
||||
Supports **ALL** Bedrock Endpoints (including streaming).
|
||||
|
|
@ -33,39 +64,235 @@ Supports **ALL** Bedrock Endpoints (including streaming).
|
|||
|
||||
Let's call the Bedrock [`/converse` endpoint](https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_Converse.html)
|
||||
|
||||
1. Add AWS Keys to your environment
|
||||
1. Create a `config.yaml` file with your Bedrock model
|
||||
|
||||
```bash
|
||||
```yaml showLineNumbers
|
||||
model_list:
|
||||
- model_name: my-bedrock-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-west-2
|
||||
custom_llm_provider: bedrock
|
||||
```
|
||||
|
||||
Set your AWS credentials:
|
||||
|
||||
```bash showLineNumbers
|
||||
export AWS_ACCESS_KEY_ID="" # Access key
|
||||
export AWS_SECRET_ACCESS_KEY="" # Secret access key
|
||||
export AWS_REGION_NAME="" # us-east-1, us-east-2, us-west-1, us-west-2
|
||||
```
|
||||
|
||||
2. Start LiteLLM Proxy
|
||||
|
||||
```bash
|
||||
litellm
|
||||
```bash showLineNumbers
|
||||
litellm --config config.yaml
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
3. Test it!
|
||||
|
||||
Let's call the Bedrock converse endpoint
|
||||
Let's call the Bedrock converse endpoint using the model name from config:
|
||||
|
||||
```bash
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse' \
|
||||
-H 'Authorization: Bearer anything' \
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/converse' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"messages": [
|
||||
{"role": "user",
|
||||
"content": [{"text": "Hello"}]
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"text": "Hello, how are you?"}]
|
||||
}
|
||||
],
|
||||
"inferenceConfig": {
|
||||
"maxTokens": 100
|
||||
}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
## Setup with config.yaml
|
||||
|
||||
Use config.yaml to define Bedrock models and use them via passthrough endpoints.
|
||||
|
||||
### 1. Define models in config.yaml
|
||||
|
||||
```yaml showLineNumbers
|
||||
model_list:
|
||||
- model_name: my-claude-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-west-2
|
||||
custom_llm_provider: bedrock
|
||||
|
||||
- model_name: my-cohere-model
|
||||
litellm_params:
|
||||
model: bedrock/cohere.command-r-v1:0
|
||||
aws_region_name: us-east-1
|
||||
custom_llm_provider: bedrock
|
||||
```
|
||||
|
||||
### 2. Start proxy with config
|
||||
|
||||
```bash showLineNumbers
|
||||
litellm --config config.yaml
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
### 3. Call Bedrock Converse endpoint
|
||||
|
||||
Use the `model_name` from config in the URL path:
|
||||
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-claude-model/converse' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"text": "Hello, how are you?"}]
|
||||
}
|
||||
],
|
||||
"inferenceConfig": {
|
||||
"temperature": 0.5,
|
||||
"maxTokens": 100
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
### 4. Call Bedrock Converse Stream endpoint
|
||||
|
||||
For streaming responses, use the `/converse-stream` endpoint:
|
||||
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-claude-model/converse-stream' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"text": "Tell me a short story"}]
|
||||
}
|
||||
],
|
||||
"inferenceConfig": {
|
||||
"temperature": 0.7,
|
||||
"maxTokens": 200
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
### Supported Bedrock Endpoints with config.yaml
|
||||
|
||||
When using models from config.yaml, you can call any Bedrock endpoint:
|
||||
|
||||
| Endpoint | Description | Example |
|
||||
|----------|-------------|---------|
|
||||
| `/model/{model_name}/converse` | Converse API | `http://0.0.0.0:4000/bedrock/model/my-claude-model/converse` |
|
||||
| `/model/{model_name}/converse-stream` | Streaming Converse | `http://0.0.0.0:4000/bedrock/model/my-claude-model/converse-stream` |
|
||||
| `/model/{model_name}/invoke` | Legacy Invoke API | `http://0.0.0.0:4000/bedrock/model/my-claude-model/invoke` |
|
||||
| `/model/{model_name}/invoke-with-response-stream` | Legacy Streaming | `http://0.0.0.0:4000/bedrock/model/my-claude-model/invoke-with-response-stream` |
|
||||
|
||||
The proxy automatically resolves the `model_name` to the actual Bedrock model ID and region configured in your `config.yaml`.
|
||||
|
||||
### Load Balancing Across Multiple Deployments
|
||||
|
||||
Define multiple Bedrock deployments with the same `model_name` to enable automatic load balancing.
|
||||
|
||||
#### 1. Define multiple deployments in config.yaml
|
||||
|
||||
```yaml showLineNumbers
|
||||
model_list:
|
||||
# First deployment - us-west-2
|
||||
- model_name: my-claude-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-west-2
|
||||
custom_llm_provider: bedrock
|
||||
|
||||
# Second deployment - us-east-1 (load balanced)
|
||||
- model_name: my-claude-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-east-1
|
||||
custom_llm_provider: bedrock
|
||||
```
|
||||
|
||||
#### 2. Start proxy with config
|
||||
|
||||
```bash showLineNumbers
|
||||
litellm --config config.yaml
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
#### 3. Call the endpoint - requests are automatically load balanced
|
||||
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-claude-model/invoke' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"max_tokens": 100,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hello, how are you?"
|
||||
}
|
||||
],
|
||||
"anthropic_version": "bedrock-2023-05-31"
|
||||
}'
|
||||
```
|
||||
|
||||
The proxy will automatically distribute requests across both `us-west-2` and `us-east-1` deployments. This works for all Bedrock endpoints: `/invoke`, `/invoke-with-response-stream`, `/converse`, and `/converse-stream`.
|
||||
|
||||
#### Using boto3 SDK with load balancing
|
||||
|
||||
You can also call the load-balanced endpoint using the boto3 SDK:
|
||||
|
||||
```python showLineNumbers
|
||||
import boto3
|
||||
import json
|
||||
import os
|
||||
|
||||
# Set dummy AWS credentials (required by boto3, but not used by LiteLLM proxy)
|
||||
os.environ['AWS_ACCESS_KEY_ID'] = 'dummy'
|
||||
os.environ['AWS_SECRET_ACCESS_KEY'] = 'dummy'
|
||||
os.environ['AWS_BEARER_TOKEN_BEDROCK'] = "sk-1234" # your litellm proxy api key
|
||||
|
||||
# Point boto3 to the LiteLLM proxy
|
||||
bedrock_runtime = boto3.client(
|
||||
service_name='bedrock-runtime',
|
||||
region_name='us-west-2',
|
||||
endpoint_url='http://0.0.0.0:4000/bedrock'
|
||||
)
|
||||
|
||||
# Call the load-balanced model
|
||||
response = bedrock_runtime.invoke_model(
|
||||
modelId='my-claude-model', # Your model_name from config.yaml
|
||||
contentType='application/json',
|
||||
accept='application/json',
|
||||
body=json.dumps({
|
||||
"max_tokens": 100,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hello, how are you?"
|
||||
}
|
||||
],
|
||||
"anthropic_version": "bedrock-2023-05-31"
|
||||
})
|
||||
)
|
||||
|
||||
# Parse response
|
||||
response_body = json.loads(response['body'].read())
|
||||
print(response_body['content'][0]['text'])
|
||||
```
|
||||
|
||||
The proxy will automatically load balance your boto3 requests across all configured deployments.
|
||||
|
||||
|
||||
## Examples
|
||||
|
||||
|
|
@ -84,7 +311,7 @@ Key Changes:
|
|||
|
||||
#### LiteLLM Proxy Call
|
||||
|
||||
```bash
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse' \
|
||||
-H 'Authorization: Bearer sk-anything' \
|
||||
-H 'Content-Type: application/json' \
|
||||
|
|
@ -99,7 +326,7 @@ curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse'
|
|||
|
||||
#### Direct Bedrock API Call
|
||||
|
||||
```bash
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'https://bedrock-runtime.us-west-2.amazonaws.com/model/cohere.command-r-v1:0/converse' \
|
||||
-H 'Authorization: AWS4-HMAC-SHA256..' \
|
||||
-H 'Content-Type: application/json' \
|
||||
|
|
@ -114,9 +341,25 @@ curl -X POST 'https://bedrock-runtime.us-west-2.amazonaws.com/model/cohere.comma
|
|||
|
||||
### **Example 2: Apply Guardrail**
|
||||
|
||||
**Setup**: Set AWS credentials for direct passthrough
|
||||
|
||||
```bash showLineNumbers
|
||||
export AWS_ACCESS_KEY_ID="your-access-key"
|
||||
export AWS_SECRET_ACCESS_KEY="your-secret-key"
|
||||
export AWS_REGION_NAME="us-west-2"
|
||||
```
|
||||
|
||||
Start proxy:
|
||||
|
||||
```bash showLineNumbers
|
||||
litellm
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
#### LiteLLM Proxy Call
|
||||
|
||||
```bash
|
||||
```bash showLineNumbers
|
||||
curl "http://0.0.0.0:4000/bedrock/guardrail/guardrailIdentifier/version/guardrailVersion/apply" \
|
||||
-H 'Authorization: Bearer sk-anything' \
|
||||
-H 'Content-Type: application/json' \
|
||||
|
|
@ -129,7 +372,7 @@ curl "http://0.0.0.0:4000/bedrock/guardrail/guardrailIdentifier/version/guardrai
|
|||
|
||||
#### Direct Bedrock API Call
|
||||
|
||||
```bash
|
||||
```bash showLineNumbers
|
||||
curl "https://bedrock-runtime.us-west-2.amazonaws.com/guardrail/guardrailIdentifier/version/guardrailVersion/apply" \
|
||||
-H 'Authorization: AWS4-HMAC-SHA256..' \
|
||||
-H 'Content-Type: application/json' \
|
||||
|
|
@ -142,7 +385,25 @@ curl "https://bedrock-runtime.us-west-2.amazonaws.com/guardrail/guardrailIdentif
|
|||
|
||||
### **Example 3: Query Knowledge Base**
|
||||
|
||||
```bash
|
||||
**Setup**: Set AWS credentials for direct passthrough
|
||||
|
||||
```bash showLineNumbers
|
||||
export AWS_ACCESS_KEY_ID="your-access-key"
|
||||
export AWS_SECRET_ACCESS_KEY="your-secret-key"
|
||||
export AWS_REGION_NAME="us-west-2"
|
||||
```
|
||||
|
||||
Start proxy:
|
||||
|
||||
```bash showLineNumbers
|
||||
litellm
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
#### LiteLLM Proxy Call
|
||||
|
||||
```bash showLineNumbers
|
||||
curl -X POST "http://0.0.0.0:4000/bedrock/knowledgebases/{knowledgeBaseId}/retrieve" \
|
||||
-H 'Authorization: Bearer sk-anything' \
|
||||
-H 'Content-Type: application/json' \
|
||||
|
|
@ -163,7 +424,7 @@ curl -X POST "http://0.0.0.0:4000/bedrock/knowledgebases/{knowledgeBaseId}/retri
|
|||
|
||||
#### Direct Bedrock API Call
|
||||
|
||||
```bash
|
||||
```bash showLineNumbers
|
||||
curl -X POST "https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases/{knowledgeBaseId}/retrieve" \
|
||||
-H 'Authorization: AWS4-HMAC-SHA256..' \
|
||||
-H 'Content-Type: application/json' \
|
||||
|
|
@ -194,7 +455,7 @@ Use this, to avoid giving developers the raw AWS Keys, but still letting them us
|
|||
|
||||
1. Setup environment
|
||||
|
||||
```bash
|
||||
```bash showLineNumbers
|
||||
export DATABASE_URL=""
|
||||
export LITELLM_MASTER_KEY=""
|
||||
export AWS_ACCESS_KEY_ID="" # Access key
|
||||
|
|
@ -202,7 +463,7 @@ export AWS_SECRET_ACCESS_KEY="" # Secret access key
|
|||
export AWS_REGION_NAME="" # us-east-1, us-east-2, us-west-1, us-west-2
|
||||
```
|
||||
|
||||
```bash
|
||||
```bash showLineNumbers
|
||||
litellm
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
|
|
@ -210,7 +471,7 @@ litellm
|
|||
|
||||
2. Generate virtual key
|
||||
|
||||
```bash
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/key/generate' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
|
|
@ -219,7 +480,7 @@ curl -X POST 'http://0.0.0.0:4000/key/generate' \
|
|||
|
||||
Expected Response
|
||||
|
||||
```bash
|
||||
```bash showLineNumbers
|
||||
{
|
||||
...
|
||||
"key": "sk-1234ewknldferwedojwojw"
|
||||
|
|
@ -229,7 +490,7 @@ Expected Response
|
|||
3. Test it!
|
||||
|
||||
|
||||
```bash
|
||||
```bash showLineNumbers
|
||||
curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse' \
|
||||
-H 'Authorization: Bearer sk-1234ewknldferwedojwojw' \
|
||||
-H 'Content-Type: application/json' \
|
||||
|
|
@ -246,46 +507,46 @@ curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse'
|
|||
|
||||
Call Bedrock Agents via LiteLLM proxy
|
||||
|
||||
```python
|
||||
**Setup**: Set AWS credentials on your LiteLLM proxy server
|
||||
|
||||
```bash showLineNumbers
|
||||
export AWS_ACCESS_KEY_ID="your-access-key"
|
||||
export AWS_SECRET_ACCESS_KEY="your-secret-key"
|
||||
export AWS_REGION_NAME="us-west-2"
|
||||
```
|
||||
|
||||
Start proxy:
|
||||
|
||||
```bash showLineNumbers
|
||||
litellm
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
**Usage from Python**:
|
||||
|
||||
```python showLineNumbers
|
||||
import os
|
||||
import boto3
|
||||
from botocore.config import Config
|
||||
|
||||
# # Define your proxy endpoint
|
||||
proxy_endpoint = "http://0.0.0.0:4000/bedrock" # 👈 your proxy base url
|
||||
|
||||
# # Create a Config object with the proxy
|
||||
# Custom headers
|
||||
custom_headers = {
|
||||
'litellm_user_api_key': 'Bearer sk-1234', # 👈 your proxy api key
|
||||
}
|
||||
|
||||
|
||||
os.environ["AWS_ACCESS_KEY_ID"] = "my-fake-key-id"
|
||||
os.environ["AWS_SECRET_ACCESS_KEY"] = "my-fake-access-key"
|
||||
import boto3
|
||||
|
||||
# Set dummy AWS credentials (required by boto3, but not used by LiteLLM proxy)
|
||||
os.environ["AWS_ACCESS_KEY_ID"] = "dummy"
|
||||
os.environ["AWS_SECRET_ACCESS_KEY"] = "dummy"
|
||||
os.environ["AWS_BEARER_TOKEN_BEDROCK"] = "sk-1234" # your litellm proxy api key
|
||||
|
||||
# Create the client
|
||||
runtime_client = boto3.client(
|
||||
service_name="bedrock-agent-runtime",
|
||||
region_name="us-west-2",
|
||||
endpoint_url=proxy_endpoint
|
||||
endpoint_url="http://0.0.0.0:4000/bedrock"
|
||||
)
|
||||
|
||||
# Custom header injection
|
||||
def inject_custom_headers(request, **kwargs):
|
||||
request.headers.update(custom_headers)
|
||||
|
||||
# Attach the event to inject custom headers before the request is sent
|
||||
runtime_client.meta.events.register('before-send.*.*', inject_custom_headers)
|
||||
|
||||
|
||||
response = runtime_client.invoke_agent(
|
||||
agentId="L1RT58GYRW",
|
||||
agentAliasId="MFPSBCXYTW",
|
||||
sessionId="12345",
|
||||
inputText="Who do you know?"
|
||||
)
|
||||
agentId="L1RT58GYRW",
|
||||
agentAliasId="MFPSBCXYTW",
|
||||
sessionId="12345",
|
||||
inputText="Who do you know?"
|
||||
)
|
||||
|
||||
completion = ""
|
||||
|
||||
|
|
@ -294,5 +555,4 @@ for event in response.get("completion"):
|
|||
completion += chunk["bytes"].decode()
|
||||
|
||||
print(completion)
|
||||
|
||||
```
|
||||
|
|
|
|||
|
|
@ -1,6 +1,10 @@
|
|||
# CometAPI
|
||||
LiteLLM supports all AI models from [CometAPI](https://www.cometapi.com/). CometAPI provides access to 500+ AI models through a unified API interface, including cutting-edge models like GPT-5, Claude Opus 4.1, and various other state-of-the-art language models.
|
||||
|
||||
<a target="_blank" href="https://colab.research.google.com/github/BerriAI/litellm/blob/main/cookbook/LiteLLM_CometAPI.ipynb">
|
||||
<img src="https://colab.research.google.com/assets/colab-badge.svg" alt="Open In Colab"/>
|
||||
</a>
|
||||
|
||||
## Authentication
|
||||
|
||||
To use CometAPI models, you need to obtain an API key from [CometAPI Token Console](https://api.cometapi.com/console/token). CometAPI offers free tokens for new users - you can get your free API key instantly by registering.
|
||||
|
|
|
|||
|
|
@ -320,6 +320,16 @@ Okta requires the `GENERIC_CLIENT_STATE` parameter:
|
|||
GENERIC_CLIENT_STATE="random-string" # Required for Okta
|
||||
```
|
||||
|
||||
### Okta PKCE
|
||||
|
||||
If your Okta application is configured to require PKCE (Proof Key for Code Exchange), enable it by setting:
|
||||
|
||||
```bash
|
||||
GENERIC_CLIENT_USE_PKCE="true"
|
||||
```
|
||||
|
||||
This is required when your Okta app settings enforce PKCE for enhanced security. LiteLLM will automatically handle PKCE parameter generation and verification during the OAuth flow.
|
||||
|
||||
### Common Configuration Issues
|
||||
|
||||
#### Missing Protocol in Base URL
|
||||
|
|
|
|||
|
|
@ -533,6 +533,7 @@ router_settings:
|
|||
| GENERIC_CLIENT_ID | Client ID for generic OAuth providers
|
||||
| GENERIC_CLIENT_SECRET | Client secret for generic OAuth providers
|
||||
| GENERIC_CLIENT_STATE | State parameter for generic client authentication
|
||||
| GENERIC_CLIENT_USE_PKCE | Enable PKCE (Proof Key for Code Exchange) for generic OAuth providers. Set to "true" when your OAuth provider requires PKCE. **Default is false**
|
||||
| GENERIC_SSO_HEADERS | Comma-separated list of additional headers to add to the request - e.g. Authorization=Bearer `<token>`, Content-Type=application/json, etc.
|
||||
| GENERIC_INCLUDE_CLIENT_ID | Include client ID in requests for OAuth
|
||||
| GENERIC_SCOPE | Scope settings for generic OAuth providers
|
||||
|
|
@ -752,6 +753,7 @@ router_settings:
|
|||
| SPEND_LOGS_URL | URL for retrieving spend logs
|
||||
| SPEND_LOG_CLEANUP_BATCH_SIZE | Number of logs deleted per batch during cleanup. Default is 1000
|
||||
| SSL_CERTIFICATE | Path to the SSL certificate file
|
||||
| SSL_ECDH_CURVE | ECDH curve for SSL/TLS key exchange (e.g., 'X25519' to disable PQC).
|
||||
| SSL_SECURITY_LEVEL | [BETA] Security level for SSL/TLS connections. E.g. `DEFAULT@SECLEVEL=1`
|
||||
| SSL_VERIFY | Flag to enable or disable SSL certificate verification
|
||||
| SSL_CERT_FILE | Path to the SSL certificate file for custom CA bundle
|
||||
|
|
|
|||
|
|
@ -38,13 +38,17 @@ model_list:
|
|||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "pillar-minitor-everything" # you can change my name
|
||||
- guardrail_name: "pillar-monitor-everything" # you can change my name
|
||||
litellm_params:
|
||||
guardrail: pillar
|
||||
mode: [pre_call, post_call] # Monitor both input and output
|
||||
api_key: os.environ/PILLAR_API_KEY # Your Pillar API key
|
||||
api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint
|
||||
on_flagged_action: "monitor" # Log threats but allow requests
|
||||
persist_session: true # Keep conversations visible in Pillar dashboard
|
||||
async_mode: false # Request synchronous verdicts
|
||||
include_scanners: true # Return scanner category breakdown
|
||||
include_evidence: true # Include detailed findings for triage
|
||||
default_on: true # Enable for all requests
|
||||
|
||||
general_settings:
|
||||
|
|
@ -104,10 +108,14 @@ guardrails:
|
|||
api_key: os.environ/PILLAR_API_KEY # Your Pillar API key
|
||||
api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint
|
||||
on_flagged_action: "block" # Block malicious requests
|
||||
persist_session: true # Keep records for investigation
|
||||
async_mode: false # Require an immediate verdict
|
||||
include_scanners: true # Understand which rule triggered
|
||||
include_evidence: true # Capture concrete evidence
|
||||
default_on: true # Enable for all requests
|
||||
|
||||
general_settings:
|
||||
master_key: "your-master-key-here"
|
||||
master_key: "YOUR_LITELLM_PROXY_MASTER_KEY"
|
||||
|
||||
litellm_settings:
|
||||
set_verbose: true
|
||||
|
|
@ -136,10 +144,14 @@ guardrails:
|
|||
api_key: os.environ/PILLAR_API_KEY # Your Pillar API key
|
||||
api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint
|
||||
on_flagged_action: "monitor" # Log threats but allow requests
|
||||
persist_session: false # Skip dashboard storage for low latency
|
||||
async_mode: false # Still receive results inline
|
||||
include_scanners: false # Minimal payload for performance
|
||||
include_evidence: false # Omit details to keep responses light
|
||||
default_on: true # Enable for all requests
|
||||
|
||||
general_settings:
|
||||
master_key: "your-secure-master-key-here"
|
||||
master_key: "YOUR_LITELLM_PROXY_MASTER_KEY"
|
||||
|
||||
litellm_settings:
|
||||
set_verbose: true # Enable detailed logging
|
||||
|
|
@ -169,10 +181,14 @@ guardrails:
|
|||
api_key: os.environ/PILLAR_API_KEY # Your Pillar API key
|
||||
api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint
|
||||
on_flagged_action: "block" # Block threats on input and output
|
||||
persist_session: true # Preserve conversations in Pillar dashboard
|
||||
async_mode: false # Require synchronous approval
|
||||
include_scanners: true # Inspect which scanners fired
|
||||
include_evidence: true # Include detailed evidence for auditing
|
||||
default_on: true # Enable for all requests
|
||||
|
||||
general_settings:
|
||||
master_key: "your-secure-master-key-here"
|
||||
master_key: "YOUR_LITELLM_PROXY_MASTER_KEY"
|
||||
|
||||
litellm_settings:
|
||||
set_verbose: true # Enable detailed logging
|
||||
|
|
@ -229,19 +245,139 @@ Logs the violation but allows the request to proceed:
|
|||
on_flagged_action: "monitor"
|
||||
```
|
||||
|
||||
## Advanced Configuration
|
||||
|
||||
**Quick takeaways**
|
||||
- Every request still runs *all* Pillar scanners; these options only change what comes back.
|
||||
- Choose richer responses when you need audit trails, lighter responses when latency or cost matters.
|
||||
- Blocking is controlled by LiteLLM’s `on_flagged_action` configuration—Pillar headers do not change block/monitor behaviour.
|
||||
|
||||
Pillar Security executes the full scanner suite on each call. The settings below tune the Protect response headers LiteLLM sends, letting you balance fidelity, retention, and latency.
|
||||
|
||||
### Response Control
|
||||
|
||||
#### Data Retention (`persist_session`)
|
||||
```yaml
|
||||
persist_session: false # Default: true
|
||||
```
|
||||
- **Why**: Controls whether Pillar stores session data for dashboard visibility.
|
||||
- **Set false for**: Ephemeral testing, privacy-sensitive interactions.
|
||||
- **Set true for**: Production monitoring, compliance, historical review (default behaviour).
|
||||
- **Impact**: `false` means the conversation will *not* appear in the Pillar dashboard.
|
||||
|
||||
#### Response Detail Level
|
||||
The following toggles grow the payload size without changing detection behaviour.
|
||||
|
||||
```yaml
|
||||
include_scanners: true # → plr_scanners (default true in LiteLLM)
|
||||
include_evidence: true # → plr_evidence (default true in LiteLLM)
|
||||
```
|
||||
|
||||
- **Minimal response** (`include_scanners=false`, `include_evidence=false`)
|
||||
```json
|
||||
{
|
||||
"session_id": "abc-123",
|
||||
"flagged": true
|
||||
}
|
||||
```
|
||||
Use when you only care about whether Pillar detected a threat.
|
||||
|
||||
> **📝 Note:** `flagged: true` means Pillar’s scanners recommend blocking. Pillar only reports this verdict—LiteLLM enforces your policy via the `on_flagged_action` configuration (no Pillar header controls it):
|
||||
> - `on_flagged_action: "block"` → LiteLLM raises a 400 guardrail error
|
||||
> - `on_flagged_action: "monitor"` → LiteLLM logs the threat but still returns the LLM response
|
||||
|
||||
- **Scanner breakdown** (`include_scanners=true`)
|
||||
```json
|
||||
{
|
||||
"session_id": "abc-123",
|
||||
"flagged": true,
|
||||
"scanners": {
|
||||
"jailbreak": true,
|
||||
"prompt_injection": false,
|
||||
"pii": false,
|
||||
"secret": false,
|
||||
"toxic_language": false
|
||||
/* ... more categories ... */
|
||||
}
|
||||
}
|
||||
```
|
||||
Use when you need to know which categories triggered.
|
||||
|
||||
- **Full context** (both toggles true)
|
||||
```json
|
||||
{
|
||||
"session_id": "abc-123",
|
||||
"flagged": true,
|
||||
"scanners": { /* ... */ },
|
||||
"evidence": [
|
||||
{
|
||||
"category": "jailbreak",
|
||||
"type": "prompt_injection",
|
||||
"evidence": "Ignore previous instructions",
|
||||
"metadata": { "start_idx": 0, "end_idx": 28 }
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
Ideal for debugging, audit logs, or compliance exports.
|
||||
|
||||
### Processing Mode (`async_mode`)
|
||||
```yaml
|
||||
async_mode: true # Default: false
|
||||
```
|
||||
- **Why**: Queue the request for background processing instead of waiting for a synchronous verdict.
|
||||
- **Response shape**:
|
||||
```json
|
||||
{
|
||||
"status": "queued",
|
||||
"session_id": "abc-123",
|
||||
"position": 1
|
||||
}
|
||||
```
|
||||
- **Set true for**: Large batch jobs, latency-tolerant pipelines.
|
||||
- **Set false for**: Real-time user flows (default).
|
||||
- ⚠️ **Note**: Async mode returns only a 202 queue acknowledgment (no flagged verdict). LiteLLM treats that as “no block,” so the pre-call hook always allows the request. Use async mode only for post-call or monitor-only workflows where delayed review is acceptable.
|
||||
|
||||
### Complete Examples
|
||||
|
||||
```yaml
|
||||
guardrails:
|
||||
# Production: full fidelity & dashboard visibility
|
||||
- guardrail_name: "pillar-production"
|
||||
litellm_params:
|
||||
guardrail: pillar
|
||||
mode: [pre_call, post_call]
|
||||
persist_session: true
|
||||
include_scanners: true
|
||||
include_evidence: true
|
||||
on_flagged_action: "block"
|
||||
|
||||
# Testing: lightweight, no persistence
|
||||
- guardrail_name: "pillar-testing"
|
||||
litellm_params:
|
||||
guardrail: pillar
|
||||
mode: pre_call
|
||||
persist_session: false
|
||||
include_scanners: false
|
||||
include_evidence: false
|
||||
on_flagged_action: "monitor"
|
||||
```
|
||||
|
||||
Keep in mind that LiteLLM forwards these values as the documented `plr_*` headers, so any direct HTTP integrations outside the proxy can reuse the same guidance.
|
||||
|
||||
## Examples
|
||||
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="safe" label="Simple Safe Request">
|
||||
|
||||
**Safe requset**
|
||||
**Safe request**
|
||||
|
||||
```bash
|
||||
# Test with safe content
|
||||
curl -X POST "http://localhost:4000/v1/chat/completions" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer your-master-key-here" \
|
||||
-H "Authorization: Bearer YOUR_LITELLM_PROXY_MASTER_KEY" \
|
||||
-d '{
|
||||
"model": "gpt-4.1-mini",
|
||||
"messages": [{"role": "user", "content": "Hello! Can you tell me a joke?"}],
|
||||
|
|
@ -300,7 +436,7 @@ curl -X POST "http://localhost:4000/v1/chat/completions" \
|
|||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/chat/completions" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer your-master-key-here" \
|
||||
-H "Authorization: Bearer YOUR_LITELLM_PROXY_MASTER_KEY" \
|
||||
-d '{
|
||||
"model": "gpt-4.1-mini",
|
||||
"messages": [
|
||||
|
|
@ -350,7 +486,7 @@ curl -X POST "http://localhost:4000/v1/chat/completions" \
|
|||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/chat/completions" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer your-master-key-here" \
|
||||
-H "Authorization: Bearer YOUR_LITELLM_PROXY_MASTER_KEY" \
|
||||
-d '{
|
||||
"model": "gpt-4.1-mini",
|
||||
"messages": [
|
||||
|
|
@ -405,4 +541,4 @@ Feel free to contact us at support@pillar.security
|
|||
- [Pillar Security API Docs](https://docs.pillar.security/docs/api/introduction)
|
||||
- [Pillar Security Dashboard](https://app.pillar.security)
|
||||
- [Pillar Security Website](https://pillar.security)
|
||||
- [LiteLLM Docs](https://docs.litellm.ai)
|
||||
- [LiteLLM Docs](https://docs.litellm.ai)
|
||||
|
|
|
|||
|
|
@ -252,7 +252,7 @@ litellm --config /path/to/config.yaml
|
|||
3. Use the MCP server in Claude Code
|
||||
|
||||
```bash
|
||||
claude mcp add --transport http litellm_proxy http://0.0.0.0:4000 --header "Authorization: Bearer sk-LITELLM_VIRTUAL_KEY"
|
||||
claude mcp add --transport http litellm_proxy http://0.0.0.0:4000/github_mcp/mcp --header "Authorization: Bearer sk-LITELLM_VIRTUAL_KEY"
|
||||
```
|
||||
|
||||
4. Authenticate via Claude Code
|
||||
|
|
|
|||
|
|
@ -347,6 +347,9 @@ const sidebars = {
|
|||
]
|
||||
},
|
||||
"moderation",
|
||||
"bedrock_invoke",
|
||||
"bedrock_converse",
|
||||
"ocr",
|
||||
{
|
||||
type: "category",
|
||||
label: "Pass-through Endpoints (Anthropic SDK, etc.)",
|
||||
|
|
@ -536,6 +539,7 @@ const sidebars = {
|
|||
"providers/datarobot",
|
||||
"providers/ovhcloud",
|
||||
"providers/wandb_inference",
|
||||
"providers/cometapi",
|
||||
],
|
||||
},
|
||||
{
|
||||
|
|
|
|||
|
|
@ -263,6 +263,7 @@ use_client: bool = False
|
|||
ssl_verify: Union[str, bool] = True
|
||||
ssl_security_level: Optional[str] = None
|
||||
ssl_certificate: Optional[str] = None
|
||||
ssl_ecdh_curve: Optional[str] = None # Set to 'X25519' to disable PQC and improve performance
|
||||
disable_streaming_logging: bool = False
|
||||
disable_token_counter: bool = False
|
||||
disable_add_transform_inline_image_block: bool = False
|
||||
|
|
@ -1288,6 +1289,7 @@ from .llms.hyperbolic.chat.transformation import HyperbolicChatConfig
|
|||
from .llms.vercel_ai_gateway.chat.transformation import VercelAIGatewayConfig
|
||||
from .llms.ovhcloud.chat.transformation import OVHCloudChatConfig
|
||||
from .llms.ovhcloud.embedding.transformation import OVHCloudEmbeddingConfig
|
||||
from .llms.cometapi.embed.transformation import CometAPIEmbeddingConfig
|
||||
from .llms.lemonade.chat.transformation import LemonadeChatConfig
|
||||
from .main import * # type: ignore
|
||||
from .integrations import *
|
||||
|
|
@ -1325,6 +1327,7 @@ from .batch_completion.main import * # type: ignore
|
|||
from .rerank_api.main import *
|
||||
from .llms.anthropic.experimental_pass_through.messages.handler import *
|
||||
from .responses.main import *
|
||||
from .ocr.main import *
|
||||
from .realtime_api.main import _arealtime
|
||||
from .fine_tuning.main import *
|
||||
from .files.main import *
|
||||
|
|
|
|||
|
|
@ -525,6 +525,7 @@ openai_compatible_providers: List = [
|
|||
"vercel_ai_gateway",
|
||||
"aiml",
|
||||
"wandb",
|
||||
"cometapi",
|
||||
]
|
||||
openai_text_completion_compatible_providers: List = (
|
||||
[ # providers that support `/v1/completions`
|
||||
|
|
@ -849,6 +850,7 @@ BEDROCK_CONVERSE_MODELS = [
|
|||
"deepseek.v3-v1:0",
|
||||
"openai.gpt-oss-20b-1:0",
|
||||
"openai.gpt-oss-120b-1:0",
|
||||
"anthropic.claude-haiku-4-5-20251001-v1:0",
|
||||
"anthropic.claude-sonnet-4-5-20250929-v1:0",
|
||||
"anthropic.claude-opus-4-1-20250805-v1:0",
|
||||
"anthropic.claude-opus-4-20250514-v1:0",
|
||||
|
|
|
|||
|
|
@ -693,6 +693,15 @@ class CostCalculatorUtils:
|
|||
model=model,
|
||||
image_response=completion_response,
|
||||
)
|
||||
elif custom_llm_provider == litellm.LlmProviders.COMETAPI.value:
|
||||
from litellm.llms.cometapi.image_generation.cost_calculator import (
|
||||
cost_calculator as cometapi_image_cost_calculator,
|
||||
)
|
||||
|
||||
return cometapi_image_cost_calculator(
|
||||
model=model,
|
||||
image_response=completion_response,
|
||||
)
|
||||
elif custom_llm_provider == litellm.LlmProviders.GEMINI.value:
|
||||
from litellm.llms.gemini.image_generation.cost_calculator import (
|
||||
cost_calculator as gemini_image_cost_calculator,
|
||||
|
|
|
|||
5
litellm/llms/azure_ai/ocr/__init__.py
Normal file
5
litellm/llms/azure_ai/ocr/__init__.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
"""Azure AI OCR module."""
|
||||
from .transformation import AzureAIOCRConfig
|
||||
|
||||
__all__ = ["AzureAIOCRConfig"]
|
||||
|
||||
268
litellm/llms/azure_ai/ocr/transformation.py
Normal file
268
litellm/llms/azure_ai/ocr/transformation.py
Normal file
|
|
@ -0,0 +1,268 @@
|
|||
"""
|
||||
Azure AI OCR transformation implementation.
|
||||
"""
|
||||
from typing import Dict, Optional
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.image_handling import (
|
||||
async_convert_url_to_base64,
|
||||
convert_url_to_base64,
|
||||
)
|
||||
from litellm.llms.base_llm.ocr.transformation import DocumentType, OCRRequestData
|
||||
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
class AzureAIOCRConfig(MistralOCRConfig):
|
||||
"""
|
||||
Azure AI OCR transformation configuration.
|
||||
|
||||
Azure AI uses Mistral's OCR API but with a different endpoint format.
|
||||
Inherits transformation logic from MistralOCRConfig since they use the same format.
|
||||
|
||||
Reference: Azure AI Foundry OCR documentation
|
||||
|
||||
Important: Azure AI only supports base64 data URIs (data:image/..., data:application/pdf;base64,...).
|
||||
Regular URLs are not supported.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Dict,
|
||||
model: str,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Dict:
|
||||
"""
|
||||
Validate environment and return headers for Azure AI OCR.
|
||||
|
||||
Azure AI uses Bearer token authentication with AZURE_AI_API_KEY.
|
||||
"""
|
||||
# Get API key from environment if not provided
|
||||
if api_key is None:
|
||||
api_key = get_secret_str("AZURE_AI_API_KEY")
|
||||
|
||||
if api_key is None:
|
||||
raise ValueError(
|
||||
"Missing Azure AI API Key - A call is being made to Azure AI but no key is set either in the environment variables or via params"
|
||||
)
|
||||
|
||||
# Validate API base is provided
|
||||
if api_base is None:
|
||||
api_base = get_secret_str("AZURE_AI_API_BASE")
|
||||
|
||||
if api_base is None:
|
||||
raise ValueError(
|
||||
"Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter"
|
||||
)
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
**headers,
|
||||
}
|
||||
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
Get complete URL for Azure AI OCR endpoint.
|
||||
|
||||
Azure AI endpoint format: https://<api_base>/providers/mistral/azure/ocr
|
||||
|
||||
Args:
|
||||
api_base: Azure AI API base URL
|
||||
model: Model name (not used in URL construction)
|
||||
optional_params: Optional parameters
|
||||
|
||||
Returns: Complete URL for Azure AI OCR endpoint
|
||||
"""
|
||||
if api_base is None:
|
||||
raise ValueError(
|
||||
"Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter"
|
||||
)
|
||||
|
||||
# Ensure no trailing slash
|
||||
api_base = api_base.rstrip("/")
|
||||
|
||||
# Azure AI OCR endpoint format
|
||||
return f"{api_base}/providers/mistral/azure/ocr"
|
||||
|
||||
def _convert_url_to_data_uri_sync(self, url: str) -> str:
|
||||
"""
|
||||
Synchronously convert a URL to a base64 data URI.
|
||||
|
||||
Azure AI OCR doesn't have internet access, so we need to fetch URLs
|
||||
and convert them to base64 data URIs.
|
||||
|
||||
Args:
|
||||
url: The URL to convert
|
||||
|
||||
Returns:
|
||||
Base64 data URI string
|
||||
"""
|
||||
verbose_logger.debug(f"Azure AI OCR: Converting URL to base64 data URI (sync): {url}")
|
||||
|
||||
# Fetch and convert to base64 data URI
|
||||
# convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
|
||||
data_uri = convert_url_to_base64(url=url)
|
||||
|
||||
verbose_logger.debug(f"Azure AI OCR: Converted URL to data URI (length: {len(data_uri)})")
|
||||
|
||||
return data_uri
|
||||
|
||||
async def _convert_url_to_data_uri_async(self, url: str) -> str:
|
||||
"""
|
||||
Asynchronously convert a URL to a base64 data URI.
|
||||
|
||||
Azure AI OCR doesn't have internet access, so we need to fetch URLs
|
||||
and convert them to base64 data URIs.
|
||||
|
||||
Args:
|
||||
url: The URL to convert
|
||||
|
||||
Returns:
|
||||
Base64 data URI string
|
||||
"""
|
||||
verbose_logger.debug(f"Azure AI OCR: Converting URL to base64 data URI (async): {url}")
|
||||
|
||||
# Fetch and convert to base64 data URI asynchronously
|
||||
# async_convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
|
||||
data_uri = await async_convert_url_to_base64(url=url)
|
||||
|
||||
verbose_logger.debug(f"Azure AI OCR: Converted URL to data URI (length: {len(data_uri)})")
|
||||
|
||||
return data_uri
|
||||
|
||||
def transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
"""
|
||||
Transform OCR request for Azure AI, converting URLs to base64 data URIs (sync).
|
||||
|
||||
Azure AI OCR doesn't have internet access, so we automatically fetch
|
||||
any URLs and convert them to base64 data URIs synchronously.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
document: Document dict from user
|
||||
optional_params: Already mapped optional parameters
|
||||
headers: Request headers
|
||||
**kwargs: Additional arguments
|
||||
|
||||
Returns:
|
||||
OCRRequestData with JSON data
|
||||
"""
|
||||
verbose_logger.debug(f"Azure AI OCR transform_ocr_request (sync) - model: {model}")
|
||||
|
||||
if not isinstance(document, dict):
|
||||
raise ValueError(f"Expected document dict, got {type(document)}")
|
||||
|
||||
# Check if we need to convert URL to base64
|
||||
doc_type = document.get("type")
|
||||
transformed_document = document.copy()
|
||||
|
||||
if doc_type == "document_url":
|
||||
document_url = document.get("document_url", "")
|
||||
# If it's not already a data URI, convert it
|
||||
if document_url and not document_url.startswith("data:"):
|
||||
verbose_logger.debug(
|
||||
"Azure AI OCR: Converting document URL to base64 data URI (sync)"
|
||||
)
|
||||
data_uri = self._convert_url_to_data_uri_sync(url=document_url)
|
||||
transformed_document["document_url"] = data_uri
|
||||
elif doc_type == "image_url":
|
||||
image_url = document.get("image_url", "")
|
||||
# If it's not already a data URI, convert it
|
||||
if image_url and not image_url.startswith("data:"):
|
||||
verbose_logger.debug(
|
||||
"Azure AI OCR: Converting image URL to base64 data URI (sync)"
|
||||
)
|
||||
data_uri = self._convert_url_to_data_uri_sync(url=image_url)
|
||||
transformed_document["image_url"] = data_uri
|
||||
|
||||
# Call parent's transform to build the request
|
||||
return super().transform_ocr_request(
|
||||
model=model,
|
||||
document=transformed_document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def async_transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
"""
|
||||
Transform OCR request for Azure AI, converting URLs to base64 data URIs (async).
|
||||
|
||||
Azure AI OCR doesn't have internet access, so we automatically fetch
|
||||
any URLs and convert them to base64 data URIs asynchronously.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
document: Document dict from user
|
||||
optional_params: Already mapped optional parameters
|
||||
headers: Request headers
|
||||
**kwargs: Additional arguments
|
||||
|
||||
Returns:
|
||||
OCRRequestData with JSON data
|
||||
"""
|
||||
verbose_logger.debug(f"Azure AI OCR async_transform_ocr_request - model: {model}")
|
||||
|
||||
if not isinstance(document, dict):
|
||||
raise ValueError(f"Expected document dict, got {type(document)}")
|
||||
|
||||
# Check if we need to convert URL to base64
|
||||
doc_type = document.get("type")
|
||||
transformed_document = document.copy()
|
||||
|
||||
if doc_type == "document_url":
|
||||
document_url = document.get("document_url", "")
|
||||
# If it's not already a data URI, convert it
|
||||
if document_url and not document_url.startswith("data:"):
|
||||
verbose_logger.debug(
|
||||
"Azure AI OCR: Converting document URL to base64 data URI (async)"
|
||||
)
|
||||
data_uri = await self._convert_url_to_data_uri_async(url=document_url)
|
||||
transformed_document["document_url"] = data_uri
|
||||
elif doc_type == "image_url":
|
||||
image_url = document.get("image_url", "")
|
||||
# If it's not already a data URI, convert it
|
||||
if image_url and not image_url.startswith("data:"):
|
||||
verbose_logger.debug(
|
||||
"Azure AI OCR: Converting image URL to base64 data URI (async)"
|
||||
)
|
||||
data_uri = await self._convert_url_to_data_uri_async(url=image_url)
|
||||
transformed_document["image_url"] = data_uri
|
||||
|
||||
# Call parent's transform to build the request
|
||||
return super().transform_ocr_request(
|
||||
model=model,
|
||||
document=transformed_document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
22
litellm/llms/base_llm/ocr/__init__.py
Normal file
22
litellm/llms/base_llm/ocr/__init__.py
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
"""Base OCR transformation module."""
|
||||
from .transformation import (
|
||||
BaseOCRConfig,
|
||||
DocumentType,
|
||||
OCRPage,
|
||||
OCRPageDimensions,
|
||||
OCRPageImage,
|
||||
OCRRequestData,
|
||||
OCRResponse,
|
||||
OCRUsageInfo,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"BaseOCRConfig",
|
||||
"DocumentType",
|
||||
"OCRResponse",
|
||||
"OCRPage",
|
||||
"OCRPageDimensions",
|
||||
"OCRPageImage",
|
||||
"OCRUsageInfo",
|
||||
"OCRRequestData",
|
||||
]
|
||||
207
litellm/llms/base_llm/ocr/transformation.py
Normal file
207
litellm/llms/base_llm/ocr/transformation.py
Normal file
|
|
@ -0,0 +1,207 @@
|
|||
"""
|
||||
Base OCR transformation configuration.
|
||||
"""
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
# DocumentType for OCR - Mistral format document dict
|
||||
DocumentType = Dict[str, str]
|
||||
|
||||
|
||||
class OCRPageDimensions(BaseModel):
|
||||
"""Page dimensions from OCR response."""
|
||||
dpi: Optional[int] = None
|
||||
height: Optional[int] = None
|
||||
width: Optional[int] = None
|
||||
|
||||
|
||||
class OCRPageImage(BaseModel):
|
||||
"""Image extracted from OCR page."""
|
||||
image_base64: Optional[str] = None
|
||||
bbox: Optional[Dict[str, Any]] = None
|
||||
|
||||
model_config = {"extra": "allow"}
|
||||
|
||||
|
||||
class OCRPage(BaseModel):
|
||||
"""Single page from OCR response."""
|
||||
index: int
|
||||
markdown: str
|
||||
images: Optional[List[OCRPageImage]] = None
|
||||
dimensions: Optional[OCRPageDimensions] = None
|
||||
|
||||
model_config = {"extra": "allow"}
|
||||
|
||||
|
||||
class OCRUsageInfo(BaseModel):
|
||||
"""Usage information from OCR response."""
|
||||
pages_processed: Optional[int] = None
|
||||
doc_size_bytes: Optional[int] = None
|
||||
|
||||
model_config = {"extra": "allow"}
|
||||
|
||||
|
||||
class OCRResponse(BaseModel):
|
||||
"""
|
||||
Standard OCR response format.
|
||||
Standardized to Mistral OCR format - other providers should transform to this format.
|
||||
"""
|
||||
pages: List[OCRPage]
|
||||
model: str
|
||||
document_annotation: Optional[Any] = None
|
||||
usage_info: Optional[OCRUsageInfo] = None
|
||||
object: str = "ocr"
|
||||
|
||||
model_config = {"extra": "allow"}
|
||||
|
||||
|
||||
class OCRRequestData(BaseModel):
|
||||
"""OCR request data structure."""
|
||||
data: Optional[Union[Dict, bytes]] = None
|
||||
files: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class BaseOCRConfig:
|
||||
"""
|
||||
Base configuration for OCR transformations.
|
||||
Handles provider-agnostic OCR operations.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def get_supported_ocr_params(self, model: str) -> list:
|
||||
"""
|
||||
Get supported OCR parameters for this provider.
|
||||
Override this method in provider-specific implementations.
|
||||
"""
|
||||
return []
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
) -> dict:
|
||||
"""Map OCR parameters to provider-specific parameters."""
|
||||
return optional_params
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Dict,
|
||||
model: str,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Dict:
|
||||
"""
|
||||
Validate environment and return headers.
|
||||
Override in provider-specific implementations.
|
||||
"""
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
Get complete URL for OCR endpoint.
|
||||
Override in provider-specific implementations.
|
||||
"""
|
||||
raise NotImplementedError("get_complete_url must be implemented by provider")
|
||||
|
||||
def transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
"""
|
||||
Transform OCR request to provider-specific format.
|
||||
Override in provider-specific implementations.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
document: Document to process (Mistral format dict, or file path, bytes, etc.)
|
||||
optional_params: Optional parameters for the request
|
||||
headers: Request headers
|
||||
|
||||
Returns:
|
||||
OCRRequestData with data and files fields
|
||||
"""
|
||||
raise NotImplementedError("transform_ocr_request must be implemented by provider")
|
||||
|
||||
async def async_transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
"""
|
||||
Async transform OCR request to provider-specific format.
|
||||
Optional method - providers can override if they need async transformations
|
||||
(e.g., Azure AI for URL-to-base64 conversion).
|
||||
|
||||
Default implementation falls back to sync transform_ocr_request.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
document: Document to process (Mistral format dict, or file path, bytes, etc.)
|
||||
optional_params: Optional parameters for the request
|
||||
headers: Request headers
|
||||
|
||||
Returns:
|
||||
OCRRequestData with data and files fields
|
||||
"""
|
||||
# Default implementation: call sync version
|
||||
return self.transform_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Transform provider-specific OCR response to standard format.
|
||||
Override in provider-specific implementations.
|
||||
"""
|
||||
raise NotImplementedError("transform_ocr_response must be implemented by provider")
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: dict,
|
||||
) -> Exception:
|
||||
"""Get appropriate error class for the provider."""
|
||||
return BaseLLMException(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
3
litellm/llms/cometapi/embed/__init__.py
Normal file
3
litellm/llms/cometapi/embed/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .transformation import CometAPIEmbeddingConfig
|
||||
|
||||
__all__ = ["CometAPIEmbeddingConfig"]
|
||||
157
litellm/llms/cometapi/embed/transformation.py
Normal file
157
litellm/llms/cometapi/embed/transformation.py
Normal file
|
|
@ -0,0 +1,157 @@
|
|||
"""
|
||||
CometAPI Embedding API support - OpenAI compatible
|
||||
"""
|
||||
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
|
||||
from litellm.types.utils import EmbeddingResponse, Usage
|
||||
|
||||
from ..common_utils import CometAPIException
|
||||
|
||||
|
||||
class CometAPIEmbeddingConfig(BaseEmbeddingConfig):
|
||||
"""
|
||||
Configuration class for CometAPI Embedding API.
|
||||
|
||||
Since CometAPI is OpenAI-compatible, this class provides OpenAI-standard
|
||||
embedding functionality with CometAPI-specific authentication and endpoints.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Get the complete URL for the CometAPI embedding endpoint.
|
||||
"""
|
||||
api_base = (
|
||||
"https://api.cometapi.com/v1" if api_base is None else api_base.rstrip("/")
|
||||
)
|
||||
complete_url = f"{api_base}/embeddings"
|
||||
return complete_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:
|
||||
"""
|
||||
Validate and set up authentication headers for CometAPI.
|
||||
"""
|
||||
if api_key is None:
|
||||
api_key = get_secret_str("COMETAPI_KEY")
|
||||
|
||||
default_headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"accept": "application/json",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
if "Authorization" in headers:
|
||||
default_headers["Authorization"] = headers["Authorization"]
|
||||
|
||||
return {**default_headers, **headers}
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
"""
|
||||
Get the supported OpenAI parameters for embedding requests.
|
||||
CometAPI supports standard OpenAI embedding parameters.
|
||||
"""
|
||||
return [
|
||||
"dimensions",
|
||||
"encoding_format",
|
||||
"user",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OpenAI parameters to CometAPI format.
|
||||
"""
|
||||
supported_openai_params = self.get_supported_openai_params(model)
|
||||
for param, value in non_default_params.items():
|
||||
if param in supported_openai_params:
|
||||
optional_params[param] = value
|
||||
return optional_params
|
||||
|
||||
def transform_embedding_request(
|
||||
self,
|
||||
model: str,
|
||||
input: AllEmbeddingInputValues,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform the embedding request into CometAPI format.
|
||||
"""
|
||||
return {"input": input, "model": model, **optional_params}
|
||||
|
||||
def transform_embedding_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: EmbeddingResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: Optional[str],
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> EmbeddingResponse:
|
||||
"""
|
||||
Transform CometAPI response into standard EmbeddingResponse format.
|
||||
"""
|
||||
try:
|
||||
raw_response_json = raw_response.json()
|
||||
except Exception:
|
||||
raise CometAPIException(
|
||||
message=raw_response.text,
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
model_response.model = raw_response_json.get("model")
|
||||
model_response.data = raw_response_json.get("data")
|
||||
model_response.object = raw_response_json.get("object")
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=raw_response_json.get("usage", {}).get("prompt_tokens", 0),
|
||||
total_tokens=raw_response_json.get("usage", {}).get("total_tokens", 0),
|
||||
)
|
||||
|
||||
model_response.usage = usage
|
||||
return model_response
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
"""
|
||||
Get the appropriate error class for CometAPI exceptions.
|
||||
"""
|
||||
return CometAPIException(
|
||||
message=error_message, status_code=status_code, headers=headers
|
||||
)
|
||||
13
litellm/llms/cometapi/image_generation/__init__.py
Normal file
13
litellm/llms/cometapi/image_generation/__init__.py
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
|
||||
from .transformation import CometAPIImageGenerationConfig
|
||||
|
||||
__all__ = [
|
||||
"CometAPIImageGenerationConfig",
|
||||
]
|
||||
|
||||
|
||||
def get_cometapi_image_generation_config(model: str) -> BaseImageGenerationConfig:
|
||||
return CometAPIImageGenerationConfig()
|
||||
25
litellm/llms/cometapi/image_generation/cost_calculator.py
Normal file
25
litellm/llms/cometapi/image_generation/cost_calculator.py
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
from typing import Any
|
||||
|
||||
import litellm
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
|
||||
def cost_calculator(
|
||||
model: str,
|
||||
image_response: Any,
|
||||
) -> float:
|
||||
"""
|
||||
CometAPI image generation cost calculator
|
||||
"""
|
||||
_model_info = litellm.get_model_info(
|
||||
model=model,
|
||||
custom_llm_provider=litellm.LlmProviders.COMETAPI.value,
|
||||
)
|
||||
output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0
|
||||
num_images: int = 0
|
||||
if isinstance(image_response, ImageResponse):
|
||||
if image_response.data:
|
||||
num_images = len(image_response.data)
|
||||
return output_cost_per_image * num_images
|
||||
else:
|
||||
raise ValueError(f"image_response must be of type ImageResponse got type={type(image_response)}")
|
||||
170
litellm/llms/cometapi/image_generation/transformation.py
Normal file
170
litellm/llms/cometapi/image_generation/transformation.py
Normal file
|
|
@ -0,0 +1,170 @@
|
|||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
OpenAIImageGenerationOptionalParams,
|
||||
)
|
||||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class CometAPIImageGenerationConfig(BaseImageGenerationConfig):
|
||||
DEFAULT_BASE_URL: str = "https://api.cometapi.com"
|
||||
IMAGE_GENERATION_ENDPOINT: str = "v1/images/generations"
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIImageGenerationOptionalParams]:
|
||||
"""
|
||||
https://api.cometapi.com/v1/images/generations
|
||||
"""
|
||||
return [
|
||||
"n",
|
||||
"quality",
|
||||
"response_format",
|
||||
"size",
|
||||
"style",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
|
||||
for k in non_default_params.keys():
|
||||
if k not in optional_params.keys():
|
||||
if k in supported_params:
|
||||
# CometAPI uses OpenAI-compatible parameters, so we can pass them directly
|
||||
optional_params[k] = non_default_params[k]
|
||||
elif drop_params:
|
||||
pass
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Parameter {k} is not supported for model {model}. Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters."
|
||||
)
|
||||
|
||||
return optional_params
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Get the complete url for the request
|
||||
"""
|
||||
complete_url: str = (
|
||||
api_base
|
||||
or get_secret_str("COMETAPI_BASE_URL")
|
||||
or get_secret_str("COMETAPI_API_BASE")
|
||||
or self.DEFAULT_BASE_URL
|
||||
)
|
||||
|
||||
complete_url = complete_url.rstrip("/")
|
||||
complete_url = f"{complete_url}/{self.IMAGE_GENERATION_ENDPOINT}"
|
||||
return complete_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:
|
||||
final_api_key: Optional[str] = (
|
||||
api_key or
|
||||
get_secret_str("COMETAPI_KEY") or
|
||||
get_secret_str("COMETAPI_API_KEY")
|
||||
)
|
||||
if not final_api_key:
|
||||
raise ValueError("COMETAPI_KEY or COMETAPI_API_KEY is not set")
|
||||
|
||||
headers["Authorization"] = f"Bearer {final_api_key}"
|
||||
headers["Content-Type"] = "application/json"
|
||||
return headers
|
||||
|
||||
def transform_image_generation_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform the image generation request to the CometAPI image generation request body
|
||||
|
||||
https://api.cometapi.com/v1/images/generations
|
||||
"""
|
||||
# CometAPI uses OpenAI-compatible format
|
||||
request_body = {
|
||||
"prompt": prompt,
|
||||
"model": model,
|
||||
**optional_params,
|
||||
}
|
||||
return request_body
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ImageResponse:
|
||||
"""
|
||||
Transform the image generation response to the litellm image response
|
||||
|
||||
https://api.cometapi.com/v1/images/generations
|
||||
"""
|
||||
try:
|
||||
response_data = raw_response.json()
|
||||
except Exception as e:
|
||||
raise self.get_error_class(
|
||||
error_message=f"Error transforming image generation response: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
if not model_response.data:
|
||||
model_response.data = []
|
||||
|
||||
# CometAPI returns OpenAI-compatible format
|
||||
# Expected format: {"created": timestamp, "data": [{"url": "...", "b64_json": "..."}]}
|
||||
if "data" in response_data:
|
||||
for image_data in response_data["data"]:
|
||||
image_obj = ImageObject(
|
||||
b64_json=image_data.get("b64_json"),
|
||||
url=image_data.get("url"),
|
||||
)
|
||||
model_response.data.append(image_obj)
|
||||
|
||||
return model_response
|
||||
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
import os
|
||||
import ssl
|
||||
import sys
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Mapping, Optional, Union
|
||||
|
||||
|
|
@ -114,6 +115,28 @@ def get_ssl_configuration(
|
|||
# but falls back to widely compatible ones
|
||||
custom_ssl_context.set_ciphers(DEFAULT_SSL_CIPHERS)
|
||||
|
||||
# Configure ECDH curve for key exchange (e.g., to disable PQC and improve performance)
|
||||
# Set SSL_ECDH_CURVE env var or litellm.ssl_ecdh_curve to 'X25519' to disable PQC
|
||||
# Common valid curves: X25519, prime256v1, secp384r1, secp521r1
|
||||
ssl_ecdh_curve = os.getenv("SSL_ECDH_CURVE", litellm.ssl_ecdh_curve)
|
||||
if ssl_ecdh_curve and isinstance(ssl_ecdh_curve, str):
|
||||
try:
|
||||
custom_ssl_context.set_ecdh_curve(ssl_ecdh_curve)
|
||||
verbose_logger.debug(f"SSL ECDH curve set to: {ssl_ecdh_curve}")
|
||||
except AttributeError:
|
||||
verbose_logger.warning(
|
||||
f"SSL ECDH curve configuration not supported. "
|
||||
f"Python version: {sys.version.split()[0]}, OpenSSL version: {ssl.OPENSSL_VERSION}. "
|
||||
f"Requested curve: {ssl_ecdh_curve}. Continuing with default curves."
|
||||
)
|
||||
except ValueError as e:
|
||||
# Invalid curve name
|
||||
verbose_logger.warning(
|
||||
f"Invalid SSL ECDH curve name: '{ssl_ecdh_curve}'. {e}. "
|
||||
f"Common valid curves: X25519, prime256v1, secp384r1, secp521r1. "
|
||||
f"Continuing with default curves (including PQC)."
|
||||
)
|
||||
|
||||
# Use our custom SSL context instead of the original ssl_verify value
|
||||
return custom_ssl_context
|
||||
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
|
|||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
|
||||
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
|
||||
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
|
|
@ -1256,6 +1257,289 @@ class BaseLLMHTTPHandler:
|
|||
api_key=api_key,
|
||||
)
|
||||
|
||||
def _prepare_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: Dict[str, str],
|
||||
optional_params: dict,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
headers: Optional[Dict[str, Any]],
|
||||
provider_config: BaseOCRConfig,
|
||||
litellm_params: dict,
|
||||
) -> Tuple[Dict[str, Any], str, Dict[str, Any], None]:
|
||||
"""
|
||||
Shared logic for preparing OCR requests.
|
||||
Returns: (headers, complete_url, data, files)
|
||||
"""
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRRequestData
|
||||
|
||||
headers = provider_config.validate_environment(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=headers or {},
|
||||
model=model,
|
||||
)
|
||||
|
||||
complete_url = provider_config.get_complete_url(
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
# Transform the request to get data and files
|
||||
transformed_result = provider_config.transform_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
if not isinstance(transformed_result, OCRRequestData):
|
||||
raise ValueError(
|
||||
f"Provider {provider_config.__class__.__name__} must return OCRRequestData"
|
||||
)
|
||||
|
||||
# Data is always a dict for Mistral OCR format
|
||||
if not isinstance(transformed_result.data, dict):
|
||||
raise ValueError(f"Expected dict data for OCR request, got {type(transformed_result.data)}")
|
||||
|
||||
data = transformed_result.data
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": complete_url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
return headers, complete_url, data, None
|
||||
|
||||
async def _async_prepare_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: Dict[str, str],
|
||||
optional_params: dict,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
headers: Optional[Dict[str, Any]],
|
||||
provider_config: BaseOCRConfig,
|
||||
litellm_params: dict,
|
||||
) -> Tuple[Dict[str, Any], str, Dict[str, Any], None]:
|
||||
"""
|
||||
Async version of _prepare_ocr_request for providers that need async transforms.
|
||||
Returns: (headers, complete_url, data, files)
|
||||
"""
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRRequestData
|
||||
|
||||
headers = provider_config.validate_environment(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=headers or {},
|
||||
model=model,
|
||||
)
|
||||
|
||||
complete_url = provider_config.get_complete_url(
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
# Use async transform (providers can override this method if they need async operations)
|
||||
transformed_result = await provider_config.async_transform_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
if not isinstance(transformed_result, OCRRequestData):
|
||||
raise ValueError(
|
||||
f"Provider {provider_config.__class__.__name__} must return OCRRequestData"
|
||||
)
|
||||
|
||||
# Data is always a dict for Mistral OCR format
|
||||
if not isinstance(transformed_result.data, dict):
|
||||
raise ValueError(f"Expected dict data for OCR request, got {type(transformed_result.data)}")
|
||||
|
||||
data = transformed_result.data
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": complete_url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
return headers, complete_url, data, None
|
||||
|
||||
def _transform_ocr_response(
|
||||
self,
|
||||
provider_config: BaseOCRConfig,
|
||||
model: str,
|
||||
response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> OCRResponse:
|
||||
"""Shared logic for transforming OCR responses."""
|
||||
return provider_config.transform_ocr_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def ocr(
|
||||
self,
|
||||
model: str,
|
||||
document: Dict[str, str],
|
||||
optional_params: dict,
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
custom_llm_provider: str,
|
||||
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
|
||||
aocr: bool = False,
|
||||
headers: Optional[Dict[str, Any]] = None,
|
||||
provider_config: Optional[BaseOCRConfig] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> Union[OCRResponse, Coroutine[Any, Any, OCRResponse]]:
|
||||
"""
|
||||
Sync OCR handler.
|
||||
"""
|
||||
if provider_config is None:
|
||||
raise ValueError(
|
||||
f"No provider config found for model: {model} and provider: {custom_llm_provider}"
|
||||
)
|
||||
|
||||
if litellm_params is None:
|
||||
litellm_params = {}
|
||||
|
||||
if aocr is True:
|
||||
return self.async_ocr(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
client=client,
|
||||
headers=headers,
|
||||
provider_config=provider_config,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
# Prepare the request
|
||||
headers, complete_url, data, files = self._prepare_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
provider_config=provider_config,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
client = _get_httpx_client()
|
||||
|
||||
try:
|
||||
# Make the POST request with JSON data (Mistral format)
|
||||
response = client.post(
|
||||
url=complete_url,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
return self._transform_ocr_response(
|
||||
provider_config=provider_config,
|
||||
model=model,
|
||||
response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_ocr(
|
||||
self,
|
||||
model: str,
|
||||
document: Dict[str, str],
|
||||
optional_params: dict,
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
custom_llm_provider: str,
|
||||
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
|
||||
headers: Optional[Dict[str, Any]] = None,
|
||||
provider_config: Optional[BaseOCRConfig] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Async OCR handler.
|
||||
"""
|
||||
if provider_config is None:
|
||||
raise ValueError(
|
||||
f"No provider config found for model: {model} and provider: {custom_llm_provider}"
|
||||
)
|
||||
|
||||
if litellm_params is None:
|
||||
litellm_params = {}
|
||||
|
||||
# Prepare the request using async prepare method
|
||||
headers, complete_url, data, files = await self._async_prepare_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
provider_config=provider_config,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
else:
|
||||
async_httpx_client = client
|
||||
|
||||
try:
|
||||
# Make the async POST request with JSON data (Mistral format)
|
||||
response = await async_httpx_client.post(
|
||||
url=complete_url,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
return self._transform_ocr_response(
|
||||
provider_config=provider_config,
|
||||
model=model,
|
||||
response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_anthropic_messages_handler(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -2995,6 +3279,7 @@ class BaseLLMHTTPHandler:
|
|||
BaseGoogleGenAIGenerateContentConfig,
|
||||
BaseAnthropicMessagesConfig,
|
||||
BaseBatchesConfig,
|
||||
BaseOCRConfig,
|
||||
"BasePassthroughConfig",
|
||||
],
|
||||
):
|
||||
|
|
|
|||
2
litellm/llms/mistral/ocr/__init__.py
Normal file
2
litellm/llms/mistral/ocr/__init__.py
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
"""Mistral OCR transformation module."""
|
||||
|
||||
223
litellm/llms/mistral/ocr/transformation.py
Normal file
223
litellm/llms/mistral/ocr/transformation.py
Normal file
|
|
@ -0,0 +1,223 @@
|
|||
"""
|
||||
Mistral OCR transformation implementation.
|
||||
"""
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.base_llm.ocr.transformation import (
|
||||
BaseOCRConfig,
|
||||
DocumentType,
|
||||
OCRRequestData,
|
||||
OCRResponse,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
class MistralOCRConfig(BaseOCRConfig):
|
||||
"""
|
||||
Mistral OCR transformation configuration.
|
||||
|
||||
Reference: https://docs.mistral.ai/api/#tag/ocr
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def get_supported_ocr_params(self, model: str) -> list:
|
||||
"""
|
||||
Get supported OCR parameters for Mistral OCR.
|
||||
|
||||
Mistral OCR supports:
|
||||
- pages: List of page numbers to process
|
||||
- include_image_base64: Whether to include base64 encoded images
|
||||
- image_limit: Maximum number of images to return
|
||||
- image_min_size: Minimum size of images to include
|
||||
- bbox_annotation_format: Format for bounding box annotations
|
||||
- document_annotation_format: Format for document annotations
|
||||
"""
|
||||
return [
|
||||
"pages",
|
||||
"include_image_base64",
|
||||
"image_limit",
|
||||
"image_min_size",
|
||||
"bbox_annotation_format",
|
||||
"document_annotation_format",
|
||||
]
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OCR parameters to Mistral-specific format.
|
||||
|
||||
Mistral accepts these parameters directly, so no transformation needed.
|
||||
Just filter out unsupported params.
|
||||
"""
|
||||
supported_params = self.get_supported_ocr_params(model=model)
|
||||
|
||||
# Only include params that are in the supported list
|
||||
mapped_params = {}
|
||||
for param, value in non_default_params.items():
|
||||
if param in supported_params:
|
||||
mapped_params[param] = value
|
||||
|
||||
return mapped_params
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Dict,
|
||||
model: str,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Dict:
|
||||
"""
|
||||
Validate environment and return headers for Mistral OCR.
|
||||
"""
|
||||
# Get API key from environment if not provided
|
||||
if api_key is None:
|
||||
api_key = (
|
||||
get_secret_str("MISTRAL_API_KEY")
|
||||
)
|
||||
|
||||
if api_key is None:
|
||||
raise ValueError(
|
||||
"Missing Mistral API Key - A call is being made to Mistral but no key is set either in the environment variables or via params"
|
||||
)
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
**headers,
|
||||
}
|
||||
|
||||
# Don't set Content-Type for multipart/form-data - httpx will handle it
|
||||
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
Get complete URL for Mistral OCR endpoint.
|
||||
|
||||
Returns: https://api.mistral.ai/v1/ocr
|
||||
"""
|
||||
if api_base is None:
|
||||
api_base = "https://api.mistral.ai/v1"
|
||||
|
||||
# Ensure no trailing slash
|
||||
api_base = api_base.rstrip("/")
|
||||
|
||||
# Remove /v1 if it's already in the base to avoid duplication
|
||||
if api_base.endswith("/v1"):
|
||||
return f"{api_base}/ocr"
|
||||
|
||||
return f"{api_base}/v1/ocr"
|
||||
|
||||
|
||||
def transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
"""
|
||||
Transform OCR request to Mistral-specific format.
|
||||
|
||||
Mistral OCR API accepts:
|
||||
{
|
||||
"model": "mistral-ocr-latest",
|
||||
"document": {
|
||||
"type": "document_url",
|
||||
"document_url": "<https-url or data-uri>"
|
||||
},
|
||||
"pages": [0], # optional
|
||||
"include_image_base64": false, # optional
|
||||
...
|
||||
}
|
||||
|
||||
Args:
|
||||
model: Model name (e.g., "mistral-ocr-latest")
|
||||
document: Document dict from user (Mistral format) - already validated in main.py
|
||||
optional_params: Already mapped optional parameters
|
||||
headers: Request headers
|
||||
|
||||
Returns:
|
||||
OCRRequestData with JSON data
|
||||
"""
|
||||
verbose_logger.debug(f"Mistral OCR transform_ocr_request - model: {model}")
|
||||
|
||||
# Document parameter is the Mistral-format dict from the user
|
||||
# Just pass it through as-is to the Mistral API
|
||||
if not isinstance(document, dict):
|
||||
raise ValueError(f"Expected document dict, got {type(document)}")
|
||||
|
||||
# Build request data - use document dict directly
|
||||
data = {
|
||||
"model": model,
|
||||
"document": document, # Pass through the Mistral-format document dict
|
||||
}
|
||||
|
||||
# Add all optional parameters from the already-mapped optional_params
|
||||
data.update(optional_params)
|
||||
|
||||
# No multipart files - using JSON
|
||||
return OCRRequestData(data=data, files=None)
|
||||
|
||||
def transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: Any,
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Return Mistral OCR response in native format.
|
||||
|
||||
Mistral OCR is the standard format for LiteLLM OCR responses.
|
||||
No transformation needed - return native response.
|
||||
|
||||
Mistral OCR returns:
|
||||
{
|
||||
"pages": [
|
||||
{
|
||||
"index": 0,
|
||||
"markdown": "extracted text content",
|
||||
"images": [...],
|
||||
"dimensions": {...}
|
||||
},
|
||||
...
|
||||
],
|
||||
"model": "mistral-ocr-2505-completion",
|
||||
"document_annotation": null,
|
||||
"usage_info": {...}
|
||||
}
|
||||
"""
|
||||
try:
|
||||
response_json = raw_response.json()
|
||||
|
||||
verbose_logger.debug(f"Mistral OCR response keys: {response_json.keys()}")
|
||||
|
||||
# Return native Mistral format - no transformation
|
||||
return OCRResponse(
|
||||
pages=response_json.get("pages", []),
|
||||
model=response_json.get("model", model),
|
||||
document_annotation=response_json.get("document_annotation"),
|
||||
usage_info=response_json.get("usage_info"),
|
||||
object="ocr",
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Error parsing Mistral OCR response: {e}")
|
||||
raise e
|
||||
|
||||
26
litellm/llms/openai/image_edit/__init__.py
Normal file
26
litellm/llms/openai/image_edit/__init__.py
Normal file
|
|
@ -0,0 +1,26 @@
|
|||
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
|
||||
|
||||
from .dalle2_transformation import DallE2ImageEditConfig
|
||||
from .transformation import OpenAIImageEditConfig
|
||||
|
||||
__all__ = ["OpenAIImageEditConfig", "DallE2ImageEditConfig", "get_openai_image_edit_config"]
|
||||
|
||||
|
||||
def get_openai_image_edit_config(model: str) -> BaseImageEditConfig:
|
||||
"""
|
||||
Get the appropriate OpenAI image edit config based on the model.
|
||||
|
||||
Args:
|
||||
model: The model name (e.g., "dall-e-2", "gpt-image-1")
|
||||
|
||||
Returns:
|
||||
The appropriate config instance for the model
|
||||
"""
|
||||
model_normalized = model.lower().replace("-", "").replace("_", "")
|
||||
|
||||
if model_normalized == "dalle2":
|
||||
return DallE2ImageEditConfig()
|
||||
else:
|
||||
# Default to standard OpenAI config for gpt-image-1 and other models
|
||||
return OpenAIImageEditConfig()
|
||||
|
||||
101
litellm/llms/openai/image_edit/dalle2_transformation.py
Normal file
101
litellm/llms/openai/image_edit/dalle2_transformation.py
Normal file
|
|
@ -0,0 +1,101 @@
|
|||
from io import BufferedReader
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Tuple, cast
|
||||
|
||||
from httpx._types import RequestFiles
|
||||
|
||||
import litellm
|
||||
from litellm.images.utils import ImageEditRequestUtils
|
||||
from litellm.types.images.main import ImageEditRequestParams
|
||||
from litellm.types.llms.openai import FileTypes
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
from .transformation import OpenAIImageEditConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class DallE2ImageEditConfig(OpenAIImageEditConfig):
|
||||
"""
|
||||
DALL-E-2 specific configuration for image edit API.
|
||||
|
||||
DALL-E-2 only supports editing a single image (not an array).
|
||||
Uses "image" field name instead of "image[]".
|
||||
"""
|
||||
|
||||
def transform_image_edit_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
image: FileTypes,
|
||||
image_edit_optional_request_params: Dict,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[Dict, RequestFiles]:
|
||||
"""
|
||||
Transform image edit request for DALL-E-2.
|
||||
|
||||
DALL-E-2 only accepts a single image with field name "image" (not "image[]").
|
||||
"""
|
||||
request = ImageEditRequestParams(
|
||||
model=model,
|
||||
image=image,
|
||||
prompt=prompt,
|
||||
**image_edit_optional_request_params,
|
||||
)
|
||||
request_dict = cast(Dict, request)
|
||||
|
||||
#########################################################
|
||||
# Separate images and masks as `files` and send other parameters as `data`
|
||||
#########################################################
|
||||
_image_list = request_dict.get("image")
|
||||
_mask = request_dict.get("mask")
|
||||
data_without_files = {
|
||||
k: v for k, v in request_dict.items() if k not in ["image", "mask"]
|
||||
}
|
||||
files_list: List[Tuple[str, Any]] = []
|
||||
|
||||
# Handle image parameter - DALL-E-2 only supports single image
|
||||
if _image_list is not None:
|
||||
image_list = (
|
||||
[_image_list] if not isinstance(_image_list, list) else _image_list
|
||||
)
|
||||
|
||||
# Validate only one image is provided
|
||||
if len(image_list) > 1:
|
||||
raise litellm.BadRequestError(
|
||||
message="DALL-E-2 only supports editing a single image. Please provide one image.",
|
||||
model=model,
|
||||
llm_provider="openai",
|
||||
)
|
||||
|
||||
# Use "image" field name (singular) for DALL-E-2
|
||||
for _image in image_list:
|
||||
if _image is not None:
|
||||
self._add_image_to_files(
|
||||
files_list=files_list,
|
||||
image=_image,
|
||||
field_name="image",
|
||||
)
|
||||
|
||||
# Handle mask parameter if provided
|
||||
if _mask is not None:
|
||||
# Handle case where mask can be a list (extract first mask)
|
||||
if isinstance(_mask, list):
|
||||
_mask = _mask[0] if _mask else None
|
||||
|
||||
if _mask is not None:
|
||||
mask_content_type: str = ImageEditRequestUtils.get_image_content_type(
|
||||
_mask
|
||||
)
|
||||
if isinstance(_mask, BufferedReader):
|
||||
files_list.append(("mask", (_mask.name, _mask, mask_content_type)))
|
||||
else:
|
||||
files_list.append(("mask", ("mask.png", _mask, mask_content_type)))
|
||||
|
||||
return data_without_files, files_list
|
||||
|
||||
|
|
@ -27,6 +27,11 @@ else:
|
|||
|
||||
|
||||
class OpenAIImageEditConfig(BaseImageEditConfig):
|
||||
"""
|
||||
Base configuration for OpenAI image edit API.
|
||||
Used for models like gpt-image-1 that support multiple images.
|
||||
"""
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
"""
|
||||
All OpenAI Image Edits params are supported
|
||||
|
|
@ -57,6 +62,20 @@ class OpenAIImageEditConfig(BaseImageEditConfig):
|
|||
"""No mapping applied since inputs are in OpenAI spec already"""
|
||||
return dict(image_edit_optional_params)
|
||||
|
||||
def _add_image_to_files(
|
||||
self,
|
||||
files_list: List[Tuple[str, Any]],
|
||||
image: Any,
|
||||
field_name: str,
|
||||
) -> None:
|
||||
"""Add an image to the files list with appropriate content type"""
|
||||
image_content_type = ImageEditRequestUtils.get_image_content_type(image)
|
||||
|
||||
if isinstance(image, BufferedReader):
|
||||
files_list.append((field_name, (image.name, image, image_content_type)))
|
||||
else:
|
||||
files_list.append((field_name, ("image.png", image, image_content_type)))
|
||||
|
||||
def transform_image_edit_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -67,9 +86,10 @@ class OpenAIImageEditConfig(BaseImageEditConfig):
|
|||
headers: dict,
|
||||
) -> Tuple[Dict, RequestFiles]:
|
||||
"""
|
||||
No transform applied since inputs are in OpenAI spec already
|
||||
Transform image edit request to OpenAI API format.
|
||||
|
||||
This handles buffered readers as images to be sent as multipart/form-data for OpenAI
|
||||
Handles multipart/form-data for images. Uses "image[]" field name
|
||||
to support multiple images (e.g., for gpt-image-1).
|
||||
"""
|
||||
request = ImageEditRequestParams(
|
||||
model=model,
|
||||
|
|
@ -94,19 +114,14 @@ class OpenAIImageEditConfig(BaseImageEditConfig):
|
|||
image_list = (
|
||||
[_image_list] if not isinstance(_image_list, list) else _image_list
|
||||
)
|
||||
|
||||
for _image in image_list:
|
||||
if _image is not None:
|
||||
image_content_type: str = (
|
||||
ImageEditRequestUtils.get_image_content_type(_image)
|
||||
self._add_image_to_files(
|
||||
files_list=files_list,
|
||||
image=_image,
|
||||
field_name="image[]",
|
||||
)
|
||||
if isinstance(_image, BufferedReader):
|
||||
files_list.append(
|
||||
("image[]", (_image.name, _image, image_content_type))
|
||||
)
|
||||
else:
|
||||
files_list.append(
|
||||
("image[]", ("image.png", _image, image_content_type))
|
||||
)
|
||||
# Handle mask parameter if provided
|
||||
if _mask is not None:
|
||||
# Handle case where mask can be a list (extract first mask)
|
||||
|
|
|
|||
|
|
@ -4754,6 +4754,33 @@ def embedding( # noqa: PLR0915
|
|||
aembedding=aembedding,
|
||||
litellm_params={},
|
||||
)
|
||||
elif custom_llm_provider == "cometapi":
|
||||
api_key = (
|
||||
api_key
|
||||
or litellm.cometapi_key
|
||||
or get_secret_str("COMETAPI_KEY")
|
||||
or litellm.api_key
|
||||
)
|
||||
api_base = (
|
||||
api_base
|
||||
or litellm.api_base
|
||||
or get_secret_str("COMETAPI_API_BASE")
|
||||
or "https://api.cometapi.com/v1"
|
||||
)
|
||||
response = base_llm_http_handler.embedding(
|
||||
model=model,
|
||||
input=input,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
logging_obj=logging,
|
||||
timeout=timeout,
|
||||
model_response=EmbeddingResponse(),
|
||||
optional_params=optional_params,
|
||||
client=client,
|
||||
aembedding=aembedding,
|
||||
litellm_params={},
|
||||
)
|
||||
elif custom_llm_provider in litellm._custom_providers:
|
||||
custom_handler: Optional[CustomLLM] = None
|
||||
for item in litellm.custom_provider_map:
|
||||
|
|
|
|||
|
|
@ -400,6 +400,44 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"anthropic.claude-haiku-4-5@20251001": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"anthropic.claude-3-5-sonnet-20240620-v1:0": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -810,6 +848,25 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"apac.anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"apac.anthropic.claude-3-sonnet-20240229-v1:0": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -4612,6 +4669,48 @@
|
|||
"supports_web_search": true,
|
||||
"tool_use_system_prompt_tokens": 264
|
||||
},
|
||||
"claude-haiku-4-5-20251001": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 2e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "anthropic",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 64000,
|
||||
"max_tokens": 64000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"claude-haiku-4-5": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 2e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "anthropic",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 64000,
|
||||
"max_tokens": 64000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"claude-3-5-sonnet-20240620": {
|
||||
"cache_creation_input_token_cost": 3.75e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 6e-06,
|
||||
|
|
@ -7741,6 +7840,25 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"eu.anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"eu.anthropic.claude-3-5-sonnet-20240620-v1:0": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -9486,6 +9604,54 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"gemini-2.5-flash-image": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.039,
|
||||
"output_cost_per_reasoning_token": 2.5e-06,
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"rpm": 100000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-flash-image",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supports_audio_output": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"tpm": 8000000
|
||||
},
|
||||
"gemini-2.5-flash-image-preview": {
|
||||
"cache_read_input_token_cost": 7.5e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
|
|
@ -10939,6 +11105,54 @@
|
|||
"supports_web_search": true,
|
||||
"tpm": 8000000
|
||||
},
|
||||
"gemini/gemini-2.5-flash-image": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.039,
|
||||
"output_cost_per_reasoning_token": 2.5e-06,
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"rpm": 100000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-flash-image",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supports_audio_output": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"tpm": 8000000
|
||||
},
|
||||
"gemini/gemini-2.5-flash-image-preview": {
|
||||
"cache_read_input_token_cost": 7.5e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
|
|
@ -13650,8 +13864,56 @@
|
|||
"lemonade/Qwen3-Coder-30B-A3B-Instruct-GGUF": {
|
||||
"input_cost_per_token": 0,
|
||||
"litellm_provider": "lemonade",
|
||||
"max_tokens": 32768,
|
||||
"max_input_tokens": 32768,
|
||||
"max_tokens": 262144,
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"lemonade/gpt-oss-20b-mxfp4-GGUF": {
|
||||
"input_cost_per_token": 0,
|
||||
"litellm_provider": "lemonade",
|
||||
"max_tokens": 131072,
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"lemonade/gpt-oss-120b-mxfp-GGUF": {
|
||||
"input_cost_per_token": 0,
|
||||
"litellm_provider": "lemonade",
|
||||
"max_tokens": 131072,
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"lemonade/Gemma-3-4b-it-GGUF": {
|
||||
"input_cost_per_token": 0,
|
||||
"litellm_provider": "lemonade",
|
||||
"max_tokens": 128000,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"lemonade/Qwen3-4B-Instruct-2507-GGUF": {
|
||||
"input_cost_per_token": 0,
|
||||
"litellm_provider": "lemonade",
|
||||
"max_tokens": 262144,
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0,
|
||||
|
|
@ -14466,6 +14728,25 @@
|
|||
"supports_vision": true,
|
||||
"tool_use_system_prompt_tokens": 346
|
||||
},
|
||||
"jp.anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"lambda_ai/deepseek-llama3.3-70b": {
|
||||
"input_cost_per_token": 2e-07,
|
||||
"litellm_provider": "lambda_ai",
|
||||
|
|
@ -17153,6 +17434,8 @@
|
|||
},
|
||||
"openrouter/anthropic/claude-opus-4": {
|
||||
"input_cost_per_image": 0.0048,
|
||||
"cache_creation_input_token_cost": 1.875e-05,
|
||||
"cache_read_input_token_cost": 1.5e-06,
|
||||
"input_cost_per_token": 1.5e-05,
|
||||
"litellm_provider": "openrouter",
|
||||
"max_input_tokens": 200000,
|
||||
|
|
@ -17163,6 +17446,7 @@
|
|||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -17170,6 +17454,9 @@
|
|||
},
|
||||
"openrouter/anthropic/claude-opus-4.1": {
|
||||
"input_cost_per_image": 0.0048,
|
||||
"cache_creation_input_token_cost": 1.875e-05,
|
||||
"cache_creation_input_token_cost_above_1hr": 3e-05,
|
||||
"cache_read_input_token_cost": 1.5e-06,
|
||||
"input_cost_per_token": 1.5e-05,
|
||||
"litellm_provider": "openrouter",
|
||||
"max_input_tokens": 200000,
|
||||
|
|
@ -17180,6 +17467,7 @@
|
|||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -17187,6 +17475,10 @@
|
|||
},
|
||||
"openrouter/anthropic/claude-sonnet-4": {
|
||||
"input_cost_per_image": 0.0048,
|
||||
"cache_creation_input_token_cost": 3.75e-06,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
|
||||
"cache_read_input_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
|
||||
"input_cost_per_token": 3e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 6e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 2.25e-05,
|
||||
|
|
@ -17199,6 +17491,31 @@
|
|||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"tool_use_system_prompt_tokens": 159
|
||||
},
|
||||
"openrouter/anthropic/claude-sonnet-4.5": {
|
||||
"input_cost_per_image": 0.0048,
|
||||
"cache_creation_input_token_cost": 3.75e-06,
|
||||
"cache_read_input_token_cost": 3e-07,
|
||||
"input_cost_per_token": 3e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 6e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 2.25e-05,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
|
||||
"litellm_provider": "openrouter",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -20097,6 +20414,25 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"us.anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"us.anthropic.claude-3-5-sonnet-20240620-v1:0": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -21368,6 +21704,25 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"vertex_ai/claude-haiku-4-5@20251001": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "vertex_ai-anthropic_models",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/haiku-4-5",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"vertex_ai/claude-3-5-sonnet": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"litellm_provider": "vertex_ai-anthropic_models",
|
||||
|
|
@ -21560,8 +21915,8 @@
|
|||
"input_cost_per_token_batches": 7.5e-06,
|
||||
"litellm_provider": "vertex_ai-anthropic_models",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"max_output_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 7.5e-05,
|
||||
"output_cost_per_token_batches": 3.75e-05,
|
||||
|
|
@ -21577,8 +21932,8 @@
|
|||
"input_cost_per_token_batches": 7.5e-06,
|
||||
"litellm_provider": "vertex_ai-anthropic_models",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"max_output_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 7.5e-05,
|
||||
"output_cost_per_token_batches": 3.75e-05,
|
||||
|
|
|
|||
5
litellm/ocr/__init__.py
Normal file
5
litellm/ocr/__init__.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
"""OCR module for LiteLLM."""
|
||||
from .main import aocr, ocr
|
||||
|
||||
__all__ = ["ocr", "aocr"]
|
||||
|
||||
301
litellm/ocr/main.py
Normal file
301
litellm/ocr/main.py
Normal file
|
|
@ -0,0 +1,301 @@
|
|||
"""
|
||||
Main OCR function for LiteLLM.
|
||||
"""
|
||||
import asyncio
|
||||
import contextvars
|
||||
from functools import partial
|
||||
from typing import Any, Coroutine, Dict, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import request_timeout
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.utils import ProviderConfigManager, client
|
||||
|
||||
####### ENVIRONMENT VARIABLES ###################
|
||||
base_llm_http_handler = BaseLLMHTTPHandler()
|
||||
#################################################
|
||||
|
||||
|
||||
@client
|
||||
async def aocr(
|
||||
model: str,
|
||||
document: Dict[str, str],
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Async OCR function.
|
||||
|
||||
Args:
|
||||
model: Model name (e.g., "mistral/mistral-ocr-latest")
|
||||
document: Document to process in Mistral format:
|
||||
{"type": "document_url", "document_url": "https://..."} for PDFs/docs or
|
||||
{"type": "image_url", "image_url": "https://..."} for images
|
||||
api_key: Optional API key
|
||||
api_base: Optional API base URL
|
||||
timeout: Optional timeout
|
||||
custom_llm_provider: Optional custom LLM provider
|
||||
extra_headers: Optional extra headers
|
||||
**kwargs: Additional parameters (e.g., include_image_base64, pages, image_limit)
|
||||
|
||||
Returns:
|
||||
OCRResponse in Mistral OCR format with pages, model, usage_info, etc.
|
||||
|
||||
Example:
|
||||
```python
|
||||
import litellm
|
||||
|
||||
# OCR with PDF
|
||||
response = await litellm.aocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": "https://arxiv.org/pdf/2201.04234"
|
||||
},
|
||||
include_image_base64=True
|
||||
)
|
||||
|
||||
# OCR with image
|
||||
response = await litellm.aocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={
|
||||
"type": "image_url",
|
||||
"image_url": "https://example.com/image.png"
|
||||
}
|
||||
)
|
||||
|
||||
# OCR with base64 encoded PDF
|
||||
response = await litellm.aocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": f"data:application/pdf;base64,{base64_pdf}"
|
||||
}
|
||||
)
|
||||
```
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["aocr"] = True
|
||||
|
||||
# Get custom llm provider
|
||||
if custom_llm_provider is None:
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model=model, api_base=api_base
|
||||
)
|
||||
|
||||
func = partial(
|
||||
ocr,
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response
|
||||
|
||||
if response is None:
|
||||
raise ValueError(
|
||||
f"Got an unexpected None response from the OCR API: {response}"
|
||||
)
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def ocr(
|
||||
model: str,
|
||||
document: Dict[str, str],
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> Union[OCRResponse, Coroutine[Any, Any, OCRResponse]]:
|
||||
"""
|
||||
Synchronous OCR function.
|
||||
|
||||
Args:
|
||||
model: Model name (e.g., "mistral/mistral-ocr-latest")
|
||||
document: Document to process in Mistral format:
|
||||
{"type": "document_url", "document_url": "https://..."} for PDFs/docs or
|
||||
{"type": "image_url", "image_url": "https://..."} for images
|
||||
api_key: Optional API key
|
||||
api_base: Optional API base URL
|
||||
timeout: Optional timeout
|
||||
custom_llm_provider: Optional custom LLM provider
|
||||
extra_headers: Optional extra headers
|
||||
**kwargs: Additional parameters (e.g., include_image_base64, pages, image_limit)
|
||||
|
||||
Returns:
|
||||
OCRResponse in Mistral OCR format with pages, model, usage_info, etc.
|
||||
|
||||
Example:
|
||||
```python
|
||||
import litellm
|
||||
|
||||
# OCR with PDF
|
||||
response = litellm.ocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": "https://arxiv.org/pdf/2201.04234"
|
||||
},
|
||||
include_image_base64=True
|
||||
)
|
||||
|
||||
# OCR with image
|
||||
response = litellm.ocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={
|
||||
"type": "image_url",
|
||||
"image_url": "https://example.com/image.png"
|
||||
}
|
||||
)
|
||||
|
||||
# OCR with base64 encoded PDF
|
||||
response = litellm.ocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": f"data:application/pdf;base64,{base64_pdf}"
|
||||
}
|
||||
)
|
||||
|
||||
# Access pages
|
||||
for page in response.pages:
|
||||
print(f"Page {page.index}: {page.markdown}")
|
||||
```
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("aocr", False) is True
|
||||
|
||||
# Validate document parameter format (Mistral spec)
|
||||
if not isinstance(document, dict):
|
||||
raise ValueError(f"document must be a dict with 'type' and URL field, got {type(document)}")
|
||||
|
||||
doc_type = document.get("type")
|
||||
if doc_type not in ["document_url", "image_url"]:
|
||||
raise ValueError(f"Invalid document type: {doc_type}. Must be 'document_url' or 'image_url'")
|
||||
|
||||
model, custom_llm_provider, dynamic_api_key, dynamic_api_base = (
|
||||
litellm.get_llm_provider(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
)
|
||||
)
|
||||
|
||||
# Update with dynamic values if available
|
||||
if dynamic_api_key:
|
||||
api_key = dynamic_api_key
|
||||
if dynamic_api_base:
|
||||
api_base = dynamic_api_base
|
||||
|
||||
# Get provider config
|
||||
ocr_provider_config: Optional[BaseOCRConfig] = (
|
||||
ProviderConfigManager.get_provider_ocr_config(
|
||||
model=model,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
)
|
||||
|
||||
if ocr_provider_config is None:
|
||||
raise ValueError(
|
||||
f"OCR is not supported for provider: {custom_llm_provider}"
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"OCR call - model: {model}, provider: {custom_llm_provider}"
|
||||
)
|
||||
|
||||
# Extract OCR-specific parameters from kwargs
|
||||
supported_params = ocr_provider_config.get_supported_ocr_params(model=model)
|
||||
non_default_params = {}
|
||||
for param in supported_params:
|
||||
if param in kwargs:
|
||||
non_default_params[param] = kwargs.pop(param)
|
||||
|
||||
# Map parameters to provider-specific format
|
||||
optional_params = ocr_provider_config.map_ocr_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params={},
|
||||
model=model,
|
||||
)
|
||||
|
||||
verbose_logger.debug(f"OCR optional_params after mapping: {optional_params}")
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
litellm_params={
|
||||
"litellm_call_id": litellm_call_id,
|
||||
"api_base": api_base,
|
||||
},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
# Call the handler - pass document dict directly
|
||||
response = base_llm_http_handler.ocr(
|
||||
model=model,
|
||||
document=document, # Pass the entire document dict
|
||||
optional_params=optional_params,
|
||||
timeout=timeout or request_timeout,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
aocr=_is_async,
|
||||
headers=extra_headers,
|
||||
provider_config=ocr_provider_config,
|
||||
litellm_params={
|
||||
"api_base": api_base,
|
||||
"api_key": api_key,
|
||||
},
|
||||
)
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
|
@ -54,12 +54,7 @@ async def allm_passthrough_route(
|
|||
cookies: Optional[CookieTypes] = None,
|
||||
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
|
||||
**kwargs,
|
||||
) -> Union[
|
||||
httpx.Response,
|
||||
Coroutine[Any, Any, httpx.Response],
|
||||
Generator[Any, Any, Any],
|
||||
AsyncGenerator[Any, Any],
|
||||
]:
|
||||
) -> Union[httpx.Response, AsyncGenerator[Any, Any]]:
|
||||
"""
|
||||
Async: Reranks a list of documents based on their relevance to the query
|
||||
"""
|
||||
|
|
@ -111,23 +106,25 @@ async def allm_passthrough_route(
|
|||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
# Since allm_passthrough_route=True, we always get a coroutine from _async_passthrough_request
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
||||
try:
|
||||
# Only call raise_for_status if it's a Response object (not a generator)
|
||||
if isinstance(response, httpx.Response):
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_text = await e.response.aread()
|
||||
error_text_str = error_text.decode("utf-8")
|
||||
raise Exception(error_text_str)
|
||||
|
||||
|
||||
return response
|
||||
else:
|
||||
response = init_response
|
||||
|
||||
return response
|
||||
# This shouldn't happen when allm_passthrough_route=True, but handle it for type safety
|
||||
raise Exception("Expected coroutine from async passthrough route")
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
# For HTTP errors, re-raise as-is to preserve the original error details
|
||||
# The caller (e.g., proxy layer) can handle conversion to appropriate response format
|
||||
raise e
|
||||
except Exception as e:
|
||||
# For passthrough routes, we need to get the provider config to properly handle errors
|
||||
# For other exceptions, use provider-specific error handling
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
|
|
@ -186,6 +183,7 @@ def llm_passthrough_route(
|
|||
) -> Union[
|
||||
httpx.Response,
|
||||
Coroutine[Any, Any, httpx.Response],
|
||||
Coroutine[Any, Any, Union[httpx.Response, AsyncGenerator[Any, Any]]],
|
||||
Generator[Any, Any, Any],
|
||||
AsyncGenerator[Any, Any],
|
||||
]:
|
||||
|
|
@ -200,8 +198,10 @@ def llm_passthrough_route(
|
|||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
_is_async = allm_passthrough_route
|
||||
|
||||
if client is None:
|
||||
if allm_passthrough_route:
|
||||
if _is_async:
|
||||
client = litellm.module_level_aclient
|
||||
else:
|
||||
client = litellm.module_level_client
|
||||
|
|
@ -302,24 +302,40 @@ def llm_passthrough_route(
|
|||
# Update logging object with streaming status
|
||||
litellm_logging_obj.stream = is_streaming_request
|
||||
|
||||
## LOGGING PRE-CALL
|
||||
request_data = data if data else json
|
||||
litellm_logging_obj.pre_call(
|
||||
input=request_data,
|
||||
api_key=provider_api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": request_data,
|
||||
"api_base": str(updated_url),
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
response = client.client.send(request=request, stream=is_streaming_request)
|
||||
if asyncio.iscoroutine(response):
|
||||
if is_streaming_request:
|
||||
return _async_streaming(response, litellm_logging_obj, provider_config)
|
||||
else:
|
||||
return response
|
||||
response.raise_for_status()
|
||||
|
||||
if (
|
||||
hasattr(response, "iter_bytes") and is_streaming_request
|
||||
): # yield the chunk, so we can store it in the logging object
|
||||
|
||||
return _sync_streaming(response, litellm_logging_obj, provider_config)
|
||||
if _is_async:
|
||||
# Return the coroutine to be awaited by the caller
|
||||
return _async_passthrough_request(
|
||||
client=client,
|
||||
request=request,
|
||||
is_streaming_request=is_streaming_request,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
provider_config=provider_config,
|
||||
)
|
||||
else:
|
||||
# Sync path - client.client.send returns Response directly
|
||||
response: httpx.Response = client.client.send(request=request, stream=is_streaming_request) # type: ignore
|
||||
response.raise_for_status()
|
||||
|
||||
# For non-streaming responses, yield the entire response
|
||||
return response
|
||||
if (
|
||||
hasattr(response, "iter_bytes") and is_streaming_request
|
||||
): # yield the chunk, so we can store it in the logging object
|
||||
return _sync_streaming(response, litellm_logging_obj, provider_config)
|
||||
else:
|
||||
# For non-streaming responses, yield the entire response
|
||||
return response
|
||||
except Exception as e:
|
||||
if provider_config is None:
|
||||
raise e
|
||||
|
|
@ -329,6 +345,39 @@ def llm_passthrough_route(
|
|||
)
|
||||
|
||||
|
||||
async def _async_passthrough_request(
|
||||
client: Union[HTTPHandler, AsyncHTTPHandler],
|
||||
request: httpx.Request,
|
||||
is_streaming_request: bool,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
provider_config: "BasePassthroughConfig",
|
||||
) -> Union[httpx.Response, AsyncGenerator[Any, Any]]:
|
||||
"""
|
||||
Handle async passthrough requests.
|
||||
Uses async client to send request and properly handles streaming.
|
||||
"""
|
||||
# client.client.send returns a coroutine for async clients
|
||||
response_result = client.client.send(request=request, stream=is_streaming_request)
|
||||
|
||||
# Check if it's a coroutine and await it
|
||||
if asyncio.iscoroutine(response_result):
|
||||
if is_streaming_request:
|
||||
# Pass the coroutine to _async_streaming which will await it
|
||||
return _async_streaming(
|
||||
response=response_result,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
provider_config=provider_config,
|
||||
)
|
||||
else:
|
||||
response = await response_result
|
||||
await response.aread()
|
||||
response.raise_for_status()
|
||||
return response
|
||||
else:
|
||||
# Fallback for sync-like behavior (shouldn't happen in async path)
|
||||
raise Exception("Expected coroutine from async client")
|
||||
|
||||
|
||||
def _sync_streaming(
|
||||
response: httpx.Response,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
|
|
|
|||
|
|
@ -308,6 +308,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"allm_passthrough_route",
|
||||
"avector_store_search",
|
||||
"avector_store_create",
|
||||
"aocr",
|
||||
],
|
||||
version: Optional[str] = None,
|
||||
user_model: Optional[str] = None,
|
||||
|
|
@ -398,6 +399,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"allm_passthrough_route",
|
||||
"avector_store_search",
|
||||
"avector_store_create",
|
||||
"aocr",
|
||||
],
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
general_settings: dict,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import json
|
||||
import re
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Collection, Dict, List, Optional
|
||||
|
||||
import orjson
|
||||
from fastapi import Request, UploadFile, status
|
||||
|
|
@ -149,7 +149,7 @@ def _safe_get_request_headers(request: Optional[Request]) -> dict:
|
|||
def check_file_size_under_limit(
|
||||
request_data: dict,
|
||||
file: UploadFile,
|
||||
router_model_names: List[str],
|
||||
router_model_names: Collection[str],
|
||||
) -> bool:
|
||||
"""
|
||||
Check if any files passed in request are under max_file_size_mb
|
||||
|
|
|
|||
|
|
@ -23,6 +23,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
if not guardrail_name:
|
||||
raise ValueError("Pillar guardrail name is required")
|
||||
|
||||
optional_params = getattr(litellm_params, "optional_params", None)
|
||||
|
||||
_pillar_callback = PillarGuardrail(
|
||||
guardrail_name=guardrail_name,
|
||||
api_key=litellm_params.api_key,
|
||||
|
|
@ -30,12 +32,34 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
on_flagged_action=getattr(litellm_params, "on_flagged_action", "monitor"),
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=litellm_params.default_on,
|
||||
async_mode=_get_config_value(
|
||||
litellm_params, optional_params, "async_mode"
|
||||
),
|
||||
persist_session=_get_config_value(
|
||||
litellm_params, optional_params, "persist_session"
|
||||
),
|
||||
include_scanners=_get_config_value(
|
||||
litellm_params, optional_params, "include_scanners"
|
||||
),
|
||||
include_evidence=_get_config_value(
|
||||
litellm_params, optional_params, "include_evidence"
|
||||
),
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_pillar_callback)
|
||||
|
||||
return _pillar_callback
|
||||
|
||||
|
||||
def _get_config_value(litellm_params, optional_params, attribute_name):
|
||||
"""Return guardrail configuration value prioritising optional params when present."""
|
||||
|
||||
if optional_params is not None:
|
||||
value = getattr(optional_params, attribute_name, None)
|
||||
if value is not None:
|
||||
return value
|
||||
return getattr(litellm_params, attribute_name, None)
|
||||
|
||||
|
||||
guardrail_initializer_registry = {
|
||||
SupportedGuardrailIntegrations.PILLAR.value: initialize_guardrail,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -69,6 +69,10 @@ class PillarGuardrail(CustomGuardrail):
|
|||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
on_flagged_action: Optional[str] = None,
|
||||
async_mode: Optional[bool] = None,
|
||||
persist_session: Optional[bool] = None,
|
||||
include_scanners: Optional[bool] = None,
|
||||
include_evidence: Optional[bool] = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
"""
|
||||
|
|
@ -110,6 +114,31 @@ class PillarGuardrail(CustomGuardrail):
|
|||
f"Pillar Guardrail: Initialized with on_flagged_action: {self.on_flagged_action}"
|
||||
)
|
||||
|
||||
self.async_mode = self._resolve_bool_config(
|
||||
provided_value=async_mode,
|
||||
env_var="PILLAR_ASYNC",
|
||||
default=None,
|
||||
setting_name="async_mode",
|
||||
)
|
||||
self.persist_session = self._resolve_bool_config(
|
||||
provided_value=persist_session,
|
||||
env_var="PILLAR_PERSIST",
|
||||
default=None,
|
||||
setting_name="persist_session",
|
||||
)
|
||||
self.include_scanners = self._resolve_bool_config(
|
||||
provided_value=include_scanners,
|
||||
env_var="PILLAR_INCLUDE_SCANNERS",
|
||||
default=True,
|
||||
setting_name="include_scanners",
|
||||
)
|
||||
self.include_evidence = self._resolve_bool_config(
|
||||
provided_value=include_evidence,
|
||||
env_var="PILLAR_INCLUDE_EVIDENCE",
|
||||
default=True,
|
||||
setting_name="include_evidence",
|
||||
)
|
||||
|
||||
# Define supported event hooks
|
||||
supported_event_hooks = [
|
||||
GuardrailEventHooks.pre_call,
|
||||
|
|
@ -347,12 +376,74 @@ class PillarGuardrail(CustomGuardrail):
|
|||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
# Add Pillar-specific headers for enhanced response data
|
||||
headers["plr_evidence"] = "true"
|
||||
headers["plr_scanners"] = "true"
|
||||
# Add Pillar-specific headers based on configuration
|
||||
self._set_bool_header(headers, "plr_scanners", self.include_scanners)
|
||||
self._set_bool_header(headers, "plr_evidence", self.include_evidence)
|
||||
self._set_bool_header(headers, "plr_async", self.async_mode)
|
||||
self._set_bool_header(headers, "plr_persist", self.persist_session)
|
||||
|
||||
return headers
|
||||
|
||||
def _set_bool_header(
|
||||
self, headers: Dict[str, str], header_name: str, value: Optional[bool]
|
||||
) -> None:
|
||||
"""Apply a boolean value as a lowercase string HTTP header when provided."""
|
||||
|
||||
if value is None:
|
||||
return
|
||||
headers[header_name] = "true" if value else "false"
|
||||
|
||||
def _resolve_bool_config(
|
||||
self,
|
||||
provided_value: Optional[Union[bool, str, int]],
|
||||
env_var: Optional[str],
|
||||
default: Optional[bool],
|
||||
setting_name: str,
|
||||
) -> Optional[bool]:
|
||||
"""Resolve configuration precedence: explicit value -> environment -> default."""
|
||||
|
||||
if provided_value is not None:
|
||||
try:
|
||||
return self._parse_bool_value(provided_value)
|
||||
except ValueError:
|
||||
verbose_proxy_logger.warning(
|
||||
"Pillar Guardrail: Invalid boolean value '%s' for %s, falling back to default.",
|
||||
provided_value,
|
||||
setting_name,
|
||||
)
|
||||
return default
|
||||
|
||||
if env_var:
|
||||
env_value = os.getenv(env_var)
|
||||
if env_value is not None:
|
||||
try:
|
||||
return self._parse_bool_value(env_value)
|
||||
except ValueError:
|
||||
verbose_proxy_logger.warning(
|
||||
"Pillar Guardrail: Invalid boolean env value '%s' for %s, falling back to default.",
|
||||
env_value,
|
||||
env_var,
|
||||
)
|
||||
return default
|
||||
|
||||
return default
|
||||
|
||||
@staticmethod
|
||||
def _parse_bool_value(value: Union[bool, str, int]) -> bool:
|
||||
"""Normalise various truthy/falsey inputs to a strict boolean."""
|
||||
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, int):
|
||||
return bool(value)
|
||||
|
||||
value_str = str(value).strip().lower()
|
||||
if value_str in {"true", "1", "yes", "y", "on"}:
|
||||
return True
|
||||
if value_str in {"false", "0", "no", "n", "off"}:
|
||||
return False
|
||||
raise ValueError(f"Unrecognised boolean value: {value}")
|
||||
|
||||
def _extract_model_and_provider(self, data: dict) -> Tuple[str, str]:
|
||||
"""
|
||||
Extract the model and provider from the request data.
|
||||
|
|
|
|||
|
|
@ -9,7 +9,10 @@ Has all /sso/* routes
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import os
|
||||
import secrets
|
||||
from copy import deepcopy
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
|
|
@ -381,7 +384,10 @@ async def get_generic_sso_response(
|
|||
try:
|
||||
result = await generic_sso.verify_and_process(
|
||||
request,
|
||||
params={"include_client_id": generic_include_client_id},
|
||||
params=SSOAuthenticationHandler.prepare_token_exchange_parameters(
|
||||
request=request,
|
||||
generic_include_client_id=generic_include_client_id,
|
||||
),
|
||||
headers=additional_generic_sso_headers_dict,
|
||||
)
|
||||
|
||||
|
|
@ -1067,30 +1073,97 @@ class SSOAuthenticationHandler:
|
|||
allow_insecure_http=True,
|
||||
scope=generic_scope,
|
||||
)
|
||||
with generic_sso:
|
||||
# TODO: state should be a random string and added to the user session with cookie
|
||||
# or a cryptographicly signed state that we can verify stateless
|
||||
# For simplification we are using a static state, this is not perfect but some
|
||||
# SSO providers do not allow stateless verification
|
||||
redirect_params = (
|
||||
SSOAuthenticationHandler._get_generic_sso_redirect_params(
|
||||
state=state,
|
||||
generic_authorization_endpoint=generic_authorization_endpoint,
|
||||
)
|
||||
)
|
||||
|
||||
return await generic_sso.get_login_redirect(**redirect_params) # type: ignore
|
||||
return await SSOAuthenticationHandler.get_generic_sso_redirect_response(
|
||||
generic_sso=generic_sso,
|
||||
state=state,
|
||||
generic_authorization_endpoint=generic_authorization_endpoint,
|
||||
)
|
||||
raise ValueError(
|
||||
"Unknown SSO provider. Please setup SSO with client IDs https://docs.litellm.ai/docs/proxy/admin_ui_sso"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def get_generic_sso_redirect_response(
|
||||
generic_sso: Any,
|
||||
state: Optional[str] = None,
|
||||
generic_authorization_endpoint: Optional[str] = None,
|
||||
) -> Optional[RedirectResponse]:
|
||||
"""
|
||||
Get the redirect response for Generic SSO
|
||||
"""
|
||||
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
|
||||
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
with generic_sso:
|
||||
# TODO: state should be a random string and added to the user session with cookie
|
||||
# or a cryptographicly signed state that we can verify stateless
|
||||
# For simplification we are using a static state, this is not perfect but some
|
||||
# SSO providers do not allow stateless verification
|
||||
redirect_params, code_verifier = (
|
||||
SSOAuthenticationHandler._get_generic_sso_redirect_params(
|
||||
state=state,
|
||||
generic_authorization_endpoint=generic_authorization_endpoint,
|
||||
)
|
||||
)
|
||||
|
||||
# Separate PKCE params from state params (fastapi-sso doesn't accept code_challenge)
|
||||
pkce_params = {}
|
||||
state_only_params = {}
|
||||
for key, value in redirect_params.items():
|
||||
if key in ("code_challenge", "code_challenge_method"):
|
||||
pkce_params[key] = value
|
||||
else:
|
||||
state_only_params[key] = value
|
||||
|
||||
# Get the redirect response from fastapi-sso with only state param
|
||||
redirect_response = await generic_sso.get_login_redirect(**state_only_params) # type: ignore
|
||||
|
||||
# If PKCE is enabled, add PKCE parameters to the redirect URL
|
||||
if code_verifier and "state" in redirect_params:
|
||||
|
||||
# Store code_verifier in cache (10 min TTL)
|
||||
cache_key = f"pkce_verifier:{redirect_params['state']}"
|
||||
user_api_key_cache.set_cache(
|
||||
key=cache_key,
|
||||
value=code_verifier,
|
||||
ttl=600,
|
||||
)
|
||||
|
||||
# Add PKCE parameters to the authorization URL
|
||||
if pkce_params:
|
||||
parsed_url = urlparse(str(redirect_response.headers["location"]))
|
||||
query_params = parse_qs(parsed_url.query)
|
||||
|
||||
# Add PKCE parameters
|
||||
for key, value in pkce_params.items():
|
||||
query_params[key] = [value]
|
||||
|
||||
# Reconstruct the URL with PKCE parameters
|
||||
new_query = urlencode(query_params, doseq=True)
|
||||
new_url = urlunparse((
|
||||
parsed_url.scheme,
|
||||
parsed_url.netloc,
|
||||
parsed_url.path,
|
||||
parsed_url.params,
|
||||
new_query,
|
||||
parsed_url.fragment
|
||||
))
|
||||
|
||||
# Update the redirect response
|
||||
redirect_response.headers["location"] = new_url
|
||||
verbose_proxy_logger.debug(
|
||||
"PKCE parameters added to authorization URL"
|
||||
)
|
||||
return redirect_response
|
||||
|
||||
@staticmethod
|
||||
def _get_generic_sso_redirect_params(
|
||||
state: Optional[str] = None,
|
||||
generic_authorization_endpoint: Optional[str] = None,
|
||||
) -> dict:
|
||||
) -> Tuple[dict, Optional[str]]:
|
||||
"""
|
||||
Get redirect parameters for Generic SSO with proper state priority handling.
|
||||
Optionally generates PKCE parameters if GENERIC_CLIENT_USE_PKCE is enabled.
|
||||
|
||||
Priority order:
|
||||
1. CLI state (if provided)
|
||||
|
|
@ -1102,9 +1175,12 @@ class SSOAuthenticationHandler:
|
|||
generic_authorization_endpoint: Authorization endpoint URL
|
||||
|
||||
Returns:
|
||||
dict: Redirect parameters for SSO login
|
||||
Tuple[dict, Optional[str]]:
|
||||
- Redirect parameters for SSO login (may include PKCE params)
|
||||
- code_verifier (if PKCE is enabled, None otherwise)
|
||||
"""
|
||||
redirect_params = {}
|
||||
code_verifier: Optional[str] = None
|
||||
|
||||
if state:
|
||||
# CLI state takes priority
|
||||
|
|
@ -1122,7 +1198,18 @@ class SSOAuthenticationHandler:
|
|||
uuid.uuid4().hex
|
||||
) # set state param for okta - required
|
||||
|
||||
return redirect_params
|
||||
# Handle PKCE (Proof Key for Code Exchange) if enabled
|
||||
# Set GENERIC_CLIENT_USE_PKCE=true to enable PKCE for enhanced OAuth security
|
||||
use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
|
||||
if use_pkce:
|
||||
code_verifier, code_challenge = SSOAuthenticationHandler.generate_pkce_params()
|
||||
redirect_params["code_challenge"] = code_challenge
|
||||
redirect_params["code_challenge_method"] = "S256"
|
||||
verbose_proxy_logger.debug(
|
||||
"PKCE enabled - code_challenge added to authorization request"
|
||||
)
|
||||
|
||||
return redirect_params, code_verifier
|
||||
|
||||
@staticmethod
|
||||
def should_use_sso_handler(
|
||||
|
|
@ -1606,6 +1693,69 @@ class SSOAuthenticationHandler:
|
|||
redirect_response = RedirectResponse(url=litellm_dashboard_ui, status_code=303)
|
||||
redirect_response.set_cookie(key="token", value=jwt_token)
|
||||
return redirect_response
|
||||
|
||||
|
||||
@staticmethod
|
||||
def prepare_token_exchange_parameters(
|
||||
request: Request,
|
||||
generic_include_client_id: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Prepare token exchange parameters for Generic SSO.
|
||||
|
||||
Args:
|
||||
request: Request object
|
||||
generic_include_client_id: Generic OAuth Client ID
|
||||
|
||||
Returns:
|
||||
dict: Token exchange parameters
|
||||
"""
|
||||
# Prepare token exchange parameters
|
||||
token_params = {"include_client_id": generic_include_client_id}
|
||||
|
||||
# Retrieve PKCE code_verifier if PKCE was used in authorization
|
||||
query_params = dict(request.query_params)
|
||||
state = query_params.get("state")
|
||||
if state:
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
cache_key = f"pkce_verifier:{state}"
|
||||
code_verifier = user_api_key_cache.get_cache(key=cache_key)
|
||||
|
||||
if code_verifier:
|
||||
# Add code_verifier to token exchange parameters
|
||||
token_params["code_verifier"] = code_verifier
|
||||
verbose_proxy_logger.debug(
|
||||
"PKCE code_verifier retrieved and will be included in token exchange"
|
||||
)
|
||||
|
||||
# Clean up the cache entry (single-use verifier)
|
||||
user_api_key_cache.delete_cache(key=cache_key)
|
||||
return token_params
|
||||
|
||||
|
||||
@staticmethod
|
||||
def generate_pkce_params() -> Tuple[str, str]:
|
||||
"""
|
||||
Generate PKCE (Proof Key for Code Exchange) parameters for OAuth 2.0.
|
||||
|
||||
Returns:
|
||||
Tuple[str, str]: (code_verifier, code_challenge)
|
||||
- code_verifier: Random 43-128 character string (we use 43 for efficiency)
|
||||
- code_challenge: Base64-URL-encoded SHA256 hash of the code_verifier
|
||||
|
||||
Reference: https://datatracker.ietf.org/doc/html/rfc7636
|
||||
"""
|
||||
# Generate a cryptographically random code_verifier (43 characters)
|
||||
# Using 32 random bytes which becomes 43 characters when base64-url-encoded
|
||||
code_verifier = base64.urlsafe_b64encode(secrets.token_bytes(32)).decode('utf-8').rstrip('=')
|
||||
|
||||
# Generate code_challenge using S256 method (SHA256)
|
||||
code_challenge_bytes = hashlib.sha256(code_verifier.encode('utf-8')).digest()
|
||||
code_challenge = base64.urlsafe_b64encode(code_challenge_bytes).decode('utf-8').rstrip('=')
|
||||
|
||||
return code_verifier, code_challenge
|
||||
|
||||
|
||||
|
||||
class MicrosoftSSOHandler:
|
||||
|
|
@ -1739,7 +1889,7 @@ class MicrosoftSSOHandler:
|
|||
Extract app roles from the Microsoft Entra ID (Azure AD) id_token JWT.
|
||||
|
||||
App roles are assigned in the Azure AD Enterprise Application and appear
|
||||
in the 'roles' claim of the id_token.
|
||||
in the 'app_roles' claim of the id_token.
|
||||
|
||||
Args:
|
||||
id_token (Optional[str]): The JWT id_token from Microsoft SSO
|
||||
|
|
@ -1758,8 +1908,9 @@ class MicrosoftSSOHandler:
|
|||
# (signature is already verified by fastapi_sso)
|
||||
decoded_token = jwt.decode(id_token, options={"verify_signature": False})
|
||||
|
||||
# Extract roles claim from the token
|
||||
roles = decoded_token.get("roles", [])
|
||||
# Extract app_roles claim from the token
|
||||
## check for both 'roles' and 'app_roles' claims
|
||||
roles = decoded_token.get("app_roles", []) or decoded_token.get("roles", [])
|
||||
|
||||
if roles and isinstance(roles, list):
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
|
|||
2
litellm/proxy/ocr_endpoints/__init__.py
Normal file
2
litellm/proxy/ocr_endpoints/__init__.py
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
# OCR Endpoints
|
||||
|
||||
97
litellm/proxy/ocr_endpoints/endpoints.py
Normal file
97
litellm/proxy/ocr_endpoints/endpoints.py
Normal file
|
|
@ -0,0 +1,97 @@
|
|||
#### OCR Endpoints #####
|
||||
|
||||
import orjson
|
||||
from fastapi import APIRouter, Depends, Request, Response
|
||||
from fastapi.responses import ORJSONResponse
|
||||
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/ocr",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse,
|
||||
tags=["ocr"],
|
||||
)
|
||||
@router.post(
|
||||
"/ocr",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse,
|
||||
tags=["ocr"],
|
||||
)
|
||||
async def ocr(
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
OCR endpoint for extracting text from documents and images.
|
||||
|
||||
Follows the Mistral OCR API spec:
|
||||
https://docs.mistral.ai/capabilities/vision/#optical-character-recognition-ocr
|
||||
|
||||
Example:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/ocr" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "mistral/mistral-ocr-latest",
|
||||
"document": {
|
||||
"type": "document_url",
|
||||
"document_url": "https://arxiv.org/pdf/2201.04234"
|
||||
}
|
||||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
select_data_generator,
|
||||
user_api_base,
|
||||
user_max_tokens,
|
||||
user_model,
|
||||
user_request_timeout,
|
||||
user_temperature,
|
||||
version,
|
||||
)
|
||||
|
||||
# Read request body
|
||||
body = await request.body()
|
||||
data = orjson.loads(body)
|
||||
|
||||
# Process request using ProxyBaseLLMRequestProcessing
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
return await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
route_type="aocr",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=select_data_generator,
|
||||
model=None,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
user_max_tokens=user_max_tokens,
|
||||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
)
|
||||
except Exception as e:
|
||||
raise await processor._handle_llm_api_exception(
|
||||
e=e,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
|
||||
|
|
@ -8,7 +8,7 @@ Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc.
|
|||
|
||||
import json
|
||||
import os
|
||||
from typing import Optional, cast
|
||||
from typing import Any, Optional, Union, cast
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
|
||||
|
|
@ -482,6 +482,172 @@ async def anthropic_proxy_route(
|
|||
return received_value
|
||||
|
||||
|
||||
# Bedrock endpoint actions - consolidated list used for model extraction and streaming detection
|
||||
BEDROCK_ENDPOINT_ACTIONS = {
|
||||
"invoke",
|
||||
"invoke-with-response-stream",
|
||||
"converse",
|
||||
"converse-stream",
|
||||
"count_tokens",
|
||||
"count-tokens",
|
||||
}
|
||||
|
||||
BEDROCK_STREAMING_ACTIONS = {"invoke-with-response-stream", "converse-stream"}
|
||||
|
||||
|
||||
def _extract_model_from_bedrock_endpoint(endpoint: str) -> str:
|
||||
"""
|
||||
Extract model name from Bedrock endpoint path.
|
||||
|
||||
Handles model names with slashes (e.g., aws/anthropic/bedrock-claude-3-5-sonnet-v1)
|
||||
by finding the action in the endpoint and extracting everything between "model" and the action.
|
||||
|
||||
Args:
|
||||
endpoint: The endpoint path (e.g., "/model/aws/anthropic/model-name/invoke")
|
||||
|
||||
Returns:
|
||||
The extracted model name (e.g., "aws/anthropic/model-name")
|
||||
|
||||
Raises:
|
||||
ValueError: If model cannot be extracted from endpoint
|
||||
"""
|
||||
try:
|
||||
endpoint_parts = endpoint.split("/")
|
||||
|
||||
if "application-inference-profile" in endpoint:
|
||||
# Format: model/application-inference-profile/{profile-id}/{action}
|
||||
return "/".join(endpoint_parts[1:3])
|
||||
|
||||
# Format: model/{modelId}/{action}
|
||||
# Find the index of the action in the endpoint parts
|
||||
action_index = None
|
||||
for idx, part in enumerate(endpoint_parts):
|
||||
if part in BEDROCK_ENDPOINT_ACTIONS:
|
||||
action_index = idx
|
||||
break
|
||||
|
||||
if action_index is not None and action_index > 1:
|
||||
# Join all parts between "model" and the action
|
||||
return "/".join(endpoint_parts[1:action_index])
|
||||
|
||||
# Fallback to taking everything after "model" if no action found
|
||||
return "/".join(endpoint_parts[1:])
|
||||
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"Model missing from endpoint. Expected format: /model/{{modelId}}/{{action}}. Got: {endpoint}"
|
||||
) from e
|
||||
|
||||
|
||||
async def handle_bedrock_passthrough_router_model(
|
||||
model: str,
|
||||
endpoint: str,
|
||||
request: Request,
|
||||
request_body: dict,
|
||||
llm_router: litellm.Router,
|
||||
) -> Union[Response, StreamingResponse]:
|
||||
"""
|
||||
Handle Bedrock passthrough for router models (models defined in config.yaml).
|
||||
|
||||
This helper delegates to llm_router.allm_passthrough_route for proper credential
|
||||
and configuration management from the router.
|
||||
|
||||
Args:
|
||||
model: The router model name (e.g., "aws/anthropic/bedrock-claude-3-5-sonnet-v1")
|
||||
endpoint: The Bedrock endpoint path (e.g., "/model/{modelId}/invoke")
|
||||
request: The FastAPI request object
|
||||
request_body: The parsed request body
|
||||
llm_router: The LiteLLM router instance
|
||||
|
||||
Returns:
|
||||
Response or StreamingResponse depending on endpoint type
|
||||
"""
|
||||
# Detect streaming based on endpoint
|
||||
is_streaming = any(action in endpoint for action in BEDROCK_STREAMING_ACTIONS)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Bedrock router passthrough: model='{model}', endpoint='{endpoint}', streaming={is_streaming}"
|
||||
)
|
||||
|
||||
# Call router passthrough
|
||||
try:
|
||||
result = await llm_router.allm_passthrough_route(
|
||||
model=model,
|
||||
method=request.method,
|
||||
endpoint=endpoint,
|
||||
request_query_params=request.query_params,
|
||||
request_headers=dict(request.headers),
|
||||
stream=is_streaming,
|
||||
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,
|
||||
)
|
||||
except httpx.HTTPStatusError as e:
|
||||
# Handle HTTP errors from the provider by converting to HTTPException
|
||||
error_body = await e.response.aread()
|
||||
error_text = error_body.decode("utf-8")
|
||||
|
||||
raise HTTPException(
|
||||
status_code=e.response.status_code,
|
||||
detail={"error": error_text},
|
||||
)
|
||||
except Exception as e:
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
# If it's a BaseLLMException (from non-HTTP errors), convert to HTTPException
|
||||
if isinstance(e, BaseLLMException):
|
||||
raise HTTPException(
|
||||
status_code=e.status_code,
|
||||
detail={"error": e.message},
|
||||
)
|
||||
# Re-raise any other exceptions
|
||||
raise e
|
||||
|
||||
# Handle streaming response
|
||||
if is_streaming:
|
||||
import inspect
|
||||
|
||||
if inspect.isasyncgen(result):
|
||||
# AsyncGenerator case
|
||||
return StreamingResponse(
|
||||
content=result,
|
||||
status_code=200,
|
||||
headers={"content-type": "application/vnd.amazon.eventstream"},
|
||||
)
|
||||
else:
|
||||
# httpx.Response case
|
||||
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,
|
||||
),
|
||||
)
|
||||
|
||||
# Handle 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,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def handle_bedrock_count_tokens(
|
||||
endpoint: str,
|
||||
request: Request,
|
||||
|
|
@ -560,6 +726,15 @@ async def bedrock_llm_proxy_route(
|
|||
):
|
||||
"""
|
||||
Handles Bedrock LLM API calls.
|
||||
|
||||
Supports both direct Bedrock models and router models from config.yaml.
|
||||
|
||||
Endpoints:
|
||||
- /model/{modelId}/invoke
|
||||
- /model/{modelId}/invoke-with-response-stream
|
||||
- /model/{modelId}/converse
|
||||
- /model/{modelId}/converse-stream
|
||||
- /model/application-inference-profile/{profileId}/{action}
|
||||
"""
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.proxy_server import (
|
||||
|
|
@ -588,24 +763,38 @@ async def bedrock_llm_proxy_route(
|
|||
request_body=request_body,
|
||||
)
|
||||
|
||||
data: Dict[str, Any] = {}
|
||||
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
# Extract model from endpoint path using helper
|
||||
try:
|
||||
endpoint_parts = endpoint.split("/")
|
||||
if "application-inference-profile" in endpoint:
|
||||
# For application-inference-profile, include the profile ID part as well
|
||||
model = "/".join(endpoint_parts[1:3])
|
||||
else:
|
||||
model = endpoint_parts[1]
|
||||
except Exception:
|
||||
model = _extract_model_from_bedrock_endpoint(endpoint=endpoint)
|
||||
except ValueError as e:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Model missing from endpoint. Expected format: /model/<Model>/<endpoint>. Got: "
|
||||
+ endpoint,
|
||||
},
|
||||
detail={"error": str(e)},
|
||||
)
|
||||
|
||||
# Check if this is a router model (from config.yaml)
|
||||
is_router_model = is_passthrough_request_using_router_model(
|
||||
request_body={"model": model}, llm_router=llm_router
|
||||
)
|
||||
|
||||
# If router model, use dedicated router passthrough handler
|
||||
if is_router_model and llm_router:
|
||||
return await handle_bedrock_passthrough_router_model(
|
||||
model=model,
|
||||
endpoint=endpoint,
|
||||
request=request,
|
||||
request_body=request_body,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
# Fall back to existing implementation for direct Bedrock models
|
||||
verbose_proxy_logger.debug(
|
||||
f"Bedrock passthrough: Using direct Bedrock model '{model}' for endpoint '{endpoint}'"
|
||||
)
|
||||
|
||||
data: Dict[str, Any] = {}
|
||||
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
|
||||
data["method"] = request.method
|
||||
data["endpoint"] = endpoint
|
||||
data["data"] = request_body
|
||||
|
|
|
|||
|
|
@ -1,16 +1,25 @@
|
|||
model_list:
|
||||
- model_name: db-openai-endpoint
|
||||
- model_name: mistral/*
|
||||
litellm_params:
|
||||
model: openai/gm
|
||||
api_key: hi
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
|
||||
|
||||
litellm_settings:
|
||||
callbacks: ["dynamic_rate_limiter_v3"]
|
||||
priority_reservation:
|
||||
"prod": 0.9 # 90% reserved for production (9 RPM)
|
||||
"dev": 0.1 # 10% reserved for development (1 RPM)
|
||||
priority_reservation_settings:
|
||||
default_priority: 0.2 # Weight (0%) assigned to keys without explicit priority metadata
|
||||
saturation_threshold: 0.50 # A model is saturated if it has hit 50% of its RPM limit
|
||||
model: mistral/*
|
||||
- model_name: special-bedrock-model
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-west-2
|
||||
custom_llm_provider: bedrock
|
||||
- model_name: aws/anthropic/bedrock-claude-3-5-sonnet-v1
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-west-2
|
||||
custom_llm_provider: bedrock
|
||||
# Load balancing test - multiple deployments with same model_name
|
||||
- model_name: load-balanced-claude
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-west-2
|
||||
custom_llm_provider: bedrock
|
||||
- model_name: load-balanced-claude
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
aws_region_name: us-east-1
|
||||
custom_llm_provider: bedrock
|
||||
|
|
@ -308,6 +308,7 @@ from litellm.proxy.management_endpoints.user_agent_analytics_endpoints import (
|
|||
)
|
||||
from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update
|
||||
from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMiddleware
|
||||
from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
router as openai_files_router,
|
||||
)
|
||||
|
|
@ -9774,6 +9775,7 @@ app.include_router(response_router)
|
|||
app.include_router(batches_router)
|
||||
app.include_router(public_endpoints_router)
|
||||
app.include_router(rerank_router)
|
||||
app.include_router(ocr_router)
|
||||
app.include_router(image_router)
|
||||
app.include_router(fine_tuning_router)
|
||||
app.include_router(vector_store_router)
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ ROUTE_ENDPOINT_MAPPING = {
|
|||
"alist_input_items": "/responses/{response_id}/input_items",
|
||||
"aimage_edit": "/images/edits",
|
||||
"acancel_responses": "/responses/{response_id}/cancel",
|
||||
"aocr": "/ocr",
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -98,6 +99,7 @@ async def route_request(
|
|||
"allm_passthrough_route",
|
||||
"avector_store_search",
|
||||
"avector_store_create",
|
||||
"aocr",
|
||||
],
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -3707,6 +3707,7 @@ def construct_database_url_from_env_vars() -> Optional[str]:
|
|||
database_username = os.getenv("DATABASE_USERNAME")
|
||||
database_password = os.getenv("DATABASE_PASSWORD")
|
||||
database_name = os.getenv("DATABASE_NAME")
|
||||
database_schema = os.getenv("DATABASE_SCHEMA")
|
||||
|
||||
if database_host and database_username and database_name:
|
||||
# Handle the problem of special character escaping in the database URL
|
||||
|
|
@ -3722,6 +3723,9 @@ def construct_database_url_from_env_vars() -> Optional[str]:
|
|||
else:
|
||||
database_url = f"postgresql://{database_username_enc}@{database_host}/{database_name_enc}"
|
||||
|
||||
if database_schema:
|
||||
database_url += f"?schema={database_schema}"
|
||||
|
||||
return database_url
|
||||
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -189,7 +189,7 @@ class RoutingArgs(enum.Enum):
|
|||
|
||||
|
||||
class Router:
|
||||
model_names: List = []
|
||||
model_names: set = set()
|
||||
cache_responses: Optional[bool] = False
|
||||
default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour
|
||||
tenacity = None
|
||||
|
|
@ -872,6 +872,14 @@ class Router:
|
|||
generate_content_stream, call_type="generate_content_stream"
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# OCR routes
|
||||
#########################################################
|
||||
from litellm.ocr import aocr, ocr
|
||||
|
||||
self.aocr = self.factory_function(aocr, call_type="aocr")
|
||||
self.ocr = self.factory_function(ocr, call_type="ocr")
|
||||
|
||||
def validate_fallbacks(self, fallback_param: Optional[List]):
|
||||
"""
|
||||
Validate the fallbacks parameter.
|
||||
|
|
@ -1057,7 +1065,7 @@ class Router:
|
|||
|
||||
self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs)
|
||||
request_priority = kwargs.get("priority") or self.default_priority
|
||||
start_time = time.time()
|
||||
start_time = time.perf_counter()
|
||||
_is_prompt_management_model = self._is_prompt_management_model(model)
|
||||
|
||||
if _is_prompt_management_model:
|
||||
|
|
@ -1070,7 +1078,7 @@ class Router:
|
|||
response = await self.schedule_acompletion(**kwargs)
|
||||
else:
|
||||
response = await self.async_function_with_fallbacks(**kwargs)
|
||||
end_time = time.time()
|
||||
end_time = time.perf_counter()
|
||||
_duration = end_time - start_time
|
||||
asyncio.create_task(
|
||||
self.service_logger_obj.async_service_success_hook(
|
||||
|
|
@ -1245,7 +1253,7 @@ class Router:
|
|||
input_kwargs_for_streaming_fallback["model"] = model
|
||||
|
||||
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs)
|
||||
start_time = time.time()
|
||||
start_time = time.perf_counter()
|
||||
deployment = await self.async_get_available_deployment(
|
||||
model=model,
|
||||
messages=messages,
|
||||
|
|
@ -1254,7 +1262,7 @@ class Router:
|
|||
)
|
||||
|
||||
_timeout_debug_deployment_dict = deployment
|
||||
end_time = time.time()
|
||||
end_time = time.perf_counter()
|
||||
_duration = end_time - start_time
|
||||
asyncio.create_task(
|
||||
self.service_logger_obj.async_service_success_hook(
|
||||
|
|
@ -1834,8 +1842,8 @@ class Router:
|
|||
await self.scheduler.add_request(request=item)
|
||||
|
||||
## POLL QUEUE
|
||||
end_time = time.time() + self.timeout
|
||||
curr_time = time.time()
|
||||
end_time = time.monotonic() + self.timeout
|
||||
curr_time = time.monotonic()
|
||||
poll_interval = self.scheduler.polling_interval # poll every 3ms
|
||||
make_request = False
|
||||
|
||||
|
|
@ -1852,7 +1860,7 @@ class Router:
|
|||
break
|
||||
else: ## ELSE -> loop till default_timeout
|
||||
await asyncio.sleep(poll_interval)
|
||||
curr_time = time.time()
|
||||
curr_time = time.monotonic()
|
||||
|
||||
if make_request:
|
||||
try:
|
||||
|
|
@ -1896,8 +1904,8 @@ class Router:
|
|||
await self.scheduler.add_request(request=item)
|
||||
|
||||
## POLL QUEUE
|
||||
end_time = time.time() + self.timeout
|
||||
curr_time = time.time()
|
||||
end_time = time.monotonic() + self.timeout
|
||||
curr_time = time.monotonic()
|
||||
poll_interval = self.scheduler.polling_interval # poll every 3ms
|
||||
make_request = False
|
||||
|
||||
|
|
@ -1914,7 +1922,7 @@ class Router:
|
|||
break
|
||||
else: ## ELSE -> loop till default_timeout
|
||||
await asyncio.sleep(poll_interval)
|
||||
curr_time = time.time()
|
||||
curr_time = time.monotonic()
|
||||
|
||||
if make_request:
|
||||
try:
|
||||
|
|
@ -2732,6 +2740,37 @@ class Router:
|
|||
)
|
||||
)
|
||||
raise e
|
||||
|
||||
def _add_deployment_model_to_endpoint_for_llm_passthrough_route(
|
||||
self, kwargs: Dict[str, Any],
|
||||
model: str,
|
||||
model_name: str
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Add the deployment model to the endpoint for LLM passthrough route.
|
||||
|
||||
e.g for bedrock invoke users can pass endpoint as /model/special-bedrock-model/invoke
|
||||
it should be actually sent as /model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke
|
||||
"""
|
||||
if "endpoint" in kwargs and kwargs["endpoint"]:
|
||||
# For provider-specific endpoints, strip the provider prefix from model_name
|
||||
# e.g., "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0" -> "us.anthropic.claude-3-5-sonnet-20240620-v1:0"
|
||||
from litellm import get_llm_provider
|
||||
|
||||
try:
|
||||
# get_llm_provider returns (model_without_prefix, provider, api_key, api_base)
|
||||
stripped_model_name, _, _, _ = get_llm_provider(
|
||||
model=model_name,
|
||||
custom_llm_provider=kwargs.get("custom_llm_provider"),
|
||||
api_base=kwargs.get("api_base"),
|
||||
)
|
||||
replacement_model_name = stripped_model_name
|
||||
except Exception:
|
||||
# If get_llm_provider fails, fall back to using model_name as-is
|
||||
replacement_model_name = model_name
|
||||
|
||||
kwargs["endpoint"] = kwargs["endpoint"].replace(model, replacement_model_name)
|
||||
return kwargs
|
||||
|
||||
async def _ageneric_api_call_with_fallbacks_helper(
|
||||
self, model: str, original_generic_function: Callable, **kwargs
|
||||
|
|
@ -2764,6 +2803,7 @@ class Router:
|
|||
model_name = data["model"]
|
||||
self.total_calls[model_name] += 1
|
||||
|
||||
self._add_deployment_model_to_endpoint_for_llm_passthrough_route(kwargs=kwargs, model=model, model_name=model_name)
|
||||
### get custom
|
||||
response = original_generic_function(
|
||||
**{
|
||||
|
|
@ -2842,6 +2882,12 @@ class Router:
|
|||
|
||||
self.total_calls[model_name] += 1
|
||||
|
||||
# For passthrough routes, use the actual model from deployment
|
||||
# and swap model name in endpoint if present
|
||||
if "endpoint" in kwargs and kwargs["endpoint"]:
|
||||
kwargs["endpoint"] = kwargs["endpoint"].replace(model, model_name)
|
||||
kwargs["model"] = model_name
|
||||
|
||||
# Perform pre-call checks for routing strategy
|
||||
self.routing_strategy_pre_call_checks(deployment=deployment)
|
||||
|
||||
|
|
@ -3537,6 +3583,9 @@ class Router:
|
|||
"avector_store_create",
|
||||
"vector_store_search",
|
||||
"vector_store_create",
|
||||
"aocr",
|
||||
"ocr",
|
||||
"aadapter_generate_content"
|
||||
] = "assistants",
|
||||
):
|
||||
"""
|
||||
|
|
@ -3553,6 +3602,7 @@ class Router:
|
|||
"generate_content_stream",
|
||||
"vector_store_search",
|
||||
"vector_store_create",
|
||||
"ocr",
|
||||
):
|
||||
|
||||
def sync_wrapper(
|
||||
|
|
@ -3595,6 +3645,8 @@ class Router:
|
|||
"aimage_edit",
|
||||
"agenerate_content",
|
||||
"agenerate_content_stream",
|
||||
"aocr",
|
||||
"ocr",
|
||||
):
|
||||
return await self._ageneric_api_call_with_fallbacks(
|
||||
original_function=original_function,
|
||||
|
|
@ -4915,22 +4967,25 @@ class Router:
|
|||
- hash
|
||||
- use hash as id
|
||||
"""
|
||||
concat_str = model_group
|
||||
# Optimized: Use list and join instead of string concatenation in loop
|
||||
# This avoids creating many temporary string objects (O(n) vs O(n²) complexity)
|
||||
parts = [model_group]
|
||||
for k, v in litellm_params.items():
|
||||
if isinstance(k, str):
|
||||
concat_str += k
|
||||
parts.append(k)
|
||||
elif isinstance(k, dict):
|
||||
concat_str += json.dumps(k)
|
||||
parts.append(json.dumps(k))
|
||||
else:
|
||||
concat_str += str(k)
|
||||
parts.append(str(k))
|
||||
|
||||
if isinstance(v, str):
|
||||
concat_str += v
|
||||
parts.append(v)
|
||||
elif isinstance(v, dict):
|
||||
concat_str += json.dumps(v)
|
||||
parts.append(json.dumps(v))
|
||||
else:
|
||||
concat_str += str(v)
|
||||
parts.append(str(v))
|
||||
|
||||
concat_str = "".join(parts)
|
||||
hash_object = hashlib.sha256(concat_str.encode())
|
||||
|
||||
return hash_object.hexdigest()
|
||||
|
|
@ -5154,7 +5209,7 @@ class Router:
|
|||
verbose_router_logger.debug(
|
||||
f"\nInitialized Model List {self.get_model_names()}"
|
||||
)
|
||||
self.model_names = [m["model_name"] for m in model_list]
|
||||
self.model_names = {m["model_name"] for m in model_list}
|
||||
|
||||
# Build model_name index for O(1) lookups
|
||||
self._build_model_name_index(self.model_list)
|
||||
|
|
@ -5360,7 +5415,7 @@ class Router:
|
|||
self._add_model_to_list_and_index_map(
|
||||
model=_deployment, model_id=deployment.model_info.id
|
||||
)
|
||||
self.model_names.append(deployment.model_name)
|
||||
self.model_names.add(deployment.model_name)
|
||||
return deployment
|
||||
|
||||
def _update_deployment_indices_after_removal(
|
||||
|
|
@ -5519,9 +5574,15 @@ class Router:
|
|||
Returns -> Deployment or None
|
||||
|
||||
Raise Exception -> if model found in invalid format
|
||||
|
||||
Optimized with O(1) index lookup instead of O(n) linear scan.
|
||||
"""
|
||||
for model in self.model_list:
|
||||
if model["model_name"] == model_group_name:
|
||||
# O(1) lookup in model_name index
|
||||
if model_group_name in self.model_name_to_deployment_indices:
|
||||
indices = self.model_name_to_deployment_indices[model_group_name]
|
||||
if indices:
|
||||
# Return first deployment for this model_name
|
||||
model = self.model_list[indices[0]]
|
||||
if isinstance(model, dict):
|
||||
return Deployment(**model)
|
||||
elif isinstance(model, Deployment):
|
||||
|
|
@ -5631,11 +5692,13 @@ class Router:
|
|||
Returns
|
||||
- dict: the model in list with 'model_name', 'litellm_params', Optional['model_info']
|
||||
- None: could not find deployment in list
|
||||
|
||||
Optimized with O(1) index lookup instead of O(n) linear scan.
|
||||
"""
|
||||
for model in self.model_list:
|
||||
if "model_info" in model and "id" in model["model_info"]:
|
||||
if id == model["model_info"]["id"]:
|
||||
return model
|
||||
# O(1) lookup via model_id_to_deployment_index_map
|
||||
if id in self.model_id_to_deployment_index_map:
|
||||
idx = self.model_id_to_deployment_index_map[id]
|
||||
return self.model_list[idx]
|
||||
return None
|
||||
|
||||
def get_model_group(self, id: str) -> Optional[List]:
|
||||
|
|
@ -6169,17 +6232,33 @@ class Router:
|
|||
if 'model_name' is none, returns all.
|
||||
|
||||
Returns list of model id's.
|
||||
|
||||
Optimized with O(1) or O(k) index lookup when model_name provided,
|
||||
instead of O(n) linear scan.
|
||||
"""
|
||||
ids = []
|
||||
for model in self.model_list:
|
||||
if "model_info" in model and "id" in model["model_info"]:
|
||||
id = model["model_info"]["id"]
|
||||
if exclude_team_models and model["model_info"].get("team_id"):
|
||||
continue
|
||||
if model_name is not None and model["model_name"] == model_name:
|
||||
ids.append(id)
|
||||
elif model_name is None:
|
||||
ids.append(id)
|
||||
|
||||
if model_name is not None:
|
||||
# O(1) lookup in model_name index, then O(k) iteration where k = deployments for this model_name
|
||||
if model_name in self.model_name_to_deployment_indices:
|
||||
indices = self.model_name_to_deployment_indices[model_name]
|
||||
for idx in indices:
|
||||
model = self.model_list[idx]
|
||||
if "model_info" in model and "id" in model["model_info"]:
|
||||
if exclude_team_models and model["model_info"].get("team_id"):
|
||||
continue
|
||||
ids.append(model["model_info"]["id"])
|
||||
else:
|
||||
# When model_name is None, return all model IDs
|
||||
# Use the index map keys for O(n) where n = total deployments
|
||||
for model_id in self.model_id_to_deployment_index_map.keys():
|
||||
idx = self.model_id_to_deployment_index_map[model_id]
|
||||
model = self.model_list[idx]
|
||||
if "model_info" in model and "id" in model["model_info"]:
|
||||
if exclude_team_models and model["model_info"].get("team_id"):
|
||||
continue
|
||||
ids.append(model_id)
|
||||
|
||||
return ids
|
||||
|
||||
def has_model_id(self, candidate_id: str) -> bool:
|
||||
|
|
@ -6257,7 +6336,9 @@ class Router:
|
|||
model_name=model_name, model=model, team_id=team_id
|
||||
):
|
||||
if model_alias is not None:
|
||||
alias_model = copy.deepcopy(model)
|
||||
# Optimized: Use shallow copy since we only modify top-level model_name
|
||||
# This is much faster than deepcopy for nested dict structures
|
||||
alias_model = model.copy()
|
||||
alias_model["model_name"] = model_alias
|
||||
returned_models.append(alias_model)
|
||||
else:
|
||||
|
|
@ -6271,7 +6352,8 @@ class Router:
|
|||
model_name=model_name, model=model, team_id=team_id
|
||||
):
|
||||
if model_alias is not None:
|
||||
alias_model = copy.deepcopy(model)
|
||||
# Optimized: Use shallow copy since we only modify top-level model_name
|
||||
alias_model = model.copy()
|
||||
alias_model["model_name"] = model_alias
|
||||
returned_models.append(alias_model)
|
||||
else:
|
||||
|
|
@ -7070,7 +7152,7 @@ class Router:
|
|||
if isinstance(healthy_deployments, dict):
|
||||
return healthy_deployments
|
||||
|
||||
start_time = time.time()
|
||||
start_time = time.perf_counter()
|
||||
if (
|
||||
self.routing_strategy == "usage-based-routing-v2"
|
||||
and self.lowesttpm_logger_v2 is not None
|
||||
|
|
@ -7137,7 +7219,7 @@ class Router:
|
|||
f"get_available_deployment for model: {model}, Selected deployment: {self.print_deployment(deployment)} for model: {model}"
|
||||
)
|
||||
|
||||
end_time = time.time()
|
||||
end_time = time.perf_counter()
|
||||
_duration = end_time - start_time
|
||||
asyncio.create_task(
|
||||
self.service_logger_obj.async_service_success_hook(
|
||||
|
|
|
|||
|
|
@ -360,6 +360,22 @@ class PillarGuardrailConfigModel(BaseModel):
|
|||
default="monitor",
|
||||
description="Action to take when content is flagged: 'block' (raise exception) or 'monitor' (log only)",
|
||||
)
|
||||
async_mode: Optional[bool] = Field(
|
||||
default=None,
|
||||
description="Set to True to request asynchronous analysis (sets `plr_async` header). Defaults to provider behaviour when omitted.",
|
||||
)
|
||||
persist_session: Optional[bool] = Field(
|
||||
default=None,
|
||||
description="Controls Pillar session persistence (sets `plr_persist` header). Set to False to disable persistence.",
|
||||
)
|
||||
include_scanners: Optional[bool] = Field(
|
||||
default=True,
|
||||
description="Include scanner category summaries in responses (sets `plr_scanners` header).",
|
||||
)
|
||||
include_evidence: Optional[bool] = Field(
|
||||
default=True,
|
||||
description="Include detailed evidence payloads in responses (sets `plr_evidence` header).",
|
||||
)
|
||||
|
||||
|
||||
class NomaGuardrailConfigModel(BaseModel):
|
||||
|
|
|
|||
|
|
@ -15,6 +15,22 @@ class PillarGuardrailConfigModelOptionalParams(BaseModel):
|
|||
default="monitor",
|
||||
description="Action to take when content is flagged: 'block' (raise exception) or 'monitor' (log only). If not provided, the `PILLAR_ON_FLAGGED_ACTION` environment variable is checked, defaults to 'monitor'.",
|
||||
)
|
||||
async_mode: Optional[bool] = Field(
|
||||
default=None,
|
||||
description="Set to True to request asynchronous analysis (sets `plr_async` header).",
|
||||
)
|
||||
persist_session: Optional[bool] = Field(
|
||||
default=None,
|
||||
description="Set to False to disable session persistence (sets `plr_persist` header).",
|
||||
)
|
||||
include_scanners: Optional[bool] = Field(
|
||||
default=True,
|
||||
description="Include scanner summaries in response payloads (sets `plr_scanners` header).",
|
||||
)
|
||||
include_evidence: Optional[bool] = Field(
|
||||
default=True,
|
||||
description="Include detailed evidence objects in response payloads (sets `plr_evidence` header).",
|
||||
)
|
||||
|
||||
|
||||
class PillarGuardrailConfigModel(
|
||||
|
|
|
|||
|
|
@ -143,8 +143,10 @@ from litellm.litellm_core_utils.token_counter import get_modified_max_tokens
|
|||
from litellm.llms.base_llm.google_genai.transformation import (
|
||||
BaseGoogleGenAIGenerateContentConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig
|
||||
from litellm.llms.bedrock.common_utils import BedrockModelInfo
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
|
||||
from litellm.router_utils.get_retry_from_policy import (
|
||||
get_num_retries_from_retry_policy,
|
||||
reset_retry_policy,
|
||||
|
|
@ -5119,9 +5121,9 @@ def json_schema_type(python_type_name: str):
|
|||
return python_to_json_schema_types.get(python_type_name, "string")
|
||||
|
||||
|
||||
def function_to_dict(input_function): # noqa: C901
|
||||
def function_to_dict(input_function) -> dict: # noqa: C901
|
||||
"""Using type hints and numpy-styled docstring,
|
||||
produce a dictionnary usable for OpenAI function calling
|
||||
produce a dictionary usable for OpenAI function calling
|
||||
|
||||
Parameters
|
||||
----------
|
||||
|
|
@ -7211,6 +7213,8 @@ class ProviderConfigManager:
|
|||
return VolcEngineEmbeddingConfig()
|
||||
elif litellm.LlmProviders.OVHCLOUD == provider:
|
||||
return litellm.OVHCloudEmbeddingConfig()
|
||||
elif litellm.LlmProviders.COMETAPI == provider:
|
||||
return litellm.CometAPIEmbeddingConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -7518,6 +7522,12 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return get_aiml_image_generation_config(model)
|
||||
elif LlmProviders.COMETAPI == provider:
|
||||
from litellm.llms.cometapi.image_generation import (
|
||||
get_cometapi_image_generation_config,
|
||||
)
|
||||
|
||||
return get_cometapi_image_generation_config(model)
|
||||
elif LlmProviders.GEMINI == provider:
|
||||
from litellm.llms.gemini.image_generation import (
|
||||
get_gemini_image_generation_config,
|
||||
|
|
@ -7549,11 +7559,9 @@ class ProviderConfigManager:
|
|||
provider: LlmProviders,
|
||||
) -> Optional[BaseImageEditConfig]:
|
||||
if LlmProviders.OPENAI == provider:
|
||||
from litellm.llms.openai.image_edit.transformation import (
|
||||
OpenAIImageEditConfig,
|
||||
)
|
||||
from litellm.llms.openai.image_edit import get_openai_image_edit_config
|
||||
|
||||
return OpenAIImageEditConfig()
|
||||
return get_openai_image_edit_config(model=model)
|
||||
elif LlmProviders.AZURE == provider:
|
||||
from litellm.llms.azure.image_edit.transformation import (
|
||||
AzureImageEditConfig,
|
||||
|
|
@ -7578,6 +7586,25 @@ class ProviderConfigManager:
|
|||
return LiteLLMProxyImageEditConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def get_provider_ocr_config(
|
||||
model: str,
|
||||
provider: LlmProviders,
|
||||
) -> Optional["BaseOCRConfig"]:
|
||||
"""
|
||||
Get OCR configuration for a given provider.
|
||||
"""
|
||||
from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig
|
||||
|
||||
PROVIDER_TO_CONFIG_MAP = {
|
||||
litellm.LlmProviders.MISTRAL: MistralOCRConfig,
|
||||
litellm.LlmProviders.AZURE_AI: AzureAIOCRConfig,
|
||||
}
|
||||
config_class = PROVIDER_TO_CONFIG_MAP.get(provider, None)
|
||||
if config_class is None:
|
||||
return None
|
||||
return config_class()
|
||||
|
||||
@staticmethod
|
||||
def get_provider_google_genai_generate_content_config(
|
||||
model: str,
|
||||
|
|
|
|||
|
|
@ -400,6 +400,44 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"anthropic.claude-haiku-4-5@20251001": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"anthropic.claude-3-5-sonnet-20240620-v1:0": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -810,6 +848,25 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"apac.anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"apac.anthropic.claude-3-sonnet-20240229-v1:0": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -4612,6 +4669,48 @@
|
|||
"supports_web_search": true,
|
||||
"tool_use_system_prompt_tokens": 264
|
||||
},
|
||||
"claude-haiku-4-5-20251001": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 2e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "anthropic",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 64000,
|
||||
"max_tokens": 64000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"claude-haiku-4-5": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 2e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "anthropic",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 64000,
|
||||
"max_tokens": 64000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"claude-3-5-sonnet-20240620": {
|
||||
"cache_creation_input_token_cost": 3.75e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 6e-06,
|
||||
|
|
@ -7741,6 +7840,25 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"eu.anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"eu.anthropic.claude-3-5-sonnet-20240620-v1:0": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -9486,6 +9604,54 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"gemini-2.5-flash-image": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.039,
|
||||
"output_cost_per_reasoning_token": 2.5e-06,
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"rpm": 100000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-flash-image",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supports_audio_output": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"tpm": 8000000
|
||||
},
|
||||
"gemini-2.5-flash-image-preview": {
|
||||
"cache_read_input_token_cost": 7.5e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
|
|
@ -10939,6 +11105,54 @@
|
|||
"supports_web_search": true,
|
||||
"tpm": 8000000
|
||||
},
|
||||
"gemini/gemini-2.5-flash-image": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.039,
|
||||
"output_cost_per_reasoning_token": 2.5e-06,
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"rpm": 100000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-flash-image",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supports_audio_output": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"tpm": 8000000
|
||||
},
|
||||
"gemini/gemini-2.5-flash-image-preview": {
|
||||
"cache_read_input_token_cost": 7.5e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
|
|
@ -13197,11 +13411,11 @@
|
|||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": false,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": false,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": false,
|
||||
"supports_tool_choice": true,
|
||||
|
|
@ -13650,8 +13864,56 @@
|
|||
"lemonade/Qwen3-Coder-30B-A3B-Instruct-GGUF": {
|
||||
"input_cost_per_token": 0,
|
||||
"litellm_provider": "lemonade",
|
||||
"max_tokens": 32768,
|
||||
"max_input_tokens": 32768,
|
||||
"max_tokens": 262144,
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"lemonade/gpt-oss-20b-mxfp4-GGUF": {
|
||||
"input_cost_per_token": 0,
|
||||
"litellm_provider": "lemonade",
|
||||
"max_tokens": 131072,
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"lemonade/gpt-oss-120b-mxfp-GGUF": {
|
||||
"input_cost_per_token": 0,
|
||||
"litellm_provider": "lemonade",
|
||||
"max_tokens": 131072,
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"lemonade/Gemma-3-4b-it-GGUF": {
|
||||
"input_cost_per_token": 0,
|
||||
"litellm_provider": "lemonade",
|
||||
"max_tokens": 128000,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"lemonade/Qwen3-4B-Instruct-2507-GGUF": {
|
||||
"input_cost_per_token": 0,
|
||||
"litellm_provider": "lemonade",
|
||||
"max_tokens": 262144,
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0,
|
||||
|
|
@ -14466,6 +14728,25 @@
|
|||
"supports_vision": true,
|
||||
"tool_use_system_prompt_tokens": 346
|
||||
},
|
||||
"jp.anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"lambda_ai/deepseek-llama3.3-70b": {
|
||||
"input_cost_per_token": 2e-07,
|
||||
"litellm_provider": "lambda_ai",
|
||||
|
|
@ -17153,6 +17434,8 @@
|
|||
},
|
||||
"openrouter/anthropic/claude-opus-4": {
|
||||
"input_cost_per_image": 0.0048,
|
||||
"cache_creation_input_token_cost": 1.875e-05,
|
||||
"cache_read_input_token_cost": 1.5e-06,
|
||||
"input_cost_per_token": 1.5e-05,
|
||||
"litellm_provider": "openrouter",
|
||||
"max_input_tokens": 200000,
|
||||
|
|
@ -17163,6 +17446,7 @@
|
|||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -17170,6 +17454,9 @@
|
|||
},
|
||||
"openrouter/anthropic/claude-opus-4.1": {
|
||||
"input_cost_per_image": 0.0048,
|
||||
"cache_creation_input_token_cost": 1.875e-05,
|
||||
"cache_creation_input_token_cost_above_1hr": 3e-05,
|
||||
"cache_read_input_token_cost": 1.5e-06,
|
||||
"input_cost_per_token": 1.5e-05,
|
||||
"litellm_provider": "openrouter",
|
||||
"max_input_tokens": 200000,
|
||||
|
|
@ -17180,6 +17467,7 @@
|
|||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -17187,6 +17475,10 @@
|
|||
},
|
||||
"openrouter/anthropic/claude-sonnet-4": {
|
||||
"input_cost_per_image": 0.0048,
|
||||
"cache_creation_input_token_cost": 3.75e-06,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
|
||||
"cache_read_input_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
|
||||
"input_cost_per_token": 3e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 6e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 2.25e-05,
|
||||
|
|
@ -17199,6 +17491,31 @@
|
|||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"tool_use_system_prompt_tokens": 159
|
||||
},
|
||||
"openrouter/anthropic/claude-sonnet-4.5": {
|
||||
"input_cost_per_image": 0.0048,
|
||||
"cache_creation_input_token_cost": 3.75e-06,
|
||||
"cache_read_input_token_cost": 3e-07,
|
||||
"input_cost_per_token": 3e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 6e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 2.25e-05,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
|
||||
"litellm_provider": "openrouter",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -20097,6 +20414,25 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"us.anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"us.anthropic.claude-3-5-sonnet-20240620-v1:0": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -21368,6 +21704,25 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"vertex_ai/claude-haiku-4-5@20251001": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "vertex_ai-anthropic_models",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-06,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/haiku-4-5",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"vertex_ai/claude-3-5-sonnet": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"litellm_provider": "vertex_ai-anthropic_models",
|
||||
|
|
@ -21560,8 +21915,8 @@
|
|||
"input_cost_per_token_batches": 7.5e-06,
|
||||
"litellm_provider": "vertex_ai-anthropic_models",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"max_output_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 7.5e-05,
|
||||
"output_cost_per_token_batches": 3.75e-05,
|
||||
|
|
@ -21577,8 +21932,8 @@
|
|||
"input_cost_per_token_batches": 7.5e-06,
|
||||
"litellm_provider": "vertex_ai-anthropic_models",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"max_output_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 7.5e-05,
|
||||
"output_cost_per_token_batches": 3.75e-05,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[tool.poetry]
|
||||
name = "litellm"
|
||||
version = "1.78.1"
|
||||
version = "1.78.3"
|
||||
description = "Library to easily interface with LLM API providers"
|
||||
authors = ["BerriAI"]
|
||||
license = "MIT"
|
||||
|
|
@ -157,7 +157,7 @@ requires = ["poetry-core", "wheel"]
|
|||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "1.78.1"
|
||||
version = "1.78.3"
|
||||
version_files = [
|
||||
"pyproject.toml:^version"
|
||||
]
|
||||
|
|
|
|||
Binary file not shown.
|
Before Width: | Height: | Size: 1.5 MiB After Width: | Height: | Size: 1.9 MiB |
|
|
@ -9,6 +9,7 @@ import base64
|
|||
from io import BytesIO
|
||||
from unittest.mock import patch, AsyncMock
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
|
|
@ -30,6 +31,72 @@ class TestCustomLogger(CustomLogger):
|
|||
self.standard_logging_payload = kwargs.get("standard_logging_object", None)
|
||||
pass
|
||||
|
||||
|
||||
class BaseLLMImageEditTest(ABC):
|
||||
"""
|
||||
Abstract base test class that enforces a common test across all image edit test classes.
|
||||
"""
|
||||
|
||||
@property
|
||||
def image_edit_function(self):
|
||||
return litellm.image_edit
|
||||
|
||||
@property
|
||||
def async_image_edit_function(self):
|
||||
return litellm.aimage_edit
|
||||
|
||||
@abstractmethod
|
||||
def get_base_image_edit_call_args(self) -> dict:
|
||||
"""Must return the base image edit call args"""
|
||||
pass
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _handle_rate_limits(self):
|
||||
"""Fixture to handle rate limit errors for all test methods"""
|
||||
try:
|
||||
yield
|
||||
except litellm.RateLimitError:
|
||||
pytest.skip("Rate limit exceeded")
|
||||
except litellm.InternalServerError:
|
||||
pytest.skip("Model is overloaded")
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.flaky(retries=3, delay=2)
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_image_edit_litellm_sdk(self, sync_mode):
|
||||
"""
|
||||
Test image edit functionality with both sync and async modes.
|
||||
"""
|
||||
litellm._turn_on_debug()
|
||||
try:
|
||||
prompt = """
|
||||
Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO.
|
||||
"""
|
||||
|
||||
call_args = self.get_base_image_edit_call_args()
|
||||
call_args["prompt"] = prompt
|
||||
|
||||
if sync_mode:
|
||||
result = self.image_edit_function(**call_args)
|
||||
else:
|
||||
result = await self.async_image_edit_function(**call_args)
|
||||
|
||||
print("result from image edit", result)
|
||||
|
||||
# Validate the response meets expected schema
|
||||
ImageResponse.model_validate(result)
|
||||
|
||||
if isinstance(result, ImageResponse) and result.data:
|
||||
image_base64 = result.data[0].b64_json
|
||||
if image_base64:
|
||||
image_bytes = base64.b64decode(image_base64)
|
||||
|
||||
# Save the image to a file
|
||||
with open("test_image_edit.png", "wb") as f:
|
||||
f.write(image_bytes)
|
||||
except litellm.ContentPolicyViolationError as e:
|
||||
pass
|
||||
|
||||
# Get the current directory of the file being run
|
||||
pwd = os.path.dirname(os.path.realpath(__file__))
|
||||
|
||||
|
|
@ -49,45 +116,31 @@ def get_test_images_as_bytesio():
|
|||
bytesio_images.append(BytesIO(image_bytes))
|
||||
return bytesio_images
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.flaky(retries=3, delay=2)
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_image_edit_litellm_sdk(sync_mode):
|
||||
from litellm import image_edit, aimage_edit
|
||||
litellm._turn_on_debug()
|
||||
try:
|
||||
prompt = """
|
||||
Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO.
|
||||
"""
|
||||
|
||||
if sync_mode:
|
||||
result = image_edit(
|
||||
prompt=prompt,
|
||||
model="gpt-image-1",
|
||||
image=TEST_IMAGES,
|
||||
)
|
||||
else:
|
||||
result = await aimage_edit(
|
||||
prompt=prompt,
|
||||
model="gpt-image-1",
|
||||
image=TEST_IMAGES,
|
||||
)
|
||||
print("result from image edit", result)
|
||||
class TestOpenAIImageEditGPTImage1(BaseLLMImageEditTest):
|
||||
"""
|
||||
Concrete implementation of BaseLLMImageEditTest for OpenAI image edits.
|
||||
"""
|
||||
|
||||
# Validate the response meets expected schema
|
||||
ImageResponse.model_validate(result)
|
||||
|
||||
if isinstance(result, ImageResponse) and result.data:
|
||||
image_base64 = result.data[0].b64_json
|
||||
if image_base64:
|
||||
image_bytes = base64.b64decode(image_base64)
|
||||
def get_base_image_edit_call_args(self) -> dict:
|
||||
"""Return base call args for OpenAI image edit"""
|
||||
return {
|
||||
"model": "gpt-image-1",
|
||||
"image": TEST_IMAGES,
|
||||
}
|
||||
|
||||
# Save the image to a file
|
||||
with open("test_image_edit.png", "wb") as f:
|
||||
f.write(image_bytes)
|
||||
except litellm.ContentPolicyViolationError as e:
|
||||
pass
|
||||
class TestOpenAIImageEditDallE2(BaseLLMImageEditTest):
|
||||
"""
|
||||
Concrete implementation of BaseLLMImageEditTest for OpenAI DALL-E-2 image edits.
|
||||
DALL-E-2 only supports a single image (not an array).
|
||||
"""
|
||||
|
||||
def get_base_image_edit_call_args(self) -> dict:
|
||||
"""Return base call args for OpenAI DALL-E-2 image edit (single image only)"""
|
||||
return {
|
||||
"model": "dall-e-2",
|
||||
"image": SINGLE_TEST_IMAGE,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.flaky(retries=3, delay=2)
|
||||
|
|
|
|||
|
|
@ -3232,6 +3232,60 @@ async def test_bedrock_passthrough(sync_mode: bool):
|
|||
assert response.status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_passthrough_router():
|
||||
"""
|
||||
Test bedrock passthrough using litellm.Router with async mode.
|
||||
Tests that the router:
|
||||
1. Resolves the router model name to the actual deployment
|
||||
2. Replaces the router model name in the endpoint with the actual deployment model
|
||||
"""
|
||||
import litellm
|
||||
from litellm import Router
|
||||
|
||||
litellm._turn_on_debug()
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "special-bedrock-model",
|
||||
"litellm_params": {
|
||||
"model": "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
data = {
|
||||
"max_tokens": 512,
|
||||
"messages": [{"role": "user", "content": "Hey"}],
|
||||
"system": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Analyze if this message indicates a new conversation topic. If it does, extract a 2-3 word title that captures the new topic. Format your response as a JSON object with two fields: 'isNewTopic' (boolean) and 'title' (string, or null if isNewTopic is false). Only include these fields, no other text.",
|
||||
}
|
||||
],
|
||||
"temperature": 0,
|
||||
"metadata": {
|
||||
"user_id": "5dd07c33da27e6d2968d94ea20bf47a7b090b6b158b82328d54da2909a108e84"
|
||||
},
|
||||
"anthropic_version": "bedrock-2023-05-31",
|
||||
"anthropic_beta": ["claude-code-20250219"],
|
||||
}
|
||||
|
||||
# Endpoint uses the router model name which should be replaced with actual deployment
|
||||
response = await router.allm_passthrough_route(
|
||||
model="special-bedrock-model",
|
||||
method="POST",
|
||||
endpoint="/model/special-bedrock-model/invoke",
|
||||
data=data,
|
||||
)
|
||||
|
||||
print(response.text)
|
||||
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_converse__streaming_passthrough(monkeypatch):
|
||||
import litellm
|
||||
|
|
|
|||
140
tests/ocr_tests/base_ocr_unit_tests.py
Normal file
140
tests/ocr_tests/base_ocr_unit_tests.py
Normal file
|
|
@ -0,0 +1,140 @@
|
|||
"""
|
||||
Base test class for OCR functionality across different providers.
|
||||
|
||||
This follows the same pattern as BaseLLMChatTest in tests/llm_translation/base_llm_unit_tests.py
|
||||
"""
|
||||
import pytest
|
||||
import litellm
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
|
||||
# Test resources
|
||||
TEST_IMAGE_PATH = "test_image_edit.png"
|
||||
TEST_PDF_URL = "https://arxiv.org/pdf/2201.04234"
|
||||
|
||||
|
||||
class BaseOCRTest(ABC):
|
||||
"""
|
||||
Abstract base test class that enforces common OCR tests across all providers.
|
||||
|
||||
Each provider-specific test class should inherit from this and implement
|
||||
get_base_ocr_call_args() to return provider-specific configuration.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def get_base_ocr_call_args(self) -> dict:
|
||||
"""Must return the base OCR call args for the specific provider"""
|
||||
pass
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _handle_rate_limits(self):
|
||||
"""Fixture to handle rate limit errors for all test methods"""
|
||||
try:
|
||||
yield
|
||||
except litellm.RateLimitError:
|
||||
pytest.skip("Rate limit exceeded")
|
||||
except litellm.InternalServerError:
|
||||
pytest.skip("Model is overloaded")
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_basic_ocr_with_url(self, sync_mode):
|
||||
"""
|
||||
Test basic OCR with a public URL.
|
||||
"""
|
||||
litellm._turn_on_debug()
|
||||
base_ocr_call_args = self.get_base_ocr_call_args()
|
||||
print("BASE OCR Call args=", base_ocr_call_args)
|
||||
|
||||
try:
|
||||
if sync_mode:
|
||||
response = litellm.ocr(
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": TEST_PDF_URL
|
||||
},
|
||||
**base_ocr_call_args,
|
||||
)
|
||||
else:
|
||||
response = await litellm.aocr(
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": TEST_PDF_URL
|
||||
},
|
||||
**base_ocr_call_args,
|
||||
)
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print(f"Sync Mode: {sync_mode}")
|
||||
print(f"Response type: {type(response)}")
|
||||
print(f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}")
|
||||
|
||||
# Check if response has expected OCR format
|
||||
assert hasattr(response, "pages"), "Response should have 'pages' attribute"
|
||||
assert hasattr(response, "model"), "Response should have 'model' attribute"
|
||||
assert hasattr(response, "object"), "Response should have 'object' attribute"
|
||||
assert response.object == "ocr", f"Expected object='ocr', got '{response.object}'"
|
||||
|
||||
# Validate pages structure
|
||||
assert isinstance(response.pages, list), "pages should be a list"
|
||||
assert len(response.pages) > 0, "Should have at least one page"
|
||||
|
||||
# Check first page structure
|
||||
first_page = response.pages[0]
|
||||
assert hasattr(first_page, "index"), "Page should have 'index' attribute"
|
||||
assert hasattr(first_page, "markdown"), "Page should have 'markdown' attribute"
|
||||
|
||||
# Extract text from all pages for validation
|
||||
total_text = "\n\n".join(page.markdown for page in response.pages if page.markdown)
|
||||
print(f"Total pages: {len(response.pages)}")
|
||||
print(f"Total extracted text length: {len(total_text)} characters")
|
||||
print(f"First 200 chars: {total_text[:200]}")
|
||||
print(f"Model: {response.model}")
|
||||
if response.usage_info:
|
||||
print(f"Pages processed: {response.usage_info.pages_processed}")
|
||||
print(f"{'='*80}\n")
|
||||
|
||||
assert len(total_text) > 0, "Should extract some text from the document"
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(f"OCR call failed: {str(e)}")
|
||||
|
||||
def test_ocr_response_structure(self):
|
||||
"""
|
||||
Test that the OCR response has the correct structure.
|
||||
"""
|
||||
litellm.set_verbose = True
|
||||
base_ocr_call_args = self.get_base_ocr_call_args()
|
||||
|
||||
response = litellm.ocr(
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": TEST_PDF_URL
|
||||
},
|
||||
**base_ocr_call_args,
|
||||
)
|
||||
|
||||
# Validate response structure
|
||||
assert hasattr(response, "pages"), "Response should have 'pages' attribute"
|
||||
assert hasattr(response, "model"), "Response should have 'model' attribute"
|
||||
assert hasattr(response, "object"), "Response should have 'object' attribute"
|
||||
assert hasattr(response, "usage_info"), "Response should have 'usage_info' attribute"
|
||||
|
||||
assert isinstance(response.pages, list), "pages should be a list"
|
||||
assert len(response.pages) > 0, "Should have at least one page"
|
||||
assert response.object == "ocr", "object should be 'ocr'"
|
||||
|
||||
# Validate first page structure
|
||||
first_page = response.pages[0]
|
||||
assert hasattr(first_page, "index"), "Page should have 'index' attribute"
|
||||
assert hasattr(first_page, "markdown"), "Page should have 'markdown' attribute"
|
||||
assert isinstance(first_page.markdown, str), "markdown should be a string"
|
||||
|
||||
print(f"\nResponse structure validated:")
|
||||
print(f" - object: {response.object}")
|
||||
print(f" - model: {response.model}")
|
||||
print(f" - pages: {len(response.pages)}")
|
||||
if response.usage_info:
|
||||
print(f" - pages_processed: {response.usage_info.pages_processed}")
|
||||
print(f" - doc_size_bytes: {response.usage_info.doc_size_bytes}")
|
||||
|
||||
27
tests/ocr_tests/test_ocr_azure_ai.py
Normal file
27
tests/ocr_tests/test_ocr_azure_ai.py
Normal file
|
|
@ -0,0 +1,27 @@
|
|||
"""
|
||||
Test OCR functionality with Azure AI API.
|
||||
|
||||
Note: Azure AI OCR automatically converts URLs to base64 data URIs since
|
||||
the Azure AI endpoint doesn't have internet access.
|
||||
"""
|
||||
import os
|
||||
from base_ocr_unit_tests import BaseOCRTest
|
||||
|
||||
class TestAzureAIOCR(BaseOCRTest):
|
||||
"""
|
||||
Test class for Azure AI OCR functionality.
|
||||
Inherits from BaseOCRTest and provides Azure AI-specific configuration.
|
||||
|
||||
Note: For Azure AI, LiteLLM will automatically convert URLs to base64 data URIs before
|
||||
sending to the API, since Azure AI OCR endpoint doesn't have internet access.
|
||||
"""
|
||||
|
||||
def get_base_ocr_call_args(self) -> dict:
|
||||
"""
|
||||
Return the base OCR call args for Azure AI.
|
||||
"""
|
||||
return {
|
||||
"model": "azure_ai/mistral-document-ai-2505",
|
||||
"api_key": os.getenv("AZURE_AI_API_KEY_MISTRAL"),
|
||||
"api_base": os.getenv("AZURE_AI_API_BASE_MISTRAL"),
|
||||
}
|
||||
88
tests/ocr_tests/test_ocr_mistral.py
Normal file
88
tests/ocr_tests/test_ocr_mistral.py
Normal file
|
|
@ -0,0 +1,88 @@
|
|||
"""
|
||||
Test OCR functionality with Mistral API.
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import pytest
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from base_ocr_unit_tests import BaseOCRTest, TEST_PDF_URL
|
||||
|
||||
|
||||
class TestMistralOCR(BaseOCRTest):
|
||||
"""
|
||||
Test class for Mistral OCR functionality.
|
||||
"""
|
||||
|
||||
def get_base_ocr_call_args(self) -> dict:
|
||||
"""Return the base OCR call args for Mistral"""
|
||||
return {
|
||||
"model": "mistral/mistral-ocr-latest",
|
||||
"api_key": os.getenv("MISTRAL_API_KEY"),
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_aocr_with_mistral():
|
||||
"""
|
||||
Test OCR with Router using Mistral OCR deployment.
|
||||
"""
|
||||
litellm.set_verbose = True
|
||||
|
||||
# Create router with Mistral OCR deployment
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "mistral-ocr",
|
||||
"litellm_params": {
|
||||
"model": "mistral/mistral-ocr-latest",
|
||||
"api_key": os.getenv("MISTRAL_API_KEY"),
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
try:
|
||||
# Call OCR through router
|
||||
response = await router.aocr(
|
||||
model="mistral-ocr",
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": TEST_PDF_URL
|
||||
},
|
||||
)
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print("Router OCR Test")
|
||||
print(f"Response type: {type(response)}")
|
||||
print(f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}")
|
||||
|
||||
# Check if response has expected Mistral OCR format
|
||||
assert hasattr(response, "pages"), "Response should have 'pages' attribute"
|
||||
assert hasattr(response, "model"), "Response should have 'model' attribute"
|
||||
assert hasattr(response, "object"), "Response should have 'object' attribute"
|
||||
assert response.object == "ocr", f"Expected object='ocr', got '{response.object}'"
|
||||
|
||||
# Validate pages structure
|
||||
assert isinstance(response.pages, list), "pages should be a list"
|
||||
assert len(response.pages) > 0, "Should have at least one page"
|
||||
|
||||
# Check first page structure
|
||||
first_page = response.pages[0]
|
||||
assert hasattr(first_page, "index"), "Page should have 'index' attribute"
|
||||
assert hasattr(first_page, "markdown"), "Page should have 'markdown' attribute"
|
||||
|
||||
# Extract text from all pages for validation
|
||||
total_text = "\n\n".join(page.markdown for page in response.pages if page.markdown)
|
||||
print(f"Total pages: {len(response.pages)}")
|
||||
print(f"Total extracted text length: {len(total_text)} characters")
|
||||
print(f"First 200 chars: {total_text[:200]}")
|
||||
print(f"Model: {response.model}")
|
||||
if response.usage_info:
|
||||
print(f"Pages processed: {response.usage_info.pages_processed}")
|
||||
print(f"{'='*80}\n")
|
||||
|
||||
assert len(total_text) > 0, "Should extract some text from the document"
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(f"Router OCR call failed: {str(e)}")
|
||||
|
||||
|
|
@ -1411,7 +1411,8 @@ def test_generate_model_id_with_deployment_model_name(model_list):
|
|||
"Expected TypeError when model_group is None - this confirms our fix is needed"
|
||||
)
|
||||
except TypeError as e:
|
||||
assert "unsupported operand type(s) for +=" in str(e)
|
||||
# After optimization, error message changed but still fails appropriately on None
|
||||
assert "unsupported operand type(s) for +=" in str(e) or "expected str instance, NoneType found" in str(e)
|
||||
print(f"✓ Correctly failed with None model_group (as expected): {e}")
|
||||
except Exception as e:
|
||||
pytest.fail(f"Unexpected error with None model_group: {e}")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
import sys
|
||||
import os
|
||||
import pytest
|
||||
import ast
|
||||
import ast
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
|
|
@ -177,3 +179,97 @@ class TestRouterIndexManagement:
|
|||
# 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]
|
||||
|
||||
def test_no_linear_scans_in_router(self):
|
||||
"""
|
||||
Static analysis test to ensure Router doesn't use O(n) linear scans.
|
||||
|
||||
Scans router.py for 'in self.model_list' pattern which indicates
|
||||
inefficient O(n) iteration instead of using index-based O(1) lookups.
|
||||
|
||||
Methods should use:
|
||||
- model_id_to_deployment_index_map for O(1) model_id lookups
|
||||
- model_name_to_deployment_indices for O(1) + O(k) model_name lookups
|
||||
"""
|
||||
# Methods that are allowed to iterate through self.model_list
|
||||
ALLOWED_METHODS = [
|
||||
"_get_deployment_by_litellm_model", # Edge case: lookup by litellm_params.model (not indexed)
|
||||
]
|
||||
|
||||
# Get path to router.py
|
||||
router_file = os.path.join(
|
||||
os.path.dirname(os.path.dirname(os.path.dirname(__file__))),
|
||||
"litellm",
|
||||
"router.py"
|
||||
)
|
||||
|
||||
# Read the file
|
||||
with open(router_file, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Parse with AST
|
||||
tree = ast.parse(content)
|
||||
|
||||
# Find violations
|
||||
violations = []
|
||||
ignore_methods = set(ALLOWED_METHODS)
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.FunctionDef):
|
||||
method_name = node.name
|
||||
|
||||
# Skip ignored methods
|
||||
if method_name in ignore_methods:
|
||||
continue
|
||||
|
||||
# Get source for this method
|
||||
try:
|
||||
method_source = ast.get_source_segment(content, node)
|
||||
if not method_source:
|
||||
continue
|
||||
|
||||
# Check for the anti-pattern: "in self.model_list"
|
||||
# This catches: for x in self.model_list, if x in self.model_list, etc.
|
||||
if "in self.model_list" in method_source:
|
||||
# Extract the specific line for better error reporting
|
||||
lines = method_source.split('\n')
|
||||
pattern_line = None
|
||||
for line in lines:
|
||||
if "in self.model_list" in line:
|
||||
pattern_line = line.strip()
|
||||
break
|
||||
|
||||
violations.append({
|
||||
"method": method_name,
|
||||
"line": node.lineno,
|
||||
"pattern": pattern_line or "in self.model_list"
|
||||
})
|
||||
except Exception:
|
||||
# Skip if we can't get source segment
|
||||
pass
|
||||
|
||||
# Assert no violations
|
||||
if violations:
|
||||
error_msg = "\n".join([
|
||||
f" - {v['method']}() at line {v['line']}: {v['pattern']}"
|
||||
for v in violations
|
||||
])
|
||||
|
||||
pytest.fail(
|
||||
f"\n{'='*70}\n"
|
||||
f"Found O(n) linear scan pattern in router.py:\n\n"
|
||||
f"{error_msg}\n\n"
|
||||
f"These methods should use index maps instead:\n"
|
||||
f" - model_id_to_deployment_index_map (for model_id lookups)\n"
|
||||
f" - model_name_to_deployment_indices (for model_name lookups)\n\n"
|
||||
f"If a method legitimately needs O(n) iteration, add it to\n"
|
||||
f"ALLOWED_METHODS in this test method.\n"
|
||||
f"{'='*70}\n"
|
||||
)
|
||||
def test_model_names_is_set(self):
|
||||
"""Verify that model_names uses a set for O(1) lookups, not a list (O(n))"""
|
||||
router = Router(model_list=[])
|
||||
|
||||
assert isinstance(router.model_names, set), (
|
||||
f"model_names should be a set for O(1) lookups, but got {type(router.model_names)}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -649,5 +649,5 @@ def test_bedrock_anthropic_prompt_caching():
|
|||
|
||||
assert prompt_cost >= 0
|
||||
assert completion_cost >= 0
|
||||
assert round(prompt_cost, 3) == 0.845
|
||||
assert round(prompt_cost, 3) == 0.111
|
||||
assert round(completion_cost, 5) == 0.00820
|
||||
|
|
|
|||
|
|
@ -399,3 +399,41 @@ async def test_session_validation():
|
|||
mock_valid_session = MockClientSession()
|
||||
transport3 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_valid_session) # type: ignore
|
||||
assert transport3.client is mock_valid_session # Should reuse session
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"env_curve,litellm_curve,expected_curve,should_call",
|
||||
[
|
||||
# env_curve: SSL_ECDH_CURVE env var | litellm_curve: litellm.ssl_ecdh_curve variable
|
||||
# expected_curve: curve that should be set | should_call: whether set_ecdh_curve() should be called
|
||||
|
||||
# Valid configurations
|
||||
("X25519", None, "X25519", True), # Env var only
|
||||
("prime256v1", None, "prime256v1", True), # Different valid curve
|
||||
(None, "secp384r1", "secp384r1", True), # litellm variable only
|
||||
("X25519", "secp521r1", "X25519", True), # Env var takes precedence
|
||||
# Empty/None configurations - should skip
|
||||
("", None, None, False), # Empty string - skip configuration
|
||||
(None, None, None, False), # None value - skip configuration
|
||||
]
|
||||
)
|
||||
def test_ssl_ecdh_curve(env_curve, litellm_curve, expected_curve, should_call, monkeypatch):
|
||||
"""Test SSL ECDH curve configuration with valid curves and precedence"""
|
||||
with patch.dict(os.environ, clear=True):
|
||||
if env_curve:
|
||||
monkeypatch.setenv("SSL_ECDH_CURVE", env_curve)
|
||||
|
||||
original_value = litellm.ssl_ecdh_curve
|
||||
try:
|
||||
litellm.ssl_ecdh_curve = litellm_curve
|
||||
|
||||
with patch.object(ssl.SSLContext, 'set_ecdh_curve') as mock_set_curve:
|
||||
ssl_context = get_ssl_configuration()
|
||||
|
||||
if should_call:
|
||||
mock_set_curve.assert_called_once_with(expected_curve)
|
||||
else:
|
||||
mock_set_curve.assert_not_called()
|
||||
assert isinstance(ssl_context, ssl.SSLContext)
|
||||
finally:
|
||||
litellm.ssl_ecdh_curve = original_value
|
||||
|
|
|
|||
|
|
@ -8,8 +8,12 @@ and following LiteLLM testing patterns and best practices.
|
|||
# Standard library imports
|
||||
import os
|
||||
import sys
|
||||
from typing import Dict
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
# Add parent directory to path for imports
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
# Third-party imports
|
||||
import pytest
|
||||
from fastapi.exceptions import HTTPException
|
||||
|
|
@ -26,9 +30,6 @@ from litellm.proxy.guardrails.guardrail_hooks.pillar import (
|
|||
)
|
||||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
|
||||
# Add parent directory to path for imports
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# FIXTURES
|
||||
|
|
@ -221,6 +222,18 @@ def mock_llm_response():
|
|||
return mock_response
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def pillar_async_response():
|
||||
"""Fixture providing an asynchronous Pillar API queue response."""
|
||||
return Response(
|
||||
json={"status": "queued", "session_id": "async-session", "position": 1},
|
||||
status_code=202,
|
||||
request=Request(
|
||||
method="POST", url="https://api.pillar.security/api/v1/protect"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_llm_response_with_tools():
|
||||
"""Fixture providing a mock LLM response with tool calls."""
|
||||
|
|
@ -440,6 +453,55 @@ async def test_post_call_hook_with_tool_calls(
|
|||
assert result == mock_llm_response_with_tools
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# HEADER CONFIGURATION TESTS
|
||||
# =========================================================================
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_hook_custom_header_overrides(
|
||||
sample_request_data,
|
||||
user_api_key_dict,
|
||||
dual_cache,
|
||||
pillar_async_response,
|
||||
):
|
||||
"""Ensure configuration values translate into correct Protect headers."""
|
||||
|
||||
guardrail = PillarGuardrail(
|
||||
guardrail_name="pillar-header-test",
|
||||
api_key="test-pillar-key",
|
||||
api_base="https://api.pillar.security",
|
||||
on_flagged_action="monitor",
|
||||
persist_session=False,
|
||||
async_mode=True,
|
||||
include_scanners=False,
|
||||
include_evidence=False,
|
||||
)
|
||||
|
||||
captured_headers: Dict[str, str] = {}
|
||||
|
||||
async def _mock_post(*args, **kwargs):
|
||||
captured_headers.update(kwargs.get("headers", {}))
|
||||
return pillar_async_response
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new=_mock_post,
|
||||
):
|
||||
result = await guardrail.async_pre_call_hook(
|
||||
data=sample_request_data,
|
||||
cache=dual_cache,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert result == sample_request_data
|
||||
assert captured_headers.get("plr_persist") == "false"
|
||||
assert captured_headers.get("plr_async") == "true"
|
||||
assert captured_headers.get("plr_scanners") == "false"
|
||||
assert captured_headers.get("plr_evidence") == "false"
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# EDGE CASE TESTS
|
||||
# ============================================================================
|
||||
|
|
@ -0,0 +1,55 @@
|
|||
"""
|
||||
Unit tests for EntraID app roles JWT claim extraction.
|
||||
|
||||
This module tests the get_app_roles_from_id_token method to ensure it correctly
|
||||
extracts app roles from Microsoft EntraID JWT tokens and prevents regressions.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import jwt
|
||||
|
||||
from litellm.proxy.management_endpoints.ui_sso import MicrosoftSSOHandler
|
||||
|
||||
|
||||
class TestEntraIDAppRoles:
|
||||
"""Test EntraID app roles extraction from JWT tokens"""
|
||||
|
||||
def test_get_app_roles_from_id_token_works_without_roles(self):
|
||||
"""Test that JWT token works fine without app_roles claim"""
|
||||
# Arrange - Token without app_roles (normal user)
|
||||
payload = {
|
||||
"sub": "user123",
|
||||
"email": "user@company.com",
|
||||
"aud": "litellm-app",
|
||||
"iss": "https://login.microsoftonline.com/tenant-id/v2.0",
|
||||
"exp": 9999999999,
|
||||
}
|
||||
no_roles_token = jwt.encode(payload, "secret", algorithm="HS256")
|
||||
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token(no_roles_token)
|
||||
|
||||
# Assert - Should return empty list, not error
|
||||
assert result == []
|
||||
assert len(result) == 0
|
||||
|
||||
def test_get_app_roles_from_id_token_assigns_roles_when_present(self):
|
||||
"""Test that valid app roles are properly assigned when present"""
|
||||
# Arrange - Token with valid roles
|
||||
payload = {
|
||||
"sub": "user123",
|
||||
"email": "admin@company.com",
|
||||
"app_roles": ["proxy_admin"],
|
||||
"aud": "litellm-app",
|
||||
"iss": "https://login.microsoftonline.com/tenant-id/v2.0",
|
||||
"exp": 9999999999,
|
||||
}
|
||||
valid_roles_token = jwt.encode(payload, "secret", algorithm="HS256")
|
||||
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token(valid_roles_token)
|
||||
|
||||
# Assert - Should extract the role
|
||||
assert result == ["proxy_admin"]
|
||||
assert len(result) == 1
|
||||
assert "proxy_admin" in result
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -18,12 +18,12 @@ import litellm
|
|||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
BaseOpenAIPassThroughHandler,
|
||||
RouteChecks,
|
||||
bedrock_llm_proxy_route,
|
||||
create_pass_through_route,
|
||||
llm_passthrough_factory_proxy_route,
|
||||
vllm_proxy_route,
|
||||
vertex_discovery_proxy_route,
|
||||
vertex_proxy_route,
|
||||
bedrock_llm_proxy_route,
|
||||
vllm_proxy_route,
|
||||
)
|
||||
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
|
||||
|
||||
|
|
@ -996,6 +996,64 @@ class TestBedrockLLMProxyRoute:
|
|||
assert call_kwargs["model"] == "anthropic.claude-3-sonnet-20240229-v1:0"
|
||||
assert result == "success"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_error_handling_returns_actual_error(self):
|
||||
"""
|
||||
Test that when Bedrock API returns an error, it is properly propagated to the user
|
||||
instead of being returned as a generic "Internal Server Error".
|
||||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
handle_bedrock_passthrough_router_model,
|
||||
)
|
||||
|
||||
mock_request = Mock()
|
||||
mock_request.method = "POST"
|
||||
mock_request.headers = {"content-type": "application/json"}
|
||||
mock_request.query_params = {}
|
||||
|
||||
mock_request_body = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"textaaa": "Hello"}]
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
bedrock_error_message = '{"message":"ContentBlock object at messages.0.content.0 must set one of the following keys: text, image, toolUse, toolResult, document, video."}'
|
||||
|
||||
# Create a mock httpx.Response for the error
|
||||
mock_error_response = Mock(spec=httpx.Response)
|
||||
mock_error_response.status_code = 400
|
||||
mock_error_response.aread = AsyncMock(return_value=bedrock_error_message.encode('utf-8'))
|
||||
|
||||
# Create the HTTPStatusError
|
||||
mock_http_error = httpx.HTTPStatusError(
|
||||
message="Bad Request",
|
||||
request=Mock(spec=httpx.Request),
|
||||
response=mock_error_response,
|
||||
)
|
||||
|
||||
mock_llm_router = Mock()
|
||||
mock_llm_router.allm_passthrough_route = AsyncMock(side_effect=mock_http_error)
|
||||
|
||||
endpoint = "model/test-model/converse"
|
||||
model = "test-model"
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handle_bedrock_passthrough_router_model(
|
||||
model=model,
|
||||
endpoint=endpoint,
|
||||
request=mock_request,
|
||||
request_body=mock_request_body,
|
||||
llm_router=mock_llm_router,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "ContentBlock object at messages.0.content.0 must set one of the following keys" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
class TestLLMPassthroughFactoryProxyRoute:
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -1548,3 +1548,74 @@ def test_get_deployment_model_info_base_model_merge_priority():
|
|||
assert result["key"] == "gpt-4"
|
||||
|
||||
print("✓ Base model merge priority test passed!")
|
||||
|
||||
|
||||
def test_add_deployment_model_to_endpoint_for_llm_passthrough_route():
|
||||
"""
|
||||
Test that _add_deployment_model_to_endpoint_for_llm_passthrough_route correctly strips bedrock provider prefix
|
||||
"""
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "special-bedrock-model",
|
||||
"litellm_params": {
|
||||
"model": "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
# Test Case 1: Bedrock model with provider prefix - should strip "bedrock/" prefix
|
||||
kwargs = {
|
||||
"endpoint": "/model/special-bedrock-model/invoke",
|
||||
"custom_llm_provider": "bedrock",
|
||||
}
|
||||
result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route(
|
||||
kwargs=kwargs,
|
||||
model="special-bedrock-model",
|
||||
model_name="bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
)
|
||||
assert (
|
||||
result["endpoint"] == "/model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke"
|
||||
), f"Expected '/model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke', got '{result['endpoint']}'"
|
||||
|
||||
# Test Case 2: Bedrock invoke-with-response-stream endpoint
|
||||
kwargs = {
|
||||
"endpoint": "/model/special-bedrock-model/invoke-with-response-stream",
|
||||
"custom_llm_provider": "bedrock",
|
||||
}
|
||||
result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route(
|
||||
kwargs=kwargs,
|
||||
model="special-bedrock-model",
|
||||
model_name="bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
)
|
||||
assert (
|
||||
result["endpoint"] == "/model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke-with-response-stream"
|
||||
), f"Expected streaming endpoint with stripped prefix, got '{result['endpoint']}'"
|
||||
|
||||
# Test Case 3: Bedrock converse endpoint
|
||||
kwargs = {
|
||||
"endpoint": "/model/bedrock-model/converse",
|
||||
"custom_llm_provider": "bedrock",
|
||||
}
|
||||
result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route(
|
||||
kwargs=kwargs,
|
||||
model="bedrock-model",
|
||||
model_name="bedrock/us.meta.llama3-8b-instruct-v1:0",
|
||||
)
|
||||
assert (
|
||||
result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/converse"
|
||||
), f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/converse', got '{result['endpoint']}'"
|
||||
|
||||
# Test Case 4: Bedrock provider prefix auto-detected from model_name
|
||||
kwargs = {
|
||||
"endpoint": "/model/router-model/invoke",
|
||||
}
|
||||
result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route(
|
||||
kwargs=kwargs,
|
||||
model="router-model",
|
||||
model_name="bedrock/us.meta.llama3-8b-instruct-v1:0",
|
||||
)
|
||||
assert (
|
||||
result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/invoke"
|
||||
), f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/invoke', got '{result['endpoint']}'"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue