Merge remote-tracking branch 'origin/main' into fix_vertex_expired_tokens

This commit is contained in:
Oz Ben-Ami 2025-08-07 14:52:04 -04:00
commit 0c85fe4b70
154 changed files with 10909 additions and 1219 deletions

View file

@ -10,6 +10,7 @@ anthropic
orjson==3.10.12 # fast /embedding responses
pydantic==2.10.2
google-cloud-aiplatform==1.43.0
google-cloud-iam==2.19.1
fastapi-sso==0.16.0
uvloop==0.21.0
mcp==1.10.1 # for MCP server

View file

@ -0,0 +1,53 @@
import base64
from openai import OpenAI
import time
client = OpenAI(
base_url="http://0.0.0.0:4001",
api_key="sk-1234"
)
# Function to encode the image
def encode_image(image_path):
with open(image_path, "rb") as image_file:
return base64.b64encode(image_file.read()).decode("utf-8")
# Path to your image
image_path = "litellm/proxy/logo.jpg"
# Getting the Base64 string
base64_image = encode_image(image_path)
response = client.responses.create(
model="bedrock/us.anthropic.claude-3-5-sonnet-20241022-v2:0",
input=[
{
"role": "user",
"content": [
{ "type": "input_text", "text": "what color is the image"},
{
"type": "input_image",
"image_url": f"data:image/jpeg;base64,{base64_image}",
},
],
}
],
)
print(response.output_text)
print("response1 id===", response.id)
print("sleeping for 20 seconds...")
time.sleep(20)
print("making follow up request for existing id")
response2 = client.responses.create(
model="bedrock/us.anthropic.claude-3-5-sonnet-20241022-v2:0",
previous_response_id=response.id,
input="ok, and what objects are in the image?"
)
print(response2.output_text)

View file

@ -1,9 +1,11 @@
{{- if .Values.migrationJob.enabled }}
# This job runs the prisma migrations for the LiteLLM DB.
# This job runs the Prisma migrations for the LiteLLM DB.
apiVersion: batch/v1
kind: Job
metadata:
name: {{ include "litellm.fullname" . }}-migrations
labels:
{{- include "litellm.labels" . | nindent 4 }}
annotations:
{{- if .Values.migrationJob.hooks.argocd.enabled }}
argocd.argoproj.io/hook: PreSync
@ -18,6 +20,8 @@ metadata:
spec:
template:
metadata:
labels:
{{- include "litellm.labels" . | nindent 8 }}
annotations:
{{- with .Values.migrationJob.annotations }}
{{- toYaml . | nindent 8 }}

View file

@ -1,7 +1,7 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Caching - In-Memory, Redis, s3, Redis Semantic Cache, Disk
# Caching - In-Memory, Redis, s3, gcs, Redis Semantic Cache, Disk
[**See Code**](https://github.com/BerriAI/litellm/blob/main/litellm/caching/caching.py)
@ -14,7 +14,7 @@ import TabItem from '@theme/TabItem';
:::
## Initialize Cache - In Memory, Redis, s3 Bucket, Redis Semantic, Disk Cache, Qdrant Semantic
## Initialize Cache - In Memory, Redis, s3 Bucket, gcs Bucket, Redis Semantic, Disk Cache, Qdrant Semantic
<Tabs>
@ -28,6 +28,8 @@ pip install redis
For the hosted version you can setup your own Redis DB here: https://redis.io/try-free/
**Basic Redis Cache**
```python
import litellm
from litellm import completion
@ -48,6 +50,91 @@ response2 = completion(
# response1 == response2, response 1 is cached
```
**GCP IAM Redis Authentication**
For GCP Memorystore Redis with IAM authentication:
```shell
pip install google-cloud-iam
```
```python
import litellm
from litellm import completion
# For Redis Cluster with GCP IAM
from litellm.caching.redis_cluster_cache import RedisClusterCache
litellm.cache = RedisClusterCache(
startup_nodes=[
{"host": "10.128.0.2", "port": 6379},
{"host": "10.128.0.2", "port": 11008},
],
gcp_service_account="projects/-/serviceAccounts/your-sa@project.iam.gserviceaccount.com",
ssl=True,
ssl_cert_reqs=None,
ssl_check_hostname=False,
)
# Make completion calls
response1 = completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "Tell me a joke."}]
)
response2 = completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "Tell me a joke."}]
)
# response1 == response2, response 1 is cached
```
**Environment Variables for GCP IAM Redis**
You can also set these as environment variables:
```shell
export REDIS_HOST="10.128.0.2"
export REDIS_PORT="6379"
export REDIS_GCP_SERVICE_ACCOUNT="projects/-/serviceAccounts/your-sa@project.iam.gserviceaccount.com"
export REDIS_SSL="False"
```
Then simply initialize:
```python
litellm.cache = Cache(type="redis")
```
</TabItem>
<TabItem value="gcs" label="gcs-cache">
Set environment variables
```shell
GCS_BUCKET_NAME="my-cache-bucket"
GCS_PATH_SERVICE_ACCOUNT="/path/to/service_account.json"
```
```python
import litellm
from litellm import completion
from litellm.caching.caching import Cache
litellm.cache = Cache(type="gcs", gcs_bucket_name="my-cache-bucket", gcs_path_service_account="/path/to/service_account.json")
response1 = completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "Tell me a joke."}]
)
response2 = completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "Tell me a joke."}]
)
# response1 == response2, response 1 is cached
```
</TabItem>
@ -405,7 +492,7 @@ Advanced Params
```python
litellm.enable_cache(
type: Optional[Literal["local", "redis", "s3", "disk"]] = "local",
type: Optional[Literal["local", "redis", "s3", "gcs", "disk"]] = "local",
host: Optional[str] = None,
port: Optional[str] = None,
password: Optional[str] = None,
@ -429,7 +516,7 @@ Update the Cache params
```python
litellm.update_cache(
type: Optional[Literal["local", "redis", "s3", "disk"]] = "local",
type: Optional[Literal["local", "redis", "s3", "gcs", "disk"]] = "local",
host: Optional[str] = None,
port: Optional[str] = None,
password: Optional[str] = None,
@ -490,7 +577,7 @@ cache.get_cache = get_cache
```python
def __init__(
self,
type: Optional[Literal["local", "redis", "redis-semantic", "s3", "disk"]] = "local",
type: Optional[Literal["local", "redis", "redis-semantic", "s3", "gcs", "disk"]] = "local",
supported_call_types: Optional[
List[Literal["completion", "acompletion", "embedding", "aembedding", "atranscription", "transcription"]]
] = ["completion", "acompletion", "embedding", "aembedding", "atranscription", "transcription"],
@ -504,6 +591,13 @@ def __init__(
namespace: Optional[str] = None,
default_in_redis_ttl: Optional[float] = None,
redis_flush_size=None,
# GCP IAM Redis authentication params
gcp_service_account: Optional[str] = None,
gcp_ssl_ca_certs: Optional[str] = None,
ssl: Optional[bool] = None,
ssl_cert_reqs: Optional[Union[str, None]] = None,
ssl_check_hostname: Optional[bool] = None,
# redis semantic cache params
similarity_threshold: Optional[float] = None,

View file

@ -15,6 +15,7 @@ import os
# set env
os.environ["BRAINTRUST_API_KEY"] = ""
os.environ["BRAINTRUST_API_BASE"] = "https://api.braintrustdata.com/v1"
os.environ['OPENAI_API_KEY']=""
# set braintrust as a callback, litellm will send the data to braintrust
@ -35,6 +36,7 @@ response = litellm.completion(
```env
BRAINTRUST_API_KEY=""
BRAINTRUST_API_BASE="https://api.braintrustdata.com/v1"
```
2. Add braintrust to callbacks
@ -157,6 +159,8 @@ For more examples, [**Click Here**](../proxy/user_keys.md#chatcompletions)
</TabItem>
</Tabs>
You can use `BRAINTRUST_API_BASE` to point to your self-hosted Braintrust data plane. Read more about this [here](https://www.braintrust.dev/docs/guides/self-hosting).
## Full API Spec
Here's everything you can pass in metadata for a braintrust request

View file

@ -4,6 +4,7 @@ import TabItem from '@theme/TabItem';
# Anthropic
LiteLLM supports all anthropic models.
- `claude-opus-4-1-20250805`
- `claude-4` (`claude-opus-4-20250514`, `claude-sonnet-4-20250514`)
- `claude-3.7` (`claude-3-7-sonnet-20250219`)
- `claude-3.5` (`claude-3-5-sonnet-20240620`)

View file

@ -0,0 +1,75 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Oracle Cloud Infrastructure (OCI)
LiteLLM supports the following models for OCI on-demand GenAI API.
Check the [OCI Models List](https://docs.oracle.com/en-us/iaas/Content/generative-ai/pretrained-models.htm) to see if the model is available for your region.
- `meta.llama-4-maverick-17b-128e-instruct-fp8`
- `meta.llama-4-scout-17b-16e-instruct`
- `meta.llama-3.3-70b-instruct`
- `meta.llama-3.2-90b-vision-instruct`
- `meta.llama-3.1-405b-instruct`
- `xai.grok-4`
- `xai.grok-3`
- `xai.grok-3-fast`
- `xai.grok-3-mini`
- `xai.grok-3-mini-fast`
## Authentication
LiteLLM uses OCI signing key authentication. Follow the [official Oracle tutorial](https://docs.oracle.com/en-us/iaas/Content/API/Concepts/apisigningkey.htm) to create a signing key and obtain the following parameters:
- `user`
- `fingerprint`
- `tenancy`
- `region`
- `key_file`
## Usage
Input the parameters obtained from the OCI signing key creation process into the `completion` function.
```python
import os
from litellm import completion
messages = [{"role": "user", "content": "Hey! how's it going?"}]
response = completion(
model="oci/xai.grok-4",
messages=messages,
oci_region=<your_oci_region>,
oci_user=<your_oci_user>,
oci_fingerprint=<your_oci_fingerprint>,
oci_tenancy=<your_oci_tenancy>,
oci_key=<string_with_content_of_oci_key>,
oci_compartment_id=<oci_compartment_id>,
)
print(response)
```
## Usage - Streaming
Just set `stream=True` when calling completion.
```python
import os
from litellm import completion
messages = [{"role": "user", "content": "Hey! how's it going?"}]
response = completion(
model="oci/xai.grok-4",
messages=messages,
stream=True,
oci_region=<your_oci_region>,
oci_user=<your_oci_user>,
oci_fingerprint=<your_oci_fingerprint>,
oci_tenancy=<your_oci_tenancy>,
oci_key=<string_with_content_of_oci_key>,
oci_compartment_id=<oci_compartment_id>,
)
for chunk in response:
print(chunk["choices"][0]["delta"]["content"]) # same as openai format
```

View file

@ -204,7 +204,71 @@ For quick testing, you can also use REDIS_URL, eg.:
REDIS_URL="rediss://.."
```
but we **don't** recommend using REDIS_URL in prod. We've noticed a performance difference between using it vs. redis_host, port, etc.
but we **don't** recommend using REDIS_URL in prod. We've noticed a performance difference between using it vs. redis_host, port, etc.
#### GCP IAM Authentication
For GCP Memorystore Redis with IAM authentication, install the required dependency:
:::info
IAM authentication for redis is only supported via GCP and only on Redis Clusters for now.
:::
```shell
pip install google-cloud-iam
```
<Tabs>
<TabItem value="gcp-iam-config" label="Set on config.yaml">
For Redis Cluster with GCP IAM:
```yaml
litellm_settings:
cache: True
cache_params:
type: redis
redis_startup_nodes: [{"host": "10.128.0.2", "port": 6379}, {"host": "10.128.0.2", "port": 11008}]
gcp_service_account: "projects/-/serviceAccounts/your-sa@project.iam.gserviceaccount.com"
ssl: true
ssl_cert_reqs: null
ssl_check_hostname: false
```
</TabItem>
<TabItem value="gcp-iam-env" label="Set on .env">
You can configure GCP IAM Redis authentication in your .env:
For Redis Cluster:
```env
REDIS_CLUSTER_NODES='[{"host": "10.128.0.2", "port": 6379}, {"host": "10.128.0.2", "port": 11008}]'
REDIS_GCP_SERVICE_ACCOUNT="projects/-/serviceAccounts/your-sa@project.iam.gserviceaccount.com"
REDIS_GCP_SSL_CA_CERTS="./server-ca.pem"
REDIS_SSL="True"
REDIS_SSL_CERT_REQS="None"
REDIS_SSL_CHECK_HOSTNAME="False"
```
**GCP Authentication Setup**
Make sure your GCP credentials are configured:
```shell
# Option 1: Service account key file
export GOOGLE_APPLICATION_CREDENTIALS="/path/to/service-account-key.json"
# Option 2: If running on GCP compute instance with service account attached
# No additional setup needed
```
</TabItem>
</Tabs>
#### Step 2: Add Redis Credentials to .env
Set either `REDIS_URL` or the `REDIS_HOST` in your os environment, to enable caching.
@ -917,6 +981,13 @@ cache_params:
password: secret_password # Redis server password
namespace: Optional[str] = None,
# GCP IAM Authentication for Redis
gcp_service_account: "projects/-/serviceAccounts/your-sa@project.iam.gserviceaccount.com" # GCP service account for IAM authentication
gcp_ssl_ca_certs: "./server-ca.pem" # Path to SSL CA certificate file for GCP Memorystore Redis
ssl: true # Enable SSL for secure connections
ssl_cert_reqs: null # Set to null for self-signed certificates
ssl_check_hostname: false # Set to false for self-signed certificates
# S3 cache parameters
s3_bucket_name: your_s3_bucket_name # Name of the S3 bucket

View file

@ -58,6 +58,13 @@ litellm_settings:
service_name: "mymaster"
sentinel_nodes: [["localhost", 26379]]
# Optional - GCP IAM Authentication for Redis
gcp_service_account: "projects/-/serviceAccounts/your-sa@project.iam.gserviceaccount.com" # GCP service account for IAM authentication
gcp_ssl_ca_certs: "./server-ca.pem" # Path to SSL CA certificate file for GCP Memorystore Redis
ssl: true # Enable SSL for secure connections
ssl_cert_reqs: null # Set to null for self-signed certificates
ssl_check_hostname: false # Set to false for self-signed certificates
# Optional - Qdrant Semantic Cache Settings
qdrant_semantic_cache_embedding_model: openai-embedding # the model should be defined on the model_list
qdrant_collection_name: test_collection
@ -362,6 +369,7 @@ router_settings:
| BEDROCK_MAX_POLICY_SIZE | Maximum size for Bedrock policy. Default is 75
| BERRISPEND_ACCOUNT_ID | Account ID for BerriSpend service
| BRAINTRUST_API_KEY | API key for Braintrust integration
| BRAINTRUST_API_BASE | Base URL for Braintrust API. Default is https://api.braintrustdata.com/v1
| CACHED_STREAMING_CHUNK_DELAY | Delay in seconds for cached streaming chunks. Default is 0.02
| CIRCLE_OIDC_TOKEN | OpenID Connect token for CircleCI
| CIRCLE_OIDC_TOKEN_V2 | Version 2 of the OpenID Connect token for CircleCI
@ -649,6 +657,8 @@ router_settings:
| REDIS_PASSWORD | Password for Redis service
| REDIS_PORT | Port number for Redis server
| REDIS_SOCKET_TIMEOUT | Timeout in seconds for Redis socket operations. Default is 0.1
| REDIS_GCP_SERVICE_ACCOUNT | GCP service account for IAM authentication with Redis. Format: "projects/-/serviceAccounts/name@project.iam.gserviceaccount.com"
| REDIS_GCP_SSL_CA_CERTS | Path to SSL CA certificate file for secure GCP Memorystore Redis connections
| REDOC_URL | The path to the Redoc Fast API documentation. **By default this is "/redoc"**
| REPEATED_STREAMING_CHUNK_LIMIT | Limit for repeated streaming chunks to detect looping. Default is 100
| REPLICATE_MODEL_NAME_WITH_ID_LENGTH | Length of Replicate model names with ID. Default is 64

View file

@ -261,7 +261,12 @@ model_list:
litellm_settings:
callbacks: ["prometheus"]
custom_prometheus_metadata_labels: ["metadata.foo", "metadata.bar"]
custom_prometheus_tags: ["prod", "staging", "batch-job"]
custom_prometheus_tags:
- "prod"
- "staging"
- "batch-job"
- "User-Agent: RooCode/*"
- "User-Agent: claude-cli/*"
```
2. Make a request with tags
@ -297,16 +302,26 @@ curl -L -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
```
**How Custom Tags Work:**
- Each configured tag becomes a boolean label in prometheus metrics
- If a tag is present in the request, the label value is `"true"`
- If a tag is not present in the request, the label value is `"false"`
- Each configured tag becomes a boolean label in prometheus metrics
- If a tag matches (exact or wildcard), the label value is `"true"`, otherwise `"false"`
- Tag names are sanitized for prometheus compatibility (e.g., `"batch-job"` becomes `"tag_batch_job"`)
- **Wildcard patterns** supported using `*` (e.g., `"User-Agent: RooCode/*"` matches `"User-Agent: RooCode/1.0.0"`)
**Example with wildcards:**
```yaml
litellm_settings:
callbacks: ["prometheus"]
custom_prometheus_tags:
- "User-Agent: RooCode/*"
- "User-Agent: claude-cli/*"
```
**Use Cases:**
- Environment tracking (`prod`, `staging`, `dev`)
- Request type classification (`batch-job`, `user-facing`, `background`)
- Feature flags (`new-feature`, `beta-users`)
- Team or service identification (`team-a`, `service-xyz`)
- User-Agent Tracking - use this to track how much Roo Code, Claude Code, Gemini CLI are used (`User-Agent: RooCode/*`, `User-Agent: claude-cli/*`, `User-Agent: gemini-cli/*`)
## Configuring Metrics and Labels

View file

@ -86,6 +86,11 @@ response = client.chat.completions.create(
print(response)
```
</TabItem>
<TabItem value="litellm_sdk" label="LiteLLM Python SDK">
[**👉 Go Here**](../providers/litellm_proxy#send-all-sdk-requests-to-litellm-proxy)
</TabItem>
<TabItem value="azureopenai" label="AzureOpenAI Python">

View file

@ -28,7 +28,7 @@ import TabItem from '@theme/TabItem';
docker run \
-e STORE_MODEL_IN_DB=True \
-p 4000:4000 \
ghcr.io/berriai/litellm:v1.74.15.rc.1
ghcr.io/berriai/litellm:1.74.15.rc.1
```
</TabItem>

View file

@ -469,7 +469,8 @@ const sidebars = {
"providers/featherless_ai",
"providers/nebius",
"providers/dashscope",
"providers/bytez"
"providers/bytez",
"providers/oci",
],
},
{

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

View file

@ -1,160 +0,0 @@
import json
from typing import TYPE_CHECKING, Any, List, Optional, Union, cast
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import SpendLogsPayload
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionResponseMessage,
GenericChatCompletionMessage,
ResponseInputParam,
)
from litellm.types.utils import ChatCompletionMessageToolCall, Message, ModelResponse
if TYPE_CHECKING:
from litellm.responses.litellm_completion_transformation.transformation import (
ChatCompletionSession,
)
else:
ChatCompletionSession = Any
class _ENTERPRISE_ResponsesSessionHandler:
@staticmethod
async def get_chat_completion_message_history_for_previous_response_id(
previous_response_id: str,
) -> ChatCompletionSession:
"""
Return the chat completion message history for a previous response id
"""
from litellm.responses.litellm_completion_transformation.transformation import (
ChatCompletionSession,
LiteLLMCompletionResponsesConfig,
)
verbose_proxy_logger.debug(
"inside get_chat_completion_message_history_for_previous_response_id"
)
all_spend_logs: List[
SpendLogsPayload
] = await _ENTERPRISE_ResponsesSessionHandler.get_all_spend_logs_for_previous_response_id(
previous_response_id
)
verbose_proxy_logger.debug(
"found %s spend logs for this response id", len(all_spend_logs)
)
litellm_session_id: Optional[str] = None
if len(all_spend_logs) > 0:
litellm_session_id = all_spend_logs[0].get("session_id")
chat_completion_message_history: List[
Union[
AllMessageValues,
GenericChatCompletionMessage,
ChatCompletionMessageToolCall,
ChatCompletionResponseMessage,
Message,
]
] = []
for spend_log in all_spend_logs:
proxy_server_request: Union[str, dict] = (
spend_log.get("proxy_server_request") or "{}"
)
proxy_server_request_dict: Optional[dict] = None
response_input_param: Optional[Union[str, ResponseInputParam]] = None
if isinstance(proxy_server_request, dict):
proxy_server_request_dict = proxy_server_request
else:
proxy_server_request_dict = json.loads(proxy_server_request)
############################################################
# Add Input messages for this Spend Log
############################################################
if proxy_server_request_dict:
_response_input_param = proxy_server_request_dict.get("input", None)
if isinstance(_response_input_param, str):
response_input_param = _response_input_param
elif isinstance(_response_input_param, dict):
response_input_param = cast(
ResponseInputParam, _response_input_param
)
if response_input_param:
chat_completion_messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=response_input_param,
responses_api_request=proxy_server_request_dict or {},
)
chat_completion_message_history.extend(chat_completion_messages)
############################################################
# Add Output messages for this Spend Log
############################################################
_response_output = spend_log.get("response", "{}")
if isinstance(_response_output, dict):
# transform `ChatCompletion Response` to `ResponsesAPIResponse`
model_response = ModelResponse(**_response_output)
for choice in model_response.choices:
if hasattr(choice, "message"):
chat_completion_message_history.append(
getattr(choice, "message")
)
verbose_proxy_logger.debug(
"chat_completion_message_history %s",
json.dumps(chat_completion_message_history, indent=4, default=str),
)
return ChatCompletionSession(
messages=chat_completion_message_history,
litellm_session_id=litellm_session_id,
)
@staticmethod
async def get_all_spend_logs_for_previous_response_id(
previous_response_id: str,
) -> List[SpendLogsPayload]:
"""
Get all spend logs for a previous response id
SQL query
SELECT session_id FROM spend_logs WHERE response_id = previous_response_id, SELECT * FROM spend_logs WHERE session_id = session_id
"""
from litellm.proxy.proxy_server import prisma_client
verbose_proxy_logger.debug("decoding response id=%s", previous_response_id)
decoded_response_id = (
ResponsesAPIRequestUtils._decode_responses_api_response_id(
previous_response_id
)
)
previous_response_id = decoded_response_id.get(
"response_id", previous_response_id
)
if prisma_client is None:
return []
query = """
WITH matching_session AS (
SELECT session_id
FROM "LiteLLM_SpendLogs"
WHERE request_id = $1
)
SELECT *
FROM "LiteLLM_SpendLogs"
WHERE session_id IN (SELECT session_id FROM matching_session)
ORDER BY "endTime" ASC;
"""
spend_logs = await prisma_client.db.query_raw(query, previous_response_id)
verbose_proxy_logger.debug(
"Found the following spend logs for previous response id %s: %s",
previous_response_id,
json.dumps(spend_logs, indent=4, default=str),
)
return spend_logs

View file

@ -122,19 +122,19 @@ class PrometheusLogger(CustomLogger):
# Counter for total_output_tokens
self.litellm_tokens_metric = self._counter_factory(
"litellm_total_tokens",
"litellm_total_tokens_metric",
"Total number of input + output tokens from LLM requests",
labelnames=self.get_labels_for_metric("litellm_total_tokens_metric"),
)
self.litellm_input_tokens_metric = self._counter_factory(
"litellm_input_tokens",
"litellm_input_tokens_metric",
"Total number of input tokens from LLM requests",
labelnames=self.get_labels_for_metric("litellm_input_tokens_metric"),
)
self.litellm_output_tokens_metric = self._counter_factory(
"litellm_output_tokens",
"litellm_output_tokens_metric",
"Total number of output tokens from LLM requests",
labelnames=self.get_labels_for_metric("litellm_output_tokens_metric"),
)
@ -2293,10 +2293,60 @@ def get_custom_labels_from_metadata(metadata: dict) -> Dict[str, str]:
return result
def _tag_matches_wildcard_configured_pattern(tags: List[str], configured_tag: str) -> bool:
"""
Check if any of the request tags matches a wildcard configured pattern
Args:
tags: List[str] - The request tags
configured_tag: str - The configured tag
Returns:
bool - True if any of the request tags matches the configured tag, False otherwise
e.g.
tags = ["User-Agent: curl/7.68.0", "User-Agent: python-requests/2.28.1", "prod"]
configured_tag = "User-Agent: curl/*"
_tag_matches_wildcard_configured_pattern(tags=tags, configured_tag=configured_tag) # True
configured_tag = "User-Agent: python-requests/*"
_tag_matches_wildcard_configured_pattern(tags=tags, configured_tag=configured_tag) # True
configured_tag = "gm"
_tag_matches_wildcard_configured_pattern(tags=tags, configured_tag=configured_tag) # False
"""
import re
from litellm.router_utils.pattern_match_deployments import PatternMatchRouter
pattern_router = PatternMatchRouter()
regex_pattern = pattern_router._pattern_to_regex(configured_tag)
return any(re.match(pattern=regex_pattern, string=tag) for tag in tags)
def get_custom_labels_from_tags(tags: List[str]) -> Dict[str, str]:
"""
Get custom labels from tags based on admin configuration
Get custom labels from tags based on admin configuration.
Supports both exact matches and wildcard patterns:
- Exact match: "prod" matches "prod" exactly
- Wildcard pattern: "User-Agent: curl/*" matches "User-Agent: curl/7.68.0"
Reuses PatternMatchRouter for wildcard pattern matching.
Returns dict of label_name: "true" if the tag matches the configured tag, "false" otherwise
{
"tag_User-Agent_curl": "true",
"tag_User-Agent_python_requests": "false",
"tag_Environment_prod": "true",
"tag_Environment_dev": "false",
"tag_Service_api_gateway_v2": "true",
"tag_Service_web_app_v1": "false",
}
"""
import re
from litellm.router_utils.pattern_match_deployments import PatternMatchRouter
from litellm.types.integrations.prometheus import _sanitize_prometheus_label_name
configured_tags = litellm.custom_prometheus_tags
@ -2304,16 +2354,22 @@ def get_custom_labels_from_tags(tags: List[str]) -> Dict[str, str]:
return {}
result: Dict[str, str] = {}
pattern_router = PatternMatchRouter()
# Map each configured tag to its presence in the request tags
for configured_tag in configured_tags:
# Create a safe prometheus label name
label_name = _sanitize_prometheus_label_name(f"tag_{configured_tag}")
# Check if this tag is present in the request tags
# Check for exact match first (backwards compatibility)
if configured_tag in tags:
result[label_name] = "true"
else:
result[label_name] = "false"
continue
# Use PatternMatchRouter for wildcard pattern matching
if "*" in configured_tag and _tag_matches_wildcard_configured_pattern(tags=tags, configured_tag=configured_tag):
result[label_name] = "true"
continue
# No match found
result[label_name] = "false"
return result

View file

@ -20,7 +20,6 @@ class EnterpriseRouteChecks:
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"🚨🚨🚨 DISABLING LLM API ENDPOINTS is an Enterprise feature\n🚨 {CommonProxyErrors.not_premium_user.value}",
)
return False
return get_secret_bool("DISABLE_LLM_API_ENDPOINTS") is True

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm-enterprise"
version = "0.1.16"
version = "0.1.19"
description = "Package for LiteLLM Enterprise features"
authors = ["BerriAI"]
readme = "README.md"
@ -22,7 +22,7 @@ requires = ["poetry-core"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "0.1.16"
version = "0.1.19"
version_files = [
"pyproject.toml:version",
"../requirements.txt:litellm-enterprise==",

View file

@ -15,13 +15,3 @@ CREATE TABLE "LiteLLM_MCPServerTable" (
CONSTRAINT "LiteLLM_MCPServerTable_pkey" PRIMARY KEY ("server_id")
);
-- Migration for existing tables: rename alias to server_name if upgrading
DO $$
BEGIN
IF EXISTS (SELECT 1 FROM information_schema.columns WHERE table_name = 'LiteLLM_MCPServerTable' AND column_name = 'alias') THEN
ALTER TABLE "LiteLLM_MCPServerTable" RENAME COLUMN "alias" TO "server_name";
END IF;
END $$;
-- Migration for existing tables: add alias column if upgrading
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "alias" TEXT;

View file

@ -0,0 +1,10 @@
-- Migration for existing tables: rename alias to server_name if upgrading
DO $$
BEGIN
IF EXISTS (SELECT 1 FROM information_schema.columns WHERE table_name = 'LiteLLM_MCPServerTable' AND column_name = 'alias') THEN
ALTER TABLE "LiteLLM_MCPServerTable" RENAME COLUMN "alias" TO "server_name";
END IF;
END $$;
-- Migration for existing tables: add alias column if upgrading
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "alias" TEXT;

View file

@ -269,6 +269,7 @@ blocked_user_list: Optional[Union[str, List]] = None
banned_keywords_list: Optional[Union[str, List]] = None
llm_guard_mode: Literal["all", "key-specific", "request-specific"] = "all"
guardrail_name_config_map: Dict[str, GuardrailItem] = {}
include_cost_in_streaming_usage: bool = False
### PROMPTS ###
from litellm.types.prompts.init_prompts import PromptSpec
@ -429,6 +430,9 @@ project = None
config_path = None
vertex_ai_safety_settings: Optional[dict] = None
BEDROCK_CONVERSE_MODELS = [
"openai.gpt-oss-20b-1:0",
"openai.gpt-oss-120b-1:0",
"anthropic.claude-opus-4-1-20250805-v1:0",
"anthropic.claude-opus-4-20250514-v1:0",
"anthropic.claude-sonnet-4-20250514-v1:0",
"anthropic.claude-3-7-sonnet-20250219-v1:0",
@ -1199,6 +1203,7 @@ from .llms.nebius.chat.transformation import NebiusConfig
from .llms.dashscope.chat.transformation import DashScopeChatConfig
from .llms.moonshot.chat.transformation import MoonshotChatConfig
from .llms.v0.chat.transformation import V0ChatConfig
from .llms.oci.chat.transformation import OCIChatConfig
from .llms.morph.chat.transformation import MorphChatConfig
from .llms.lambda_ai.chat.transformation import LambdaAIChatConfig
from .llms.hyperbolic.chat.transformation import HyperbolicChatConfig

View file

@ -108,9 +108,20 @@ verbose_router_logger.addHandler(handler)
verbose_proxy_logger.addHandler(handler)
verbose_logger.addHandler(handler)
# Suppress httpx request logging at INFO level
httpx_logger = logging.getLogger("httpx")
httpx_logger.setLevel(logging.WARNING)
def _suppress_loggers():
"""Suppress noisy loggers at INFO level"""
# Suppress httpx request logging at INFO level
httpx_logger = logging.getLogger("httpx")
httpx_logger.setLevel(logging.WARNING)
# Suppress APScheduler logging at INFO level
apscheduler_executors_logger = logging.getLogger("apscheduler.executors.default")
apscheduler_executors_logger.setLevel(logging.WARNING)
apscheduler_scheduler_logger = logging.getLogger("apscheduler.scheduler")
apscheduler_scheduler_logger.setLevel(logging.WARNING)
# Call the suppression function
_suppress_loggers()
ALL_LOGGERS = [
logging.getLogger(),

View file

@ -12,7 +12,7 @@ import json
# s/o [@Frank Colson](https://www.linkedin.com/in/frank-colson-422b9b183/) for this redis implementation
import os
from typing import List, Optional, Union
from typing import Callable, List, Optional, Union
import redis # type: ignore
import redis.asyncio as async_redis # type: ignore
@ -34,7 +34,7 @@ def _get_redis_kwargs():
"retry",
}
include_args = ["url"]
include_args = ["url", "redis_connect_func", "gcp_service_account", "gcp_ssl_ca_certs"]
available_args = [x for x in arg_spec.args if x not in exclude_args] + include_args
@ -72,6 +72,12 @@ def _get_redis_cluster_kwargs(client=None):
available_args.append("password")
available_args.append("username")
available_args.append("ssl")
available_args.append("ssl_cert_reqs")
available_args.append("ssl_check_hostname")
available_args.append("ssl_ca_certs")
available_args.append("redis_connect_func") # Needed for sync clusters and IAM detection
available_args.append("gcp_service_account")
available_args.append("gcp_ssl_ca_certs")
return available_args
@ -93,6 +99,73 @@ def _redis_kwargs_from_environment():
return return_dict
def _generate_gcp_iam_access_token(service_account: str) -> str:
"""
Generate GCP IAM access token for Redis authentication.
Args:
service_account: GCP service account in format 'projects/-/serviceAccounts/name@project.iam.gserviceaccount.com'
Returns:
Access token string for GCP IAM authentication
"""
try:
from google.cloud import iam_credentials_v1
except ImportError:
raise ImportError(
"google-cloud-iam is required for GCP IAM Redis authentication. "
"Install it with: pip install google-cloud-iam"
)
client = iam_credentials_v1.IAMCredentialsClient()
request = iam_credentials_v1.GenerateAccessTokenRequest(
name=service_account,
scope=['https://www.googleapis.com/auth/cloud-platform'],
)
response = client.generate_access_token(request=request)
return str(response.access_token)
def create_gcp_iam_redis_connect_func(
service_account: str,
ssl_ca_certs: Optional[str] = None,
) -> Callable:
"""
Creates a custom Redis connection function for GCP IAM authentication.
Args:
service_account: GCP service account in format 'projects/-/serviceAccounts/name@project.iam.gserviceaccount.com'
ssl_ca_certs: Path to SSL CA certificate file for secure connections
Returns:
A connection function that can be used with Redis clients
"""
def iam_connect(self):
"""Initialize the connection and authenticate using GCP IAM"""
from redis.exceptions import AuthenticationError, AuthenticationWrongNumberOfArgsError
from redis.utils import str_if_bytes
self._parser.on_connect(self)
auth_args = (_generate_gcp_iam_access_token(service_account),)
self.send_command("AUTH", *auth_args, check_health=False)
try:
auth_response = self.read_response()
except AuthenticationWrongNumberOfArgsError:
# Fallback to password auth if IAM fails
if hasattr(self, 'password') and self.password:
self.send_command("AUTH", self.password, check_health=False)
auth_response = self.read_response()
else:
raise
if str_if_bytes(auth_response) != "OK":
raise AuthenticationError("GCP IAM authentication failed")
return iam_connect
def get_redis_url_from_environment():
if "REDIS_URL" in os.environ:
return os.environ["REDIS_URL"]
@ -156,6 +229,27 @@ def _get_redis_client_logic(**env_overrides):
if _service_name is not None:
redis_kwargs["service_name"] = _service_name
# Handle GCP IAM authentication
_gcp_service_account = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT")
_gcp_ssl_ca_certs = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str("REDIS_GCP_SSL_CA_CERTS")
if _gcp_service_account is not None:
verbose_logger.debug("Setting up GCP IAM authentication for Redis with service account.")
redis_kwargs["redis_connect_func"] = create_gcp_iam_redis_connect_func(
service_account=_gcp_service_account,
ssl_ca_certs=_gcp_ssl_ca_certs
)
# Store GCP service account in redis_connect_func for async cluster access
redis_kwargs["redis_connect_func"]._gcp_service_account = _gcp_service_account
# Remove GCP-specific kwargs that shouldn't be passed to Redis client
redis_kwargs.pop("gcp_service_account", None)
redis_kwargs.pop("gcp_ssl_ca_certs", None)
# Only enable SSL if explicitly requested AND SSL CA certs are provided
if _gcp_ssl_ca_certs and redis_kwargs.get("ssl", False):
redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
redis_kwargs.pop("host", None)
redis_kwargs.pop("port", None)
@ -198,7 +292,7 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
for item in redis_kwargs["startup_nodes"]:
new_startup_nodes.append(ClusterNode(**item))
redis_kwargs.pop("startup_nodes")
cluster_kwargs.pop("startup_nodes", None)
return redis.RedisCluster(startup_nodes=new_startup_nodes, **cluster_kwargs) # type: ignore
@ -273,7 +367,7 @@ def get_redis_client(**env_overrides):
def get_redis_async_client(
**env_overrides,
) -> async_redis.Redis:
) -> Union[async_redis.Redis, async_redis.RedisCluster]:
redis_kwargs = _get_redis_client_logic(**env_overrides)
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
args = _get_redis_url_kwargs(client=async_redis.Redis.from_url)
@ -298,14 +392,46 @@ def get_redis_async_client(
if arg in args:
cluster_kwargs[arg] = redis_kwargs[arg]
# Handle GCP IAM authentication for async clusters
redis_connect_func = cluster_kwargs.pop("redis_connect_func", None)
from litellm import get_secret_str
# Get GCP service account - first try from redis_connect_func, then from environment
gcp_service_account = None
if redis_connect_func and hasattr(redis_connect_func, '_gcp_service_account'):
gcp_service_account = redis_connect_func._gcp_service_account
else:
gcp_service_account = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT")
verbose_logger.info(f"DEBUG: Redis cluster kwargs: redis_connect_func={redis_connect_func is not None}, gcp_service_account_provided={gcp_service_account is not None}")
# If GCP IAM is configured (indicated by redis_connect_func), generate access token and use as password
if redis_connect_func and gcp_service_account:
verbose_logger.info("DEBUG: Generating IAM token for service account (value not logged for security reasons)")
try:
# Generate IAM access token using the helper function
access_token = _generate_gcp_iam_access_token(gcp_service_account)
cluster_kwargs["password"] = access_token
verbose_logger.info("DEBUG: Successfully generated GCP IAM access token for async Redis cluster")
except Exception as e:
verbose_logger.error(f"Failed to generate GCP IAM access token: {e}")
from redis.exceptions import AuthenticationError
raise AuthenticationError("Failed to generate GCP IAM access token")
else:
verbose_logger.info(f"DEBUG: Not using GCP IAM auth - redis_connect_func={redis_connect_func is not None}, gcp_service_account={gcp_service_account}")
new_startup_nodes: List[ClusterNode] = []
for item in redis_kwargs["startup_nodes"]:
new_startup_nodes.append(ClusterNode(**item))
redis_kwargs.pop("startup_nodes")
return async_redis.RedisCluster(
cluster_kwargs.pop("startup_nodes", None)
# Create async RedisCluster with IAM token as password if available
cluster_client = async_redis.RedisCluster(
startup_nodes=new_startup_nodes, **cluster_kwargs # type: ignore
)
return cluster_client
# Check for Redis Sentinel
if "sentinel_nodes" in redis_kwargs and "service_name" in redis_kwargs:

View file

@ -7,4 +7,5 @@ from .qdrant_semantic_cache import QdrantSemanticCache
from .redis_cache import RedisCache
from .redis_cluster_cache import RedisClusterCache
from .redis_semantic_cache import RedisSemanticCache
from .s3_cache import S3Cache
from .s3_cache import S3Cache
from .gcs_cache import GCSCache

View file

@ -34,6 +34,7 @@ from .redis_cache import RedisCache
from .redis_cluster_cache import RedisClusterCache
from .redis_semantic_cache import RedisSemanticCache
from .s3_cache import S3Cache
from .gcs_cache import GCSCache
def print_verbose(print_statement):
@ -92,6 +93,9 @@ class Cache:
s3_aws_session_token: Optional[str] = None,
s3_config: Optional[Any] = None,
s3_path: Optional[str] = None,
gcs_bucket_name: Optional[str] = None,
gcs_path_service_account: Optional[str] = None,
gcs_path: Optional[str] = None,
redis_semantic_cache_embedding_model: str = "text-embedding-ada-002",
redis_semantic_cache_index_name: Optional[str] = None,
redis_flush_size: Optional[int] = None,
@ -102,6 +106,9 @@ class Cache:
qdrant_collection_name: Optional[str] = None,
qdrant_quantization_config: Optional[str] = None,
qdrant_semantic_cache_embedding_model: str = "text-embedding-ada-002",
# GCP IAM authentication parameters
gcp_service_account: Optional[str] = None,
gcp_ssl_ca_certs: Optional[str] = None,
**kwargs,
):
"""
@ -140,6 +147,11 @@ class Cache:
s3_aws_session_token (str, optional): The aws session token for the s3 cache. Defaults to None.
s3_config (dict, optional): The config for the s3 cache. Defaults to None.
# GCS Cache Args
gcs_bucket_name (str, optional): The bucket name for the gcs cache. Defaults to None.
gcs_path_service_account (str, optional): Path to the service account json.
gcs_path (str, optional): Folder path inside the bucket to store cache files.
# Common Cache Args
supported_call_types (list, optional): List of call types to cache for. Defaults to cache == on for all call types.
**kwargs: Additional keyword arguments for redis.Redis() cache
@ -152,14 +164,21 @@ class Cache:
"""
if type == LiteLLMCacheType.REDIS:
if redis_startup_nodes:
self.cache: BaseCache = RedisClusterCache(
host=host,
port=port,
password=password,
redis_flush_size=redis_flush_size,
startup_nodes=redis_startup_nodes,
# Only pass GCP parameters if they are provided
cluster_kwargs = {
"host": host,
"port": port,
"password": password,
"redis_flush_size": redis_flush_size,
"startup_nodes": redis_startup_nodes,
**kwargs,
)
}
if gcp_service_account is not None:
cluster_kwargs["gcp_service_account"] = gcp_service_account
if gcp_ssl_ca_certs is not None:
cluster_kwargs["gcp_ssl_ca_certs"] = gcp_ssl_ca_certs
self.cache: BaseCache = RedisClusterCache(**cluster_kwargs)
else:
self.cache = RedisCache(
host=host,
@ -204,6 +223,12 @@ class Cache:
s3_path=s3_path,
**kwargs,
)
elif type == LiteLLMCacheType.GCS:
self.cache = GCSCache(
bucket_name=gcs_bucket_name,
path_service_account=gcs_path_service_account,
gcs_path=gcs_path,
)
elif type == LiteLLMCacheType.AZURE_BLOB:
self.cache = AzureBlobCache(
account_url=azure_account_url,

View file

@ -0,0 +1,97 @@
"""GCS Cache implementation
Supports syncing responses to Google Cloud Storage Buckets using HTTP requests.
"""
import json
import asyncio
from typing import Optional
from litellm._logging import print_verbose, verbose_logger
from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
_get_httpx_client,
httpxSpecialProvider,
)
from .base_cache import BaseCache
class GCSCache(BaseCache):
def __init__(self, bucket_name: Optional[str] = None, path_service_account: Optional[str] = None, gcs_path: Optional[str] = None) -> None:
super().__init__()
self.bucket_name = bucket_name or GCSBucketBase(bucket_name=None).BUCKET_NAME
self.path_service_account = path_service_account or GCSBucketBase(bucket_name=None).path_service_account_json
self.key_prefix = gcs_path.rstrip("/") + "/" if gcs_path else ""
# create httpx clients
self.async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
self.sync_client = _get_httpx_client()
def _construct_headers(self) -> dict:
base = GCSBucketBase(bucket_name=self.bucket_name)
base.path_service_account_json = self.path_service_account
base.BUCKET_NAME = self.bucket_name
return base.sync_construct_request_headers()
def set_cache(self, key, value, **kwargs):
try:
print_verbose(f"LiteLLM SET Cache - GCS. Key={key}. Value={value}")
headers = self._construct_headers()
object_name = self.key_prefix + key
bucket_name = self.bucket_name
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}"
data = json.dumps(value)
self.sync_client.post(url=url, data=data, headers=headers)
except Exception as e:
print_verbose(f"GCS Caching: set_cache() - Got exception from GCS: {e}")
async def async_set_cache(self, key, value, **kwargs):
try:
headers = self._construct_headers()
object_name = self.key_prefix + key
bucket_name = self.bucket_name
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}"
data = json.dumps(value)
await self.async_client.post(url=url, data=data, headers=headers)
except Exception as e:
print_verbose(f"GCS Caching: async_set_cache() - Got exception from GCS: {e}")
def get_cache(self, key, **kwargs):
try:
headers = self._construct_headers()
object_name = self.key_prefix + key
bucket_name = self.bucket_name
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{object_name}?alt=media"
response = self.sync_client.get(url=url, headers=headers)
if response.status_code == 200:
cached_response = json.loads(response.text)
verbose_logger.debug(
f"Got GCS Cache: key: {key}, cached_response {cached_response}. Type Response {type(cached_response)}"
)
return cached_response
return None
except Exception as e:
verbose_logger.error(f"GCS Caching: get_cache() - Got exception from GCS: {e}")
async def async_get_cache(self, key, **kwargs):
try:
headers = self._construct_headers()
object_name = self.key_prefix + key
bucket_name = self.bucket_name
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{object_name}?alt=media"
response = await self.async_client.get(url=url, headers=headers)
if response.status_code == 200:
return json.loads(response.text)
return None
except Exception as e:
verbose_logger.error(f"GCS Caching: async_get_cache() - Got exception from GCS: {e}")
def flush_cache(self):
pass
async def disconnect(self):
pass
async def async_set_cache_pipeline(self, cache_list, **kwargs):
tasks = []
for val in cache_list:
tasks.append(self.async_set_cache(val[0], val[1], **kwargs))
await asyncio.gather(*tasks)

View file

@ -279,6 +279,7 @@ LITELLM_CHAT_PROVIDERS = [
"dashscope",
"moonshot",
"v0",
"oci",
"morph",
"lambda_ai",
]
@ -765,6 +766,7 @@ MAXIMUM_TRACEBACK_LINES_TO_LOG = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG",
X_LITELLM_DISABLE_CALLBACKS = "x-litellm-disable-callbacks"
LITELLM_METADATA_FIELD = "litellm_metadata"
OLD_LITELLM_METADATA_FIELD = "metadata"
LITELLM_TRUNCATED_PAYLOAD_FIELD = "litellm_truncated"
########################### LiteLLM Proxy Specific Constants ###########################
########################################################################################

View file

@ -1265,7 +1265,7 @@ class BaseTokenUsageProcessor:
Combine multiple Usage objects into a single Usage object, checking model keys for nested values.
"""
from litellm.types.utils import (
CompletionTokensDetails,
CompletionTokensDetailsWrapper,
PromptTokensDetailsWrapper,
Usage,
)
@ -1320,7 +1320,7 @@ class BaseTokenUsageProcessor:
not hasattr(combined, "completion_tokens_details")
or not combined.completion_tokens_details
):
combined.completion_tokens_details = CompletionTokensDetails()
combined.completion_tokens_details = CompletionTokensDetailsWrapper()
# Check what keys exist in the model's completion_tokens_details
for attr in usage.completion_tokens_details.model_fields:

View file

@ -42,7 +42,7 @@ class BraintrustLogger(CustomLogger):
) -> None:
super().__init__()
self.validate_environment(api_key=api_key)
self.api_base = api_base or API_BASE
self.api_base = api_base or os.getenv("BRAINTRUST_API_BASE") or API_BASE
self.default_project_id = None
self.api_key: str = api_key or os.getenv("BRAINTRUST_API_KEY") # type: ignore
self.headers = {

View file

@ -34,8 +34,6 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp import (
MCPDuringCallRequestObject,
MCPDuringCallResponseObject,
MCPPostCallResponseObject,
MCPPreCallRequestObject,
MCPPreCallResponseObject,
@ -412,59 +410,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
#########################################################
# MCP TOOL CALL HOOKS
#########################################################
async def async_pre_mcp_tool_call_hook(
self,
kwargs,
request_obj: MCPPreCallRequestObject,
start_time,
end_time
) -> Optional[MCPPreCallResponseObject]:
"""
This hook gets called before the MCP tool call is made.
Useful for:
- Validating tool calls before execution
- Modifying arguments before they are sent to the MCP server
- Implementing access control and rate limiting
- Adding custom metadata or tracking information
Args:
kwargs: The logging kwargs containing model call details
request_obj: MCPPreCallRequestObject containing tool name, arguments, and metadata
start_time: Start time of the request
end_time: End time of the request
Returns:
MCPPreCallResponseObject with validation results and any modifications
"""
return None
async def async_during_mcp_tool_call_hook(
self,
kwargs,
request_obj: MCPDuringCallRequestObject,
start_time,
end_time
) -> Optional[MCPDuringCallResponseObject]:
"""
This hook gets called during the MCP tool call execution.
Useful for:
- Concurrent monitoring and validation during tool execution
- Implementing timeouts and cancellation logic
- Real-time cost tracking and billing
- Performance monitoring and metrics collection
Args:
kwargs: The logging kwargs containing model call details
request_obj: MCPDuringCallRequestObject containing tool execution context
start_time: Start time of the request
end_time: End time of the request
Returns:
MCPDuringCallResponseObject with execution control decisions
"""
return None
async def async_post_mcp_tool_call_hook(
self, kwargs, response_obj: MCPPostCallResponseObject, start_time, end_time
@ -595,3 +541,14 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
model_call_details_copy["standard_logging_object"] = standard_logging_object_copy
return model_call_details_copy
async def get_proxy_server_request_from_cold_storage_with_object_key(
self,
object_key: str,
) -> Optional[dict]:
"""
Get the proxy server request from cold storage using the object key directly.
"""
pass

View file

@ -304,7 +304,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
data=prepped.body,
headers=prepped.headers,
)
SigV4Auth(credentials, "s3", self.s3_region_name).add_auth(aws_request)
aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(
aws_region_name=self.s3_region_name
)
SigV4Auth(credentials, "s3", aws_region_name).add_auth(aws_request)
# Prepare the signed headers
signed_headers = dict(aws_request.headers.items())
@ -444,7 +447,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
data=prepped.body,
headers=prepped.headers,
)
SigV4Auth(credentials, "s3", self.s3_region_name).add_auth(aws_request)
aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(
aws_region_name=self.s3_region_name
)
SigV4Auth(credentials, "s3", aws_region_name).add_auth(aws_request)
# Prepare the signed headers
signed_headers = dict(aws_request.headers.items())
@ -455,3 +461,108 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
response.raise_for_status()
except Exception as e:
verbose_logger.exception(f"Error uploading to s3: {str(e)}")
async def _download_object_from_s3(self, s3_object_key: str) -> Optional[dict]:
"""
Download and parse JSON object from S3.
Args:
s3_object_key: The S3 object key to download
Returns:
Optional[dict]: The parsed JSON object or None if not found/error
"""
try:
import hashlib
import requests
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError("Missing boto3 to call S3. Run 'pip install boto3'.")
try:
from litellm.litellm_core_utils.asyncify import asyncify
# Get AWS credentials
asyncified_get_credentials = asyncify(self.get_credentials)
credentials = await asyncified_get_credentials(
aws_access_key_id=self.s3_aws_access_key_id,
aws_secret_access_key=self.s3_aws_secret_access_key,
aws_session_token=self.s3_aws_session_token,
aws_region_name=self.s3_region_name,
aws_session_name=self.s3_aws_session_name,
aws_profile_name=self.s3_aws_profile_name,
aws_role_name=self.s3_aws_role_name,
aws_web_identity_token=self.s3_aws_web_identity_token,
aws_sts_endpoint=self.s3_aws_sts_endpoint,
)
verbose_logger.debug(
f"s3_v2 logger - downloading data from s3 - {s3_object_key}"
)
# Prepare the URL
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{s3_object_key}"
if self.s3_endpoint_url:
url = self.s3_endpoint_url + "/" + s3_object_key
# Prepare the request for GET operation
# For GET requests, we need x-amz-content-sha256 with hash of empty string
empty_string_hash = hashlib.sha256(b"").hexdigest()
headers = {
"x-amz-content-sha256": empty_string_hash,
}
req = requests.Request("GET", url, headers=headers)
prepped = req.prepare()
# Sign the request
aws_request = AWSRequest(
method=prepped.method,
url=prepped.url,
headers=prepped.headers,
)
SigV4Auth(credentials, "s3", self.s3_region_name).add_auth(aws_request)
# Prepare the signed headers
signed_headers = dict(aws_request.headers.items())
# Make the request
response = await self.async_httpx_client.get(url, headers=signed_headers)
if response.status_code != 200:
verbose_logger.exception("S3 object not found, saw response=", response.text)
return None
# Parse JSON response
return response.json()
except Exception as e:
verbose_logger.exception(f"Error downloading from S3: {str(e)}")
return None
async def get_proxy_server_request_from_cold_storage_with_object_key(
self,
object_key: str,
) -> Optional[dict]:
"""
Get the proxy server request from cold storage
Allows fetching a dict of the proxy server request from s3 or GCS bucket.
Args:
request_id: The unique request ID to search for
start_time: The start time of the request (datetime or ISO string)
Returns:
Optional[dict]: The request data dictionary or None if not found
"""
try:
# Download and return the object from S3
downloaded_object = await self._download_object_from_s3(object_key)
return downloaded_object
except Exception as e:
verbose_logger.exception(f"Error retrieving object {object_key} from cold storage: {str(e)}")
return None

View file

@ -10,6 +10,7 @@ Example:
from typing import Union
from litellm import _custom_logger_compatible_callbacks_literal
from litellm.integrations.agentops import AgentOps
from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheControlHook
from litellm.integrations.argilla import ArgillaLogger
@ -150,3 +151,14 @@ class CustomLoggerRegistry:
if callback_class == class_type:
callback_strs.append(callback_str)
return callback_strs
@classmethod
def get_class_type_for_custom_logger_name(
cls,
custom_logger_name: _custom_logger_compatible_callbacks_literal,
) -> type:
"""
Get the class type for a given custom logger name
"""
return cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE[custom_logger_name]

View file

@ -356,6 +356,8 @@ def get_llm_provider( # noqa: PLR0915
# bytez models
elif model.startswith("bytez/"):
custom_llm_provider = "bytez"
elif model.startswith("oci/"):
custom_llm_provider = "oci"
if not custom_llm_provider:
if litellm.suppress_debug_info is False:
print() # noqa

View file

@ -146,7 +146,9 @@ def get_supported_openai_params( # noqa: PLR0915
return litellm.HuggingFaceChatConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "jina_ai":
if request_type == "embeddings":
return litellm.JinaAIEmbeddingConfig().get_supported_openai_params()
return litellm.JinaAIEmbeddingConfig().get_supported_openai_params(
model=model
)
elif custom_llm_provider == "together_ai":
return litellm.TogetherAIConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "databricks":

View file

@ -3830,6 +3830,8 @@ class StandardLoggingPayloadSetup:
] = None,
usage_object: Optional[dict] = None,
proxy_server_request: Optional[dict] = None,
start_time: Optional[dt_object] = None,
response_id: Optional[str] = None,
) -> StandardLoggingMetadata:
"""
Clean and filter the metadata dictionary to include only the specified keys in StandardLoggingMetadata.
@ -3881,6 +3883,7 @@ class StandardLoggingPayloadSetup:
usage_object=usage_object,
requester_custom_headers=None,
user_api_key_request_route=None,
cold_storage_object_key=None,
)
if isinstance(metadata, dict):
# Filter the metadata dictionary to include only the specified keys
@ -3913,6 +3916,16 @@ class StandardLoggingPayloadSetup:
proxy_server_request=proxy_server_request,
)
# Generate cold storage object key if cold storage is configured
if start_time is not None and response_id is not None:
cold_storage_object_key = StandardLoggingPayloadSetup._generate_cold_storage_object_key(
start_time=start_time,
response_id=response_id,
team_alias=clean_metadata.get("user_api_key_team_alias"),
)
if cold_storage_object_key:
clean_metadata["cold_storage_object_key"] = cold_storage_object_key
return clean_metadata
@staticmethod
@ -4071,6 +4084,49 @@ class StandardLoggingPayloadSetup:
return api_base.rstrip("/")
return api_base
@staticmethod
def _generate_cold_storage_object_key(
start_time: dt_object,
response_id: str,
team_alias: Optional[str] = None,
) -> Optional[str]:
"""
Generate cold storage object key in the same format as S3Logger.
Args:
start_time: The start time of the request
response_id: The response ID
team_alias: Optional team alias for team-based prefixing
Returns:
Optional[str]: The generated object key or None if cold storage not configured
"""
# Generate object key in same format as S3Logger
from litellm.integrations.s3 import get_s3_object_key
from litellm.proxy.spend_tracking.cold_storage_handler import ColdStorageHandler
# Only generate object key if cold storage is configured
configured_cold_storage_logger = ColdStorageHandler._get_configured_cold_storage_custom_logger()
if configured_cold_storage_logger is None:
return None
try:
# Generate file name in same format as litellm.utils.get_logging_id
s3_file_name = f"time-{start_time.strftime('%H-%M-%S-%f')}_{response_id}"
s3_object_key = get_s3_object_key(
s3_path="", # Use empty path as default
team_alias_prefix="", # Don't split by team alias for cold storage
start_time=start_time,
s3_file_name=s3_file_name,
)
return s3_object_key
except Exception:
# If any error occurs in generating the key, return None
return None
@staticmethod
def get_error_information(
original_exception: Optional[Exception],
@ -4322,6 +4378,8 @@ def get_standard_logging_object_payload(
),
usage_object=usage.model_dump(),
proxy_server_request=proxy_server_request,
start_time=start_time,
response_id=id,
)
_request_body = proxy_server_request.get("body", {})
@ -4469,6 +4527,7 @@ def get_standard_logging_metadata(
usage_object=None,
requester_custom_headers=None,
user_api_key_request_route=None,
cold_storage_object_key=None,
)
if isinstance(metadata, dict):
# Filter the metadata dictionary to include only the specified keys

