mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge remote-tracking branch 'origin/main' into fix_vertex_expired_tokens
This commit is contained in:
commit
0c85fe4b70
154 changed files with 10909 additions and 1219 deletions
|
|
@ -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
|
||||
|
|
|
|||
53
cookbook/misc/test_responses_api.py
Normal file
53
cookbook/misc/test_responses_api.py
Normal 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)
|
||||
|
||||
|
||||
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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`)
|
||||
|
|
|
|||
75
docs/my-website/docs/providers/oci.md
Normal file
75
docs/my-website/docs/providers/oci.md
Normal 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
|
||||
```
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
||||
|
|
|
|||
|
|
@ -469,7 +469,8 @@ const sidebars = {
|
|||
"providers/featherless_ai",
|
||||
"providers/nebius",
|
||||
"providers/dashscope",
|
||||
"providers/bytez"
|
||||
"providers/bytez",
|
||||
"providers/oci",
|
||||
],
|
||||
},
|
||||
{
|
||||
|
|
|
|||
BIN
enterprise/dist/litellm_enterprise-0.1.17-py3-none-any.whl
vendored
Normal file
BIN
enterprise/dist/litellm_enterprise-0.1.17-py3-none-any.whl
vendored
Normal file
Binary file not shown.
BIN
enterprise/dist/litellm_enterprise-0.1.17.tar.gz
vendored
Normal file
BIN
enterprise/dist/litellm_enterprise-0.1.17.tar.gz
vendored
Normal file
Binary file not shown.
BIN
enterprise/dist/litellm_enterprise-0.1.19-py3-none-any.whl
vendored
Normal file
BIN
enterprise/dist/litellm_enterprise-0.1.19-py3-none-any.whl
vendored
Normal file
Binary file not shown.
BIN
enterprise/dist/litellm_enterprise-0.1.19.tar.gz
vendored
Normal file
BIN
enterprise/dist/litellm_enterprise-0.1.19.tar.gz
vendored
Normal file
Binary file not shown.
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
97
litellm/caching/gcs_cache.py
Normal file
97
litellm/caching/gcs_cache.py
Normal 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)
|
||||
|
|
@ -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 ###########################
|
||||
########################################################################################
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
6
litellm/llms/jina_ai/common_utils.py
Normal file
6
litellm/llms/jina_ai/common_utils.py
Normal 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)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
868
litellm/llms/oci/chat/transformation.py
Normal file
868
litellm/llms/oci/chat/transformation.py
Normal 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,
|
||||
)
|
||||
]
|
||||
)
|
||||
19
litellm/llms/oci/common_utils.py
Normal file
19
litellm/llms/oci/common_utils.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)}"}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
73
litellm/proxy/spend_tracking/cold_storage_handler.py
Normal file
73
litellm/proxy/spend_tracking/cold_storage_handler.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
183
litellm/types/llms/oci.py
Normal 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
|
||||
|
|
@ -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]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
512
poetry.lock
generated
File diff suppressed because it is too large
Load diff
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
),
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
6
tests/local_testing/test_gcs_cache_unit_tests.py
Normal file
6
tests/local_testing/test_gcs_cache_unit_tests.py
Normal 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
Loading…
Add table
Reference in a new issue