View file

@ -1,4 +1,4 @@
from typing import Callable, List, Set, Type, Union
from typing import TYPE_CHECKING, Callable, List, Optional, Set, Type, Union
import litellm
from litellm._logging import verbose_logger
@ -6,6 +6,11 @@ from litellm.integrations.additional_logging_utils import AdditionalLoggingUtils
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import CallbacksByType
if TYPE_CHECKING:
from litellm import _custom_logger_compatible_callbacks_literal
else:
_custom_logger_compatible_callbacks_literal = str
class LoggingCallbackManager:
"""
@ -343,3 +348,26 @@ class LoggingCallbackManager:
elif callable(callback):
return getattr(callback, "__name__", str(callback))
return str(callback)
def get_active_custom_logger_for_callback_name(
self,
callback_name: _custom_logger_compatible_callbacks_literal,
) -> Optional[CustomLogger]:
"""
Get the active custom logger for a given callback name
"""
from litellm.litellm_core_utils.custom_logger_registry import (
CustomLoggerRegistry,
)
# get the custom logger class type
custom_logger_class_type = CustomLoggerRegistry.get_class_type_for_custom_logger_name(callback_name)
# get the active custom logger
custom_logger = self.get_custom_loggers_for_type(custom_logger_class_type)
if len(custom_logger) == 0:
raise ValueError(f"No active custom logger found for callback name: {callback_name}")
return custom_logger[0]

View file

@ -519,25 +519,25 @@ def unpack_defs(schema: dict, defs: dict) -> None:
}
# Use iterative approach with queue to avoid recursion
# Each item in queue is (node, parent_container, key/index, active_defs, seen_ids)
# Each item in queue is (node, parent_container, key/index, active_defs, ref_chain)
queue: deque[
tuple[Any, Union[dict, list, None], Union[str, int, None], dict, set]
] = deque([(schema, None, None, root_defs, set())])
while queue:
node, parent, key, active_defs, seen = queue.popleft()
# Avoid infinite loops on self-referential schemas
if id(node) in seen:
continue
seen = seen.copy() # Create new set for this branch
seen.add(id(node))
node, parent, key, active_defs, ref_chain = queue.popleft()
# ----------------------------- dict -----------------------------
if isinstance(node, dict):
# --- Case 1: this node *is* a reference ---
if "$ref" in node:
ref_name = node["$ref"].split("/")[-1]
# Check for circular reference in the resolution chain
if ref_name in ref_chain:
# Circular reference detected - leave as-is to prevent infinite recursion
continue
target_schema = active_defs.get(ref_name)
# Unknown reference – leave untouched
if target_schema is None:
@ -563,8 +563,12 @@ def unpack_defs(schema: dict, defs: dict) -> None:
schema.update(resolved)
resolved = schema
# Add to ref chain to track circular references
new_ref_chain = ref_chain.copy()
new_ref_chain.add(ref_name)
# Add resolved node to queue for further processing
queue.append((resolved, parent, key, child_defs, seen))
queue.append((resolved, parent, key, child_defs, new_ref_chain))
continue
# --- Case 2: regular dict – process its values ---
@ -577,13 +581,13 @@ def unpack_defs(schema: dict, defs: dict) -> None:
# Add all dict values to queue
for k, v in node.items():
queue.append((v, node, k, current_defs, seen))
queue.append((v, node, k, current_defs, ref_chain))
# ---------------------------- list ------------------------------
elif isinstance(node, list):
# Add all list items to queue
for idx, item in enumerate(node):
queue.append((item, node, idx, active_defs, seen))
queue.append((item, node, idx, active_defs, ref_chain))
def _get_image_mime_type_from_url(url: str) -> Optional[str]:

View file

@ -1,5 +1,6 @@
import copy
import json
import mimetypes
import re
import uuid
import xml.etree.ElementTree as ET
@ -13,6 +14,7 @@ import litellm.types
import litellm.types.llms
from litellm import verbose_logger
from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_async_httpx_client
from litellm.types.files import get_file_extension_from_mime_type
from litellm.types.llms.anthropic import *
from litellm.types.llms.bedrock import MessageBlock as BedrockMessageBlock
from litellm.types.llms.custom_http import httpxSpecialProvider
@ -2351,7 +2353,6 @@ def stringify_json_tool_call_content(messages: List) -> List:
###### AMAZON BEDROCK #######
import base64
import mimetypes
from email.message import Message
import httpx
@ -2479,20 +2480,11 @@ class BedrockImageProcessor:
)
if is_document:
potential_extensions = mimetypes.guess_all_extensions(mime_type)
valid_extensions = [
ext[1:]
for ext in potential_extensions
if ext[1:] in supported_doc_formats
]
return BedrockImageProcessor._get_document_format(
mime_type=mime_type,
supported_doc_formats=supported_doc_formats
)
if not valid_extensions:
raise ValueError(
f"No supported extensions for MIME type: {mime_type}. Supported formats: {supported_doc_formats}"
)
# Use first valid extension instead of provided image_format
return valid_extensions[0]
else:
#########################################################
# Check if image_format is an image or video
@ -2502,6 +2494,60 @@ class BedrockImageProcessor:
f"Unsupported image format: {image_format}. Supported formats: {supported_image_and_video_formats}"
)
return image_format
@staticmethod
def _get_document_format(
mime_type: str,
supported_doc_formats: List[str]
) -> str:
"""
Get the document format from the mime type
- Primary method - uses `mimetypes.guess_all_extensions`
- Fallback method - uses `get_file_extension_from_mime_type`
Relevant Issue: https://github.com/BerriAI/litellm/issues/12260
`mimetypes` is not available in docker containers, so we fallback to `get_file_extension_from_mime_type`
Args:
mime_type: The mime type of the document
supported_doc_formats: The supported document formats for the current model
Returns:
The document format
"""
valid_extensions: Optional[List[str]] = None
potential_extensions = mimetypes.guess_all_extensions(
mime_type, strict=False
)
valid_extensions = [
ext[1:]
for ext in potential_extensions
if ext[1:] in supported_doc_formats
]
# Fallback to types/files.py if mimetypes doesn't return valid extensions
#################
# litellm runs on docker containers and `mimetypes` depends on the installed mimetypes of the OS
# we fallback to well known mime types in types/files.py if mimetypes doesn't return valid extensions
if not valid_extensions:
try:
fallback_extension = get_file_extension_from_mime_type(mime_type)
if fallback_extension in supported_doc_formats:
valid_extensions = [fallback_extension]
except ValueError:
# Neither mimetypes nor files.py could handle this MIME type
# get_file_extension_from_mime_type raises ValueError if the mime type is not supported
pass
if not valid_extensions:
raise ValueError(
f"No supported extensions for MIME type: {mime_type}. Supported formats: {supported_doc_formats}"
)
# Use first valid extension instead of provided image_format
return valid_extensions[0]
@staticmethod
def _create_bedrock_block(
@ -2950,7 +2996,10 @@ def process_empty_text_blocks(
]
modified_message = message.copy()
modified_message["content"] = modified_content_block
modified_message["content"] = cast(
Union[List[ChatCompletionTextObject], List[ChatCompletionThinkingBlock]],
modified_content_block,
)
return modified_message

View file

@ -527,7 +527,12 @@ class ChunkProcessor:
returned_usage, "cache_read_input_tokens", cache_read_input_tokens
) # for anthropic
if completion_tokens_details is not None:
returned_usage.completion_tokens_details = completion_tokens_details
if isinstance(completion_tokens_details, CompletionTokensDetails):
returned_usage.completion_tokens_details = CompletionTokensDetailsWrapper(
**completion_tokens_details.model_dump()
)
else:
returned_usage.completion_tokens_details = completion_tokens_details
if reasoning_tokens is not None:
if returned_usage.completion_tokens_details is None:

View file

@ -1584,7 +1584,9 @@ class CustomStreamWrapper:
except StopIteration:
if self.sent_last_chunk is True:
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks, messages=self.messages
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
response = self.model_response_creator()
@ -1768,7 +1770,9 @@ class CustomStreamWrapper:
if self.sent_last_chunk is True:
# log the final chunk with accurate streaming values
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks, messages=self.messages
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
response = self.model_response_creator()
if complete_streaming_response is not None:

View file

@ -2,7 +2,7 @@
This file contains common utils for anthropic calls.
"""
from typing import Dict, List, Optional, Union
from typing import Any, Dict, List, Optional, Union
import httpx
@ -229,6 +229,60 @@ class AnthropicModelInfo(BaseLLMModelInfo):
litellm_model_names.append(litellm_model_name)
return litellm_model_names
def get_token_counter(self) -> Optional["AnthropicTokenCounter"]:
"""
Factory method to create an Anthropic token counter.
Returns:
AnthropicTokenCounter instance for this provider.
"""
return AnthropicTokenCounter()
class AnthropicTokenCounter:
"""Token counter implementation for Anthropic provider."""
def supports_provider(
self,
deployment: Optional[Dict[str, Any]] = None,
from_endpoint: bool = False
) -> bool:
if not from_endpoint:
return False
if deployment is None:
return False
full_model = deployment.get("litellm_params", {}).get("model", "")
is_anthropic_provider = full_model.startswith("anthropic/") or "anthropic" in full_model.lower()
return is_anthropic_provider
async def count_tokens(
self,
model_to_use: str,
messages: Optional[List[Dict[str, Any]]],
deployment: Optional[Dict[str, Any]] = None,
request_model: str = "",
) -> Optional[Dict[str, Any]]:
from litellm.proxy.utils import count_tokens_with_anthropic_api
result = await count_tokens_with_anthropic_api(
model_to_use=model_to_use,
messages=messages,
deployment=deployment,
)
if result is not None:
return {
"total_tokens": result["total_tokens"],
"request_model": request_model,
"model_used": model_to_use,
"tokenizer_type": result["tokenizer_used"],
}
return None
def process_anthropic_headers(headers: Union[httpx.Headers, dict]) -> dict:
openai_headers = {}

View file

@ -70,6 +70,16 @@ class BaseLLMModelInfo(ABC):
"""
pass
def get_token_counter(self):
"""
Factory method to create a token counter for this provider.
Returns:
Optional TokenCounterInterface implementation for this provider,
or None if token counting is not supported.
"""
return None
def _convert_tool_response_to_message(
tool_calls: List[ChatCompletionToolCallChunk],

View file

@ -371,8 +371,8 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
reasoning_content += sum["text"]
thinking_block = ChatCompletionThinkingBlock(
type="thinking",
thinking=sum["text"],
signature=sum["signature"],
thinking=sum.get("text", ""),
signature=sum.get("signature", ""),
)
if thinking_blocks is None:
thinking_blocks = []

View file

@ -0,0 +1,6 @@
from litellm.llms.base_llm.chat.transformation import BaseLLMException
class JinaAIError(BaseLLMException):
def __init__(self, status_code, message):
super().__init__(status_code=status_code, message=message)

View file

@ -1,5 +1,5 @@
"""
Transformation logic from OpenAI /v1/embeddings format to Jina AI's `/v1/embeddings` format.
Transformation logic from OpenAI /v1/embeddings format to Jina AI's `/v1/embeddings` format.
Why separate file? Make it easy to see how transformation works
@ -7,13 +7,23 @@ Docs - https://jina.ai/embeddings/
"""
import types
from typing import List, Optional, Tuple
from typing import List, Optional, Tuple, Union, cast
import httpx
from litellm import LlmProviders
from litellm.secret_managers.main import get_secret_str
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm import BaseEmbeddingConfig
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
from litellm.types.utils import EmbeddingResponse
from litellm.utils import is_base64_encoded
from ..common_utils import JinaAIError
class JinaAIEmbeddingConfig:
class JinaAIEmbeddingConfig(BaseEmbeddingConfig):
"""
Reference: https://jina.ai/embeddings/
"""
@ -44,11 +54,15 @@ class JinaAIEmbeddingConfig:
and v is not None
}
def get_supported_openai_params(self) -> List[str]:
def get_supported_openai_params(self, model: str) -> List[str]:
return ["dimensions"]
def map_openai_params(
self, non_default_params: dict, optional_params: dict
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
if "dimensions" in non_default_params:
optional_params["dimensions"] = non_default_params["dimensions"]
@ -76,3 +90,88 @@ class JinaAIEmbeddingConfig:
or get_secret_str("JINA_AI_TOKEN")
)
return LlmProviders.JINA_AI.value, api_base, dynamic_api_key
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:
return (
f"{api_base}/embeddings"
if api_base
else "https://api.jina.ai/v1/embeddings"
)
def transform_embedding_request(
self,
model: str,
input: AllEmbeddingInputValues,
optional_params: dict,
headers: dict,
) -> dict:
data = {"model": model, **optional_params}
input = cast(List[str], input) if isinstance(input, List) else [input]
if any((is_base64_encoded(x) for x in input)):
transformed_input = []
for value in input:
if isinstance(value, str):
if is_base64_encoded(value):
img_data = value.split(",")[1]
transformed_input.append({"image": img_data})
else:
transformed_input.append({"text": value})
data["input"] = transformed_input
else:
data["input"] = input
return data
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:
response_json = raw_response.json()
## LOGGING
logging_obj.post_call(
input=input,
api_key=api_key,
additional_args={"complete_input_dict": request_data},
original_response=response_json,
)
return EmbeddingResponse(**response_json)
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:
default_headers = {
"Content-Type": "application/json",
}
if api_key:
default_headers["Authorization"] = f"Bearer {api_key}"
headers = {**default_headers, **headers}
return headers
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BaseLLMException:
return JinaAIError(
status_code=status_code,
message=error_message,
)

View file

@ -0,0 +1,868 @@
import base64
import datetime
import hashlib
import json
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
from urllib.parse import urlparse
import httpx
import litellm
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
_get_httpx_client,
get_async_httpx_client,
version,
)
from litellm.llms.oci.common_utils import OCIError
from litellm.types.llms.oci import (
OCIChatRequestPayload,
OCICompletionPayload,
OCICompletionResponse,
OCIContentPartUnion,
OCIImageContentPart,
OCIMessage,
OCIRoles,
OCIServingMode,
OCIStreamChunk,
OCITextContentPart,
OCIToolCall,
OCIToolDefinition,
OCIVendors,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import (
Delta,
LlmProviders,
ModelResponseStream,
StreamingChoices,
)
from litellm.utils import (
ChatCompletionMessageToolCall,
CustomStreamWrapper,
ModelResponse,
Usage,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
def sha256_base64(data: bytes) -> str:
digest = hashlib.sha256(data).digest()
return base64.b64encode(digest).decode()
def build_signature_string(method, path, headers, signed_headers):
lines = []
for header in signed_headers:
if header == "(request-target)":
value = f"{method.lower()} {path}"
else:
value = headers[header]
lines.append(f"{header}: {value}")
return "\n".join(lines)
def load_private_key_from_str(key_str: str):
try:
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
except ImportError as e:
raise ImportError(
"cryptography package is required for OCI authentication. "
"Please install it with: pip install cryptography"
) from e
key = serialization.load_pem_private_key(
key_str.encode("utf-8"),
password=None,
)
if not isinstance(key, rsa.RSAPrivateKey):
raise TypeError(
"The provided private key is not an RSA key, which is required for OCI signing."
)
return key
def get_vendor_from_model(model: str) -> OCIVendors:
"""
Extracts the vendor from the model name.
Args:
model (str): The model name.
Returns:
str: The vendor name.
"""
vendor = model.split(".")[0].lower()
if vendor == "cohere":
return OCIVendors.COHERE
else:
return OCIVendors.GENERIC
# 5 minute timeout (models may need to load)
STREAMING_TIMEOUT = 60 * 5
class OCIChatConfig(BaseConfig):
"""
Configuration class for OCI's API interface.
"""
def __init__(
self,
) -> None:
locals_ = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)
# mark the class as using a custom stream wrapper because the default only iterates on lines
setattr(self.__class__, "has_custom_stream_wrapper", True)
self.openai_to_oci_generic_param_map = {
"stream": "isStream",
"max_tokens": "maxTokens",
"max_completion_tokens": "maxTokens",
"temperature": "temperature",
"tools": "tools",
"frequency_penalty": "frequencyPenalty",
"logprobs": "logProbs",
"logit_bias": "logitBias",
"n": "numGenerations",
"presence_penalty": "presencePenalty",
"seed": "seed",
"stop": "stop",
"tool_choice": "toolChoice",
"top_p": "topP",
"max_retries": False,
"top_logprobs": False,
"modalities": False,
"prediction": False,
"stream_options": False,
"function_call": False,
"functions": False,
"extra_headers": False,
"parallel_tool_calls": False,
"audio": False,
"web_search_options": False,
}
def get_supported_openai_params(self, model: str) -> List[str]:
supported_params = []
vendor = get_vendor_from_model(model)
if vendor == OCIVendors.COHERE:
raise ValueError(
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
)
else:
open_ai_to_oci_param_map = self.openai_to_oci_generic_param_map
for key, value in open_ai_to_oci_param_map.items():
if value:
supported_params.append(key)
return supported_params
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
adapted_params = {}
vendor = get_vendor_from_model(model)
if vendor == OCIVendors.COHERE:
raise ValueError(
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
)
else:
open_ai_to_oci_param_map = self.openai_to_oci_generic_param_map
all_params = {**non_default_params, **optional_params}
for key, value in all_params.items():
alias = open_ai_to_oci_param_map.get(key)
if alias is False:
if drop_params:
continue
raise Exception(f"param `{key}` is not supported on OCI")
if alias is None:
adapted_params[key] = value
continue
adapted_params[alias] = value
return adapted_params
def sign_request(
self,
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
api_key: Optional[str] = None,
model: Optional[str] = None,
stream: Optional[bool] = None,
fake_stream: Optional[bool] = None,
) -> Tuple[dict, Optional[bytes]]:
"""
Some providers like Bedrock require signing the request. The sign request funtion needs access to `request_data` and `complete_url`
Args:
headers: dict
optional_params: dict
request_data: dict - the request body being sent in http request
api_base: str - the complete url being sent in http request
Returns:
dict - the signed headers
"""
import json
oci_region = optional_params.get("oci_region", "us-ashburn-1")
api_base = (
api_base
or litellm.api_base
or f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com"
)
oci_user = optional_params.get("oci_user")
oci_fingerprint = optional_params.get("oci_fingerprint")
oci_tenancy = optional_params.get("oci_tenancy")
oci_key = optional_params.get("oci_key")
if not oci_user or not oci_fingerprint or not oci_tenancy or not oci_key:
raise Exception(
"Missing one of the following parameters: oci_user, oci_fingerprint, oci_tenancy, oci_key"
)
method = str(optional_params.get("method", "POST")).upper()
body = json.dumps(request_data).encode("utf-8")
parsed = urlparse(api_base)
path = parsed.path or "/"
host = parsed.netloc
date = datetime.datetime.utcnow().strftime("%a, %d %b %Y %H:%M:%S GMT")
content_type = headers.get("content-type", "application/json")
content_length = str(len(body))
x_content_sha256 = sha256_base64(body)
headers_to_sign = {
"date": date,
"host": host,
"content-type": content_type,
"content-length": content_length,
"x-content-sha256": x_content_sha256,
}
signed_headers = [
"date",
"(request-target)",
"host",
"content-length",
"content-type",
"x-content-sha256",
]
signing_string = build_signature_string(
method, path, headers_to_sign, signed_headers
)
try:
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.asymmetric import padding
except ImportError as e:
raise ImportError(
"cryptography package is required for OCI authentication. "
"Please install it with: pip install cryptography"
) from e
private_key = load_private_key_from_str(oci_key)
signature = private_key.sign(
signing_string.encode("utf-8"),
padding.PKCS1v15(),
hashes.SHA256(),
)
signature_b64 = base64.b64encode(signature).decode()
key_id = f"{oci_tenancy}/{oci_user}/{oci_fingerprint}"
authorization = (
'Signature version="1",'
f'keyId="{key_id}",'
'algorithm="rsa-sha256",'
f'headers="{" ".join(signed_headers)}",'
f'signature="{signature_b64}"'
)
headers.update(
{
"authorization": authorization,
"date": date,
"host": host,
"content-type": content_type,
"content-length": content_length,
"x-content-sha256": x_content_sha256,
}
)
return headers, None
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:
oci_region = optional_params.get("oci_region", "us-ashburn-1")
api_base = (
api_base
or litellm.api_base
or f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com"
)
oci_user = optional_params.get("oci_user")
oci_fingerprint = optional_params.get("oci_fingerprint")
oci_tenancy = optional_params.get("oci_tenancy")
oci_key = optional_params.get("oci_key")
oci_compartment_id = optional_params.get("oci_compartment_id")
if (
not oci_user
or not oci_fingerprint
or not oci_tenancy
or not oci_key
or not oci_compartment_id
):
raise Exception(
"Missing one of the following parameters: oci_user, oci_fingerprint, oci_tenancy, oci_key, oci_compartment_id"
)
if not api_base:
raise Exception(
"Either `api_base` must be provided or `litellm.api_base` must be set. Alternatively, you can set the `oci_region` optional parameter to use the default OCI region."
)
headers.update(
{
"content-type": "application/json",
"user-agent": f"litellm/{version}",
}
)
if not messages:
raise Exception(
"kwarg `messages` must be an array of messages that follow the openai chat standard"
)
return headers
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:
oci_region = optional_params.get("oci_region", "us-ashburn-1")
return f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com/20231130/actions/chat"
def _get_optional_params(self, vendor: OCIVendors, optional_params: dict) -> Dict:
selected_params = {}
if vendor == OCIVendors.COHERE:
raise ValueError(
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
)
else:
open_ai_to_oci_param_map = self.openai_to_oci_generic_param_map
for value in open_ai_to_oci_param_map.values():
if value in optional_params:
selected_params[value] = optional_params[value]
if "tools" in selected_params:
selected_params["tools"] = adapt_tool_definition_to_oci_standard(
selected_params["tools"], vendor
)
return selected_params
def transform_request(
self,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
oci_compartment_id = optional_params.get("oci_compartment_id", None)
if not oci_compartment_id:
raise Exception("kwarg `oci_compartment_id` is required for OCI requests")
vendor = get_vendor_from_model(model)
if vendor == OCIVendors.COHERE:
raise Exception(
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
)
else:
data = OCICompletionPayload(
compartmentId=oci_compartment_id,
servingMode=OCIServingMode(
servingType="ON_DEMAND",
modelId=model,
),
chatRequest=OCIChatRequestPayload(
apiFormat=vendor.value,
messages=adapt_messages_to_generic_oci_standard(messages),
**self._get_optional_params(vendor, optional_params),
),
)
return data.model_dump(exclude_none=True)
def transform_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ModelResponse:
json = raw_response.json() # noqa: F811
error = json.get("error")
if error is not None:
raise OCIError(
message=str(json["error"]),
status_code=raw_response.status_code,
)
if not isinstance(json, dict):
raise OCIError(
message="Invalid response format from OCI",
status_code=raw_response.status_code,
)
try:
completion_response = OCICompletionResponse(**json)
except TypeError as e:
raise OCIError(
message=f"Response cannot be casted to OCICompletionResponse: {str(e)}",
status_code=raw_response.status_code,
)
vendor = get_vendor_from_model(model)
if vendor == OCIVendors.COHERE:
raise ValueError(
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
)
else:
iso_str = completion_response.chatResponse.timeCreated
dt = datetime.datetime.fromisoformat(iso_str.replace("Z", "+00:00"))
model_response.created = int(dt.timestamp())
model_response.model = completion_response.modelId
message = model_response.choices[0].message # type: ignore
if vendor == OCIVendors.COHERE:
raise ValueError(
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
)
else:
response_message = completion_response.chatResponse.choices[0].message
if response_message.content and response_message.content[0].type == "TEXT":
message.content = response_message.content[0].text
if response_message.toolCalls:
message.tool_calls = adapt_tools_to_openai_standard(
response_message.toolCalls
)
usage = Usage(
prompt_tokens=completion_response.chatResponse.usage.promptTokens,
completion_tokens=completion_response.chatResponse.usage.completionTokens,
total_tokens=completion_response.chatResponse.usage.totalTokens,
)
model_response.usage = usage # type: ignore
model_response._hidden_params["additional_headers"] = raw_response.headers
return model_response
@track_llm_api_timing()
def get_sync_custom_stream_wrapper(
self,
model: str,
custom_llm_provider: str,
logging_obj: LiteLLMLoggingObj,
api_base: str,
headers: dict,
data: dict,
messages: list,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
json_mode: Optional[bool] = None,
signed_json_body: Optional[bytes] = None,
) -> "OCIStreamWrapper":
if "stream" in data:
del data["stream"]
if client is None or isinstance(client, AsyncHTTPHandler):
client = _get_httpx_client(params={})
try:
response = client.post(
api_base,
headers=headers,
data=json.dumps(data),
stream=True,
logging_obj=logging_obj,
timeout=STREAMING_TIMEOUT,
)
except httpx.HTTPStatusError as e:
raise OCIError(status_code=e.response.status_code, message=e.response.text)
if response.status_code != 200:
raise OCIError(status_code=response.status_code, message=response.text)
completion_stream = response.iter_text()
streaming_response = OCIStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
)
return streaming_response
@track_llm_api_timing()
async def get_async_custom_stream_wrapper(
self,
model: str,
custom_llm_provider: str,
logging_obj: LiteLLMLoggingObj,
api_base: str,
headers: dict,
data: dict,
messages: list,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
json_mode: Optional[bool] = None,
signed_json_body: Optional[bytes] = None,
) -> "OCIStreamWrapper":
if "stream" in data:
del data["stream"]
if client is None or isinstance(client, HTTPHandler):
client = get_async_httpx_client(llm_provider=LlmProviders.BYTEZ, params={})
try:
response = await client.post(
api_base,
headers=headers,
data=json.dumps(data),
stream=True,
logging_obj=logging_obj,
timeout=STREAMING_TIMEOUT,
)
except httpx.HTTPStatusError as e:
raise OCIError(status_code=e.response.status_code, message=e.response.text)
if response.status_code != 200:
raise OCIError(status_code=response.status_code, message=response.text)
completion_stream = response.aiter_text()
streaming_response = OCIStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
)
return streaming_response
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BaseLLMException:
return OCIError(status_code=status_code, message=error_message)
open_ai_to_generic_oci_role_map: Dict[str, OCIRoles] = {
"system": "SYSTEM",
"user": "USER",
"assistant": "ASSISTANT",
"tool": "TOOL",
}
def adapt_messages_to_generic_oci_standard_content_message(
role: str, content: Union[str, list]
) -> OCIMessage:
new_content: List[OCIContentPartUnion] = []
if isinstance(content, str):
return OCIMessage(
role=open_ai_to_generic_oci_role_map[role],
content=[OCITextContentPart(text=content)],
toolCalls=None,
toolCallId=None,
)
# content is a list of content items:
# [
# {"type": "text", "text": "Hello"},
# {"type": "image_url", "image_url": "https://example.com/image.png"}
# ]
for content_item in content:
if not isinstance(content_item, dict):
raise Exception("Each content item must be a dictionary")
type = content_item.get("type")
if not isinstance(type, str):
raise Exception("Prop `type` is not a string")
if type not in ["text", "image_url"]:
raise Exception(f"Prop `{type}` is not supported")
if type == "text":
text = content_item.get("text")
if not isinstance(text, str):
raise Exception("Prop `text` is not a string")
new_content.append(OCITextContentPart(text=text))
elif type == "image_url":
image_url = content_item.get("image_url")
if not isinstance(image_url, str):
raise Exception("Prop `image_url` is not a string")
new_content.append(OCIImageContentPart(imageUrl=image_url))
return OCIMessage(
role=open_ai_to_generic_oci_role_map[role],
content=new_content,
toolCalls=None,
toolCallId=None,
)
def adapt_messages_to_generic_oci_standard_tool_call(
role: str, tool_calls: list
) -> OCIMessage:
tool_calls_formated = []
for tool_call in tool_calls:
if not isinstance(tool_call, dict):
raise Exception("Each tool call must be a dictionary")
if tool_call.get("type") != "function":
raise Exception("OCI only supports function tools")
tool_call_id = tool_call.get("id")
if not isinstance(tool_call_id, str):
raise Exception("Prop `id` is not a string")
tool_function = tool_call.get("function")
if not isinstance(tool_function, dict):
raise Exception("Prop `function` is not a dictionary")
function_name = tool_function.get("name")
if not isinstance(function_name, str):
raise Exception("Prop `name` is not a string")
arguments = tool_call["function"].get("arguments", "{}")
if not isinstance(arguments, str):
raise Exception("Prop `arguments` is not a string")
# tool_calls_formated.append(OCIToolCall(
# id=tool_call_id,
# type="FUNCTION",
# function=OCIFunction(
# name=function_name,
# arguments=arguments
# )
# ))
tool_calls_formated.append(
OCIToolCall(
id=tool_call_id,
type="FUNCTION",
name=function_name,
arguments=arguments,
)
)
return OCIMessage(
role=open_ai_to_generic_oci_role_map[role],
content=None,
toolCalls=tool_calls_formated,
toolCallId=None,
)
def adapt_messages_to_generic_oci_standard_tool_response(
role: str, tool_call_id: str, content: str
) -> OCIMessage:
return OCIMessage(
role=open_ai_to_generic_oci_role_map[role],
content=[OCITextContentPart(text=content)],
toolCalls=None,
toolCallId=tool_call_id,
)
def adapt_messages_to_generic_oci_standard(
messages: List[AllMessageValues],
) -> List[OCIMessage]:
new_messages = []
for message in messages:
role = message["role"]
content = message.get("content")
tool_calls = message.get("tool_calls")
tool_call_id = message.get("tool_call_id")
if role in ["system", "user", "assistant"] and content is not None:
if not isinstance(content, (str, list)):
raise Exception(
"Prop `content` must be a string or a list of content items"
)
new_messages.append(
adapt_messages_to_generic_oci_standard_content_message(role, content)
)
elif role == "assistant" and tool_calls is not None:
if not isinstance(tool_calls, list):
raise Exception("Prop `tool_calls` must be a list of tool calls")
new_messages.append(
adapt_messages_to_generic_oci_standard_tool_call(role, tool_calls)
)
elif role == "tool":
if not isinstance(tool_call_id, str):
raise Exception("Prop `tool_call_id` is required and must be a string")
if not isinstance(content, str):
raise Exception("Prop `content` is not a string")
new_messages.append(
adapt_messages_to_generic_oci_standard_tool_response(
role, tool_call_id, content
)
)
return new_messages
def adapt_tool_definition_to_oci_standard(tools: List[Dict], vendor: OCIVendors):
new_tools = []
if vendor == OCIVendors.COHERE:
raise ValueError(
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
)
else:
for tool in tools:
if tool["type"] != "function":
raise Exception("OCI only supports function tools")
tool_function = tool.get("function")
if not isinstance(tool_function, dict):
raise Exception("Prop `function` is not a dictionary")
new_tool = OCIToolDefinition(
type="FUNCTION",
name=tool_function.get("name"),
description=tool_function.get("description", ""),
parameters=tool_function.get("parameters", {}),
)
new_tools.append(new_tool)
return new_tools
def adapt_tools_to_openai_standard(
tools: List[OCIToolCall],
) -> List[ChatCompletionMessageToolCall]:
new_tools = []
for tool in tools:
new_tool = ChatCompletionMessageToolCall(
id=tool.id,
type="function",
function={
"name": tool.name,
"arguments": tool.arguments,
},
)
new_tools.append(new_tool)
return new_tools
class OCIStreamWrapper(CustomStreamWrapper):
"""
Custom stream wrapper for OCI responses.
This class is used to handle streaming responses from OCI's API.
"""
def __init__(
self,
**kwargs: Any,
):
super().__init__(**kwargs)
def chunk_creator(self, chunk: Any):
if not isinstance(chunk, str):
raise ValueError(f"Chunk is not a string: {chunk}")
if not chunk.startswith("data:"):
raise ValueError(f"Chunk does not start with 'data:': {chunk}")
dict_chunk = json.loads(chunk[5:]) # Remove 'data: ' prefix and parse JSON
try:
typed_chunk = OCIStreamChunk(**dict_chunk)
except TypeError as e:
raise ValueError(f"Chunk cannot be casted to OCIStreamChunk: {str(e)}")
if typed_chunk.index is None:
typed_chunk.index = 0
text = ""
if typed_chunk.message and typed_chunk.message.content:
for item in typed_chunk.message.content:
if isinstance(item, OCITextContentPart):
text += item.text
elif isinstance(item, OCIImageContentPart):
raise ValueError(
"OCI does not support image content in streaming responses"
)
else:
raise ValueError(
f"Unsupported content type in OCI response: {item.type}"
)
tool_calls = None
if typed_chunk.message and typed_chunk.message.toolCalls:
tool_calls = adapt_tools_to_openai_standard(typed_chunk.message.toolCalls)
return ModelResponseStream(
choices=[
StreamingChoices(
index=typed_chunk.index if typed_chunk.index else 0,
delta=Delta(
content=text,
tool_calls=(
[tool.model_dump() for tool in tool_calls]
if tool_calls
else None
),
provider_specific_fields=None, # OCI does not have provider specific fields in the response
thinking_blocks=None, # OCI does not have thinking blocks in the response
reasoning_content=None, # OCI does not have reasoning content in the response
),
finish_reason=typed_chunk.finishReason,
)
]
)

View file

@ -0,0 +1,19 @@
from typing import Optional
import httpx
from litellm.llms.base_llm.chat.transformation import BaseLLMException
class OCIError(BaseLLMException):
def __init__(
self,
status_code: int,
message: str,
headers: Optional[httpx.Headers] = None,
):
super().__init__(
status_code=status_code,
message=message,
headers=headers,
)

View file

@ -4,6 +4,9 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
_safe_convert_created_field,
)
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import *
@ -11,7 +14,6 @@ from litellm.types.responses.main import *
from litellm.types.router import GenericLiteLLMParams
from ..common_utils import OpenAIError
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import _safe_convert_created_field
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -47,6 +49,8 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
"top_p",
"truncation",
"user",
"service_tier",
"safety_identifier",
"extra_headers",
"extra_query",
"extra_body",

View file

@ -14,6 +14,7 @@ from litellm.types.vector_stores import (
VectorStoreSearchRequest,
VectorStoreSearchResponse,
)
from litellm.utils import add_openai_metadata
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -119,12 +120,13 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig):
api_base: str,
) -> Tuple[str, Dict]:
url = api_base # Base URL for creating vector stores
metadata = vector_store_create_optional_params.get("metadata", None)
typed_request_body = VectorStoreCreateRequest(
name=vector_store_create_optional_params.get("name", None),
file_ids=vector_store_create_optional_params.get("file_ids", None),
expires_after=vector_store_create_optional_params.get("expires_after", None),
chunking_strategy=vector_store_create_optional_params.get("chunking_strategy", None),
metadata=vector_store_create_optional_params.get("metadata", None),
metadata=add_openai_metadata(metadata) if metadata is not None else None,
)
dict_request_body = cast(dict, typed_request_body)

View file

@ -107,6 +107,7 @@ from litellm.utils import (
supports_httpx_timeout,
token_counter,
validate_and_fix_openai_messages,
validate_and_fix_openai_tools,
validate_chat_completion_tool_choice,
)
@ -151,6 +152,7 @@ from .llms.gemini.common_utils import get_api_key_from_env
from .llms.groq.chat.handler import GroqChatCompletion
from .llms.huggingface.embedding.handler import HuggingFaceEmbedding
from .llms.nlp_cloud.chat.handler import completion as nlp_cloud_chat_completion
from .llms.oci.chat.transformation import OCIChatConfig
from .llms.ollama.completion import handler as ollama
from .llms.oobabooga.chat import oobabooga
from .llms.openai.completion.handler import OpenAITextCompletion
@ -252,6 +254,7 @@ base_llm_http_handler = BaseLLMHTTPHandler()
base_llm_aiohttp_handler = BaseLLMAIOHTTPHandler()
sagemaker_chat_completion = SagemakerChatHandler()
bytez_transformation = BytezChatConfig()
oci_transformation = OCIChatConfig()
####### COMPLETION ENDPOINTS ################
@ -963,6 +966,7 @@ def completion( # type: ignore # noqa: PLR0915
raise ValueError("model param not passed in.")
# validate messages
messages = validate_and_fix_openai_messages(messages=messages)
tools = validate_and_fix_openai_tools(tools=tools)
# validate tool_choice
tool_choice = validate_chat_completion_tool_choice(tool_choice=tool_choice)
######### unpacking kwargs #####################
@ -2399,6 +2403,24 @@ def completion( # type: ignore # noqa: PLR0915
encoding=encoding,
stream=stream,
)
elif custom_llm_provider == "oci":
response = base_llm_http_handler.completion(
model=model,
messages=messages,
headers=headers,
model_response=model_response,
api_key=api_key,
api_base=api_base,
acompletion=acompletion,
logging_obj=logging,
optional_params=optional_params,
litellm_params=litellm_params,
timeout=timeout, # type: ignore
client=client,
custom_llm_provider=custom_llm_provider,
encoding=encoding,
stream=stream,
)
elif custom_llm_provider == "oobabooga":
custom_llm_provider = "oobabooga"
model_response = oobabooga.completion(
@ -3880,7 +3902,6 @@ def embedding( # noqa: PLR0915
)
elif (
custom_llm_provider == "openai_like"
or custom_llm_provider == "jina_ai"
or custom_llm_provider == "hosted_vllm"
or custom_llm_provider == "llamafile"
or custom_llm_provider == "lm_studio"
@ -4307,6 +4328,25 @@ def embedding( # noqa: PLR0915
client=client,
aembedding=aembedding,
)
elif custom_llm_provider == "jina_ai":
if isinstance(input, str):
transformed_input = [input]
else:
transformed_input = input
response = base_llm_http_handler.embedding(
model=model,
input=transformed_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,
litellm_params={},
client=client,
aembedding=aembedding,
)
elif custom_llm_provider in litellm._custom_providers:
custom_handler: Optional[CustomLLM] = None
for item in litellm.custom_provider_map:
@ -5682,7 +5722,11 @@ def stream_chunk_builder_text_completion(
def stream_chunk_builder( # noqa: PLR0915
chunks: list, messages: Optional[list] = None, start_time=None, end_time=None
chunks: list,
messages: Optional[list] = None,
start_time=None,
end_time=None,
logging_obj: Optional[Logging] = None,
) -> Optional[Union[ModelResponse, TextCompletionResponse]]:
try:
if chunks is None:
@ -5807,6 +5851,12 @@ def stream_chunk_builder( # noqa: PLR0915
setattr(response, "usage", usage)
# Add cost to usage object if include_cost_in_streaming_usage is True
if litellm.include_cost_in_streaming_usage and logging_obj is not None:
setattr(
usage, "cost", logging_obj._response_cost_calculator(result=response)
)
return response
except Exception as e:
verbose_logger.exception(

View file

@ -607,9 +607,9 @@
"supports_system_messages": true,
"supports_tool_choice": true,
"search_context_cost_per_query": {
"search_context_size_low": 30.0,
"search_context_size_medium": 35.0,
"search_context_size_high": 50.0
"search_context_size_low": 0.025,
"search_context_size_medium": 0.0275,
"search_context_size_high": 0.03
}
},
"codex-mini-latest": {
@ -3750,7 +3750,7 @@
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 3.3e-06,
"output_cost_per_token": 16.5e-06,
"output_cost_per_token": 1.65e-05,
"litellm_provider": "azure_ai",
"mode": "chat",
"supports_function_calling": true,
@ -3764,7 +3764,7 @@
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 3e-06,
"output_cost_per_token": 15e-06,
"output_cost_per_token": 1.5e-05,
"litellm_provider": "azure_ai",
"mode": "chat",
"supports_function_calling": true,
@ -3777,7 +3777,7 @@
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 0.25e-06,
"input_cost_per_token": 2.5e-07,
"output_cost_per_token": 1.27e-06,
"litellm_provider": "azure_ai",
"mode": "chat",
@ -3792,7 +3792,7 @@
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 0.275e-06,
"input_cost_per_token": 2.75e-07,
"output_cost_per_token": 1.38e-06,
"litellm_provider": "azure_ai",
"mode": "chat",
@ -5486,6 +5486,36 @@
"litellm_provider": "groq",
"mode": "audio_transcription"
},
"groq/openai/gpt-oss-20b": {
"max_tokens": 32768,
"max_input_tokens": 131072,
"max_output_tokens": 32768,
"input_cost_per_token": 1e-07,
"output_cost_per_token": 5e-07,
"litellm_provider": "groq",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_response_schema": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_web_search": true
},
"groq/openai/gpt-oss-120b": {
"max_tokens": 32766,
"max_input_tokens": 131072,
"max_output_tokens": 32766,
"input_cost_per_token": 1.5e-07,
"output_cost_per_token": 7.5e-07,
"litellm_provider": "groq",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_response_schema": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_web_search": true
},
"cerebras/llama3.1-8b": {
"max_tokens": 128000,
"max_input_tokens": 128000,
@ -5741,6 +5771,32 @@
"supports_reasoning": true,
"supports_computer_use": true
},
"claude-opus-4-1-20250805": {
"max_tokens": 32000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 1.5e-05,
"output_cost_per_token": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01,
"search_context_size_high": 0.01
},
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "anthropic",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 159,
"supports_assistant_prefill": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_computer_use": true
},
"claude-sonnet-4-20250514": {
"max_tokens": 64000,
"max_input_tokens": 200000,
@ -7337,12 +7393,12 @@
"max_audio_length_hours": 8.4,
"max_audio_per_prompt": 1,
"max_pdf_size_mb": 30,
"input_cost_per_token": 3.5e-07,
"input_cost_per_token": 3.5e-07,
"input_cost_per_audio_token": 2.1e-06,
"input_cost_per_image": 2.1e-06,
"input_cost_per_video_per_second": 2.1e-06,
"output_cost_per_token": 1.5e-06,
"output_cost_per_audio_token": 8.5e-06,
"output_cost_per_audio_token": 8.5e-06,
"litellm_provider": "gemini",
"mode": "chat",
"rpm": 10,
@ -8690,6 +8746,40 @@
"source": "https://aistudio.google.com",
"supports_tool_choice": true
},
"vertex_ai/claude-opus-4-1": {
"max_tokens": 4096,
"max_input_tokens": 200000,
"max_output_tokens": 4096,
"input_cost_per_token": 15e-06,
"output_cost_per_token": 75e-06,
"input_cost_per_token_batches": 7.5e-06,
"output_cost_per_token_batches": 37.5e-06,
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"supports_assistant_prefill": true,
"supports_tool_choice": true
},
"vertex_ai/claude-opus-4-1@20250805": {
"max_tokens": 4096,
"max_input_tokens": 200000,
"max_output_tokens": 4096,
"input_cost_per_token": 15e-06,
"output_cost_per_token": 75e-06,
"input_cost_per_token_batches": 7.5e-06,
"output_cost_per_token_batches": 37.5e-06,
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"supports_assistant_prefill": true,
"supports_tool_choice": true
},
"vertex_ai/claude-3-sonnet": {
"max_tokens": 4096,
"max_input_tokens": 200000,
@ -9039,9 +9129,9 @@
"supports_tool_choice": true
},
"vertex_ai/meta/llama-4-scout-17b-16e-instruct-maas": {
"max_tokens": 10000000.0,
"max_input_tokens": 10000000.0,
"max_output_tokens": 10000000.0,
"max_tokens": 10000000,
"max_input_tokens": 10000000,
"max_output_tokens": 10000000,
"input_cost_per_token": 2.5e-07,
"output_cost_per_token": 7e-07,
"litellm_provider": "vertex_ai-llama_models",
@ -9059,9 +9149,9 @@
]
},
"vertex_ai/meta/llama-4-scout-17b-128e-instruct-maas": {
"max_tokens": 10000000.0,
"max_input_tokens": 10000000.0,
"max_output_tokens": 10000000.0,
"max_tokens": 10000000,
"max_input_tokens": 10000000,
"max_output_tokens": 10000000,
"input_cost_per_token": 2.5e-07,
"output_cost_per_token": 7e-07,
"litellm_provider": "vertex_ai-llama_models",
@ -9079,9 +9169,9 @@
]
},
"vertex_ai/meta/llama-4-maverick-17b-128e-instruct-maas": {
"max_tokens": 1000000.0,
"max_input_tokens": 1000000.0,
"max_output_tokens": 1000000.0,
"max_tokens": 1000000,
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"input_cost_per_token": 3.5e-07,
"output_cost_per_token": 1.15e-06,
"litellm_provider": "vertex_ai-llama_models",
@ -9099,9 +9189,9 @@
]
},
"vertex_ai/meta/llama-4-maverick-17b-16e-instruct-maas": {
"max_tokens": 1000000.0,
"max_input_tokens": 1000000.0,
"max_output_tokens": 1000000.0,
"max_tokens": 1000000,
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"input_cost_per_token": 3.5e-07,
"output_cost_per_token": 1.15e-06,
"litellm_provider": "vertex_ai-llama_models",
@ -9174,7 +9264,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 2048,
"input_cost_per_token": 5e-06,
"output_cost_per_token": 16e-06,
"output_cost_per_token": 1.6e-05,
"litellm_provider": "vertex_ai-llama_models",
"mode": "chat",
"supports_system_messages": true,
@ -10480,7 +10570,7 @@
"supports_tool_choice": true,
"supports_prompt_caching": true
},
"openrouter/x-ai/grok-4":{
"openrouter/x-ai/grok-4": {
"max_tokens": 256000,
"max_input_tokens": 256000,
"max_output_tokens": 256000,
@ -10494,12 +10584,12 @@
"source": "https://openrouter.ai/x-ai/grok-4",
"supports_web_search": true
},
"openrouter/bytedance/ui-tars-1.5-7b":{
"openrouter/bytedance/ui-tars-1.5-7b": {
"max_tokens": 2048,
"max_input_tokens": 131072,
"max_output_tokens": 2048,
"input_cost_per_token": 0.1e-06,
"output_cost_per_token": 0.2e-06,
"input_cost_per_token": 1e-07,
"output_cost_per_token": 2e-07,
"litellm_provider": "openrouter",
"mode": "chat",
"source": "https://openrouter.ai/api/v1/models/bytedance/ui-tars-1.5-7b",
@ -11178,8 +11268,8 @@
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 2048,
"input_cost_per_token": 0.21e-06,
"output_cost_per_token": 0.63e-06,
"input_cost_per_token": 2.1e-07,
"output_cost_per_token": 6.3e-07,
"litellm_provider": "openrouter",
"mode": "chat",
"supports_tool_choice": true
@ -11891,6 +11981,60 @@
"supports_pdf_input": true,
"supports_tool_choice": true
},
"openai.gpt-oss-20b-1:0": {
"max_tokens": 128000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 7e-08,
"output_cost_per_token": 3e-07,
"litellm_provider": "bedrock_converse",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true
},
"openai.gpt-oss-120b-1:0": {
"max_tokens": 128000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 1.5e-07,
"output_cost_per_token": 6e-07,
"litellm_provider": "bedrock_converse",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true
},
"anthropic.claude-opus-4-1-20250805-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 1.5e-05,
"output_cost_per_token": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01,
"search_context_size_high": 0.01
},
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "bedrock_converse",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 159,
"supports_assistant_prefill": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_computer_use": true
},
"anthropic.claude-opus-4-20250514-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
@ -12093,6 +12237,32 @@
"supports_tool_choice": true,
"supports_reasoning": true
},
"us.anthropic.claude-opus-4-1-20250805-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 1.5e-05,
"output_cost_per_token": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01,
"search_context_size_high": 0.01
},
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "bedrock_converse",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 159,
"supports_assistant_prefill": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_computer_use": true
},
"us.anthropic.claude-opus-4-20250514-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
@ -12266,6 +12436,32 @@
"supports_pdf_input": true,
"supports_tool_choice": true
},
"eu.anthropic.claude-opus-4-1-20250805-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 1.5e-05,
"output_cost_per_token": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01,
"search_context_size_high": 0.01
},
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "bedrock_converse",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 159,
"supports_assistant_prefill": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_computer_use": true
},
"eu.anthropic.claude-opus-4-20250514-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
@ -14763,7 +14959,7 @@
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 16384,
"input_cost_per_token": 0.6e-06,
"input_cost_per_token": 6e-07,
"output_cost_per_token": 2.5e-06,
"litellm_provider": "fireworks_ai",
"mode": "chat",
@ -14809,6 +15005,58 @@
"source": "https://fireworks.ai/pricing",
"supports_tool_choice": false
},
"fireworks_ai/accounts/fireworks/models/glm-4p5": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 96000,
"input_cost_per_token": 5.5e-07,
"output_cost_per_token": 2.19e-06,
"litellm_provider": "fireworks_ai",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"source": "https://fireworks.ai/models/fireworks/glm-4p5"
},
"fireworks_ai/accounts/fireworks/models/glm-4p5-air": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 96000,
"input_cost_per_token": 2.2e-07,
"output_cost_per_token": 8.8e-07,
"litellm_provider": "fireworks_ai",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"source": "https://artificialanalysis.ai/models/glm-4-5-air"
},
"fireworks_ai/accounts/fireworks/models/gpt-oss-120b": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 1.5e-07,
"output_cost_per_token": 6e-07,
"litellm_provider": "fireworks_ai",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"source": "https://fireworks.ai/pricing"
},
"fireworks_ai/accounts/fireworks/models/gpt-oss-20b": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 5e-08,
"output_cost_per_token": 2e-07,
"litellm_provider": "fireworks_ai",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"source": "https://fireworks.ai/pricing"
},
"fireworks_ai/nomic-ai/nomic-embed-text-v1.5": {
"max_tokens": 8192,
"max_input_tokens": 8192,

View file

@ -492,7 +492,10 @@ class MCPServerManager:
)
tools = await self._fetch_tools_with_timeout(client, server.name)
return self._create_prefixed_tools(tools, server)
prefixed_tools = self._create_prefixed_tools(tools, server)
return prefixed_tools
except Exception as e:
verbose_logger.warning(
@ -523,6 +526,7 @@ class MCPServerManager:
async def _list_tools_task():
try:
await client.connect()
tools = await client.list_tools()
verbose_logger.debug(f"Tools from {server_name}: {tools}")
return tools
@ -644,6 +648,7 @@ class MCPServerManager:
#########################################################
# Pre MCP Tool Call Hook
# Allow validation and modification of tool calls before execution
# Using standard pre_call_hook with call_type="mcp_call"
#########################################################
if proxy_logging_obj:
pre_hook_kwargs = {
@ -651,24 +656,32 @@ class MCPServerManager:
"arguments": arguments,
"server_name": server_name_from_prefix,
"user_api_key_auth": user_api_key_auth,
"user_api_key_user_id": getattr(user_api_key_auth, 'user_id', None) if user_api_key_auth else None,
"user_api_key_team_id": getattr(user_api_key_auth, 'team_id', None) if user_api_key_auth else None,
"user_api_key_end_user_id": getattr(user_api_key_auth, 'end_user_id', None) if user_api_key_auth else None,
"user_api_key_hash": getattr(user_api_key_auth, 'api_key_hash', None) if user_api_key_auth else None,
}
# Create MCP request object for processing
mcp_request_obj = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs)
# Convert to LLM format for existing guardrail compatibility
synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs)
try:
pre_hook_result = await proxy_logging_obj.async_pre_mcp_tool_call_hook(
kwargs=pre_hook_kwargs,
request_obj=None, # Will be created in the hook
start_time=start_time,
end_time=start_time,
# Use standard pre_call_hook with call_type="mcp_call"
modified_data = await proxy_logging_obj.pre_call_hook(
user_api_key_dict=user_api_key_auth, #type: ignore
data=synthetic_llm_data,
call_type="mcp_call" #type: ignore
)
if pre_hook_result:
# Apply any argument modifications
if pre_hook_result.get("modified_arguments"):
arguments = pre_hook_result["modified_arguments"]
except (
BlockedPiiEntityError,
GuardrailRaisedException,
HTTPException,
) as e:
if modified_data:
# Convert response back to MCP format and apply modifications
modified_kwargs = proxy_logging_obj._convert_mcp_hook_response_to_kwargs(modified_data, pre_hook_kwargs)
if modified_kwargs.get("arguments") != arguments:
arguments = modified_kwargs["arguments"]
except (BlockedPiiEntityError, GuardrailRaisedException, HTTPException) as e:
# Re-raise guardrail exceptions to properly fail the MCP call
verbose_logger.error(
f"Guardrail blocked MCP tool call pre call: {str(e)}"
@ -699,22 +712,34 @@ class MCPServerManager:
name=original_tool_name,
arguments=arguments,
)
# Initialize during_hook_task as None
during_hook_task = None
tasks = []
# Start during hook if proxy_logging_obj is available
if proxy_logging_obj:
# Create synthetic LLM data for during hook processing
from litellm.types.mcp import MCPDuringCallRequestObject
from litellm.types.llms.base import HiddenParams
request_obj = MCPDuringCallRequestObject(
tool_name=name,
arguments=arguments,
server_name=server_name_from_prefix,
start_time=start_time.timestamp() if start_time else None,
hidden_params=HiddenParams(),
)
during_hook_kwargs = {
"name": name,
"arguments": arguments,
"server_name": server_name_from_prefix,
"user_api_key_auth": user_api_key_auth,
}
synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(request_obj, during_hook_kwargs)
during_hook_task = asyncio.create_task(
proxy_logging_obj.async_during_mcp_tool_call_hook(
kwargs={
"name": name,
"arguments": arguments,
"server_name": server_name_from_prefix,
},
request_obj=None, # Will be created in the hook
start_time=start_time,
end_time=start_time,
proxy_logging_obj.during_call_hook(
user_api_key_dict=user_api_key_auth,
data=synthetic_llm_data,
call_type="mcp_call" #type: ignore
)
)
tasks.append(during_hook_task)
@ -809,20 +834,28 @@ class MCPServerManager:
get_prisma_client_or_throw,
)
verbose_logger.info("Loading MCP servers from database into registry...")
# perform authz check to filter the mcp servers user has access to
prisma_client = get_prisma_client_or_throw(
"Database not connected. Connect a database to your proxy"
)
db_mcp_servers = await get_all_mcp_servers(prisma_client)
verbose_logger.info(f"Found {len(db_mcp_servers)} MCP servers in database")
# ensure the global_mcp_server_manager is up to date with the db
for server in db_mcp_servers:
verbose_logger.debug(f"Adding server to registry: {server.server_id} ({server.server_name})")
self.add_update_server(server)
verbose_logger.info(f"Registry now contains {len(self.get_registry())} servers")
def get_mcp_server_by_id(self, server_id: str) -> Optional[MCPServer]:
"""
Get the MCP Server from the server id
"""
for server in self.get_registry().values():
registry = self.get_registry()
for server in registry.values():
if server.server_id == server_id:
return server
return None
@ -1096,5 +1129,12 @@ class MCPServerManager:
for server in list_mcp_servers
]
async def reload_servers_from_database(self):
"""
Public method to reload all MCP servers from database into registry.
This can be called from management endpoints to ensure registry is up to date.
"""
await self._add_mcp_servers_from_db_to_in_memory_registry()
global_mcp_server_manager: MCPServerManager = MCPServerManager()

View file

@ -1,5 +1,5 @@
import importlib
from typing import Optional
from typing import Optional, Dict
from fastapi import APIRouter, Depends, Query, Request
@ -32,9 +32,49 @@ if MCP_AVAILABLE:
########################################################
############ MCP Server REST API Routes #################
def _get_server_auth_header(
server, mcp_server_auth_headers: Optional[Dict[str, str]], mcp_auth_header: Optional[str]
) -> Optional[str]:
"""Helper function to get server-specific auth header with case-insensitive matching."""
if mcp_server_auth_headers and server.alias:
normalized_server_alias = server.alias.lower()
normalized_headers = {k.lower(): v for k, v in mcp_server_auth_headers.items()}
server_auth = normalized_headers.get(normalized_server_alias)
if server_auth is not None:
return server_auth
elif mcp_server_auth_headers and server.server_name:
normalized_server_name = server.server_name.lower()
normalized_headers = {k.lower(): v for k, v in mcp_server_auth_headers.items()}
server_auth = normalized_headers.get(normalized_server_name)
if server_auth is not None:
return server_auth
return mcp_auth_header
def _create_tool_response_objects(tools, server_mcp_info):
"""Helper function to create tool response objects."""
return [
ListMCPToolsRestAPIResponseObject(
name=tool.name,
description=tool.description,
inputSchema=tool.inputSchema,
mcp_info=server_mcp_info,
)
for tool in tools
]
async def _get_tools_for_single_server(server, server_auth_header, mcp_protocol_version):
"""Helper function to get tools for a single server."""
tools = await global_mcp_server_manager._get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
mcp_protocol_version=mcp_protocol_version,
)
return _create_tool_response_objects(tools, server.mcp_info)
########################################################
@router.get("/tools/list", dependencies=[Depends(user_api_key_auth)])
async def list_tool_rest_api(
request: Request,
server_id: Optional[str] = Query(
None, description="The server id to list tools for"
),
@ -60,7 +100,15 @@ if MCP_AVAILABLE:
"message": "Successfully retrieved tools"
}
"""
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
try:
# Extract auth headers from request
headers = request.headers
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers)
mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
mcp_protocol_version = headers.get(MCPRequestHandler.MCP_PROTOCOL_VERSION_HEADER_NAME)
list_tools_result = []
error_message = None
@ -73,19 +121,11 @@ if MCP_AVAILABLE:
"error": "server_not_found",
"message": f"Server with id {server_id} not found"
}
server_auth_header = _get_server_auth_header(server, mcp_server_auth_headers, mcp_auth_header)
try:
tools = await global_mcp_server_manager._get_tools_from_server(
server=server,
)
for tool in tools:
list_tools_result.append(
ListMCPToolsRestAPIResponseObject(
name=tool.name,
description=tool.description,
inputSchema=tool.inputSchema,
mcp_info=server.mcp_info,
)
)
list_tools_result = await _get_tools_for_single_server(server, server_auth_header, mcp_protocol_version)
except Exception as e:
verbose_logger.exception(f"Error getting tools from {server.name}: {e}")
return {
@ -97,19 +137,11 @@ if MCP_AVAILABLE:
# Query all servers
errors = []
for server in global_mcp_server_manager.get_registry().values():
server_auth_header = _get_server_auth_header(server, mcp_server_auth_headers, mcp_auth_header)
try:
tools = await global_mcp_server_manager._get_tools_from_server(
server=server,
)
for tool in tools:
list_tools_result.append(
ListMCPToolsRestAPIResponseObject(
name=tool.name,
description=tool.description,
inputSchema=tool.inputSchema,
mcp_info=server.mcp_info,
)
)
tools_result = await _get_tools_for_single_server(server, server_auth_header, mcp_protocol_version)
list_tools_result.extend(tools_result)
except Exception as e:
verbose_logger.exception(f"Error getting tools from {server.name}: {e}")
errors.append(f"{server.name}: {str(e)}")

File diff suppressed because one or more lines are too long

View file

@ -2,27 +2,4 @@ model_list:
- model_name: openai-test
litellm_params:
model: gpt-3.5-turbo
api_key: os.environ/OPENAI_API_KEY
guardrails:
- guardrail_name: azure-text-moderation
litellm_params:
guardrail: azure/text_moderations
mode: "post_call"
api_key: os.environ/AZURE_GUARDRAIL_API_KEY
api_base: os.environ/AZURE_GUARDRAIL_API_BASE
prompts:
- prompt_id: test_my_json_prompt
litellm_params:
prompt_integration: dotprompt
prompt_id: test_hello_world_prompt
prompt_file: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts/test_hello_world_prompt.prompt
- prompt_id: test_hello_world_prompt_2
litellm_params:
prompt_integration: dotprompt
prompt_id: test_hello_world_prompt_2
prompt_file: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts/test_hello_world_prompt.prompt
litellm_settings:
callbacks: ["datadog_llm_observability"]
api_key: os.environ/OPENAI_API_KEY

View file

@ -380,6 +380,8 @@ class LiteLLMRoutes(enum.Enum):
"/health",
"/key/list",
"/user/filter/ui",
"/models",
"/v1/models",
]
# NOTE: ROUTES ONLY FOR MASTER KEY - only the Master Key should be able to Reset Spend
@ -492,6 +494,8 @@ class LiteLLMRoutes(enum.Enum):
"/global/spend/end_users",
"/global/activity",
"/global/activity/model",
"/v1/models/{model_id}",
"/models/{model_id}",
]
+ spend_tracking_routes
+ key_management_routes
@ -2206,7 +2210,7 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
braintrust: CallbackOnUI = CallbackOnUI(
litellm_callback_name="braintrust",
litellm_callback_params=["BRAINTRUST_API_KEY"],
litellm_callback_params=["BRAINTRUST_API_KEY","BRAINTRUST_API_BASE"],
ui_callback_name="Braintrust",
)
@ -2260,6 +2264,7 @@ class SpendLogsMetadata(TypedDict):
error_information: Optional[StandardLoggingPayloadErrorInformation]
usage_object: Optional[dict]
model_map_information: Optional[StandardLoggingModelInformation]
cold_storage_object_key: Optional[str] # S3/GCS object key for cold storage retrieval
class SpendLogsPayload(TypedDict):

View file

@ -200,3 +200,80 @@ async def anthropic_response( # noqa: PLR0915
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
)
@router.post(
"/v1/messages/count_tokens",
tags=["[beta] Anthropic Messages Token Counting"],
dependencies=[Depends(user_api_key_auth)],
)
async def count_tokens(
request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # Used for auth
):
"""
Count tokens for Anthropic Messages API format.
This endpoint follows the Anthropic Messages API token counting specification.
It accepts the same parameters as the /v1/messages endpoint but returns
token counts instead of generating a response.
Example usage:
```
curl -X POST "http://localhost:4000/v1/messages/count_tokens?beta=true" \
-H "Content-Type: application/json" \
-H "Authorization: Bearer your-key" \
-d '{
"model": "claude-3-sonnet-20240229",
"messages": [{"role": "user", "content": "Hello Claude!"}]
}'
```
Returns: {"input_tokens": <number>}
"""
from litellm.proxy.proxy_server import token_counter as internal_token_counter
try:
request_data = await _read_request_body(request=request)
data: dict = {**request_data}
# Extract required fields
model_name = data.get("model")
messages = data.get("messages", [])
if not model_name:
raise HTTPException(
status_code=400,
detail={"error": "model parameter is required"}
)
if not messages:
raise HTTPException(
status_code=400,
detail={"error": "messages parameter is required"}
)
# Create TokenCountRequest for the internal endpoint
from litellm.proxy._types import TokenCountRequest
token_request = TokenCountRequest(
model=model_name,
messages=messages
)
# Call the internal token counter function with direct request flag set to False
token_response = await internal_token_counter(token_request, is_direct_request=False)
# Convert the internal response to Anthropic API format
return {"input_tokens": token_response.total_tokens}
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.anthropic_endpoints.count_tokens(): Exception occurred - {}".format(str(e))
)
raise HTTPException(
status_code=500,
detail={"error": f"Internal server error: {str(e)}"}
)

View file

@ -288,9 +288,6 @@ def _is_api_route_allowed(
if valid_token is None:
raise Exception("Invalid proxy server token passed. valid_token=None.")
# Check if management routes are disabled and raise exception if they are
RouteChecks.should_call_route(route=route, valid_token=valid_token)
if not _is_user_proxy_admin(user_obj=user_obj): # if non-admin
RouteChecks.non_proxy_admin_allowed_routes_check(
user_obj=user_obj,
@ -440,6 +437,7 @@ async def get_end_user_object(
end_user_id: Optional[str],
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
route: str,
parent_otel_span: Optional[Span] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> Optional[LiteLLM_EndUserTable]:
@ -456,6 +454,8 @@ async def get_end_user_object(
_key = "end_user_id:{}".format(end_user_id)
def check_in_budget(end_user_obj: LiteLLM_EndUserTable):
if route in LiteLLMRoutes.info_routes.value: # allow calling info routes
return
if end_user_obj.litellm_budget_table is None:
return
end_user_budget = end_user_obj.litellm_budget_table.max_budget

View file

@ -83,7 +83,7 @@ class JWTHandler:
self.user_api_key_cache = user_api_key_cache
self.litellm_jwtauth = litellm_jwtauth
self.leeway = leeway
@staticmethod
def is_jwt(token: str):
parts = token.split(".")
@ -844,6 +844,7 @@ class JWTAuthManager:
user_api_key_cache: DualCache,
parent_otel_span: Optional[Span],
proxy_logging_obj: ProxyLogging,
route: str,
) -> Tuple[
Optional[LiteLLM_UserTable],
Optional[LiteLLM_OrganizationTable],
@ -892,6 +893,7 @@ class JWTAuthManager:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
)
if end_user_id
else None
@ -1133,6 +1135,7 @@ class JWTAuthManager:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
)
await JWTAuthManager.sync_user_role_and_teams(

View file

@ -25,6 +25,8 @@ class RouteChecks:
from litellm_enterprise.proxy.auth.route_checks import EnterpriseRouteChecks
EnterpriseRouteChecks.should_call_route(route=route)
except HTTPException as e:
raise e
except Exception:
pass
@ -208,7 +210,7 @@ class RouteChecks:
route=route, allowed_routes=LiteLLMRoutes.self_managed_routes.value
): # routes that manage their own allowed/disallowed logic
pass
elif route.startswith("/v1/mcp/"):
elif route.startswith("/v1/mcp/") or route.startswith("/mcp-rest/"):
pass # authN/authZ handled by api itself
else:
user_role = "unknown"
@ -386,7 +388,7 @@ class RouteChecks:
if "thread" in request.url.path or "assistant" in request.url.path:
return True
return False
@staticmethod
def is_generate_content_route(route: str) -> bool:
"""

View file

@ -49,6 +49,7 @@ from litellm.proxy.auth.auth_utils import (
from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler
from litellm.proxy.auth.oauth2_check import check_oauth2_token
from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
@ -259,6 +260,7 @@ def get_api_key(
from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_query_params,
)
api_key = api_key
passed_in_key: Optional[str] = None
if isinstance(custom_litellm_key_header, str):
@ -279,7 +281,11 @@ def get_api_key(
elif isinstance(azure_apim_header, str):
passed_in_key = azure_apim_header
api_key = azure_apim_header
elif RouteChecks.is_generate_content_route(route=route) and request is not None and _safe_get_request_query_params(request).get("key"):
elif (
RouteChecks.is_generate_content_route(route=route)
and request is not None
and _safe_get_request_query_params(request).get("key")
):
google_auth_key: str = _safe_get_request_query_params(request).get("key") or ""
passed_in_key = google_auth_key
api_key = google_auth_key
@ -609,6 +615,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
)
if _end_user_object is not None:
end_user_params["allowed_model_region"] = (
@ -1141,6 +1148,8 @@ async def user_api_key_auth(
request_data = await _read_request_body(request=request)
route: str = get_request_route(request=request)
## CHECK IF ROUTE IS ALLOWED
user_api_key_auth_obj = await _user_api_key_auth_builder(
request=request,
api_key=api_key,
@ -1152,6 +1161,9 @@ async def user_api_key_auth(
custom_litellm_key_header=custom_litellm_key_header,
)
## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ##
RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj)
end_user_id = get_end_user_id_from_request_body(
request_data, _safe_get_request_headers(request)
)

View file

@ -5,6 +5,7 @@
# +-------------------------------------------------------------+
# Thank you users! We ❤️ you! - Krrish & Ishaan
import copy
import os
import sys
@ -50,6 +51,51 @@ from litellm.types.utils import (
GUARDRAIL_NAME = "bedrock"
def _redact_pii_matches(response_json: dict) -> dict:
try:
# Create a deep copy to avoid modifying the original response
redacted_response = copy.deepcopy(response_json)
# Get assessments from the response
assessments = redacted_response.get("assessments", [])
if not assessments:
return redacted_response
for assessment in assessments:
# Redact PII entities in sensitive information policy
sensitive_info_policy = assessment.get("sensitiveInformationPolicy")
if sensitive_info_policy:
pii_entities = sensitive_info_policy.get("piiEntities", [])
for pii_entity in pii_entities:
if "match" in pii_entity:
pii_entity["match"] = "[REDACTED]"
# Redact regex matches
regexes = sensitive_info_policy.get("regexes", [])
for regex_match in regexes:
if "match" in regex_match:
regex_match["match"] = "[REDACTED]"
# Redact custom word matches in word policy
word_policy = assessment.get("wordPolicy")
if word_policy:
custom_words = word_policy.get("customWords", [])
for custom_word in custom_words:
if "match" in custom_word:
custom_word["match"] = "[REDACTED]"
managed_words = word_policy.get("managedWordLists", [])
for managed_word in managed_words:
if "match" in managed_word:
managed_word["match"] = "[REDACTED]"
return redacted_response
except Exception as e:
# We do not want to fail in any case so this is just a warning
verbose_proxy_logger.warning("Guardrail log redaction failed: %s", str(e))
return response_json
class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
def __init__(
self,
@ -271,10 +317,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
data=prepared_request.body, # type: ignore
headers=prepared_request.headers, # type: ignore
)
verbose_proxy_logger.debug("Bedrock AI response: %s", response.text)
if response.status_code == 200:
# check if the response was flagged
_json_response = response.json()
redacted_response = _redact_pii_matches(_json_response)
verbose_proxy_logger.debug("Bedrock AI response : %s", redacted_response)
bedrock_guardrail_response = BedrockGuardrailResponse(**_json_response)
if self._should_raise_guardrail_blocked_exception(
bedrock_guardrail_response

View file

@ -403,6 +403,9 @@ if MCP_AVAILABLE:
touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
)
global_mcp_server_manager.add_update_server(new_mcp_server)
# Ensure registry is up to date by reloading from database
await global_mcp_server_manager.reload_servers_from_database()
except Exception as e:
verbose_proxy_logger.exception(f"Error creating mcp server: {str(e)}")
raise HTTPException(
@ -461,6 +464,9 @@ if MCP_AVAILABLE:
detail={"error": f"MCP Server not found, passed server_id={server_id}"},
)
global_mcp_server_manager.remove_server(mcp_server_record_deleted)
# Ensure registry is up to date by reloading from database
await global_mcp_server_manager.reload_servers_from_database()
# TODO: Enterprise: Finish audit log trail
if litellm.store_audit_logs:
@ -533,6 +539,9 @@ if MCP_AVAILABLE:
},
)
global_mcp_server_manager.add_update_server(mcp_server_record_updated)
# Ensure registry is up to date by reloading from database
await global_mcp_server_manager.reload_servers_from_database()
# TODO: Enterprise: Finish audit log trail
if litellm.store_audit_logs:

View file

@ -116,6 +116,14 @@ async def serve_login_page(
missing_env_vars = show_missing_vars_in_env()
if missing_env_vars is not None:
return missing_env_vars
#########################################################
# Construct Redirect URL
base_url_to_redirect_to: Optional[str] = None
base_url_to_redirect_to = os.getenv("PROXY_BASE_URL", "")
server_root_path = os.getenv("SERVER_ROOT_PATH", "")
if server_root_path != "":
base_url_to_redirect_to += server_root_path
#########################################################
# Build the unified login page HTML
error_message = ""
@ -137,7 +145,13 @@ async def serve_login_page(
sso_button = ""
if sso_available:
sso_button = """
sso_login_url = base_url_to_redirect_to
if sso_login_url.endswith("/"):
sso_login_url += "sso/login"
else:
sso_login_url += "/sso/login"
sso_button = f"""
<div style="
margin-top: 20px;
padding-top: 20px;
@ -149,7 +163,7 @@ async def serve_login_page(
font-size: 14px;
margin-bottom: 16px;
">or</p>
<a href="/sso/login" style="
<a href="{sso_login_url}" style="
display: inline-block;
background-color: #f8fafc;
border: 1px solid #e2e8f0;
@ -167,12 +181,10 @@ async def serve_login_page(
</div>
"""
# Get the base URL for form action using proper URL construction
url_to_redirect_to = os.getenv("PROXY_BASE_URL", "")
server_root_path = os.getenv("SERVER_ROOT_PATH", "")
if server_root_path != "":
url_to_redirect_to += server_root_path
url_to_redirect_to += "/login"
if base_url_to_redirect_to.endswith("/"):
url_to_redirect_to = base_url_to_redirect_to + "login"
else:
url_to_redirect_to = base_url_to_redirect_to + "/login"
unified_login_html = f"""
<!DOCTYPE html>

View file

@ -1,6 +1,15 @@
model_list:
- model_name: vertex_ai/*
- model_name: bedrock/*
litellm_params:
model: vertex_ai/*
model: bedrock/*
litellm_settings:
callbacks: ["s3_v2"]
s3_callback_params:
s3_bucket_name: litellm-logs # AWS Bucket Name for S3
s3_region_name: us-west-2
general_settings:
cold_storage_custom_logger: s3_v2
store_prompts_in_cold_storage: true

View file

@ -126,6 +126,7 @@ import litellm
from litellm import Router
from litellm._logging import verbose_proxy_logger, verbose_router_logger
from litellm.caching.caching import DualCache, RedisCache
from litellm.caching.redis_cluster_cache import RedisClusterCache
from litellm.constants import (
DAYS_IN_A_MONTH,
DEFAULT_HEALTH_CHECK_INTERVAL,
@ -1591,7 +1592,9 @@ class ProxyConfig:
litellm.cache = Cache(**cache_params)
if litellm.cache is not None and isinstance(litellm.cache.cache, RedisCache):
if litellm.cache is not None and isinstance(
litellm.cache.cache, (RedisCache, RedisClusterCache)
):
## INIT PROXY REDIS USAGE CLIENT ##
redis_usage_cache = litellm.cache.cache
@ -1728,7 +1731,7 @@ class ProxyConfig:
self._load_environment_variables(config=config)
## Callback settings
callback_settings = config.get("callback_settings", None)
callback_settings = config.get("callback_settings", {})
## LITELLM MODULE SETTINGS (e.g. litellm.drop_params=True,..)
litellm_settings = config.get("litellm_settings", None)
@ -2670,7 +2673,7 @@ class ProxyConfig:
proxy_logging_obj: ProxyLogging
"""
_general_settings = config_data.get("general_settings", {})
if "alerting" in _general_settings:
if _general_settings is not None and "alerting" in _general_settings:
if (
general_settings is not None
and general_settings.get("alerting", None) is not None
@ -2702,14 +2705,14 @@ class ProxyConfig:
"alerting"
]
if "alert_types" in _general_settings:
if _general_settings is not None and "alert_types" in _general_settings:
general_settings["alert_types"] = _general_settings["alert_types"]
proxy_logging_obj.alert_types = general_settings["alert_types"]
proxy_logging_obj.slack_alerting_instance.update_values(
alert_types=general_settings["alert_types"], llm_router=llm_router
)
if "alert_to_webhook_url" in _general_settings:
if _general_settings is not None and "alert_to_webhook_url" in _general_settings:
general_settings["alert_to_webhook_url"] = _general_settings[
"alert_to_webhook_url"
]
@ -3766,100 +3769,37 @@ async def model_list(
Defaults to "general" when include_metadata=true
"""
global llm_model_list, general_settings, llm_router, prisma_client, user_api_key_cache, proxy_logging_obj
all_models = []
model_access_groups: Dict[str, List[str]] = defaultdict(list)
## CHECK IF MODEL RESTRICTIONS ARE SET AT KEY/TEAM LEVEL ##
if llm_router is None:
proxy_model_list = []
else:
proxy_model_list = llm_router.get_model_names()
model_access_groups = llm_router.get_model_access_groups()
## if only_model_access_groups is True,
"""
1. Get all models key/user/team has access to
2. Filter out models that are not model access groups
3. Return the models
"""
if only_model_access_groups is True:
include_model_access_groups = True
from litellm.proxy.utils import (
create_model_info_response,
get_available_models_for_user,
)
key_models = get_key_models(
# Get available models for the user
all_models = await get_available_models_for_user(
user_api_key_dict=user_api_key_dict,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
include_model_access_groups=include_model_access_groups,
)
team_models: List[str] = user_api_key_dict.team_models
if team_id:
key_models = []
team_object = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
validate_membership(user_api_key_dict=user_api_key_dict, team_table=team_object)
team_models = team_object.models
team_models = get_team_models(
team_models=team_models,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
include_model_access_groups=include_model_access_groups,
)
all_models = get_complete_model_list(
key_models=key_models,
team_models=team_models,
proxy_model_list=proxy_model_list,
user_model=user_model,
infer_model_from_keys=general_settings.get("infer_model_from_keys", False),
return_wildcard_routes=return_wildcard_routes,
llm_router=llm_router,
model_access_groups=model_access_groups,
include_model_access_groups=include_model_access_groups,
only_model_access_groups=only_model_access_groups,
general_settings=general_settings,
user_model=user_model,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
team_id=team_id,
include_model_access_groups=include_model_access_groups or False,
only_model_access_groups=only_model_access_groups or False,
return_wildcard_routes=return_wildcard_routes or False,
user_api_key_cache=user_api_key_cache,
)
# Build response data
model_data = []
for model in all_models:
model_info = {
"id": model,
"object": "model",
"created": DEFAULT_MODEL_CREATED_AT_TIME,
"owned_by": "openai",
}
# Add metadata if requested
if include_metadata:
metadata = {}
# Default fallback_type to "general" if include_metadata is true
effective_fallback_type = (
fallback_type if fallback_type is not None else "general"
)
# Validate fallback_type
valid_fallback_types = ["general", "context_window", "content_policy"]
if effective_fallback_type not in valid_fallback_types:
raise HTTPException(
status_code=400,
detail=f"Invalid fallback_type. Must be one of: {valid_fallback_types}",
)
fallbacks = get_all_fallbacks(
model=model,
llm_router=llm_router,
fallback_type=effective_fallback_type,
)
metadata["fallbacks"] = fallbacks
model_info["metadata"] = metadata
model_info = create_model_info_response(
model_id=model,
provider="openai",
include_metadata=include_metadata or False,
fallback_type=fallback_type,
llm_router=llm_router,
)
model_data.append(model_info)
return dict(
@ -3868,6 +3808,68 @@ async def model_list(
)
@router.get(
"/v1/models/{model_id}",
dependencies=[Depends(user_api_key_auth)],
tags=["model management"],
)
@router.get(
"/models/{model_id}",
dependencies=[Depends(user_api_key_auth)],
tags=["model management"],
)
async def model_info(
model_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Retrieve information about a specific model accessible to your API key.
Returns model details only if the model is available to your API key/team.
Returns 404 if the model doesn't exist or is not accessible.
Follows OpenAI API specification for individual model retrieval.
https://platform.openai.com/docs/api-reference/models/retrieve
"""
global llm_model_list, general_settings, llm_router, prisma_client, user_api_key_cache, proxy_logging_obj
from litellm.proxy.utils import (
create_model_info_response,
get_available_models_for_user,
validate_model_access,
)
# Get available models for the user
all_models = await get_available_models_for_user(
user_api_key_dict=user_api_key_dict,
llm_router=llm_router,
general_settings=general_settings,
user_model=user_model,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
team_id=None,
include_model_access_groups=False,
only_model_access_groups=False,
return_wildcard_routes=False,
user_api_key_cache=user_api_key_cache,
)
# Validate that the requested model is accessible
validate_model_access(model_id=model_id, available_models=all_models)
# Get provider information
_, provider, _, _ = litellm.get_llm_provider(model=model_id)
# Return the model information in the same format as the list endpoint
return create_model_info_response(
model_id=model_id,
provider=provider,
include_metadata=False,
fallback_type=None,
llm_router=llm_router,
)
@router.post(
"/v1/chat/completions",
dependencies=[Depends(user_api_key_auth)],
@ -3923,7 +3925,7 @@ async def chat_completion( # noqa: PLR0915
data = await _read_request_body(request=request)
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await base_llm_response_processor.base_process_llm_request(
result = await base_llm_response_processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -3941,6 +3943,10 @@ async def chat_completion( # noqa: PLR0915
user_api_base=user_api_base,
version=version,
)
if isinstance(result, BaseModel):
return result.model_dump(exclude_none=True, exclude_unset=True)
else:
return result
except RejectedRequestError as e:
_data = e.request_data
await proxy_logging_obj.post_call_failure_hook(
@ -5614,13 +5620,62 @@ async def run_thread(
# async def get_available_routes(user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth)):
def _get_provider_token_counter(deployment: dict, model_to_use: str):
"""
Auto-route to the correct provider's token counter based on model/deployment.
Uses the existing get_provider_model_info infrastructure with switch-case pattern.
"""
if deployment is None:
return None
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
full_model = deployment.get("litellm_params", {}).get("model", "")
try:
# Use existing LiteLLM logic to determine provider
model, provider, dynamic_api_key, api_base = get_llm_provider(
model=full_model,
custom_llm_provider=deployment.get("litellm_params", {}).get(
"custom_llm_provider"
),
api_base=deployment.get("litellm_params", {}).get("api_base"),
api_key=deployment.get("litellm_params", {}).get("api_key"),
)
# Switch case pattern using existing get_provider_model_info
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
# Convert string provider to LlmProviders enum
llm_provider_enum = LlmProviders(provider)
# Add more provider mappings as needed
if llm_provider_enum:
provider_model_info = ProviderConfigManager.get_provider_model_info(
model=full_model, provider=llm_provider_enum
)
if provider_model_info is not None:
return provider_model_info.get_token_counter()
except Exception:
# If provider detection fails, fall back to manual checks
if full_model.startswith("anthropic/") or "anthropic" in full_model.lower():
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
anthropic_model_info = AnthropicModelInfo()
return anthropic_model_info.get_token_counter()
return None
@router.post(
"/utils/token_counter",
tags=["llm utils"],
dependencies=[Depends(user_api_key_auth)],
response_model=TokenCountResponse,
)
async def token_counter(request: TokenCountRequest):
async def token_counter(request: TokenCountRequest, is_direct_request: bool = True):
""" """
from litellm import token_counter
@ -5653,6 +5708,31 @@ async def token_counter(request: TokenCountRequest):
litellm_model_name or request.model
) # use litellm model name, if it's not avalable then fallback to request.model
# Try provider-specific token counting first - only for non-direct requests (from provider endpoints)
provider_counter = None
if deployment is not None and not is_direct_request:
# Auto-route to the correct provider based on model
provider_counter = _get_provider_token_counter(deployment, model_to_use)
if provider_counter is not None and provider_counter.supports_provider(
deployment=deployment, from_endpoint=not is_direct_request
):
result = await provider_counter.count_tokens(
model_to_use=model_to_use,
messages=messages, # type: ignore
deployment=deployment,
request_model=request.model,
)
if result is not None:
return TokenCountResponse(
total_tokens=result["total_tokens"],
request_model=result["request_model"],
model_used=result["model_used"],
tokenizer_type=result["tokenizer_type"],
)
# Default LiteLLM token counting
custom_tokenizer: Optional[CustomHuggingfaceTokenizer] = None
if model_info is not None:
custom_tokenizer = cast(
@ -6930,56 +7010,21 @@ async def model_group_info(
status_code=500, detail={"error": "LLM Router is not loaded in"}
)
## CHECK IF MODEL RESTRICTIONS ARE SET AT KEY/TEAM LEVEL ##
model_access_groups: Dict[str, List[str]] = defaultdict(list)
if llm_router is None:
proxy_model_list = []
else:
proxy_model_list = llm_router.get_model_names()
model_access_groups = llm_router.get_model_access_groups()
from litellm.proxy.utils import get_available_models_for_user
key_models = get_key_models(
# Get available models for the user
all_models_str = await get_available_models_for_user(
user_api_key_dict=user_api_key_dict,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
)
team_models = []
if (
not user_api_key_dict.team_id
and user_api_key_dict.user_id is not None
and not _user_has_admin_view(user_api_key_dict)
):
if prisma_client is None:
raise HTTPException(
status_code=500,
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
user_object = await prisma_client.db.litellm_usertable.find_first(
where={"user_id": user_api_key_dict.user_id}
)
user_object_typed = LiteLLM_UserTable(**user_object.model_dump())
user_models = []
if user_object is not None:
user_models = get_team_models(
team_models=user_object_typed.models,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
)
team_models = user_models
else:
team_models = get_team_models(
team_models=user_api_key_dict.team_models,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
)
all_models_str = get_complete_model_list(
key_models=key_models,
team_models=team_models,
proxy_model_list=proxy_model_list,
user_model=user_model,
infer_model_from_keys=general_settings.get("infer_model_from_keys", False),
llm_router=llm_router,
general_settings=general_settings,
user_model=user_model,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
team_id=None,
include_model_access_groups=False,
only_model_access_groups=False,
return_wildcard_routes=False,
user_api_key_cache=user_api_key_cache,
)
model_groups: List[ModelGroupInfoProxy] = _get_model_group_info(
llm_router=llm_router, all_models_str=all_models_str, model_group=model_group
@ -7727,6 +7772,13 @@ async def claim_onboarding_link(data: InvitationClaim):
return user_obj
@app.get("/get_logo_url", include_in_schema=False)
def get_logo_url():
"""Get the current logo URL from environment"""
logo_path = os.getenv("UI_LOGO_PATH", "")
return {"logo_url": logo_path}
@app.get("/get_image", include_in_schema=False)
def get_image():
"""Get logo to show on admin UI"""

View file

@ -153,6 +153,8 @@ async def route_request(
"aget_responses",
"adelete_responses",
"alist_input_items",
"avector_store_create",
"avector_store_search",
]:
# moderation endpoint does not require `model` parameter
return getattr(llm_router, f"{route_type}")(**data)

View file

@ -0,0 +1,73 @@
"""
This module is responsible for handling Getting/Setting the proxy server request from cold storage.
It allows fetching a dict of the proxy server request from s3 or GCS bucket.
"""
from typing import Optional, cast
import litellm
from litellm import _custom_logger_compatible_callbacks_literal
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
class ColdStorageHandler:
"""
This class is responsible for handling Getting/Setting the proxy server request from cold storage.
It allows fetching a dict of the proxy server request from s3 or GCS bucket.
"""
async def get_proxy_server_request_from_cold_storage_with_object_key(
self,
object_key: str,
) -> Optional[dict]:
"""
Get the proxy server request from cold storage using the object key directly.
Args:
object_key: The S3/GCS object key to retrieve
Returns:
Optional[dict]: The proxy server request dict or None if not found
"""
# select the custom logger to use for cold storage
custom_logger_name: Optional[_custom_logger_compatible_callbacks_literal] = self._select_custom_logger_for_cold_storage()
# if no custom logger name is configured, return None
if custom_logger_name is None:
return None
# get the active/initialized custom logger
custom_logger: Optional[CustomLogger] = litellm.logging_callback_manager.get_active_custom_logger_for_callback_name(custom_logger_name)
# if no custom logger is found, return None
if custom_logger is None:
return None
proxy_server_request = await custom_logger.get_proxy_server_request_from_cold_storage_with_object_key(
object_key=object_key,
)
return proxy_server_request
def _select_custom_logger_for_cold_storage(
self,
) -> Optional[_custom_logger_compatible_callbacks_literal]:
cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = ColdStorageHandler._get_configured_cold_storage_custom_logger()
return cold_storage_custom_logger
@staticmethod
def _get_configured_cold_storage_custom_logger() -> Optional[_custom_logger_compatible_callbacks_literal]:
from litellm.proxy.proxy_server import general_settings
cold_storage_custom_logger: Optional[str] = general_settings.get("cold_storage_custom_logger")
if not cold_storage_custom_logger:
verbose_proxy_logger.debug("No cold storage custom logger found in general settings")
return None
return cast(_custom_logger_compatible_callbacks_literal, cold_storage_custom_logger)

View file

@ -53,6 +53,7 @@ def _get_spend_logs_metadata(
guardrail_information: Optional[StandardLoggingGuardrailInformation] = None,
usage_object: Optional[dict] = None,
model_map_information: Optional[StandardLoggingModelInformation] = None,
cold_storage_object_key: Optional[str] = None
) -> SpendLogsMetadata:
if metadata is None:
return SpendLogsMetadata(
@ -75,6 +76,7 @@ def _get_spend_logs_metadata(
model_map_information=None,
usage_object=None,
guardrail_information=None,
cold_storage_object_key=cold_storage_object_key,
)
verbose_proxy_logger.debug(
"getting payload for SpendLogs, available keys in metadata: "
@ -98,6 +100,8 @@ def _get_spend_logs_metadata(
clean_metadata["guardrail_information"] = guardrail_information
clean_metadata["usage_object"] = usage_object
clean_metadata["model_map_information"] = model_map_information
clean_metadata["cold_storage_object_key"] = cold_storage_object_key
return clean_metadata
@ -267,6 +271,11 @@ def get_logging_payload( # noqa: PLR0915
if standard_logging_payload is not None
else None
),
cold_storage_object_key=(
standard_logging_payload["metadata"].get("cold_storage_object_key", None)
if standard_logging_payload is not None
else None
),
)
special_usage_fields = ["completion_tokens", "prompt_tokens", "total_tokens"]
@ -474,6 +483,7 @@ def _sanitize_request_body_for_spend_logs_payload(
Recursively sanitize request body to prevent logging large base64 strings or other large values.
Truncates strings longer than 1000 characters and handles nested dictionaries.
"""
from litellm.constants import LITELLM_TRUNCATED_PAYLOAD_FIELD
MAX_STRING_LENGTH = 1000
if visited is None:
@ -492,7 +502,7 @@ def _sanitize_request_body_for_spend_logs_payload(
return [_sanitize_value(item) for item in value]
elif isinstance(value, str):
if len(value) > MAX_STRING_LENGTH:
return f"{value[:MAX_STRING_LENGTH]}... (truncated {len(value) - MAX_STRING_LENGTH} chars)"
return f"{value[:MAX_STRING_LENGTH]}... ({LITELLM_TRUNCATED_PAYLOAD_FIELD} {len(value) - MAX_STRING_LENGTH} chars)"
return value
return value

View file

@ -1,7 +1,7 @@
#### CRUD ENDPOINTS for UI Settings #####
from typing import Any, Dict, List, Union
from typing import Any, Dict, List, Union, Optional
from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, Depends, HTTPException, UploadFile, File
import litellm
from litellm._logging import verbose_proxy_logger
@ -19,6 +19,16 @@ class IPAddress(BaseModel):
ip: str
class UIThemeConfig(BaseModel):
"""Configuration for UI theme customization"""
# Logo configuration
logo_url: Optional[str] = Field(
default=None,
description="URL or path to custom logo image. Can be a local file path or HTTP/HTTPS URL"
)
class SettingsResponse(BaseModel):
"""Base response model for settings with values and schema information"""
@ -47,6 +57,12 @@ class DefaultTeamSettingsResponse(SettingsResponse):
pass
class UIThemeSettingsResponse(SettingsResponse):
"""Response model for UI theme settings"""
pass
@router.get(
"/get/allowed_ips",
tags=["Budget & Spend Tracking"],
@ -507,3 +523,155 @@ async def update_sso_settings(sso_config: SSOConfig):
"status": "success",
"settings": sso_data,
}
@router.get(
"/get/ui_theme_settings",
tags=["UI Theme Settings"],
dependencies=[Depends(user_api_key_auth)],
response_model=UIThemeSettingsResponse,
)
async def get_ui_theme_settings():
"""
Get UI theme configuration from the litellm_settings.
Returns current logo settings for UI customization.
"""
from litellm.proxy.proxy_server import proxy_config
# Load existing config
config = await proxy_config.get_config()
return await _get_settings_with_schema(
settings_key="ui_theme_config",
settings_class=UIThemeConfig,
config=config,
)
@router.patch(
"/update/ui_theme_settings",
tags=["UI Theme Settings"],
dependencies=[Depends(user_api_key_auth)],
)
async def update_ui_theme_settings(theme_config: UIThemeConfig):
"""
Update UI theme configuration.
Updates logo settings for the admin UI.
"""
from litellm.proxy.proxy_server import proxy_config, store_model_in_db
import os
if store_model_in_db is not True:
raise HTTPException(
status_code=500,
detail={
"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."
},
)
# Load existing config
config = await proxy_config.get_config()
# Update config with UI theme settings
if "general_settings" not in config:
config["general_settings"] = {}
if "environment_variables" not in config:
config["environment_variables"] = {}
# Convert theme config to dict
theme_data = theme_config.model_dump(exclude_none=True)
# Store UI theme config in litellm_settings (where it's retrieved from)
if "litellm_settings" not in config:
config["litellm_settings"] = {}
config["litellm_settings"]["ui_theme_config"] = theme_data
# Update UI_LOGO_PATH environment variable if logo_url is provided
# If logo_url is empty string, None, or null, remove the environment variable to use default
logo_url = theme_data.get("logo_url")
verbose_proxy_logger.debug(f"Updating logo_url: {logo_url}")
if logo_url and isinstance(logo_url, str) and logo_url.strip(): # Check if logo_url exists and is not empty/whitespace
config["environment_variables"]["UI_LOGO_PATH"] = logo_url
os.environ["UI_LOGO_PATH"] = logo_url
verbose_proxy_logger.debug(f"Set UI_LOGO_PATH to: {logo_url}")
else:
# Remove the environment variable to restore default logo
if "UI_LOGO_PATH" in config.get("environment_variables", {}):
del config["environment_variables"]["UI_LOGO_PATH"]
verbose_proxy_logger.debug("Removed UI_LOGO_PATH from config")
if "UI_LOGO_PATH" in os.environ:
del os.environ["UI_LOGO_PATH"]
verbose_proxy_logger.debug("Removed UI_LOGO_PATH from environment")
# Handle environment variable encryption if needed
stored_config = config.copy()
if "environment_variables" in stored_config and len(stored_config["environment_variables"]) > 0:
# Only encrypt if there are environment variables to encrypt
stored_config["environment_variables"] = proxy_config._encrypt_env_variables(
environment_variables=stored_config["environment_variables"]
)
# Save the updated config
await proxy_config.save_config(new_config=stored_config)
return {
"message": "Logo settings updated successfully.",
"status": "success",
"theme_config": theme_data,
}
@router.post(
"/upload/logo",
tags=["UI Theme Settings"],
dependencies=[Depends(user_api_key_auth)],
)
async def upload_logo(file: UploadFile = File(...)):
"""
Upload a custom logo for the admin UI.
Accepts image files (PNG, JPG, JPEG, SVG) and stores them for use in the UI.
"""
import os
from pathlib import Path
# Validate file type
allowed_extensions = {".png", ".jpg", ".jpeg", ".svg"}
file_extension = Path(file.filename or "").suffix.lower()
if file_extension not in allowed_extensions:
raise HTTPException(
status_code=400,
detail=f"Invalid file type. Allowed types: {', '.join(allowed_extensions)}"
)
# Validate file size (max 5MB)
file_content = await file.read()
if len(file_content) > 5 * 1024 * 1024: # 5MB
raise HTTPException(
status_code=400,
detail="File size too large. Maximum size is 5MB."
)
# Create uploads directory if it doesn't exist
current_dir = os.path.dirname(os.path.abspath(__file__))
upload_dir = os.path.join(current_dir, "..", "uploads")
os.makedirs(upload_dir, exist_ok=True)
# Generate unique filename
import uuid
unique_filename = f"logo_{uuid.uuid4().hex}{file_extension}"
file_path = os.path.join(upload_dir, unique_filename)
# Save the file
with open(file_path, "wb") as buffer:
buffer.write(file_content)
return {
"message": "Logo uploaded successfully",
"status": "success",
"file_path": file_path,
"filename": unique_filename,
"file_size": len(file_content),
}

View file

@ -22,7 +22,7 @@ from typing import (
overload,
)
from litellm.constants import MAX_TEAM_LIST_LIMIT
from litellm.constants import MAX_TEAM_LIST_LIMIT, DEFAULT_MODEL_CREATED_AT_TIME
from litellm.proxy._types import (
DB_CONNECTION_ERROR_TYPES,
CommonProxyErrors,
@ -55,11 +55,7 @@ from litellm import (
from litellm._logging import verbose_proxy_logger
from litellm._service_logger import ServiceLogging, ServiceTypes
from litellm.caching.caching import DualCache, RedisCache
from litellm.exceptions import (
BlockedPiiEntityError,
GuardrailRaisedException,
RejectedRequestError,
)
from litellm.exceptions import RejectedRequestError
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
@ -452,108 +448,6 @@ class ProxyLogging:
litellm_parent_otel_span=None,
)
async def async_pre_mcp_tool_call_hook(
self,
kwargs: dict,
request_obj: Any,
start_time: datetime,
end_time: datetime,
) -> Optional[Any]:
"""
Pre MCP Tool Call Hook
Use this to validate and modify MCP tool calls before execution.
Reuses existing LLM guardrail logic by converting MCP calls to message format.
"""
from litellm.types.llms.base import HiddenParams
from litellm.types.mcp import MCPPreCallRequestObject
callbacks = self.get_combined_callback_list(
dynamic_success_callbacks=getattr(self, "dynamic_success_callbacks", None),
global_callbacks=litellm.success_callback,
)
# Create the request object if it's not already one
if not isinstance(request_obj, MCPPreCallRequestObject):
# Convert UserAPIKeyAuth object to dict if needed
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(
kwargs.get("user_api_key_auth")
)
request_obj = MCPPreCallRequestObject(
tool_name=kwargs.get("name", ""),
arguments=kwargs.get("arguments", {}),
server_name=kwargs.get("server_name"),
user_api_key_auth=user_api_key_auth_dict,
hidden_params=HiddenParams(),
)
for callback in callbacks:
try:
_callback: Optional[CustomLogger] = None
if isinstance(callback, str):
from typing import cast
from litellm import _custom_logger_compatible_callbacks_literal
_callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class(
cast(_custom_logger_compatible_callbacks_literal, callback)
)
else:
_callback = callback # type: ignore
if _callback is not None and isinstance(_callback, CustomGuardrail):
from litellm.types.guardrails import GuardrailEventHooks
# Check if guardrail should be run for pre_call hook (reusing existing logic)
if (
_callback.should_run_guardrail(
data=kwargs, event_type=GuardrailEventHooks.pre_mcp_call
)
is not True
):
continue
# Convert MCP tool call to LLM message format for existing guardrail logic
synthetic_llm_data = self._convert_mcp_to_llm_format(
request_obj, kwargs
)
# Reuse existing LLM guardrail logic
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(
kwargs.get("user_api_key_auth")
)
result = await _callback.async_pre_call_hook(
user_api_key_dict=user_api_key_auth_dict, # type: ignore
cache=self.call_details["user_api_key_cache"],
data=synthetic_llm_data,
call_type="mcp_call",
)
# Convert result back to MCP response format if blocked/modified
if result is not None:
mcp_response = self._convert_llm_result_to_mcp_response(
result, request_obj
)
if mcp_response is not None:
return self._parse_pre_mcp_call_hook_response(
response=mcp_response, original_request=request_obj
)
except (
BlockedPiiEntityError,
GuardrailRaisedException,
HTTPException,
) as e:
# Re-raise guardrail exceptions so they can be properly handled
raise e
except Exception as e:
verbose_proxy_logger.exception(
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {}".format(
str(e)
)
)
return None
def _convert_user_api_key_auth_to_dict(self, user_api_key_auth_obj):
"""
@ -567,7 +461,7 @@ class ProxyLogging:
elif hasattr(user_api_key_auth_obj, "__dict__"):
# If it's a regular object, convert to dict
return user_api_key_auth_obj.__dict__
return user_api_key_auth_obj
return {}
def _convert_mcp_to_llm_format(self, request_obj, kwargs: dict) -> dict:
"""
@ -765,8 +659,6 @@ class ProxyLogging:
"""
Convert LLM guardrail result back to MCP during call response format.
"""
from litellm.types.mcp import MCPDuringCallResponseObject
# If result is an exception, it means the guardrail wants to stop execution
if isinstance(llm_result, Exception):
return MCPDuringCallResponseObject(
@ -836,112 +728,39 @@ class ProxyLogging:
}
return result
async def async_during_mcp_tool_call_hook(
self,
kwargs: dict,
request_obj: Any,
start_time: datetime,
end_time: datetime,
) -> Optional[Any]:
def _create_mcp_request_object_from_kwargs(self, kwargs: dict) -> "MCPPreCallRequestObject":
"""
During MCP Tool Call Hook
Use this for concurrent monitoring and validation during tool execution.
Reuses existing LLM guardrail logic by converting MCP calls to message format.
Helper function to create MCPPreCallRequestObject from kwargs for standard pre_call_hook.
"""
from litellm.types.llms.base import HiddenParams
from litellm.types.mcp import MCPDuringCallRequestObject
from litellm.types.mcp import MCPPreCallRequestObject
callbacks = self.get_combined_callback_list(
dynamic_success_callbacks=getattr(self, "dynamic_success_callbacks", None),
global_callbacks=litellm.success_callback,
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(kwargs.get("user_api_key_auth"))
return MCPPreCallRequestObject(
tool_name=kwargs.get("name", ""),
arguments=kwargs.get("arguments", {}),
server_name=kwargs.get("server_name"),
user_api_key_auth=user_api_key_auth_dict,
hidden_params=HiddenParams(),
)
# Create the request object if it's not already one
if not isinstance(request_obj, MCPDuringCallRequestObject):
request_obj = MCPDuringCallRequestObject(
tool_name=kwargs.get("name", ""),
arguments=kwargs.get("arguments", {}),
server_name=kwargs.get("server_name"),
start_time=start_time.timestamp() if start_time else None,
hidden_params=HiddenParams(),
)
for callback in callbacks:
try:
_callback: Optional[CustomLogger] = None
if isinstance(callback, str):
from typing import cast
from litellm import _custom_logger_compatible_callbacks_literal
_callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class(
cast(_custom_logger_compatible_callbacks_literal, callback)
)
else:
_callback = callback # type: ignore
if _callback is not None and isinstance(_callback, CustomGuardrail):
from litellm.types.guardrails import GuardrailEventHooks
# Check if guardrail should be run for during_call hook (reusing existing logic)
if (
_callback.should_run_guardrail(
data=kwargs, event_type=GuardrailEventHooks.during_mcp_call
)
is not True
):
continue
# Convert MCP tool call to LLM message format for existing guardrail logic
synthetic_llm_data = self._convert_mcp_to_llm_format(
request_obj, kwargs
)
# Reuse existing LLM guardrail logic for during call
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(
kwargs.get("user_api_key_auth")
)
result = await _callback.async_moderation_hook(
data=synthetic_llm_data,
user_api_key_dict=user_api_key_auth_dict, # type: ignore
call_type="mcp_call",
)
# Convert result back to MCP response format if blocked/modified
if result is not None:
mcp_response = self._convert_llm_result_to_mcp_during_response(
result, request_obj
)
if mcp_response is not None:
return self._parse_during_mcp_call_hook_response(
response=mcp_response
)
except Exception as e:
raise e
verbose_proxy_logger.exception(
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {}".format(
str(e)
)
)
return None
def _parse_during_mcp_call_hook_response(
self, response: MCPDuringCallResponseObject
) -> Dict[str, Any]:
def _convert_mcp_hook_response_to_kwargs(self, response_data: Optional[dict], original_kwargs: dict) -> dict:
"""
Parse the response from the during_mcp_tool_call_hook
1. Check if execution should continue
2. Handle any error messages
3. Apply any hidden parameter updates
Helper function to convert pre_call_hook response back to kwargs for MCP usage.
"""
result = {
"should_continue": response.should_continue,
"error_message": response.error_message,
"hidden_params": response.hidden_params,
}
return result
if not response_data:
return original_kwargs
# Apply any argument modifications from the hook response
modified_kwargs = original_kwargs.copy()
# If the response contains modified arguments, apply them
if response_data.get("modified_arguments"):
modified_kwargs["arguments"] = response_data["modified_arguments"]
return modified_kwargs
async def process_pre_call_hook_response(self, response, data, call_type):
if isinstance(response, Exception):
@ -975,6 +794,7 @@ class ProxyLogging:
"audio_transcription",
"pass_through_endpoint",
"rerank",
"mcp_call",
],
) -> None:
pass
@ -993,6 +813,7 @@ class ProxyLogging:
"audio_transcription",
"pass_through_endpoint",
"rerank",
"mcp_call",
],
) -> dict:
pass
@ -1010,6 +831,7 @@ class ProxyLogging:
"audio_transcription",
"pass_through_endpoint",
"rerank",
"mcp_call",
],
) -> Optional[dict]:
"""
@ -1081,10 +903,14 @@ class ProxyLogging:
_callback = callback # type: ignore
if _callback is not None and isinstance(_callback, CustomGuardrail):
from litellm.types.guardrails import GuardrailEventHooks
event_type = GuardrailEventHooks.pre_call
if call_type == "mcp_call":
event_type = GuardrailEventHooks.pre_mcp_call
if (
_callback.should_run_guardrail(
data=data, event_type=GuardrailEventHooks.pre_call
data=data, event_type=event_type
)
is not True
):
@ -1108,6 +934,9 @@ class ProxyLogging:
and _callback.__class__.async_pre_call_hook
!= CustomLogger.async_pre_call_hook
):
if call_type == "mcp_call" and user_api_key_dict is None:
continue
response = await _callback.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=self.call_details["user_api_key_cache"],
@ -1126,7 +955,7 @@ class ProxyLogging:
async def during_call_hook(
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
user_api_key_dict: Optional[UserAPIKeyAuth],
call_type: Literal[
"completion",
"responses",
@ -1134,6 +963,7 @@ class ProxyLogging:
"image_generation",
"moderation",
"audio_transcription",
"mcp_call",
],
):
"""
@ -1156,16 +986,26 @@ class ProxyLogging:
# Main - V2 Guardrails implementation
from litellm.types.guardrails import GuardrailEventHooks
event_type = GuardrailEventHooks.during_call
if call_type == "mcp_call":
event_type = GuardrailEventHooks.during_mcp_call
if (
callback.should_run_guardrail(
data=data, event_type=GuardrailEventHooks.during_call
data=data, event_type=event_type
)
is not True
):
continue
# Convert user_api_key_dict to proper format for async_moderation_hook
if call_type == "mcp_call":
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(user_api_key_dict)
else:
user_api_key_auth_dict = user_api_key_dict
await callback.async_moderation_hook(
data=data,
user_api_key_dict=user_api_key_dict,
user_api_key_dict=user_api_key_auth_dict, # type: ignore
call_type=call_type,
)
except Exception as e:
@ -3802,7 +3642,6 @@ def is_valid_api_key(key: str) -> bool:
def construct_database_url_from_env_vars() -> Optional[str]:
"""
Construct a DATABASE_URL from individual environment variables.
Returns:
Optional[str]: The constructed DATABASE_URL or None if required variables are missing
"""
@ -3829,5 +3668,234 @@ def construct_database_url_from_env_vars() -> Optional[str]:
database_url = f"postgresql://{database_username_enc}@{database_host}/{database_name_enc}"
return database_url
return None
async def count_tokens_with_anthropic_api(
model_to_use: str,
messages: Optional[List[Dict[str, Any]]],
deployment: Optional[Dict[str, Any]] = None,
) -> Optional[Dict[str, Any]]:
"""
Helper function to count tokens using Anthropic API directly.
Args:
model_to_use: The model name to use for token counting
messages: The messages to count tokens for
deployment: Optional deployment configuration containing API key
Returns:
Optional dict with token count and tokenizer info, or None if failed
"""
if not messages:
return None
try:
import anthropic
import os
# Get Anthropic API key from deployment config
anthropic_api_key = None
if deployment is not None:
anthropic_api_key = deployment.get("litellm_params", {}).get("api_key")
# Fallback to environment variable
if not anthropic_api_key:
anthropic_api_key = os.getenv("ANTHROPIC_API_KEY")
if anthropic_api_key and messages:
# Call Anthropic API directly for more accurate token counting
client = anthropic.Anthropic(api_key=anthropic_api_key)
# Call with explicit parameters to satisfy type checking
# Type ignore for now since messages come from generic dict input
response = client.beta.messages.count_tokens(
model=model_to_use,
messages=messages, # type: ignore
betas=["token-counting-2024-11-01"]
)
total_tokens = response.input_tokens
tokenizer_used = "anthropic_api"
return {
"total_tokens": total_tokens,
"tokenizer_used": tokenizer_used,
}
except ImportError:
verbose_proxy_logger.warning("Anthropic library not available, falling back to LiteLLM tokenizer")
except Exception as e:
verbose_proxy_logger.warning(f"Error calling Anthropic API: {e}, falling back to LiteLLM tokenizer")
return None
async def get_available_models_for_user(
user_api_key_dict: "UserAPIKeyAuth",
llm_router: Optional["Router"],
general_settings: dict,
user_model: Optional[str],
prisma_client: Optional["PrismaClient"] = None,
proxy_logging_obj: Optional["ProxyLogging"] = None,
team_id: Optional[str] = None,
include_model_access_groups: bool = False,
only_model_access_groups: bool = False,
return_wildcard_routes: bool = False,
user_api_key_cache: Optional["DualCache"] = None,
) -> List[str]:
"""
Get the list of models available to a user based on their API key and team permissions.
Args:
user_api_key_dict: User API key authentication object
llm_router: LiteLLM router instance
general_settings: General settings from config
user_model: User-specific model
prisma_client: Prisma client for database operations
proxy_logging_obj: Proxy logging object
team_id: Specific team ID to check (optional)
include_model_access_groups: Whether to include model access groups
only_model_access_groups: Whether to only return model access groups
return_wildcard_routes: Whether to return wildcard routes
Returns:
List of model names available to the user
"""
from litellm.proxy.auth.model_checks import (
get_key_models,
get_team_models,
get_complete_model_list,
)
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.management_endpoints.team_endpoints import validate_membership
# Get proxy model list and access groups
if llm_router is None:
proxy_model_list = []
model_access_groups = {}
else:
proxy_model_list = llm_router.get_model_names()
model_access_groups = llm_router.get_model_access_groups()
# Get key models
key_models = get_key_models(
user_api_key_dict=user_api_key_dict,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
include_model_access_groups=include_model_access_groups,
)
# Get team models
team_models: List[str] = user_api_key_dict.team_models
# If specific team_id is provided, validate and get team models
if team_id and prisma_client and proxy_logging_obj and user_api_key_cache:
key_models = []
team_object = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
validate_membership(user_api_key_dict=user_api_key_dict, team_table=team_object)
team_models = team_object.models
team_models = get_team_models(
team_models=team_models,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
include_model_access_groups=include_model_access_groups,
)
# Get complete model list
all_models = get_complete_model_list(
key_models=key_models,
team_models=team_models,
proxy_model_list=proxy_model_list,
user_model=user_model,
infer_model_from_keys=general_settings.get("infer_model_from_keys", False),
return_wildcard_routes=return_wildcard_routes,
llm_router=llm_router,
model_access_groups=model_access_groups,
include_model_access_groups=include_model_access_groups,
only_model_access_groups=only_model_access_groups,
)
return all_models
def create_model_info_response(
model_id: str,
provider: str,
include_metadata: bool = False,
fallback_type: Optional[str] = None,
llm_router: Optional["Router"] = None,
) -> dict:
"""
Create a standardized model info response.
Args:
model_id: The model ID
provider: The model provider
include_metadata: Whether to include metadata
fallback_type: Type of fallbacks to include
llm_router: LiteLLM router instance
Returns:
Dictionary containing model information
"""
from litellm.proxy.auth.model_checks import get_all_fallbacks
model_info = {
"id": model_id,
"object": "model",
"created": DEFAULT_MODEL_CREATED_AT_TIME,
"owned_by": provider,
}
# Add metadata if requested
if include_metadata:
metadata = {}
# Default fallback_type to "general" if include_metadata is true
effective_fallback_type = (
fallback_type if fallback_type is not None else "general"
)
# Validate fallback_type
valid_fallback_types = ["general", "context_window", "content_policy"]
if effective_fallback_type not in valid_fallback_types:
raise HTTPException(
status_code=400,
detail=f"Invalid fallback_type. Must be one of: {valid_fallback_types}",
)
fallbacks = get_all_fallbacks(
model=model_id,
llm_router=llm_router,
fallback_type=effective_fallback_type,
)
metadata["fallbacks"] = fallbacks
model_info["metadata"] = metadata
return model_info
def validate_model_access(
model_id: str,
available_models: List[str],
) -> None:
"""
Validate that a model is accessible to the user.
Args:
model_id: The model ID to validate
available_models: List of models available to the user
Raises:
HTTPException: If the model is not accessible
"""
if model_id not in available_models:
raise HTTPException(
status_code=404,
detail="The model `{}` does not exist or is not accessible".format(model_id)
)

View file

@ -0,0 +1,301 @@
import json
from typing import TYPE_CHECKING, Any, List, Optional, Union, cast
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import SpendLogsPayload
from litellm.proxy.spend_tracking.cold_storage_handler import ColdStorageHandler
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionResponseMessage,
GenericChatCompletionMessage,
ResponseInputParam,
)
from litellm.types.utils import ChatCompletionMessageToolCall, Message, ModelResponse
if TYPE_CHECKING:
from litellm.responses.litellm_completion_transformation.transformation import (
ChatCompletionSession,
)
else:
ChatCompletionSession = Any
########################################################
# Cold Storage Handler
########################################################
COLD_STORAGE_HANDLER = ColdStorageHandler()
########################################################
class ResponsesSessionHandler:
@staticmethod
async def get_chat_completion_message_history_for_previous_response_id(
previous_response_id: str,
) -> ChatCompletionSession:
"""
Return the chat completion message history for a previous response id
"""
from litellm.responses.litellm_completion_transformation.transformation import (
ChatCompletionSession,
)
verbose_proxy_logger.debug(
"inside get_chat_completion_message_history_for_previous_response_id"
)
all_spend_logs: List[
SpendLogsPayload
] = await ResponsesSessionHandler.get_all_spend_logs_for_previous_response_id(
previous_response_id
)
verbose_proxy_logger.debug(
"found %s spend logs for this response id", len(all_spend_logs)
)
litellm_session_id: Optional[str] = None
if len(all_spend_logs) > 0:
litellm_session_id = all_spend_logs[0].get("session_id")
chat_completion_message_history: List[
Union[
AllMessageValues,
GenericChatCompletionMessage,
ChatCompletionMessageToolCall,
ChatCompletionResponseMessage,
Message,
]
] = []
for spend_log in all_spend_logs:
chat_completion_message_history = await ResponsesSessionHandler.extend_chat_completion_message_with_spend_log_payload(
spend_log=spend_log,
chat_completion_message_history=chat_completion_message_history,
)
verbose_proxy_logger.debug(
"chat_completion_message_history %s",
json.dumps(chat_completion_message_history, indent=4, default=str),
)
return ChatCompletionSession(
messages=chat_completion_message_history,
litellm_session_id=litellm_session_id,
)
@staticmethod
async def extend_chat_completion_message_with_spend_log_payload(
spend_log: SpendLogsPayload,
chat_completion_message_history: List[
Union[
AllMessageValues,
GenericChatCompletionMessage,
ChatCompletionMessageToolCall,
ChatCompletionResponseMessage,
Message,
]
]
):
"""
Extend the chat completion message history with the spend log payload
"""
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
proxy_server_request_dict = await ResponsesSessionHandler.get_proxy_server_request_from_spend_log(
spend_log=spend_log,
)
response_input_param: Optional[Union[str, ResponseInputParam]] = None
_messages: Optional[Union[str, ResponseInputParam]] = None
############################################################
# Add Input messages for this Spend Log
############################################################
if proxy_server_request_dict:
_response_input_param = proxy_server_request_dict.get("input", None)
_messages = proxy_server_request_dict.get("messages", None)
if isinstance(_response_input_param, str):
response_input_param = _response_input_param
elif isinstance(_response_input_param, dict):
response_input_param = cast(
ResponseInputParam, _response_input_param
)
if response_input_param:
chat_completion_messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=response_input_param,
responses_api_request=proxy_server_request_dict or {},
)
chat_completion_message_history.extend(chat_completion_messages)
############################################################
# Check if `messages` field is present in the proxy server request dict
############################################################
elif _messages:
# ensure all messages are /chat/completions/messages
# certain requests can be stored as Responses API format - this ensures they are transformed to /chat/completions/messages
chat_completion_messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=_messages,
responses_api_request=proxy_server_request_dict or {},
)
chat_completion_message_history.extend(chat_completion_messages)
############################################################
# Add Output messages for this Spend Log
############################################################
_response_output = spend_log.get("response", "{}")
if isinstance(_response_output, dict):
# transform `ChatCompletion Response` to `ResponsesAPIResponse`
model_response = ModelResponse(**_response_output)
for choice in model_response.choices:
if hasattr(choice, "message"):
chat_completion_message_history.append(
getattr(choice, "message")
)
return chat_completion_message_history
@staticmethod
async def get_proxy_server_request_from_spend_log(
spend_log: SpendLogsPayload,
) -> Optional[dict]:
"""
Get the parsed proxy server request from the spend log
"""
proxy_server_request: Union[str, dict] = (
spend_log.get("proxy_server_request") or "{}"
)
proxy_server_request_dict: Optional[dict] = None
if isinstance(proxy_server_request, dict):
proxy_server_request_dict = proxy_server_request
else:
proxy_server_request_dict = json.loads(proxy_server_request)
############################################################
# Check if user has setup cold storage for session handling
############################################################
if ResponsesSessionHandler._should_check_cold_storage_for_full_payload(proxy_server_request_dict):
# Try to get cold storage object key from spend log metadata
_proxy_server_request_dict: Optional[dict] = None
cold_storage_object_key = ResponsesSessionHandler._get_cold_storage_object_key_from_spend_log(spend_log)
if cold_storage_object_key:
# Use the object key directly from metadata
_proxy_server_request_dict = await ResponsesSessionHandler.get_proxy_server_request_from_cold_storage_with_object_key(
object_key=cold_storage_object_key,
)
if _proxy_server_request_dict:
proxy_server_request_dict = _proxy_server_request_dict
return proxy_server_request_dict
@staticmethod
def _get_cold_storage_object_key_from_spend_log(spend_log: SpendLogsPayload) -> Optional[str]:
"""
Extract the cold storage object key from spend log metadata.
Args:
spend_log: The spend log payload containing metadata
Returns:
Optional[str]: The cold storage object key if found, None otherwise
"""
try:
metadata_str = spend_log.get("metadata", "{}")
if isinstance(metadata_str, str):
metadata_dict = json.loads(metadata_str)
return metadata_dict.get("cold_storage_object_key")
elif isinstance(metadata_str, dict):
return metadata_str.get("cold_storage_object_key")
return None
except (json.JSONDecodeError, TypeError, AttributeError):
verbose_proxy_logger.debug("Failed to parse metadata from spend log to extract cold storage object key")
return None
@staticmethod
async def get_proxy_server_request_from_cold_storage_with_object_key(
object_key: str,
) -> Optional[dict]:
"""
Get the proxy server request from cold storage using the object key directly.
Args:
object_key: The S3/GCS object key to retrieve
Returns:
Optional[dict]: The proxy server request dict or None if not found
"""
verbose_proxy_logger.debug("inside get_proxy_server_request_from_cold_storage_with_object_key...")
proxy_server_request_dict = await COLD_STORAGE_HANDLER.get_proxy_server_request_from_cold_storage_with_object_key(
object_key=object_key,
)
return proxy_server_request_dict
@staticmethod
def _should_check_cold_storage_for_full_payload(
proxy_server_request_dict: Optional[dict],
) -> bool:
"""
Only check cold storage when both are true
1. `LITELLM_TRUNCATED_PAYLOAD_FIELD` is in the proxy server request dict
2. `ColdStorageHandler._get_configured_cold_storage_custom_logger()` is not None
"""
from litellm.constants import LITELLM_TRUNCATED_PAYLOAD_FIELD
configured_cold_storage_custom_logger = ColdStorageHandler._get_configured_cold_storage_custom_logger()
if configured_cold_storage_custom_logger is None:
return False
if proxy_server_request_dict is None:
return True
if len(proxy_server_request_dict) == 0:
return True
if LITELLM_TRUNCATED_PAYLOAD_FIELD in proxy_server_request_dict:
return True
return False
@staticmethod
async def get_all_spend_logs_for_previous_response_id(
previous_response_id: str,
) -> List[SpendLogsPayload]:
"""
Get all spend logs for a previous response id
SQL query
SELECT session_id FROM spend_logs WHERE response_id = previous_response_id, SELECT * FROM spend_logs WHERE session_id = session_id
"""
from litellm.proxy.proxy_server import prisma_client
verbose_proxy_logger.debug("decoding response id=%s", previous_response_id)
decoded_response_id = (
ResponsesAPIRequestUtils._decode_responses_api_response_id(
previous_response_id
)
)
previous_response_id = decoded_response_id.get(
"response_id", previous_response_id
)
if prisma_client is None:
return []
query = """
WITH matching_session AS (
SELECT session_id
FROM "LiteLLM_SpendLogs"
WHERE request_id = $1
)
SELECT *
FROM "LiteLLM_SpendLogs"
WHERE session_id IN (SELECT session_id FROM matching_session)
ORDER BY "endTime" ASC;
"""
spend_logs = await prisma_client.db.query_raw(query, previous_response_id)
verbose_proxy_logger.debug(
"Found the following spend logs for previous response id %s: %s",
previous_response_id,
json.dumps(spend_logs, indent=4, default=str),
)
return spend_logs

View file

@ -7,21 +7,15 @@ from typing import Any, Dict, List, Literal, Optional, Tuple, Union, cast
from openai.types.responses.tool_param import FunctionToolParam
from typing_extensions import TypedDict
from litellm._logging import verbose_logger
try:
from litellm_enterprise.enterprise_callbacks.session_handler import (
_ENTERPRISE_ResponsesSessionHandler,
)
except Exception as e:
verbose_logger.debug(
f"[Non-Blocking] Unable to import _ENTERPRISE_ResponsesSessionHandler - LiteLLM Enterprise Feature - {str(e)}"
)
_ENTERPRISE_ResponsesSessionHandler = None
from litellm.caching import InMemoryCache
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.responses.litellm_completion_transformation.session_handler import (
ResponsesSessionHandler,
)
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionImageObject,
ChatCompletionImageUrlObject,
ChatCompletionResponseMessage,
ChatCompletionSystemMessage,
ChatCompletionToolCallChunk,
@ -40,7 +34,6 @@ from litellm.types.llms.openai import (
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
ResponseTextConfig,
ChatCompletionImageUrlObject,
)
from litellm.types.responses.main import (
GenericResponseOutputItem,
@ -75,13 +68,6 @@ class ChatCompletionSession(TypedDict, total=False):
litellm_session_id: Optional[str]
class ChatCompletionImageItem(TypedDict):
"""TypedDict for image items in chat completion content"""
type: Literal["image"]
image_url: ChatCompletionImageUrlObject
########### End of Initialize Classes used for Responses API ###########
@ -210,20 +196,19 @@ class LiteLLMCompletionResponsesConfig:
"""
Async hook to get the chain of previous input and output pairs and return a list of Chat Completion messages
"""
if _ENTERPRISE_ResponsesSessionHandler is not None:
chat_completion_session = ChatCompletionSession(
messages=[], litellm_session_id=None
chat_completion_session = ChatCompletionSession(
messages=[], litellm_session_id=None
)
if previous_response_id:
chat_completion_session = await ResponsesSessionHandler.get_chat_completion_message_history_for_previous_response_id(
previous_response_id=previous_response_id
)
if previous_response_id:
chat_completion_session = await _ENTERPRISE_ResponsesSessionHandler.get_chat_completion_message_history_for_previous_response_id(
previous_response_id=previous_response_id
)
_messages = litellm_completion_request.get("messages") or []
session_messages = chat_completion_session.get("messages") or []
litellm_completion_request["messages"] = session_messages + _messages
litellm_completion_request[
"litellm_trace_id"
] = chat_completion_session.get("litellm_session_id")
_messages = litellm_completion_request.get("messages") or []
session_messages = chat_completion_session.get("messages") or []
litellm_completion_request["messages"] = session_messages + _messages
litellm_completion_request[
"litellm_trace_id"
] = chat_completion_session.get("litellm_session_id")
return litellm_completion_request
@staticmethod
@ -264,6 +249,10 @@ class LiteLLMCompletionResponsesConfig:
chat_completion_messages = LiteLLMCompletionResponsesConfig._transform_responses_api_input_item_to_chat_completion_message(
input_item=_input
)
#########################################################
# If Input Item is a Tool Call Output, add it to the tool_call_output_messages list
#########################################################
if LiteLLMCompletionResponsesConfig._is_input_item_tool_call_output(
input_item=_input
):
@ -316,6 +305,11 @@ class LiteLLMCompletionResponsesConfig:
return LiteLLMCompletionResponsesConfig._transform_responses_api_tool_call_output_to_chat_completion_message(
tool_call_output=input_item
)
elif LiteLLMCompletionResponsesConfig._is_input_item_function_call(input_item):
# handle function call input items
return LiteLLMCompletionResponsesConfig._transform_responses_api_function_call_to_chat_completion_message(
function_call=input_item
)
else:
return [
GenericChatCompletionMessage(
@ -337,6 +331,13 @@ class LiteLLMCompletionResponsesConfig:
"computer_call_output",
]
@staticmethod
def _is_input_item_function_call(input_item: Any) -> bool:
"""
Check if the input item is a function call
"""
return input_item.get("type") == "function_call"
@staticmethod
def _transform_responses_api_tool_call_output_to_chat_completion_message(
tool_call_output: Dict[str, Any],
@ -402,6 +403,52 @@ class LiteLLMCompletionResponsesConfig:
return [tool_output_message]
@staticmethod
def _transform_responses_api_function_call_to_chat_completion_message(
function_call: Dict[str, Any],
) -> List[
Union[
AllMessageValues,
GenericChatCompletionMessage,
ChatCompletionResponseMessage,
]
]:
"""
Transform a Responses API function_call into a Chat Completion message with tool calls
Handles Input items of this type:
function_call:
```json
{
"type": "function_call",
"arguments":"{\"location\": \"São Paulo, Brazil\"}",
"call_id": "call_v2wlBzrlTIFl9FxPeY774GHZ",
"name": "get_weather",
"id": "fc_685c42deefc0819a822b6936faaa30be0c76bc1491ab6619",
"status": "completed"
}
```
"""
# Create a tool call for the function call
tool_call = ChatCompletionToolCallChunk(
id=function_call.get("call_id") or function_call.get("id") or "",
type="function",
function=ChatCompletionToolCallFunctionChunk(
name=function_call.get("name") or "",
arguments=function_call.get("arguments") or "",
),
index=0,
)
# Create an assistant message with the tool call
chat_completion_response_message = ChatCompletionResponseMessage(
tool_calls=[tool_call],
role="assistant",
content=None, # Function calls don't have content
)
return [chat_completion_response_message]
@staticmethod
def _transform_input_file_item_to_file_item(item: Dict[str, Any]) -> Dict[str, Any]:
"""
@ -423,7 +470,7 @@ class LiteLLMCompletionResponsesConfig:
return new_item
@staticmethod
def _transform_input_image_item_to_image_item(item: Dict[str, Any]) -> ChatCompletionImageItem:
def _transform_input_image_item_to_image_item(item: Dict[str, Any]) -> ChatCompletionImageObject:
"""
Transform a Responses API input_image item to a Chat Completion image item
"""
@ -432,8 +479,8 @@ class LiteLLMCompletionResponsesConfig:
detail=item.get("detail") or "auto"
)
return ChatCompletionImageItem(
type="image",
return ChatCompletionImageObject(
type="image_url",
image_url=image_url_obj
)
@ -444,7 +491,6 @@ class LiteLLMCompletionResponsesConfig:
"""
Transform a Responses API content into a Chat Completion content
"""
if isinstance(content, str):
return content
elif isinstance(content, list):

View file

@ -243,6 +243,8 @@ async def aresponses(
top_p: Optional[float] = None,
truncation: Optional[Literal["auto", "disabled"]] = None,
user: Optional[str] = None,
service_tier: Optional[str] = None,
safety_identifier: Optional[str] = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: Optional[Dict[str, Any]] = None,
@ -294,6 +296,8 @@ async def aresponses(
extra_body=extra_body,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
service_tier=service_tier,
safety_identifier=safety_identifier,
**kwargs,
)
@ -350,6 +354,8 @@ def responses(
top_p: Optional[float] = None,
truncation: Optional[Literal["auto", "disabled"]] = None,
user: Optional[str] = None,
service_tier: Optional[str] = None,
safety_identifier: Optional[str] = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: Optional[Dict[str, Any]] = None,

View file

@ -406,6 +406,9 @@ class Router:
self.default_max_parallel_requests = default_max_parallel_requests
self.provider_default_deployment_ids: List[str] = []
self.pattern_router = PatternMatchRouter()
self.team_pattern_routers: Dict[str, PatternMatchRouter] = (
{}
) # {"TEAM_ID": PatternMatchRouter}
self.auto_routers: Dict[str, "AutoRouter"] = {}
if model_list is not None:
@ -1207,7 +1210,7 @@ class Router:
verbose_router_logger.error(
f"Fallback also failed: {fallback_error}"
)
raise fallback_error
raise fallback_error
return FallbackStreamWrapper(stream_with_fallbacks())
@ -1403,7 +1406,7 @@ class Router:
kwargs.setdefault(metadata_variable_name, {}).update(metadata_defaults)
def _handle_clientside_credential(
self, deployment: dict, kwargs: dict
self, deployment: dict, kwargs: dict, function_name: Optional[str] = None
) -> Deployment:
"""
Handle clientside credential
@ -1413,8 +1416,11 @@ class Router:
dynamic_litellm_params = get_dynamic_litellm_params(
litellm_params=litellm_params, request_kwargs=kwargs
)
metadata = kwargs.get("metadata", {})
model_group = cast(str, metadata.get("model_group"))
# Use deployment model_name as model_group for generating model_id
metadata_variable_name = _get_router_metadata_variable_name(
function_name=function_name,
)
model_group = kwargs.get(metadata_variable_name, {}).get("model_group")
_model_id = self._generate_model_id(
model_group=model_group, litellm_params=dynamic_litellm_params
)
@ -1448,7 +1454,7 @@ class Router:
deployment_model_name = deployment["model_name"]
if is_clientside_credential(request_kwargs=kwargs):
deployment_pydantic_obj = self._handle_clientside_credential(
deployment=deployment, kwargs=kwargs
deployment=deployment, kwargs=kwargs, function_name=function_name
)
model_info = deployment_pydantic_obj.model_info.model_dump()
deployment_litellm_model_name = deployment_pydantic_obj.litellm_params.model
@ -5109,6 +5115,19 @@ class Router:
if deployment.model_info.id:
self.provider_default_deployment_ids.append(deployment.model_info.id)
_team_id = deployment.model_info.get("team_id")
_team_public_model_name = deployment.model_info.get("team_public_model_name")
if (
_team_id is not None
and _team_public_model_name is not None
and "*" in _team_public_model_name
):
if _team_id not in self.team_pattern_routers:
self.team_pattern_routers[_team_id] = PatternMatchRouter()
self.team_pattern_routers[_team_id].add_pattern(
_team_public_model_name, deployment.to_json(exclude_none=True)
)
# Azure GPT-Vision Enhancements, users can pass os.environ/
data_sources = deployment.litellm_params.get("dataSources", []) or []
@ -5920,19 +5939,17 @@ class Router:
Map a team model name to a team-specific model name.
Returns:
- team_model_name: str - the team-specific model name
- deployment id: str - the deployment id of the team-specific model
- None: if no team-specific model name is found
"""
for model in self.model_list:
model_team_id = model["model_info"].get("team_id")
model_team_public_model_name = model["model_info"].get(
"team_public_model_name"
)
if (
model_team_id == team_id
and model_team_public_model_name == team_model_name
):
return model["model_name"]
models = self.get_model_list(model_name=team_model_name, team_id=team_id)
if not models:
return None
for model in models:
if model.get("model_info", {}).get("team_id") == team_id:
return model.get("model_name")
## wildcard models
return None
def should_include_deployment(
@ -6073,6 +6090,7 @@ class Router:
if team_id specified, returns matching team-specific models
"""
if hasattr(self, "model_list"):
returned_models: List[DeploymentTypedDict] = []
@ -6087,7 +6105,17 @@ class Router:
)
if len(returned_models) == 0: # check if wildcard route
potential_wildcard_models = self.pattern_router.route(model_name)
potential_wildcard_models = self.pattern_router.route(model_name) or []
## check for team-specific wildcard models
if team_id is not None and team_id in self.team_pattern_routers:
potential_team_only_wildcard_models = (
self.team_pattern_routers[team_id].route(model_name) or []
)
potential_wildcard_models.extend(
potential_team_only_wildcard_models
)
if model_name is not None and potential_wildcard_models is not None:
for m in potential_wildcard_models:
deployment_typed_dict = DeploymentTypedDict(**m) # type: ignore
@ -6519,6 +6547,7 @@ class Router:
messages: Optional[List[Dict[str, str]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
request_kwargs: Optional[Dict] = None,
) -> Tuple[str, Union[List, Dict]]:
"""
Common checks for 'get_available_deployment' across sync + async call.
@ -6530,6 +6559,14 @@ class Router:
- List, if multiple models chosen
- Dict, if specific model chosen
"""
request_team_id: Optional[str] = None
if request_kwargs is not None:
metadata = request_kwargs.get("metadata") or {}
litellm_metadata = request_kwargs.get("litellm_metadata") or {}
request_team_id = metadata.get(
"user_api_key_team_id"
) or litellm_metadata.get("user_api_key_team_id")
# check if aliases set on litellm model alias map
if specific_deployment is True:
return model, self._get_deployment_by_litellm_model(model=model)
@ -6552,9 +6589,22 @@ class Router:
pattern_deployments = self.pattern_router.get_deployments_by_pattern(
model=model,
)
if pattern_deployments:
return model, pattern_deployments
if (
request_team_id is not None
and request_team_id in self.team_pattern_routers
):
pattern_deployments = self.team_pattern_routers[
request_team_id
].get_deployments_by_pattern(
model=model,
)
if pattern_deployments:
return model, pattern_deployments
# check if default deployment is set
if self.default_deployment is not None:
updated_deployment = copy.deepcopy(
@ -6622,6 +6672,7 @@ class Router:
messages=messages,
input=input,
specific_deployment=specific_deployment,
request_kwargs=request_kwargs,
) # type: ignore
# IF TEAM ID SPECIFIED ON MODEL, AND REQUEST CONTAINS USER_API_KEY_TEAM_ID, FILTER OUT MODELS THAT ARE NOT IN THE TEAM

View file

@ -12,6 +12,7 @@ class LiteLLMCacheType(str, Enum):
DISK = "disk"
QDRANT_SEMANTIC = "qdrant-semantic"
AZURE_BLOB = "azure-blob"
GCS = "gcs"
CachingSupportedCallTypes = Literal[

183
litellm/types/llms/oci.py Normal file
View file

@ -0,0 +1,183 @@
from __future__ import annotations
from enum import Enum
from typing import Any, Dict, List, Literal, Optional, Union
from pydantic import BaseModel
OCIRoles = Literal["SYSTEM", "USER", "ASSISTANT", "TOOL"]
class OCIVendors(Enum):
"""
A class to hold the vendor names for OCI models.
This is used to map model names to their respective vendors.
"""
COHERE = "COHERE"
GENERIC = "GENERIC"
# --- Base Models and Content Parts ---
class OCIContentPart(BaseModel):
"""Base model for content parts in an OCI message."""
pass
class OCITextContentPart(OCIContentPart):
"""Text content part for the OCI API."""
type: Literal["TEXT"] = "TEXT"
text: str
class OCIImageContentPart(OCIContentPart):
"""Image content part for the OCI API."""
type: Literal["IMAGE"] = "IMAGE"
imageUrl: str
OCIContentPartUnion = Union[OCITextContentPart, OCIImageContentPart]
# --- Models for Tools and Tool Calls ---
class OCIToolCall(BaseModel):
"""Represents a tool call made by the model."""
id: str
type: Literal["FUNCTION"] = "FUNCTION"
name: str
arguments: str # Arguments should be a JSON-serialized string
class OCIToolDefinition(BaseModel):
"""Defines a tool that can be used by the model."""
type: Literal["FUNCTION"] = "FUNCTION"
name: Optional[str] = None
description: Optional[str] = None
parameters: Optional[dict] = None
# --- Message Models (Request and Response) ---
class OCIMessage(BaseModel):
"""Model for a single message in the request/response payload."""
role: OCIRoles
content: Optional[List[OCIContentPartUnion]] = None
toolCalls: Optional[List[OCIToolCall]] = None
toolCallId: Optional[str] = None
# --- Request Payload Models ---
class OCIChatRequestPayload(BaseModel):
"""Internal 'chatRequest' payload for the OCI API."""
apiFormat: str
messages: List[OCIMessage]
tools: Optional[List[OCIToolDefinition]] = None
isStream: bool = False
numGenerations: Optional[int] = None
maxTokens: Optional[int] = None
temperature: Optional[float] = None
topP: Optional[float] = None
stop: Optional[List[str]] = None
seed: Optional[int] = None
frequencyPenalty: Optional[float] = None
presencePenalty: Optional[float] = None
class OCIServingMode(BaseModel):
"""Defines the serving mode and the model to be used."""
servingType: str
modelId: str
class OCICompletionPayload(BaseModel):
"""Pydantic model for the complete OCI chat request body."""
compartmentId: str
servingMode: OCIServingMode
chatRequest: OCIChatRequestPayload
# --- API Response Models (Non-streaming) ---
class OCICompletionTokenDetails(BaseModel):
"""Completion token details in the OCI response."""
acceptedPredictionTokens: int
reasoningTokens: int
class OCIPropmtTokensDetails(BaseModel):
"""Prompt token details in the OCI response."""
cachedTokens: int
class OCIResponseUsage(BaseModel):
"""Token usage in the OCI response."""
promptTokens: int
completionTokens: int
totalTokens: int
completionTokensDetails: OCICompletionTokenDetails
promptTokensDetails: OCIPropmtTokensDetails
class OCIResponseChoice(BaseModel):
"""A completion choice in the OCI response."""
index: int
message: OCIMessage
finishReason: Optional[str] = None
logprobs: Optional[Dict[str, Any]] = None
class OCIChatResponse(BaseModel):
"""The 'chatResponse' object in the OCI response."""
apiFormat: str
timeCreated: str
choices: List[OCIResponseChoice]
usage: OCIResponseUsage
class OCICompletionResponse(BaseModel):
"""Model for the complete non-streaming OCI response body."""
modelId: str
modelVersion: str
chatResponse: OCIChatResponse
# --- API Response Models (Streaming) ---
class OCIStreamDelta(BaseModel):
"""The content delta in a streaming chunk."""
content: Optional[List[OCIContentPartUnion]] = None
role: Optional[str] = None
toolCalls: Optional[List[OCIToolCall]] = None
class OCIStreamChunk(BaseModel):
"""Model for a single SSE event chunk from OCI."""
finishReason: Optional[str] = None
message: Optional[OCIStreamDelta] = None
pad: Optional[str] = None
index: Optional[int] = None

View file

@ -965,6 +965,8 @@ class ResponsesAPIOptionalRequestParams(TypedDict, total=False):
top_p: Optional[float]
truncation: Optional[Literal["auto", "disabled"]]
user: Optional[str]
service_tier: Optional[str]
safety_identifier: Optional[str]
prompt: Optional[PromptObject]

View file

@ -887,6 +887,10 @@ class Usage(CompletionUsage):
) # hidden param for prompt caching. Might change, once openai introduces their equivalent.
server_tool_use: Optional[ServerToolUse] = None
cost: Optional[float] = None
completion_tokens_details: Optional[CompletionTokensDetailsWrapper] = None
"""Breakdown of tokens used in a completion."""
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None
"""Breakdown of tokens used in the prompt."""
@ -904,6 +908,7 @@ class Usage(CompletionUsage):
Union[CompletionTokensDetailsWrapper, dict]
] = None,
server_tool_use: Optional[ServerToolUse] = None,
cost: Optional[float] = None,
**params,
):
# handle reasoning_tokens
@ -975,6 +980,11 @@ class Usage(CompletionUsage):
else: # maintain openai compatibility in usage object if possible
del self.server_tool_use
if cost is not None:
self.cost = cost
else:
del self.cost
## ANTHROPIC MAPPING ##
if "cache_creation_input_tokens" in params and isinstance(
params["cache_creation_input_tokens"], int
@ -1900,6 +1910,7 @@ class StandardLoggingMetadata(StandardLoggingUserAPIKeyMetadata):
vector_store_request_metadata: Optional[List[StandardLoggingVectorStoreRequest]]
applied_guardrails: Optional[List[str]]
usage_object: Optional[dict]
cold_storage_object_key: Optional[str] # S3/GCS object key for cold storage retrieval
class StandardLoggingAdditionalHeaders(TypedDict, total=False):
@ -2318,6 +2329,7 @@ class LlmProviders(str, Enum):
PG_VECTOR = "pg_vector"
HYPERBOLIC = "hyperbolic"
RECRAFT = "recraft"
OCI = "oci"
AUTO_ROUTER = "auto_router"
DOTPROMPT = "dotprompt"

View file

@ -2790,7 +2790,10 @@ def get_optional_params_embeddings( # noqa: PLR0915
)
_check_valid_arg(supported_params=supported_params)
optional_params = litellm.JinaAIEmbeddingConfig().map_openai_params(
non_default_params=non_default_params, optional_params={}
non_default_params=non_default_params,
optional_params={},
model=model,
drop_params=drop_params if drop_params is not None else False,
)
elif custom_llm_provider == "voyage":
supported_params = get_supported_openai_params(
@ -6604,6 +6607,19 @@ def validate_and_fix_openai_messages(messages: List):
new_messages.append(cleaned_message)
return validate_chat_completion_user_messages(messages=new_messages)
def validate_and_fix_openai_tools(tools: Optional[List]) -> Optional[List[dict]]:
"""
Ensure tools is List[dict] and not List[BaseModel]
"""
new_tools = []
if tools is None:
return tools
for tool in tools:
if isinstance(tool, BaseModel):
new_tools.append(tool.model_dump())
elif isinstance(tool, dict):
new_tools.append(tool)
return new_tools
def cleanup_none_field_in_message(message: AllMessageValues):
"""
@ -6916,6 +6932,8 @@ class ProviderConfigManager:
return litellm.OpenAIGPTConfig()
elif litellm.LlmProviders.NSCALE == provider:
return litellm.NscaleConfig()
elif litellm.LlmProviders.OCI == provider:
return litellm.OCIChatConfig()
elif litellm.LlmProviders.HYPERBOLIC == provider:
return litellm.HyperbolicChatConfig()
return None
@ -6940,6 +6958,12 @@ class ProviderConfigManager:
from litellm.llms.cohere.embed.transformation import CohereEmbeddingConfig
return CohereEmbeddingConfig()
elif litellm.LlmProviders.JINA_AI == provider:
from litellm.llms.jina_ai.embedding.transformation import (
JinaAIEmbeddingConfig,
)
return JinaAIEmbeddingConfig()
return None
@staticmethod

View file

@ -296,6 +296,60 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-5-2025-08-07": {
"max_tokens": 128000,
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"input_cost_per_token": 1.25e-06,
"output_cost_per_token": 1e-05,
"cache_read_input_token_cost": 1.25e-07,
"litellm_provider": "openai",
"mode": "chat",
"supports_pdf_input": true,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_prompt_caching": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-5-mini-2025-08-07": {
"max_tokens": 128000,
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"input_cost_per_token": 2.5e-07,
"output_cost_per_token": 2e-06,
"cache_read_input_token_cost": 2.5e-08,
"litellm_provider": "openai",
"mode": "chat",
"supports_pdf_input": true,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_prompt_caching": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-5-nano-2025-08-07": {
"max_tokens": 128000,
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"input_cost_per_token": 5e-08,
"output_cost_per_token": 4e-07,
"cache_read_input_token_cost": 5e-09,
"litellm_provider": "openai",
"mode": "chat",
"supports_pdf_input": true,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_prompt_caching": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"watsonx/ibm/granite-3-8b-instruct": {
"max_tokens": 8192,
"max_input_tokens": 8192,
@ -607,9 +661,9 @@
"supports_system_messages": true,
"supports_tool_choice": true,
"search_context_cost_per_query": {
"search_context_size_low": 30.0,
"search_context_size_medium": 35.0,
"search_context_size_high": 50.0
"search_context_size_low": 0.025,
"search_context_size_medium": 0.0275,
"search_context_size_high": 0.03
}
},
"codex-mini-latest": {
@ -3750,7 +3804,7 @@
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 3.3e-06,
"output_cost_per_token": 16.5e-06,
"output_cost_per_token": 1.65e-05,
"litellm_provider": "azure_ai",
"mode": "chat",
"supports_function_calling": true,
@ -3764,7 +3818,7 @@
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 3e-06,
"output_cost_per_token": 15e-06,
"output_cost_per_token": 1.5e-05,
"litellm_provider": "azure_ai",
"mode": "chat",
"supports_function_calling": true,
@ -3777,7 +3831,7 @@
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 0.25e-06,
"input_cost_per_token": 2.5e-07,
"output_cost_per_token": 1.27e-06,
"litellm_provider": "azure_ai",
"mode": "chat",
@ -3792,7 +3846,7 @@
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 0.275e-06,
"input_cost_per_token": 2.75e-07,
"output_cost_per_token": 1.38e-06,
"litellm_provider": "azure_ai",
"mode": "chat",
@ -5486,6 +5540,36 @@
"litellm_provider": "groq",
"mode": "audio_transcription"
},
"groq/openai/gpt-oss-20b": {
"max_tokens": 32768,
"max_input_tokens": 131072,
"max_output_tokens": 32768,
"input_cost_per_token": 1e-07,
"output_cost_per_token": 5e-07,
"litellm_provider": "groq",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_response_schema": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_web_search": true
},
"groq/openai/gpt-oss-120b": {
"max_tokens": 32766,
"max_input_tokens": 131072,
"max_output_tokens": 32766,
"input_cost_per_token": 1.5e-07,
"output_cost_per_token": 7.5e-07,
"litellm_provider": "groq",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_response_schema": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_web_search": true
},
"cerebras/llama3.1-8b": {
"max_tokens": 128000,
"max_input_tokens": 128000,
@ -5741,6 +5825,32 @@
"supports_reasoning": true,
"supports_computer_use": true
},
"claude-opus-4-1-20250805": {
"max_tokens": 32000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 1.5e-05,
"output_cost_per_token": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01,
"search_context_size_high": 0.01
},
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "anthropic",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 159,
"supports_assistant_prefill": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_computer_use": true
},
"claude-sonnet-4-20250514": {
"max_tokens": 64000,
"max_input_tokens": 200000,
@ -7337,12 +7447,12 @@
"max_audio_length_hours": 8.4,
"max_audio_per_prompt": 1,
"max_pdf_size_mb": 30,
"input_cost_per_token": 3.5e-07,
"input_cost_per_token": 3.5e-07,
"input_cost_per_audio_token": 2.1e-06,
"input_cost_per_image": 2.1e-06,
"input_cost_per_video_per_second": 2.1e-06,
"output_cost_per_token": 1.5e-06,
"output_cost_per_audio_token": 8.5e-06,
"output_cost_per_audio_token": 8.5e-06,
"litellm_provider": "gemini",
"mode": "chat",
"rpm": 10,
@ -8690,6 +8800,40 @@
"source": "https://aistudio.google.com",
"supports_tool_choice": true
},
"vertex_ai/claude-opus-4-1": {
"max_tokens": 4096,
"max_input_tokens": 200000,
"max_output_tokens": 4096,
"input_cost_per_token": 15e-06,
"output_cost_per_token": 75e-06,
"input_cost_per_token_batches": 7.5e-06,
"output_cost_per_token_batches": 37.5e-06,
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"supports_assistant_prefill": true,
"supports_tool_choice": true
},
"vertex_ai/claude-opus-4-1@20250805": {
"max_tokens": 4096,
"max_input_tokens": 200000,
"max_output_tokens": 4096,
"input_cost_per_token": 15e-06,
"output_cost_per_token": 75e-06,
"input_cost_per_token_batches": 7.5e-06,
"output_cost_per_token_batches": 37.5e-06,
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"supports_assistant_prefill": true,
"supports_tool_choice": true
},
"vertex_ai/claude-3-sonnet": {
"max_tokens": 4096,
"max_input_tokens": 200000,
@ -9039,9 +9183,9 @@
"supports_tool_choice": true
},
"vertex_ai/meta/llama-4-scout-17b-16e-instruct-maas": {
"max_tokens": 10000000.0,
"max_input_tokens": 10000000.0,
"max_output_tokens": 10000000.0,
"max_tokens": 10000000,
"max_input_tokens": 10000000,
"max_output_tokens": 10000000,
"input_cost_per_token": 2.5e-07,
"output_cost_per_token": 7e-07,
"litellm_provider": "vertex_ai-llama_models",
@ -9059,9 +9203,9 @@
]
},
"vertex_ai/meta/llama-4-scout-17b-128e-instruct-maas": {
"max_tokens": 10000000.0,
"max_input_tokens": 10000000.0,
"max_output_tokens": 10000000.0,
"max_tokens": 10000000,
"max_input_tokens": 10000000,
"max_output_tokens": 10000000,
"input_cost_per_token": 2.5e-07,
"output_cost_per_token": 7e-07,
"litellm_provider": "vertex_ai-llama_models",
@ -9079,9 +9223,9 @@
]
},
"vertex_ai/meta/llama-4-maverick-17b-128e-instruct-maas": {
"max_tokens": 1000000.0,
"max_input_tokens": 1000000.0,
"max_output_tokens": 1000000.0,
"max_tokens": 1000000,
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"input_cost_per_token": 3.5e-07,
"output_cost_per_token": 1.15e-06,
"litellm_provider": "vertex_ai-llama_models",
@ -9099,9 +9243,9 @@
]
},
"vertex_ai/meta/llama-4-maverick-17b-16e-instruct-maas": {
"max_tokens": 1000000.0,
"max_input_tokens": 1000000.0,
"max_output_tokens": 1000000.0,
"max_tokens": 1000000,
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"input_cost_per_token": 3.5e-07,
"output_cost_per_token": 1.15e-06,
"litellm_provider": "vertex_ai-llama_models",
@ -9174,7 +9318,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 2048,
"input_cost_per_token": 5e-06,
"output_cost_per_token": 16e-06,
"output_cost_per_token": 1.6e-05,
"litellm_provider": "vertex_ai-llama_models",
"mode": "chat",
"supports_system_messages": true,
@ -10480,7 +10624,7 @@
"supports_tool_choice": true,
"supports_prompt_caching": true
},
"openrouter/x-ai/grok-4":{
"openrouter/x-ai/grok-4": {
"max_tokens": 256000,
"max_input_tokens": 256000,
"max_output_tokens": 256000,
@ -10494,12 +10638,12 @@
"source": "https://openrouter.ai/x-ai/grok-4",
"supports_web_search": true
},
"openrouter/bytedance/ui-tars-1.5-7b":{
"openrouter/bytedance/ui-tars-1.5-7b": {
"max_tokens": 2048,
"max_input_tokens": 131072,
"max_output_tokens": 2048,
"input_cost_per_token": 0.1e-06,
"output_cost_per_token": 0.2e-06,
"input_cost_per_token": 1e-07,
"output_cost_per_token": 2e-07,
"litellm_provider": "openrouter",
"mode": "chat",
"source": "https://openrouter.ai/api/v1/models/bytedance/ui-tars-1.5-7b",
@ -11178,8 +11322,8 @@
"max_tokens": 8192,
"max_input_tokens": 8192,
"max_output_tokens": 2048,
"input_cost_per_token": 0.21e-06,
"output_cost_per_token": 0.63e-06,
"input_cost_per_token": 2.1e-07,
"output_cost_per_token": 6.3e-07,
"litellm_provider": "openrouter",
"mode": "chat",
"supports_tool_choice": true
@ -11891,6 +12035,60 @@
"supports_pdf_input": true,
"supports_tool_choice": true
},
"openai.gpt-oss-20b-1:0": {
"max_tokens": 128000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 7e-08,
"output_cost_per_token": 3e-07,
"litellm_provider": "bedrock_converse",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true
},
"openai.gpt-oss-120b-1:0": {
"max_tokens": 128000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 1.5e-07,
"output_cost_per_token": 6e-07,
"litellm_provider": "bedrock_converse",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true
},
"anthropic.claude-opus-4-1-20250805-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 1.5e-05,
"output_cost_per_token": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01,
"search_context_size_high": 0.01
},
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "bedrock_converse",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 159,
"supports_assistant_prefill": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_computer_use": true
},
"anthropic.claude-opus-4-20250514-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
@ -12093,6 +12291,32 @@
"supports_tool_choice": true,
"supports_reasoning": true
},
"us.anthropic.claude-opus-4-1-20250805-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 1.5e-05,
"output_cost_per_token": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01,
"search_context_size_high": 0.01
},
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "bedrock_converse",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 159,
"supports_assistant_prefill": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_computer_use": true
},
"us.anthropic.claude-opus-4-20250514-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
@ -12266,6 +12490,32 @@
"supports_pdf_input": true,
"supports_tool_choice": true
},
"eu.anthropic.claude-opus-4-1-20250805-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
"max_output_tokens": 32000,
"input_cost_per_token": 1.5e-05,
"output_cost_per_token": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01,
"search_context_size_high": 0.01
},
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
"litellm_provider": "bedrock_converse",
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 159,
"supports_assistant_prefill": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_computer_use": true
},
"eu.anthropic.claude-opus-4-20250514-v1:0": {
"max_tokens": 32000,
"max_input_tokens": 200000,
@ -14763,7 +15013,7 @@
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 16384,
"input_cost_per_token": 0.6e-06,
"input_cost_per_token": 6e-07,
"output_cost_per_token": 2.5e-06,
"litellm_provider": "fireworks_ai",
"mode": "chat",
@ -14809,6 +15059,58 @@
"source": "https://fireworks.ai/pricing",
"supports_tool_choice": false
},
"fireworks_ai/accounts/fireworks/models/glm-4p5": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 96000,
"input_cost_per_token": 5.5e-07,
"output_cost_per_token": 2.19e-06,
"litellm_provider": "fireworks_ai",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"source": "https://fireworks.ai/models/fireworks/glm-4p5"
},
"fireworks_ai/accounts/fireworks/models/glm-4p5-air": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 96000,
"input_cost_per_token": 2.2e-07,
"output_cost_per_token": 8.8e-07,
"litellm_provider": "fireworks_ai",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"source": "https://artificialanalysis.ai/models/glm-4-5-air"
},
"fireworks_ai/accounts/fireworks/models/gpt-oss-120b": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 1.5e-07,
"output_cost_per_token": 6e-07,
"litellm_provider": "fireworks_ai",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"source": "https://fireworks.ai/pricing"
},
"fireworks_ai/accounts/fireworks/models/gpt-oss-20b": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 5e-08,
"output_cost_per_token": 2e-07,
"litellm_provider": "fireworks_ai",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"source": "https://fireworks.ai/pricing"
},
"fireworks_ai/nomic-ai/nomic-embed-text-v1.5": {
"max_tokens": 8192,
"max_input_tokens": 8192,
@ -17445,5 +17747,136 @@
"supports_vision": false,
"supports_system_messages": true,
"supports_tool_choice": false
},
"oci/meta.llama-4-maverick-17b-128e-instruct-fp8": {
"max_tokens": 512000,
"max_input_tokens": 512000,
"max_output_tokens": 4000,
"input_cost_per_token": 7.2e-07,
"output_cost_per_token": 7.2e-07,
"litellm_provider": "oci",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_tool_choice": false,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing"
},
"oci/meta.llama-4-scout-17b-16e-instruct": {
"max_tokens": 192000,
"max_input_tokens": 192000,
"max_output_tokens": 4000,
"input_cost_per_token": 7.2e-07,
"output_cost_per_token": 7.2e-07,
"litellm_provider": "oci",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_tool_choice": false,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing"
},
"oci/meta.llama-3.3-70b-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 4000,
"input_cost_per_token": 7.2e-07,
"output_cost_per_token": 7.2e-07,
"litellm_provider": "oci",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_tool_choice": false,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing"
},
"oci/meta.llama-3.2-90b-vision-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 4000,
"input_cost_per_token": 2.0e-06,
"output_cost_per_token": 2.0e-06,
"litellm_provider": "oci",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_tool_choice": false,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing"
},
"oci/meta.llama-3.1-405b-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 4000,
"input_cost_per_token": 1.068e-05,
"output_cost_per_token": 1.068e-05,
"litellm_provider": "oci",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_tool_choice": false,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing"
},
"oci/xai.grok-4": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 3.0e-06,
"output_cost_per_token": 1.5e-07,
"litellm_provider": "oci",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_tool_choice": false,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing"
},
"oci/xai.grok-3": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 3.0e-06,
"output_cost_per_token": 1.5e-07,
"litellm_provider": "oci",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_tool_choice": false,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing"
},
"oci/xai.grok-3-mini": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 3.0e-07,
"output_cost_per_token": 5.0e-07,
"litellm_provider": "oci",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_tool_choice": false,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing"
},
"oci/xai.grok-3-fast": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 5.0e-06,
"output_cost_per_token": 2.5e-05,
"litellm_provider": "oci",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_tool_choice": false,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing"
},
"oci/xai.grok-3-mini-fast": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 6.0e-07,
"output_cost_per_token": 4.0e-06,
"litellm_provider": "oci",
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_tool_choice": false,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing"
}
}

512
poetry.lock generated

File diff suppressed because it is too large Load diff

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm"
version = "1.74.15"
version = "1.75.2"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI"]
license = "MIT"
@ -51,6 +51,7 @@ azure-identity = {version = "^1.15.0", optional = true}
azure-keyvault-secrets = {version = "^4.8.0", optional = true}
azure-storage-blob = {version="^12.25.1", optional=true}
google-cloud-kms = {version = "^2.21.3", optional = true}
google-cloud-iam = {version = "^2.19.1", optional = true}
resend = {version = "^0.8.0", optional = true}
pynacl = {version = "^1.5.0", optional = true}
websockets = {version = "^13.1.0", optional = true}
@ -59,7 +60,7 @@ redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.
mcp = {version = "^1.10.0", optional = true, python = ">=3.10"}
litellm-proxy-extras = {version = "0.2.15", optional = true}
rich = {version = "13.7.1", optional = true}
litellm-enterprise = {version = "0.1.16", optional = true}
litellm-enterprise = {version = "0.1.19", optional = true}
diskcache = {version = "^5.6.1", optional = true}
polars = {version = "^1.31.0", optional = true, python = ">=3.10"}
semantic-router = {version = "*", optional = true, python = ">=3.9"}
@ -97,6 +98,7 @@ extra_proxy = [
"azure-identity",
"azure-keyvault-secrets",
"google-cloud-kms",
"google-cloud-iam",
"resend",
"redisvl"
]
@ -152,7 +154,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "1.74.15"
version = "1.75.2"
version_files = [
"pyproject.toml:^version"
]

View file

@ -14,6 +14,7 @@ prisma==0.11.0 # for db
mangum==0.17.0 # for aws lambda functions
pynacl==1.5.0 # for encrypting keys
google-cloud-aiplatform==1.47.0 # for vertex ai calls
google-cloud-iam==2.19.1 # for GCP IAM Redis authentication
google-genai==1.22.0
anthropic[vertex]==0.54.0
mcp==1.10.1 # for MCP server
@ -59,4 +60,4 @@ websockets==13.1.0 # for realtime API
########################
# LITELLM ENTERPRISE DEPENDENCIES
########################
litellm-enterprise==0.1.16
litellm-enterprise==0.1.19

View file

@ -7,6 +7,8 @@ sys.path.insert(0, os.path.abspath("../.."))
import asyncio
import logging
import uuid
from datetime import datetime, timedelta, timezone
from unittest.mock import MagicMock, call, patch
import pytest
from prometheus_client import REGISTRY, CollectorRegistry
@ -16,16 +18,18 @@ from litellm import completion
from litellm._logging import verbose_logger
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.types.utils import (
StandardLoggingPayload,
StandardLoggingMetadata,
StandardLoggingHiddenParams,
StandardLoggingMetadata,
StandardLoggingModelInformation,
StandardLoggingPayload,
)
import pytest
from unittest.mock import MagicMock, patch, call
from datetime import datetime, timedelta, timezone
try:
from litellm_enterprise.integrations.prometheus import PrometheusLogger, UserAPIKeyLabelValues, get_custom_labels_from_metadata
from litellm_enterprise.integrations.prometheus import (
PrometheusLogger,
UserAPIKeyLabelValues,
get_custom_labels_from_metadata,
)
except Exception:
PrometheusLogger = None
from litellm.proxy._types import UserAPIKeyAuth
@ -1054,6 +1058,7 @@ def test_increment_deployment_cooled_down(prometheus_logger):
@pytest.mark.parametrize("enable_end_user_cost_tracking_prometheus_only", [True, False])
def test_prometheus_factory(monkeypatch, enable_end_user_cost_tracking_prometheus_only):
from litellm_enterprise.integrations.prometheus import prometheus_label_factory
from litellm.types.integrations.prometheus import UserAPIKeyLabelValues
monkeypatch.setattr(
@ -1130,6 +1135,119 @@ def test_get_custom_labels_from_tags_no_tags(monkeypatch):
}
def test_get_custom_labels_from_tags_wildcard_patterns(monkeypatch):
"""Test wildcard pattern matching for custom labels from tags."""
from litellm_enterprise.integrations.prometheus import get_custom_labels_from_tags
# Configure tags with wildcard patterns
monkeypatch.setattr(
"litellm.custom_prometheus_tags",
["User-Agent: curl/*", "User-Agent: python-requests/*", "Environment: prod*", "Service: api-gateway*", "exact-match"]
)
# Test tags that should match the wildcard patterns
tags = [
"User-Agent: curl/7.68.0",
"User-Agent: python-requests/2.28.1",
"Environment: production",
"Service: api-gateway-v2",
"exact-match",
"other-tag"
]
result = get_custom_labels_from_tags(tags)
expected = {
"tag_User_Agent__curl__": "true", # matches "User-Agent: curl/*"
"tag_User_Agent__python_requests__": "true", # matches "User-Agent: python-requests/*"
"tag_Environment__prod_": "true", # matches "Environment: prod*"
"tag_Service__api_gateway_": "true", # matches "Service: api-gateway*"
"tag_exact_match": "true", # exact match
}
assert result == expected
def test_get_custom_labels_from_tags_wildcard_no_matches(monkeypatch):
"""Test wildcard patterns that don't match any tags."""
from litellm_enterprise.integrations.prometheus import get_custom_labels_from_tags
# Configure tags with wildcard patterns
monkeypatch.setattr(
"litellm.custom_prometheus_tags",
["User-Agent: firefox/*", "Environment: dev*", "Service: web-app*"]
)
# Test tags that should NOT match the wildcard patterns
tags = [
"User-Agent: curl/7.68.0", # doesn't match "User-Agent: firefox/*"
"Environment: production", # doesn't match "Environment: dev*"
"Service: api-gateway-v2", # doesn't match "Service: web-app*"
"other-tag"
]
result = get_custom_labels_from_tags(tags)
expected = {
"tag_User_Agent__firefox__": "false", # no match for "User-Agent: firefox/*"
"tag_Environment__dev_": "false", # no match for "Environment: dev*"
"tag_Service__web_app_": "false", # no match for "Service: web-app*"
}
assert result == expected
def test_tag_matches_wildcard_configured_pattern():
"""Test the helper function for wildcard pattern matching."""
from litellm_enterprise.integrations.prometheus import (
_tag_matches_wildcard_configured_pattern,
)
# Test cases that should match
assert _tag_matches_wildcard_configured_pattern(
tags=["User-Agent: curl/7.68.0", "prod", "other"],
configured_tag="User-Agent: curl/*"
) is True
assert _tag_matches_wildcard_configured_pattern(
tags=["User-Agent: python-requests/2.28.1", "test"],
configured_tag="User-Agent: python-requests/*"
) is True
assert _tag_matches_wildcard_configured_pattern(
tags=["Environment: production", "debug"],
configured_tag="Environment: prod*"
) is True
# Test exact match (no wildcard)
assert _tag_matches_wildcard_configured_pattern(
tags=["prod", "test"],
configured_tag="prod"
) is True
# Test cases that should NOT match
assert _tag_matches_wildcard_configured_pattern(
tags=["User-Agent: firefox/98.0", "prod"],
configured_tag="User-Agent: curl/*"
) is False
assert _tag_matches_wildcard_configured_pattern(
tags=["Environment: development", "test"],
configured_tag="Environment: prod*"
) is False
assert _tag_matches_wildcard_configured_pattern(
tags=["staging", "test"],
configured_tag="prod"
) is False
# Test with empty tags
assert _tag_matches_wildcard_configured_pattern(
tags=[],
configured_tag="User-Agent: curl/*"
) is False
@pytest.mark.asyncio(scope="session")
async def test_initialize_remaining_budget_metrics(prometheus_logger):
"""
@ -1532,7 +1650,11 @@ def test_prometheus_label_factory_with_custom_tags(monkeypatch):
"""
Test that prometheus_label_factory correctly handles custom tags
"""
from litellm_enterprise.integrations.prometheus import get_custom_labels_from_tags, prometheus_label_factory
from litellm_enterprise.integrations.prometheus import (
get_custom_labels_from_tags,
prometheus_label_factory,
)
from litellm.types.integrations.prometheus import UserAPIKeyLabelValues
# Set custom tags configuration
@ -1567,7 +1689,11 @@ def test_prometheus_label_factory_with_no_custom_tags(monkeypatch):
"""
Test that prometheus_label_factory works when no custom tags are configured
"""
from litellm_enterprise.integrations.prometheus import get_custom_labels_from_tags, prometheus_label_factory
from litellm_enterprise.integrations.prometheus import (
get_custom_labels_from_tags,
prometheus_label_factory,
)
from litellm.types.integrations.prometheus import UserAPIKeyLabelValues
# Set empty custom tags configuration
@ -1776,3 +1902,154 @@ def test_set_llm_deployment_success_metrics_with_label_filtering():
)
prometheus_logger.litellm_deployment_success_responses.labels().inc.assert_called_once()
prometheus_logger.litellm_deployment_total_requests.labels().inc.assert_called_once()
@pytest.mark.asyncio
async def test_prometheus_token_metrics_with_prometheus_config():
"""
Test that validates the renamed token metrics are incremented correctly with a prometheus config.
This test ensures that after the metric renaming (git diff):
- litellm_total_tokens -> litellm_total_tokens_metric
- litellm_input_tokens -> litellm_input_tokens_metric
- litellm_output_tokens -> litellm_output_tokens_metric
All three metrics should be properly incremented when making a successful completion request.
"""
from prometheus_client import CollectorRegistry, Counter
import litellm
from litellm.types.integrations.prometheus import PrometheusMetricsConfig
# Clear registry before test
collectors = list(REGISTRY._collector_to_names.keys())
for collector in collectors:
REGISTRY.unregister(collector)
# Set up prometheus configuration that includes the token metrics
config = [
PrometheusMetricsConfig(
group="token_metrics_test",
metrics=[
"litellm_total_tokens_metric",
"litellm_input_tokens_metric",
"litellm_output_tokens_metric",
"litellm_requests_metric"
],
include_labels=[
"model",
"hashed_api_key",
"api_key_alias",
"team",
"team_alias"
],
)
]
# Mock litellm.prometheus_metrics_config
with patch("litellm.prometheus_metrics_config", config):
# Create PrometheusLogger with the configuration
prometheus_logger = PrometheusLogger()
# Test data with specific token counts
standard_logging_payload = create_standard_logging_payload()
standard_logging_payload["total_tokens"] = 1500
standard_logging_payload["prompt_tokens"] = 900
standard_logging_payload["completion_tokens"] = 600
standard_logging_payload["response_cost"] = 0.075
kwargs = {
"model": "gpt-3.5-turbo",
"stream": False,
"litellm_params": {
"metadata": {
"user_api_key": "test_key_hash",
"user_api_key_user_id": "test_user",
"user_api_key_team_id": "test_team",
"user_api_key_alias": "test_alias",
"user_api_key_team_alias": "test_team_alias",
}
},
"start_time": datetime.now() - timedelta(seconds=2),
"completion_start_time": datetime.now() - timedelta(seconds=1),
"api_call_start_time": datetime.now() - timedelta(seconds=1.5),
"end_time": datetime.now(),
"standard_logging_object": standard_logging_payload,
}
response_obj = MagicMock()
# Make the completion call through the logger
await prometheus_logger.async_log_success_event(
kwargs, response_obj, kwargs["start_time"], kwargs["end_time"]
)
await asyncio.sleep(2)
print("final registry values", REGISTRY._collector_to_names)
# Get metric collectors directly from registry
metric_collectors = {}
for collector, names in REGISTRY._collector_to_names.items():
metric_name = names[0] # First name is the base metric name
metric_collectors[metric_name] = collector
print("=== Final Metric Values (Direct Access) ===")
# Expected values
expected_values = {
"litellm_total_tokens_metric": 1500.0,
"litellm_input_tokens_metric": 900.0,
"litellm_output_tokens_metric": 600.0,
"litellm_requests_metric": 1.0
}
expected_label_values = {
'api_key_alias': 'test_alias',
'hashed_api_key': 'test_hash',
'model': 'gpt-3.5-turbo',
'team': 'test_team',
'team_alias': 'test_team_alias'
}
# Validate each metric directly
for metric_name, expected_value in expected_values.items():
if metric_name in metric_collectors:
collector = metric_collectors[metric_name]
# Get all samples for this metric
samples = list(collector.collect())[0].samples
# Find the _total sample (the actual counter value)
total_sample = None
for sample in samples:
if sample.name.endswith('_total'):
total_sample = sample
break
if total_sample:
actual_value = total_sample.value
actual_labels = total_sample.labels
print(f"✓ {metric_name}: expected={expected_value}, actual={actual_value}")
print(f" Labels: {actual_labels}")
# Validate the value
assert actual_value == expected_value, f"Expected {expected_value}, got {actual_value} for {metric_name}"
# Validate the labels
for label_key, expected_label_value in expected_label_values.items():
actual_label_value = actual_labels.get(label_key)
assert actual_label_value == expected_label_value, f"Expected label {label_key}={expected_label_value}, got {actual_label_value}"
print(f" ✓ {metric_name} VALIDATED")
else:
raise AssertionError(f"No _total sample found for {metric_name}")
else:
raise AssertionError(f"Metric {metric_name} not found in registry")
print("✓ All token metrics validated successfully!")
# check final value of metrics in registry

View file

@ -25,6 +25,9 @@ from litellm.types.llms.openai import (
ResponseAPIUsage,
IncompleteDetails,
)
from openai.types.responses.response_create_params import (
ResponseInputParam,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
@ -184,8 +187,8 @@ class BaseResponsesAPITest(ABC):
# basic test assert the usage seems reasonable
print("response_completed_event.response.usage=", response_completed_event.response.usage)
assert response_completed_event.response.usage.input_tokens > 0 and response_completed_event.response.usage.input_tokens < 100
assert response_completed_event.response.usage.output_tokens > 0 and response_completed_event.response.usage.output_tokens < 1000
assert response_completed_event.response.usage.total_tokens > 0 and response_completed_event.response.usage.total_tokens < 1000
assert response_completed_event.response.usage.output_tokens > 0 and response_completed_event.response.usage.output_tokens < 2000
assert response_completed_event.response.usage.total_tokens > 0 and response_completed_event.response.usage.total_tokens < 2000
# total tokens should be the sum of input and output tokens
assert response_completed_event.response.usage.total_tokens == response_completed_event.response.usage.input_tokens + response_completed_event.response.usage.output_tokens
@ -229,6 +232,7 @@ class BaseResponsesAPITest(ABC):
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.flaky(retries=3, delay=2)
@pytest.mark.asyncio
async def test_basic_openai_responses_streaming_delete_endpoint(self, sync_mode):
#litellm._turn_on_debug()
@ -278,6 +282,7 @@ class BaseResponsesAPITest(ABC):
)
@pytest.mark.parametrize("sync_mode", [False, True])
@pytest.mark.flaky(retries=3, delay=2)
@pytest.mark.asyncio
async def test_basic_openai_responses_get_endpoint(self, sync_mode):
litellm._turn_on_debug()
@ -318,6 +323,7 @@ class BaseResponsesAPITest(ABC):
raise ValueError("response is not a ResponsesAPIResponse")
@pytest.mark.asyncio
@pytest.mark.flaky(retries=3, delay=2)
async def test_basic_openai_list_input_items_endpoint(self):
"""Test that calls the OpenAI List Input Items endpoint"""
litellm._turn_on_debug()
@ -364,3 +370,73 @@ class BaseResponsesAPITest(ABC):
# assert the response is not None
assert response_1 is not None
assert response_2 is not None
@pytest.mark.asyncio
async def test_responses_api_with_tool_calls(self):
"""Test that calls the Responses API with tool calls including function call and output"""
litellm._turn_on_debug()
litellm.set_verbose = True
base_completion_call_args = self.get_base_completion_call_args()
# Define the input with message, function call, and function call output
input_data: ResponseInputParam = [
{
"type": "message",
"role": "user",
"content": "How is the weather in São Paulo today ?"
},
{
"type": "function_call",
"arguments": "{\"location\": \"São Paulo, Brazil\"}",
"call_id": "fc_1fe70e2a-a596-45ef-b72c-9b8567c460e5",
"name": "get_weather",
"id": "fc_1fe70e2a-a596-45ef-b72c-9b8567c460e5",
"status": "completed"
},
{
"type": "function_call_output",
"call_id": "fc_1fe70e2a-a596-45ef-b72c-9b8567c460e5",
"output": "Rainy"
}
]
# Define the tools
tools = [
{
"type": "function",
"name": "get_weather",
"description": "Get current temperature for a given location.",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "City and country e.g. Bogotá, Colombia"
}
},
"required": ["location"],
"additionalProperties": False
}
}
]
try:
# Make the responses API call
response = await litellm.aresponses(
input=input_data,
store=False,
tools=tools,
**base_completion_call_args
)
except litellm.InternalServerError:
pytest.skip("Skipping test due to litellm.InternalServerError")
print("litellm response=", json.dumps(response, indent=4, default=str))
# Validate the response structure
validate_responses_api_response(response, final_chunk=True)
# Additional assertions specific to tool calls
assert response is not None
assert "output" in response
assert len(response["output"]) > 0

View file

@ -5,7 +5,7 @@ from unittest.mock import patch, AsyncMock
sys.path.insert(0, os.path.abspath("../.."))
import litellm
import json
from base_responses_api import BaseResponsesAPITest
@pytest.mark.asyncio
async def test_basic_google_ai_studio_responses_api_with_tools():
litellm._turn_on_debug()
@ -85,10 +85,22 @@ async def test_mock_basic_google_ai_studio_responses_api_with_tools():
assert call_kwargs["messages"][0]["content"] == "what is the latest version of supabase python package and when was it released?"
assert call_kwargs["tools"] == [] # web search tools are converted to web_search_options, not kept as tools
class TestGoogleAIStudioResponsesAPITest(BaseResponsesAPITest):
def get_base_completion_call_args(self):
#litellm._turn_on_debug()
return {
"model": "gemini/gemini-2.5-flash-lite"
}
async def test_basic_openai_responses_delete_endpoint(self, sync_mode=False):
pass
async def test_basic_openai_responses_streaming_delete_endpoint(self, sync_mode=False):
pass
async def test_basic_openai_responses_get_endpoint(self, sync_mode=False):
pass

View file

@ -1308,3 +1308,90 @@ async def test_store_field_transformation():
assert response.created_at == 1751443898, "created_at should maintain the same value after conversion"
@pytest.mark.asyncio
async def test_aresponses_service_tier_and_safety_identifier():
"""
Test that service_tier and safety_identifier parameters are correctly sent in the request body
when using litellm.aresponses.
"""
mock_response = {
"id": "resp_01234567890abcdef",
"object": "response",
"created_at": 1753060947,
"status": "completed",
"error": None,
"incomplete_details": None,
"instructions": None,
"max_output_tokens": None,
"model": "gpt-4o-2024-05-13",
"output": [
{
"type": "text",
"id": "out_01234567890abcdef",
"text": "This is a test response with service tier and safety identifier.",
}
],
"parallel_tool_calls": True,
"previous_response_id": None,
"reasoning": None,
"store": True,
"temperature": 1.0,
"text": {"format": {"type": "text"}},
"tool_choice": "auto",
"tools": [],
"top_p": 1.0,
"truncation": "disabled",
"usage": {
"input_tokens": 15,
"input_tokens_details": {"cached_tokens": 0},
"output_tokens": 25,
"output_tokens_details": {"reasoning_tokens": 0},
"total_tokens": 40,
},
"user": None,
"metadata": {},
}
class MockResponse:
def __init__(self, json_data, status_code):
self._json_data = json_data
self.status_code = status_code
self.text = json.dumps(json_data)
def json(self):
return self._json_data
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
# Configure the mock to return our response
mock_post.return_value = MockResponse(mock_response, 200)
litellm._turn_on_debug()
litellm.set_verbose = True
# Call aresponses with service_tier and safety_identifier
response = await litellm.aresponses(
model="openai/gpt-4o",
input="Test with service tier and safety identifier",
service_tier="flex",
safety_identifier="123",
)
# Verify the request was made correctly
mock_post.assert_called_once()
request_body = mock_post.call_args.kwargs["json"]
print("request_body=", json.dumps(request_body, indent=4, default=str))
# Validate that both parameters are present in the request body
assert request_body["service_tier"] == "flex", "service_tier should be 'flex' in request body"
assert request_body["safety_identifier"] == "123", "safety_identifier should be '123' in request body"
assert request_body["model"] == "gpt-4o"
assert request_body["input"] == "Test with service tier and safety identifier"
# Validate the response
print("Response:", json.dumps(response, indent=4, default=str))

View file

@ -1200,96 +1200,99 @@ class BaseLLMChatTest(ABC):
from litellm.utils import supports_function_calling
from litellm import completion
litellm._turn_on_debug()
try:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
base_completion_call_args = self.get_base_completion_call_args()
if not supports_function_calling(base_completion_call_args["model"], None):
print("Model does not support function calling")
pytest.skip("Model does not support function calling")
def get_weather(city: str):
return f"City: {city}, Weather: Sunny with 34 degree Celcius"
base_completion_call_args = self.get_base_completion_call_args()
if not supports_function_calling(base_completion_call_args["model"], None):
print("Model does not support function calling")
pytest.skip("Model does not support function calling")
def get_weather(city: str):
return f"City: {city}, Weather: Sunny with 34 degree Celcius"
TOOLS = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the weather in a city",
"parameters": {
"$id": "https://some/internal/name",
"$schema": "https://json-schema.org/draft-07/schema",
"type": "object",
"properties": {
"city": {
"type": "string",
"description": "The city to get the weather for",
}
TOOLS = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the weather in a city",
"parameters": {
"$id": "https://some/internal/name",
"$schema": "https://json-schema.org/draft-07/schema",
"type": "object",
"properties": {
"city": {
"type": "string",
"description": "The city to get the weather for",
}
},
"required": ["city"],
"additionalProperties": False,
},
"required": ["city"],
"additionalProperties": False,
"strict": True,
},
"strict": True,
},
}
]
}
]
messages = [{ "content": "How is the weather in Mumbai?","role": "user"}]
response, iteration = "", 0
while True:
if response:
break
# Create a streaming response with tool calling enabled
stream = completion(
**base_completion_call_args,
messages=messages,
tools=TOOLS,
stream=True,
)
messages = [{ "content": "How is the weather in Mumbai?","role": "user"}]
response, iteration = "", 0
while True:
if response:
break
# Create a streaming response with tool calling enabled
stream = completion(
**base_completion_call_args,
messages=messages,
tools=TOOLS,
stream=True,
)
final_tool_calls = {}
for chunk in stream:
delta = chunk.choices[0].delta
print(delta)
if delta.content:
response += delta.content
elif delta.tool_calls:
for tool_call in chunk.choices[0].delta.tool_calls or []:
index = tool_call.index
if index not in final_tool_calls:
final_tool_calls[index] = tool_call
else:
final_tool_calls[
index
].function.arguments += tool_call.function.arguments
if final_tool_calls:
for tool_call in final_tool_calls.values():
if tool_call.function.name == "get_weather":
city = json.loads(tool_call.function.arguments)["city"]
tool_response = get_weather(city)
messages.append(
{
"role": "assistant",
"tool_calls": [tool_call],
"content": None,
}
)
messages.append(
{
"role": "tool",
"tool_call_id": tool_call.id,
"content": tool_response,
}
)
iteration += 1
if iteration > 2:
print("Something went wrong!")
break
final_tool_calls = {}
for chunk in stream:
delta = chunk.choices[0].delta
print(delta)
if delta.content:
response += delta.content
elif delta.tool_calls:
for tool_call in chunk.choices[0].delta.tool_calls or []:
index = tool_call.index
if index not in final_tool_calls:
final_tool_calls[index] = tool_call
else:
final_tool_calls[
index
].function.arguments += tool_call.function.arguments
if final_tool_calls:
for tool_call in final_tool_calls.values():
if tool_call.function.name == "get_weather":
city = json.loads(tool_call.function.arguments)["city"]
tool_response = get_weather(city)
messages.append(
{
"role": "assistant",
"tool_calls": [tool_call],
"content": None,
}
)
messages.append(
{
"role": "tool",
"tool_call_id": tool_call.id,
"content": tool_response,
}
)
iteration += 1
if iteration > 2:
print("Something went wrong!")
break
print(response)
print(response)
except litellm.ServiceUnavailableError:
pass
def test_reasoning_effort(self):
"""Test that reasoning_effort is passed correctly to the model"""

View file

@ -603,3 +603,54 @@ def test_openai_deepresearch_model_bridge():
)
print("response: ", response)
def test_openai_tool_calling():
from pydantic import BaseModel
from typing import Any, Literal
class OpenAIFunction(BaseModel):
description: Optional[str] = None
name: str
parameters: Optional[dict[str, Any]] = None
class OpenAITool(BaseModel):
type: Literal["function"]
function: OpenAIFunction
completion_params = {
"model": "openai/gpt-4.1",
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": "What is TSLA stock price at today?"}
],
}
],
"stream": False,
"temperature": 0.5,
"stop": None,
"max_tokens": 1600,
"tools": [
OpenAITool(
type="function",
function=OpenAIFunction(
description="Get the current stock price for a given ticker symbol.",
name="get_stock_price",
parameters={
"type": "object",
"properties": {
"ticker": {
"type": "string",
"description": "The stock ticker symbol, e.g. AAPL for Apple Inc.",
}
},
"required": ["ticker"],
},
),
)
],
}
response = litellm.completion(**completion_params)

View file

@ -58,10 +58,10 @@ VERTEX_MODELS_TO_NOT_TEST = [
"gemini-1.5-pro-preview-0215",
"gemini-pro-experimental",
"gemini-flash-experimental",
"gemini-1.5-flash-exp-0827",
"gemini-2.5-flash-lite-exp-0827",
"gemini-2.0-pro-exp-02-05",
"gemini-pro-flash",
"gemini-1.5-flash-exp-0827",
"gemini-2.5-flash-lite-exp-0827",
"gemini-2.0-flash-exp",
"gemini-2.0-flash-thinking-exp",
"gemini-2.0-flash-thinking-exp-01-21",
@ -149,7 +149,7 @@ async def test_get_response():
prompt = '\ndef count_nums(arr):\n """\n Write a function count_nums which takes an array of integers and returns\n the number of elements which has a sum of digits > 0.\n If a number is negative, then its first signed digit will be negative:\n e.g. -123 has signed digits -1, 2, and 3.\n >>> count_nums([]) == 0\n >>> count_nums([-1, 11, -11]) == 1\n >>> count_nums([1, 1, 2]) == 3\n """\n'
try:
response = await acompletion(
model="gemini-1.5-flash",
model="gemini-2.5-flash-lite",
messages=[
{
"role": "system",
@ -518,7 +518,7 @@ async def test_gemini_pro_vision(provider, sync_mode):
litellm.num_retries = 3
if sync_mode:
resp = litellm.completion(
model="{}/gemini-1.5-flash-preview-0514".format(provider),
model="{}/gemini-2.5-flash-lite".format(provider),
messages=[
{"role": "system", "content": "Be a good bot"},
{
@ -537,7 +537,7 @@ async def test_gemini_pro_vision(provider, sync_mode):
)
else:
resp = await litellm.acompletion(
model="{}/gemini-1.5-flash-preview-0514".format(provider),
model="{}/gemini-2.5-flash-lite".format(provider),
messages=[
{"role": "system", "content": "Be a good bot"},
{
@ -605,7 +605,7 @@ def test_completion_function_plus_pdf(load_pdf):
image_message = {"role": "user", "content": image_content}
response = completion(
model="vertex_ai_beta/gemini-1.5-flash-preview-0514",
model="vertex_ai_beta/gemini-2.5-flash-lite",
messages=[image_message],
stream=False,
)
@ -1194,7 +1194,7 @@ Using this JSON schema:
with patch.object(client, "post", side_effect=_side_effect) as mock_call:
response = completion(
model="vertex_ai_beta/gemini-1.5-flash",
model="vertex_ai_beta/gemini-2.5-flash-lite",
messages=messages,
response_format={"type": "json_object"},
client=client,
@ -1383,7 +1383,7 @@ def vertex_httpx_mock_post_invalid_schema_response_anthropic(*args, **kwargs):
[
("vertex_ai_beta/gemini-1.5-pro-001", "us-central1", True),
("gemini/gemini-1.5-pro", None, True),
("vertex_ai_beta/gemini-1.5-flash", "us-central1", True),
("vertex_ai_beta/gemini-2.5-flash-lite", "us-central1", True),
("vertex_ai/claude-3-5-sonnet@20240620", "us-east5", False),
],
)
@ -1572,7 +1572,7 @@ async def test_anthropic_message_via_anthropic_messages():
[
("vertex_ai_beta/gemini-1.5-pro-001", "us-central1", True),
("gemini/gemini-1.5-pro", None, True),
("vertex_ai_beta/gemini-1.5-flash", "us-central1", True),
("vertex_ai_beta/gemini-2.5-flash-lite", "us-central1", True),
("vertex_ai/claude-3-5-sonnet@20240620", "us-east5", False),
],
)
@ -1680,7 +1680,7 @@ async def test_gemini_pro_json_schema_args_sent_httpx_openai_schema(
@pytest.mark.parametrize(
"model", ["gemini-1.5-flash", "claude-3-5-sonnet@20240620"]
"model", ["gemini-2.5-flash-lite", "claude-3-5-sonnet@20240620"]
) # "vertex_ai",
@pytest.mark.asyncio
async def test_gemini_pro_httpx_custom_api_base(model):
@ -1820,7 +1820,7 @@ async def test_gemini_pro_function_calling_streaming(sync_mode):
load_vertex_ai_credentials()
litellm.set_verbose = True
data = {
"model": "vertex_ai/gemini-1.5-flash",
"model": "vertex_ai/gemini-2.5-flash-lite",
"messages": [
{
"role": "user",
@ -2541,7 +2541,7 @@ def mock_gemini_request(*args, **kwargs):
if "cachedContents" in kwargs["url"]:
mock_response.json.return_value = {
"name": "cachedContents/4d2kd477o3pg",
"model": "models/gemini-1.5-flash-001",
"model": "models/gemini-2.5-flash-lite-001",
"createTime": "2024-08-26T22:31:16.147190Z",
"updateTime": "2024-08-26T22:31:16.147190Z",
"expireTime": "2024-08-26T22:36:15.548934784Z",
@ -2671,7 +2671,7 @@ async def test_gemini_context_caching_anthropic_format(sync_mode):
try:
if sync_mode:
response = litellm.completion(
model="gemini/gemini-1.5-flash-001",
model="gemini/gemini-2.5-flash-lite-001",
messages=gemini_context_caching_messages,
temperature=0.2,
max_tokens=10,
@ -2679,7 +2679,7 @@ async def test_gemini_context_caching_anthropic_format(sync_mode):
)
else:
response = await litellm.acompletion(
model="gemini/gemini-1.5-flash-001",
model="gemini/gemini-2.5-flash-lite-001",
messages=gemini_context_caching_messages,
temperature=0.2,
max_tokens=10,

View file

@ -72,7 +72,7 @@ def test_batch_completions_models():
def test_batch_completion_models_all_responses():
try:
responses = batch_completion_models_all_responses(
models=["gemini/gemini-1.5-flash", "claude-3-haiku-20240307"],
models=["gemini/gemini-2.5-flash-lite", "claude-3-haiku-20240307"],
messages=[{"role": "user", "content": "write a poem"}],
max_tokens=10,
)

View file

@ -2155,8 +2155,8 @@ async def test_caching_kwargs_input(sync_mode):
Message,
ModelResponse,
Usage,
CompletionTokensDetails,
PromptTokensDetails,
CompletionTokensDetailsWrapper,
PromptTokensDetailsWrapper,
)
from datetime import datetime
@ -2187,10 +2187,10 @@ async def test_caching_kwargs_input(sync_mode):
completion_tokens=31,
prompt_tokens=16,
total_tokens=47,
completion_tokens_details=CompletionTokensDetails(
completion_tokens_details=CompletionTokensDetailsWrapper(
audio_tokens=None, reasoning_tokens=0
),
prompt_tokens_details=PromptTokensDetails(
prompt_tokens_details=PromptTokensDetailsWrapper(
audio_tokens=None, cached_tokens=0
),
),

View file

@ -3696,7 +3696,7 @@ def test_completion_volcengine():
[
# "gemini-1.0-pro",
"gemini-1.5-pro",
# "gemini-1.5-flash",
# "gemini-2.5-flash-lite",
],
)
@pytest.mark.flaky(retries=3, delay=1)
@ -3750,7 +3750,7 @@ def test_completion_gemini(model):
@pytest.mark.asyncio
async def test_acompletion_gemini():
litellm.set_verbose = True
model_name = "gemini/gemini-1.5-flash"
model_name = "gemini/gemini-2.5-flash-lite"
messages = [{"role": "user", "content": "Hey, how's it going?"}]
try:
response = await litellm.acompletion(model=model_name, messages=messages)

View file

@ -1187,3 +1187,93 @@ async def test_embedding_with_extra_headers(sync_mode):
mock_post.assert_called_once()
assert "my-test-param" in mock_post.call_args.kwargs["headers"]
@pytest.mark.parametrize(
"input_data, expected_payload_input",
[
# Case 1: Input with only text strings
(
["hello world", "foo bar"],
["hello world", "foo bar"],
),
# Case 2: Input with a mix of text and a base64 encoded image
(
[
"A picture of a cat",
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNkYAAAAAYAAjCB0C8AAAAASUVORK5CYII=",
],
[
{"text": "A picture of a cat"},
{
"image": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNkYAAAAAYAAjCB0C8AAAAASUVORK5CYII="
},
],
),
# Case 3: Input with only a base64 encoded image
(
[
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNkYAAAAAYAAjCB0C8AAAAASUVORK5CYII="
],
[
{
"image": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNkYAAAAAYAAjCB0C8AAAAASUVORK5CYII="
}
],
),
],
)
def test_jina_ai_img_embeddings(input_data, expected_payload_input):
"""
Tests the input transformation logic for Jina AI embeddings using mocks.
This test verifies that when litellm.embedding is called with a jina_ai model,
the 'input' field in the request payload is formatted correctly based on whether
the input contains text or base64 encoded images.
"""
# We patch the `post` method of the HTTPHandler. This intercepts the network
# request before it's actually sent.
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post:
# Configure the mock to return a successful, minimal valid response.
# This prevents litellm from raising an error when processing the response.
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {
"object": "list",
"data": [
{
"object": "embedding",
"index": 0,
"embedding": [0.1] * 768, # Dummy embedding vector
}
],
"model": "jina-embeddings-v4",
}
mock_post.return_value = mock_response
# Call the function we want to test
try:
litellm.embedding(
model="jina_ai/jina-embeddings-v4", input=input_data
)
except Exception as e:
pytest.fail(
f"litellm.embedding call failed with an unexpected exception: {e}"
)
# --- Assertions ---
# 1. Check that our mock `post` method was called exactly once.
mock_post.assert_called_once()
# 2. Extract the keyword arguments passed to the mock call.
# The request payload is in the 'data' keyword argument.
kwargs = mock_post.call_args.kwargs
assert "data" in kwargs
# 3. Parse the JSON payload string into a Python dictionary.
sent_data = json.loads(kwargs["data"])
# 4. This is the core of our test:
# Assert that the 'input' field in the payload matches our expectation.
assert "input" in sent_data
assert sent_data["input"] == expected_payload_input

View file

@ -0,0 +1,6 @@
from cache_unit_tests import LLMCachingUnitTests
from litellm.caching import LiteLLMCacheType
class TestGCSCacheUnitTests(LLMCachingUnitTests):
def get_cache_type(self) -> LiteLLMCacheType:
return LiteLLMCacheType.GCS

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