Merge branch 'BerriAI:main' into wandb-inference

This commit is contained in:
Anubhav Singh 2025-09-16 16:46:58 +05:30 • committed by GitHub
commit a9667e5930
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
52 changed files with 2762 additions and 248 deletions

View file

@ -316,6 +316,7 @@ curl 'http://0.0.0.0:4000/key/generate' \
| [google AI Studio - gemini](https://docs.litellm.ai/docs/providers/gemini) | ✅ | ✅ | ✅ | ✅ | | |
| [mistral ai api](https://docs.litellm.ai/docs/providers/mistral) | ✅ | ✅ | ✅ | ✅ | ✅ | |
| [cloudflare AI Workers](https://docs.litellm.ai/docs/providers/cloudflare_workers) | ✅ | ✅ | ✅ | ✅ | | |
| [CompactifAI](https://docs.litellm.ai/docs/providers/compactifai) | ✅ | ✅ | ✅ | ✅ | | |
| [cohere](https://docs.litellm.ai/docs/providers/cohere) | ✅ | ✅ | ✅ | ✅ | ✅ | |
| [anthropic](https://docs.litellm.ai/docs/providers/anthropic) | ✅ | ✅ | ✅ | ✅ | | |
| [empower](https://docs.litellm.ai/docs/providers/empower) | ✅ | ✅ | ✅ | ✅ |

View file

@ -0,0 +1,223 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# CompactifAI
https://docs.compactif.ai/
CompactifAI offers highly compressed versions of leading language models, delivering up to **70% lower inference costs**, **4x throughput gains**, and **low-latency inference** with minimal quality loss (<5%). CompactifAI's OpenAI-compatible API makes integration straightforward, enabling developers to build ultra-efficient, scalable AI applications with superior concurrency and resource efficiency.
| Property | Details |
|-------|-------|
| Description | CompactifAI offers compressed versions of leading language models with up to 70% cost reduction and 4x throughput gains |
| Provider Route on LiteLLM | `compactifai/` (add this prefix to the model name - e.g. `compactifai/cai-llama-3-1-8b-slim`) |
| Provider Doc | [CompactifAI ↗](https://docs.compactif.ai/) |
| API Endpoint for Provider | https://api.compactif.ai/v1 |
| Supported Endpoints | `/chat/completions`, `/completions` |
## Supported OpenAI Parameters
CompactifAI is fully OpenAI-compatible and supports the following parameters:
```
"stream",
"stop",
"temperature",
"top_p",
"max_tokens",
"presence_penalty",
"frequency_penalty",
"logit_bias",
"user",
"response_format",
"seed",
"tools",
"tool_choice",
"parallel_tool_calls",
"extra_headers"
```
## API Key Setup
CompactifAI API keys are available through AWS Marketplace subscription:
1. Subscribe via [AWS Marketplace](https://aws.amazon.com/marketplace)
2. Complete subscription verification (24-hour review process)
3. Access MultiverseIAM dashboard with provided credentials
4. Retrieve your API key from the dashboard
```python
import os
os.environ["COMPACTIFAI_API_KEY"] = "your-api-key"
```
## Usage
<Tabs>
<TabItem value="sdk" label="SDK">
```python
from litellm import completion
import os
os.environ['COMPACTIFAI_API_KEY'] = "your-api-key"
response = completion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[
{"role": "user", "content": "Hello from LiteLLM!"}
],
)
print(response)
```
</TabItem>
<TabItem value="proxy" label="Proxy">
```yaml
model_list:
- model_name: llama-2-compressed
litellm_params:
model: compactifai/cai-llama-3-1-8b-slim
api_key: os.environ/COMPACTIFAI_API_KEY
```
</TabItem>
</Tabs>
## Streaming
```python
from litellm import completion
import os
os.environ['COMPACTIFAI_API_KEY'] = "your-api-key"
response = completion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[
{"role": "user", "content": "Write a short story"}
],
stream=True
)
for chunk in response:
print(chunk)
```
## Advanced Usage
### Custom Parameters
```python
from litellm import completion
response = completion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[{"role": "user", "content": "Explain quantum computing"}],
temperature=0.7,
max_tokens=500,
top_p=0.9,
stop=["Human:", "AI:"]
)
```
### Function Calling
CompactifAI supports OpenAI-compatible function calling:
```python
from litellm import completion
functions = [
{
"name": "get_weather",
"description": "Get current weather information",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state"
}
},
"required": ["location"]
}
}
]
response = completion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[{"role": "user", "content": "What's the weather in San Francisco?"}],
tools=[{"type": "function", "function": f} for f in functions],
tool_choice="auto"
)
```
### Async Usage
```python
import asyncio
from litellm import acompletion
async def async_call():
response = await acompletion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[{"role": "user", "content": "Hello async world!"}]
)
return response
# Run async function
response = asyncio.run(async_call())
print(response)
```
## Available Models
CompactifAI offers compressed versions of popular models. Use the `/models` endpoint to get the latest list:
```python
import httpx
headers = {"Authorization": f"Bearer {your_api_key}"}
response = httpx.get("https://api.compactif.ai/v1/models", headers=headers)
models = response.json()
```
Common model formats:
- `compactifai/cai-llama-3-1-8b-slim`
- `compactifai/mistral-7b-compressed`
- `compactifai/codellama-7b-compressed`
## Benefits
- **Cost Efficient**: Up to 70% lower inference costs compared to standard models
- **High Performance**: 4x throughput gains with minimal quality loss (<5%)
- **Low Latency**: Optimized for fast response times
- **Drop-in Replacement**: Full OpenAI API compatibility
- **Scalable**: Superior concurrency and resource efficiency
## Error Handling
CompactifAI returns standard OpenAI-compatible error responses:
```python
from litellm import completion
from litellm.exceptions import AuthenticationError, RateLimitError
try:
response = completion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[{"role": "user", "content": "Hello"}]
)
except AuthenticationError:
print("Invalid API key")
except RateLimitError:
print("Rate limit exceeded")
```
## Support
- Documentation: https://docs.compactif.ai/
- LinkedIn: [MultiverseComputing](https://www.linkedin.com/company/multiversecomputing)
- Analysis: [Artificial Analysis Provider Comparison](https://artificialanalysis.ai/providers/compactifai)

View file

@ -10,8 +10,30 @@ import TabItem from '@theme/TabItem';
- You must set up a Postgres database (e.g. Supabase, Neon, etc.)
- To enable team member rate limits, set the environment variable `EXPERIMENTAL_MULTI_INSTANCE_RATE_LIMITING=true` **before starting the proxy server**. Without this, team member rate limits will not be enforced.
## Default Budget for Auto-Generated JWT Teams
When using JWT authentication with `team_id_upsert: true`, you can automatically assign a default budget to any newly created team.
This is configured in `default_team_settings` in your `config.yaml`.
**Example:**
```yaml
# in your config.yaml
litellm_jwtauth:
team_id_upsert: true
team_id_jwt_field: "team_id"
# ... other jwt settings
litellm_settings:
default_team_settings:
- team_id: "default-settings"
max_budget: 100.0
```
Track spend, set budgets for your Internal Team
## Setting Monthly Team Budgets
### 1. Create a team

View file

@ -1,5 +1,5 @@
---
title: "v1.77.2-stable - Bedrock Batches API"
title: "[Pre-Release] v1.77.2-stable - Bedrock Batches API"
slug: "v1-77-2"
date: 2025-09-13T10:00:00
authors:
@ -21,21 +21,22 @@ import TabItem from '@theme/TabItem';
## Deploy this version
:::info
This release is not yet live.
:::
<Tabs>
<TabItem value="docker" label="Docker">
``` showLineNumbers title="docker run litellm"
docker run \
-e STORE_MODEL_IN_DB=True \
-p 4000:4000 \
ghcr.io/berriai/litellm:v1.77.2
```
</TabItem>
<TabItem value="pip" label="Pip">
``` showLineNumbers title="pip install litellm"
pip install litellm==1.77.2
```
</TabItem>

View file

@ -453,6 +453,7 @@ const sidebars = {
"providers/elevenlabs",
"providers/fireworks_ai",
"providers/clarifai",
"providers/compactifai",
"providers/vllm",
"providers/llamafile",
"providers/infinity",

View file

@ -1030,6 +1030,7 @@ from .llms.openai_like.chat.handler import OpenAILikeChatConfig
from .llms.aiohttp_openai.chat.transformation import AiohttpOpenAIChatConfig
from .llms.galadriel.chat.transformation import GaladrielChatConfig
from .llms.github.chat.transformation import GithubChatConfig
from .llms.compactifai.chat.transformation import CompactifAIChatConfig
from .llms.empower.chat.transformation import EmpowerChatConfig
from .llms.huggingface.chat.transformation import HuggingFaceChatConfig
from .llms.huggingface.embedding.transformation import HuggingFaceEmbeddingConfig

View file

@ -203,7 +203,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
start_time=start_time,
end_time=end_time,
)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
await self._async_log_event_base(
kwargs=kwargs,
@ -212,7 +212,6 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
end_time=end_time,
)
pass
async def _async_log_event_base(self, kwargs, response_obj, start_time, end_time):
try:
@ -242,7 +241,6 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
verbose_logger.exception(f"s3 Layer Error - {str(e)}")
pass
async def async_upload_data_to_s3(
self, batch_logging_element: s3BatchLoggingElement
):
@ -277,8 +275,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Prepare the URL
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{batch_logging_element.s3_object_key}"
if self.s3_endpoint_url:
url = self.s3_endpoint_url + "/" + batch_logging_element.s3_object_key
if self.s3_endpoint_url and self.s3_bucket_name:
url = (
self.s3_endpoint_url
+ "/"
+ self.s3_bucket_name
+ "/"
+ batch_logging_element.s3_object_key
)
# Convert JSON to string
json_string = safe_dumps(batch_logging_element.payload)
@ -420,8 +424,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Prepare the URL
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{batch_logging_element.s3_object_key}"
if self.s3_endpoint_url:
url = self.s3_endpoint_url + "/" + batch_logging_element.s3_object_key
if self.s3_endpoint_url and self.s3_bucket_name:
url = (
self.s3_endpoint_url
+ "/"
+ self.s3_bucket_name
+ "/"
+ batch_logging_element.s3_object_key
)
# Convert JSON to string
json_string = safe_dumps(batch_logging_element.payload)
@ -462,14 +472,13 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
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
"""
@ -481,7 +490,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
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
@ -506,8 +515,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# 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
if self.s3_endpoint_url and self.s3_bucket_name:
url = (
self.s3_endpoint_url
+ "/"
+ self.s3_bucket_name
+ "/"
+ s3_object_key
)
# Prepare the request for GET operation
# For GET requests, we need x-amz-content-sha256 with hash of empty string
@ -533,12 +548,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
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)
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
@ -551,11 +568,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
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
"""
@ -564,5 +581,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
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
verbose_logger.exception(
f"Error retrieving object {object_key} from cold storage: {str(e)}"
)
return None

View file

@ -375,6 +375,8 @@ def get_llm_provider( # noqa: PLR0915
custom_llm_provider = "cometapi"
elif model.startswith("oci/"):
custom_llm_provider = "oci"
elif model.startswith("compactifai/"):
custom_llm_provider = "compactifai"
elif model.startswith("ovhcloud/"):
custom_llm_provider = "ovhcloud"
if not custom_llm_provider:

View file

@ -2680,7 +2680,10 @@ def _convert_to_bedrock_tool_call_invoke(
id = tool["id"]
name = tool["function"].get("name", "")
arguments = tool["function"].get("arguments", "")
arguments_dict = json.loads(arguments) if arguments else {}
if not arguments or not arguments.strip():
arguments_dict = {}
else:
arguments_dict = json.loads(arguments)
bedrock_tool = BedrockToolUseBlock(
input=arguments_dict, name=name, toolUseId=id
)

View file

@ -1,5 +1,7 @@
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
import httpx
from litellm._logging import verbose_logger
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
@ -194,3 +196,66 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
params["order"] = order
verbose_logger.debug(f"list input items url={url}")
return url, params
#########################################################
########## CANCEL RESPONSE API TRANSFORMATION ##########
#########################################################
def transform_cancel_response_api_request(
self,
response_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[str, Dict]:
"""
Transform the cancel response API request into a URL and data
Azure OpenAI API expects the following request:
- POST /openai/responses/{response_id}/cancel?api-version=xxx
This function handles URLs with query parameters by inserting the response_id
at the correct location (before any query parameters).
"""
from urllib.parse import urlparse, urlunparse
# Parse the URL to separate its components
parsed_url = urlparse(api_base)
# Insert the response_id and /cancel at the end of the path component
# Remove trailing slash if present to avoid double slashes
path = parsed_url.path.rstrip("/")
new_path = f"{path}/{response_id}/cancel"
# Reconstruct the URL with all original components but with the modified path
cancel_url = urlunparse(
(
parsed_url.scheme, # http, https
parsed_url.netloc, # domain name, port
new_path, # path with response_id and /cancel added
parsed_url.params, # parameters
parsed_url.query, # query string
parsed_url.fragment, # fragment
)
)
data: Dict = {}
verbose_logger.debug(f"cancel response url={cancel_url}")
return cancel_url, data
def transform_cancel_response_api_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> ResponsesAPIResponse:
"""
Transform the cancel response API response into a ResponsesAPIResponse
"""
try:
raw_response_json = raw_response.json()
except Exception:
from litellm.llms.azure.chat.gpt_transformation import AzureOpenAIError
raise AzureOpenAIError(
message=raw_response.text, status_code=raw_response.status_code
)
return ResponsesAPIResponse(**raw_response_json)

View file

@ -217,3 +217,28 @@ class BaseResponsesAPIConfig(ABC):
) -> bool:
"""Returns True if litellm should fake a stream for the given model and stream value"""
return False
#########################################################
########## CANCEL RESPONSE API TRANSFORMATION ##########
#########################################################
@abstractmethod
def transform_cancel_response_api_request(
self,
response_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[str, Dict]:
pass
@abstractmethod
def transform_cancel_response_api_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> ResponsesAPIResponse:
pass
#########################################################
########## END CANCEL RESPONSE API TRANSFORMATION #######
#########################################################

View file

@ -66,6 +66,7 @@ class BaseAWSLLM:
"aws_web_identity_token",
"aws_sts_endpoint",
"aws_bedrock_runtime_endpoint",
"aws_external_id",
]
def get_cache_key(self, credential_args: Dict[str, Optional[str]]) -> str:
@ -88,6 +89,7 @@ class BaseAWSLLM:
aws_role_name: Optional[str] = None,
aws_web_identity_token: Optional[str] = None,
aws_sts_endpoint: Optional[str] = None,
aws_external_id: Optional[str] = None,
):
"""
Return a boto3.Credentials object
@ -103,6 +105,7 @@ class BaseAWSLLM:
aws_role_name,
aws_web_identity_token,
aws_sts_endpoint,
aws_external_id,
]
# Iterate over parameters and update if needed
@ -127,6 +130,7 @@ class BaseAWSLLM:
aws_role_name,
aws_web_identity_token,
aws_sts_endpoint,
aws_external_id,
) = params_to_check
verbose_logger.debug(
@ -139,7 +143,8 @@ class BaseAWSLLM:
"aws_profile_name=%s\n"
"aws_role_name=%s\n"
"aws_web_identity_token=%s\n"
"aws_sts_endpoint=%s",
"aws_sts_endpoint=%s\n"
"aws_external_id=%s",
aws_access_key_id,
aws_secret_access_key,
aws_session_token,
@ -149,6 +154,7 @@ class BaseAWSLLM:
aws_role_name,
aws_web_identity_token,
aws_sts_endpoint,
aws_external_id,
)
# create cache key for non-expiring auth flows
@ -177,6 +183,7 @@ class BaseAWSLLM:
aws_session_name=aws_session_name,
aws_region_name=aws_region_name,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
elif aws_role_name is not None:
# Check if we're in IRSA and trying to assume the same role we already have
@ -205,6 +212,7 @@ class BaseAWSLLM:
aws_session_token=aws_session_token,
aws_role_name=aws_role_name,
aws_session_name=aws_session_name,
aws_external_id=aws_external_id,
)
elif aws_profile_name is not None: ### CHECK SESSION ###
@ -406,6 +414,7 @@ class BaseAWSLLM:
aws_session_name: str,
aws_region_name: Optional[str],
aws_sts_endpoint: Optional[str],
aws_external_id: Optional[str] = None,
) -> Tuple[Credentials, Optional[int]]:
"""
Authenticate with AWS Web Identity Token
@ -438,13 +447,19 @@ class BaseAWSLLM:
# https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html
# https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html
sts_response = sts_client.assume_role_with_web_identity(
RoleArn=aws_role_name,
RoleSessionName=aws_session_name,
WebIdentityToken=oidc_token,
DurationSeconds=3600,
Policy='{"Version":"2012-10-17","Statement":[{"Sid":"BedrockLiteLLM","Effect":"Allow","Action":["bedrock:InvokeModel","bedrock:InvokeModelWithResponseStream"],"Resource":"*","Condition":{"Bool":{"aws:SecureTransport":"true"},"StringLike":{"aws:UserAgent":"litellm/*"}}}]}',
)
assume_role_params = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name,
"WebIdentityToken": oidc_token,
"DurationSeconds": 3600,
"Policy": '{"Version":"2012-10-17","Statement":[{"Sid":"BedrockLiteLLM","Effect":"Allow","Action":["bedrock:InvokeModel","bedrock:InvokeModelWithResponseStream"],"Resource":"*","Condition":{"Bool":{"aws:SecureTransport":"true"},"StringLike":{"aws:UserAgent":"litellm/*"}}}]}',
}
# Add ExternalId parameter if provided
if aws_external_id is not None:
assume_role_params["ExternalId"] = aws_external_id
sts_response = sts_client.assume_role_with_web_identity(**assume_role_params)
iam_creds_dict = {
"aws_access_key_id": sts_response["Credentials"]["AccessKeyId"],
@ -464,8 +479,9 @@ class BaseAWSLLM:
iam_creds = session.get_credentials()
return iam_creds, self._get_default_ttl_for_boto3_credentials()
def _handle_irsa_cross_account(self, irsa_role_arn: str, aws_role_name: str,
aws_session_name: str, region: str, web_identity_token_file: str) -> dict:
def _handle_irsa_cross_account(self, irsa_role_arn: str, aws_role_name: str,
aws_session_name: str, region: str, web_identity_token_file: str,
aws_external_id: Optional[str] = None) -> dict:
"""Handle cross-account role assumption for IRSA."""
import boto3
@ -509,11 +525,19 @@ class BaseAWSLLM:
# Now assume the target role
verbose_logger.debug(f"Attempting to assume target role: {aws_role_name} with session: {aws_session_name}")
return sts_client_with_creds.assume_role(
RoleArn=aws_role_name, RoleSessionName=aws_session_name
)
assume_role_params = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name
}
def _handle_irsa_same_account(self, aws_role_name: str, aws_session_name: str, region: str) -> dict:
# Add ExternalId parameter if provided
if aws_external_id is not None:
assume_role_params["ExternalId"] = aws_external_id
return sts_client_with_creds.assume_role(**assume_role_params)
def _handle_irsa_same_account(self, aws_role_name: str, aws_session_name: str, region: str,
aws_external_id: Optional[str] = None) -> dict:
"""Handle same-account role assumption for IRSA."""
import boto3
@ -530,9 +554,16 @@ class BaseAWSLLM:
# Assume the role
verbose_logger.debug(f"Attempting to assume role: {aws_role_name} with session: {aws_session_name}")
return sts_client.assume_role(
RoleArn=aws_role_name, RoleSessionName=aws_session_name
)
assume_role_params = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name
}
# Add ExternalId parameter if provided
if aws_external_id is not None:
assume_role_params["ExternalId"] = aws_external_id
return sts_client.assume_role(**assume_role_params)
def _extract_credentials_and_ttl(self, sts_response: dict) -> Tuple[Credentials, Optional[int]]:
"""Extract credentials and TTL from STS response."""
@ -558,6 +589,7 @@ class BaseAWSLLM:
aws_session_token: Optional[str],
aws_role_name: str,
aws_session_name: str,
aws_external_id: Optional[str] = None,
) -> Tuple[Credentials, Optional[int]]:
"""
Authenticate with AWS Role
@ -584,11 +616,11 @@ class BaseAWSLLM:
# Check if we need to do cross-account role assumption
if aws_role_name != irsa_role_arn:
sts_response = self._handle_irsa_cross_account(
irsa_role_arn, aws_role_name, aws_session_name, region, web_identity_token_file
irsa_role_arn, aws_role_name, aws_session_name, region, web_identity_token_file, aws_external_id
)
else:
sts_response = self._handle_irsa_same_account(
aws_role_name, aws_session_name, region
aws_role_name, aws_session_name, region, aws_external_id
)
return self._extract_credentials_and_ttl(sts_response)
@ -619,9 +651,16 @@ class BaseAWSLLM:
aws_session_token=aws_session_token,
)
sts_response = sts_client.assume_role(
RoleArn=aws_role_name, RoleSessionName=aws_session_name
)
assume_role_params = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name
}
# Add ExternalId parameter if provided
if aws_external_id is not None:
assume_role_params["ExternalId"] = aws_external_id
sts_response = sts_client.assume_role(**assume_role_params)
# Extract the credentials from the response and convert to Session Credentials
sts_credentials = sts_response["Credentials"]
@ -800,6 +839,7 @@ class BaseAWSLLM:
aws_bedrock_runtime_endpoint = optional_params.pop(
"aws_bedrock_runtime_endpoint", None
) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_external_id = optional_params.pop("aws_external_id", None)
credentials: Credentials = self.get_credentials(
aws_access_key_id=aws_access_key_id,
@ -811,6 +851,7 @@ class BaseAWSLLM:
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
return Boto3CredentialsInfo(
@ -915,6 +956,7 @@ class BaseAWSLLM:
aws_profile_name = optional_params.get("aws_profile_name", None)
aws_web_identity_token = optional_params.get("aws_web_identity_token", None)
aws_sts_endpoint = optional_params.get("aws_sts_endpoint", None)
aws_external_id = optional_params.get("aws_external_id", None)
aws_region_name = self._get_aws_region_name(
optional_params=optional_params, model=model
)
@ -929,6 +971,7 @@ class BaseAWSLLM:
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
sigv4 = SigV4Auth(credentials, service_name, aws_region_name)

View file

@ -307,6 +307,7 @@ class BedrockConverseLLM(BaseAWSLLM):
) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_web_identity_token = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint = optional_params.pop("aws_sts_endpoint", None)
aws_external_id = optional_params.pop("aws_external_id", None)
optional_params.pop("aws_region_name", None)
litellm_params[
@ -323,6 +324,7 @@ class BedrockConverseLLM(BaseAWSLLM):
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
### SET RUNTIME ENDPOINT ###

View file

@ -0,0 +1 @@
# CompactifAI provider for LiteLLM

View file

@ -0,0 +1 @@
# CompactifAI chat completions

View file

@ -0,0 +1,100 @@
"""
CompactifAI chat completion transformation
"""
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union
import httpx
from litellm.secret_managers.main import get_secret_str
from litellm.types.utils import ModelResponse
from litellm.llms.openai.common_utils import OpenAIError
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class CompactifAIChatConfig(OpenAIGPTConfig):
"""
Configuration class for CompactifAI chat completions.
Since CompactifAI is OpenAI-compatible, we extend OpenAIGPTConfig.
"""
def _get_openai_compatible_provider_info(
self,
api_base: Optional[str],
api_key: Optional[str],
) -> Tuple[Optional[str], Optional[str]]:
"""
Get API base and key for CompactifAI provider.
"""
api_base = api_base or "https://api.compactif.ai/v1"
dynamic_api_key = api_key or get_secret_str("COMPACTIFAI_API_KEY") or ""
return api_base, dynamic_api_key
def transform_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List,
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ModelResponse:
"""
Transform CompactifAI response to LiteLLM format.
Since CompactifAI is OpenAI-compatible, we can use the standard OpenAI transformation.
"""
## LOGGING
logging_obj.post_call(
input=messages,
api_key=api_key,
original_response=raw_response.text,
additional_args={"complete_input_dict": request_data},
)
## RESPONSE OBJECT
response_json = raw_response.json()
# Handle JSON mode if needed
if json_mode:
for choice in response_json["choices"]:
message = choice.get("message")
if message and message.get("tool_calls"):
# Convert tool calls to content for JSON mode
tool_calls = message.get("tool_calls", [])
if len(tool_calls) == 1:
message["content"] = tool_calls[0]["function"].get("arguments", "")
message["tool_calls"] = None
returned_response = ModelResponse(**response_json)
# Set model name with provider prefix
returned_response.model = f"compactifai/{model}"
return returned_response
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BaseLLMException:
"""
Get the appropriate error class for CompactifAI errors.
Since CompactifAI is OpenAI-compatible, we use OpenAI error handling.
"""
return OpenAIError(
status_code=status_code,
message=error_message,
headers=headers,
)

View file

@ -2200,6 +2200,7 @@ class BaseLLMHTTPHandler:
litellm_params=litellm_params,
optional_params={},
)
if _is_async:
return self.async_create_file(
transformed_request=transformed_request,
@ -2216,7 +2217,6 @@ class BaseLLMHTTPHandler:
sync_httpx_client = _get_httpx_client()
else:
sync_httpx_client = client
if isinstance(transformed_request, dict) and "method" in transformed_request:
# Handle pre-signed requests (e.g., from Bedrock S3 uploads)
@ -2283,11 +2283,11 @@ class BaseLLMHTTPHandler:
e=e,
provider_config=provider_config,
)
# Store the upload URL in litellm_params for the transformation method
litellm_params_with_url = dict(litellm_params)
litellm_params_with_url["upload_url"] = api_base
return provider_config.transform_create_file_response(
model=None,
raw_response=upload_response,
@ -2423,7 +2423,7 @@ class BaseLLMHTTPHandler:
# get config from model, custom llm provider
if model is None:
raise ValueError("model is required for create_batch")
headers = provider_config.validate_environment(
api_key=api_key,
headers=headers,
@ -2606,6 +2606,159 @@ class BaseLLMHTTPHandler:
litellm_params=litellm_params_with_request,
)
def cancel_response_api_handler(
self,
response_id: str,
responses_api_provider_config: BaseResponsesAPIConfig,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Optional[str],
extra_headers: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
_is_async: bool = False,
) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]:
"""
Async version of the responses API handler.
Uses async HTTP client to make requests.
"""
if _is_async:
return self.async_cancel_response_api_handler(
response_id=response_id,
responses_api_provider_config=responses_api_provider_config,
litellm_params=litellm_params,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout,
client=client,
)
if client is None or not isinstance(client, HTTPHandler):
sync_httpx_client = _get_httpx_client(
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
)
else:
sync_httpx_client = client
headers = responses_api_provider_config.validate_environment(
headers=extra_headers or {}, model="None", litellm_params=litellm_params
)
if extra_headers:
headers.update(extra_headers)
api_base = responses_api_provider_config.get_complete_url(
api_base=litellm_params.api_base,
litellm_params=dict(litellm_params),
)
url, data = responses_api_provider_config.transform_cancel_response_api_request(
response_id=response_id,
api_base=api_base,
litellm_params=litellm_params,
headers=headers,
)
## LOGGING
logging_obj.pre_call(
input=response_id,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": url,
"headers": headers,
},
)
try:
response = sync_httpx_client.post(
url=url, headers=headers, json=data, timeout=timeout
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=responses_api_provider_config,
)
return responses_api_provider_config.transform_cancel_response_api_response(
raw_response=response,
logging_obj=logging_obj,
)
async def async_cancel_response_api_handler(
self,
response_id: str,
responses_api_provider_config: BaseResponsesAPIConfig,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Optional[str],
extra_headers: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
_is_async: bool = False,
) -> ResponsesAPIResponse:
"""
Async version of the cancel response API handler.
Uses async HTTP client to make requests.
"""
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders(custom_llm_provider),
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
)
else:
async_httpx_client = client
headers = responses_api_provider_config.validate_environment(
headers=extra_headers or {}, model="None", litellm_params=litellm_params
)
if extra_headers:
headers.update(extra_headers)
api_base = responses_api_provider_config.get_complete_url(
api_base=litellm_params.api_base,
litellm_params=dict(litellm_params),
)
url, data = responses_api_provider_config.transform_cancel_response_api_request(
response_id=response_id,
api_base=api_base,
litellm_params=litellm_params,
headers=headers,
)
## LOGGING
logging_obj.pre_call(
input=response_id,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": url,
"headers": headers,
},
)
try:
response = await async_httpx_client.post(
url=url, headers=headers, json=data, timeout=timeout
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=responses_api_provider_config,
)
return responses_api_provider_config.transform_cancel_response_api_response(
raw_response=response,
logging_obj=logging_obj,
)
def list_files(self):
"""
Lists all files
@ -2766,10 +2919,7 @@ class BaseLLMHTTPHandler:
_is_async: bool = False,
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
) -> Union[
ImageResponse,
Coroutine[Any, Any, ImageResponse],
]:
) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]:
"""
Handles image edit requests.
@ -2959,10 +3109,7 @@ class BaseLLMHTTPHandler:
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
api_key: Optional[str] = None,
) -> Union[
ImageResponse,
Coroutine[Any, Any, ImageResponse],
]:
) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]:
"""
Handles image generation requests.
When _is_async=True, returns a coroutine instead of making the call directly.
@ -3196,15 +3343,16 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
)
url, request_body = (
vector_store_provider_config.transform_search_vector_store_request(
vector_store_id=vector_store_id,
query=query,
vector_store_search_optional_params=vector_store_search_optional_params,
api_base=api_base,
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
)
(
url,
request_body,
) = vector_store_provider_config.transform_search_vector_store_request(
vector_store_id=vector_store_id,
query=query,
vector_store_search_optional_params=vector_store_search_optional_params,
api_base=api_base,
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
)
all_optional_params: Dict[str, Any] = dict(litellm_params)
all_optional_params.update(vector_store_search_optional_params or {})
@ -3295,15 +3443,16 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
)
url, request_body = (
vector_store_provider_config.transform_search_vector_store_request(
vector_store_id=vector_store_id,
query=query,
vector_store_search_optional_params=vector_store_search_optional_params,
api_base=api_base,
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
)
(
url,
request_body,
) = vector_store_provider_config.transform_search_vector_store_request(
vector_store_id=vector_store_id,
query=query,
vector_store_search_optional_params=vector_store_search_optional_params,
api_base=api_base,
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
)
all_optional_params: Dict[str, Any] = dict(litellm_params)
@ -3377,11 +3526,12 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
)
url, request_body = (
vector_store_provider_config.transform_create_vector_store_request(
vector_store_create_optional_params=vector_store_create_optional_params,
api_base=api_base,
)
(
url,
request_body,
) = vector_store_provider_config.transform_create_vector_store_request(
vector_store_create_optional_params=vector_store_create_optional_params,
api_base=api_base,
)
logging_obj.pre_call(
@ -3452,11 +3602,12 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
)
url, request_body = (
vector_store_provider_config.transform_create_vector_store_request(
vector_store_create_optional_params=vector_store_create_optional_params,
api_base=api_base,
)
(
url,
request_body,
) = vector_store_provider_config.transform_create_vector_store_request(
vector_store_create_optional_params=vector_store_create_optional_params,
api_base=api_base,
)
logging_obj.pre_call(
@ -3535,13 +3686,14 @@ class BaseLLMHTTPHandler:
sync_httpx_client = client
# Get headers and URL from the provider config
headers, api_base = (
generate_content_provider_config.sync_get_auth_token_and_url(
api_base=litellm_params.api_base,
model=model,
litellm_params=dict(litellm_params),
stream=stream,
)
(
headers,
api_base,
) = generate_content_provider_config.sync_get_auth_token_and_url(
api_base=litellm_params.api_base,
model=model,
litellm_params=dict(litellm_params),
stream=stream,
)
if extra_headers:
@ -3641,13 +3793,14 @@ class BaseLLMHTTPHandler:
async_httpx_client = client
# Get headers and URL from the provider config
headers, api_base = (
await generate_content_provider_config.get_auth_token_and_url(
model=model,
litellm_params=dict(litellm_params),
stream=stream,
api_base=litellm_params.api_base,
)
(
headers,
api_base,
) = await generate_content_provider_config.get_auth_token_and_url(
model=model,
litellm_params=dict(litellm_params),
stream=stream,
api_base=litellm_params.api_base,
)
if extra_headers:

View file

@ -425,3 +425,39 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
raise OpenAIError(
message=raw_response.text, status_code=raw_response.status_code
)
#########################################################
########## CANCEL RESPONSE API TRANSFORMATION ##########
#########################################################
def transform_cancel_response_api_request(
self,
response_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[str, Dict]:
"""
Transform the cancel response API request into a URL and data
OpenAI API expects the following request
- POST /v1/responses/{response_id}/cancel
"""
url = f"{api_base}/{response_id}/cancel"
data: Dict = {}
return url, data
def transform_cancel_response_api_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> ResponsesAPIResponse:
"""
Transform the cancel response API response into a ResponsesAPIResponse
"""
try:
raw_response_json = raw_response.json()
except Exception:
raise OpenAIError(
message=raw_response.text, status_code=raw_response.status_code
)
return ResponsesAPIResponse(**raw_response_json)

View file

@ -28,6 +28,7 @@ from litellm.types.files import (
get_file_type_from_extension,
is_gemini_1_5_accepted_file_type,
)
from litellm.types.utils import LlmProviders
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionAssistantMessage,
@ -492,7 +493,8 @@ def _transform_request_body(
data["generationConfig"] = generation_config
if cached_content is not None:
data["cachedContent"] = cached_content
if labels is not None:
# Only add labels for Vertex AI endpoints (not Google GenAI/AI Studio) and only if non-empty
if labels and custom_llm_provider != LlmProviders.GEMINI:
data["labels"] = labels
except Exception as e:
raise e
@ -647,3 +649,5 @@ def _transform_system_message(
return SystemInstructions(parts=system_content_blocks), messages
return None, messages

View file

@ -4,6 +4,9 @@ from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig
class VolcEngineChatConfig(OpenAILikeChatConfig):
"""
Reference: https://www.volcengine.com/docs/82379/1494384
"""
frequency_penalty: Optional[int] = None
function_call: Optional[Union[str, dict]] = None
functions: Optional[list] = None
@ -81,20 +84,22 @@ class VolcEngineChatConfig(OpenAILikeChatConfig):
)
if "thinking" in optional_params:
"""
The `thinking` parameters of VolcEngine model has different default values.
See the docs for details.
Refrence: https://www.volcengine.com/docs/82379/1449737#0002
"""
thinking_value = optional_params.pop("thinking")
# Handle disabled thinking case - don't add to extra_body if disabled
# Handle using thinking params case - add to extra_body if value is legal
if (
thinking_value is not None
and isinstance(thinking_value, dict)
and thinking_value.get("type") == "disabled"
and thinking_value.get("type", None) in ["enabled", "disabled", "auto"] # legal values, see docs
):
# Skip adding thinking parameter when it's disabled
pass
# Add thinking parameter to extra_body for all legal cases
optional_params.setdefault("extra_body", {})["thinking"] = thinking_value
else:
# Add thinking parameter to extra_body for all other cases
optional_params.setdefault("extra_body", {})[
"thinking"
] = thinking_value
# Skip adding thinking parameter when it's not set or has invalid value
pass
return optional_params

View file

@ -80,6 +80,8 @@ class XAIChatConfig(OpenAIGPTConfig):
return False
elif "grok-4" in model:
return False
elif "grok-code-fast" in model:
return False
return True
def _supports_frequency_penalty(self, model: str) -> bool:

View file

@ -2550,6 +2550,37 @@ def completion( # type: ignore # noqa: PLR0915
encoding=encoding,
stream=stream,
)
elif custom_llm_provider == "compactifai":
api_key = (
api_key
or get_secret_str("COMPACTIFAI_API_KEY")
or litellm.api_key
)
api_base = (
api_base
or "https://api.compactif.ai/v1"
)
## COMPLETION CALL
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,
client=client,
custom_llm_provider=custom_llm_provider,
encoding=encoding,
stream=stream,
provider_config=provider_config,
)
elif custom_llm_provider == "oobabooga":
custom_llm_provider = "oobabooga"
model_response = oobabooga.completion(

View file

@ -13477,6 +13477,23 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"bedrock/us.anthropic.claude-3-5-haiku-20241022-v1:0": {
"max_tokens": 8192,
"max_input_tokens": 200000,
"max_output_tokens": 8192,
"input_cost_per_token": 8e-07,
"output_cost_per_token": 4e-06,
"cache_creation_input_token_cost": 1e-06,
"cache_read_input_token_cost": 8e-08,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_assistant_prefill": true,
"supports_pdf_input": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"us.anthropic.claude-3-opus-20240229-v1:0": {
"max_tokens": 4096,
"max_input_tokens": 200000,

View file

@ -578,7 +578,7 @@ if MCP_AVAILABLE:
"""
import re
mcp_servers_from_path: Optional[List[str]] = None
mcp_path_match = re.match(r"^/mcp/([^/]+)(/.*)?$", path)
mcp_path_match = re.match(r"^/mcp/([^/]+/[^/]+|[^/]+)(/.*)?$", path)
if mcp_path_match:
mcp_servers_str = mcp_path_match.group(1)
if mcp_servers_str:

View file

@ -312,6 +312,8 @@ class LiteLLMRoutes(enum.Enum):
"/v1/responses/{response_id}",
"/responses/{response_id}/input_items",
"/v1/responses/{response_id}/input_items",
"/responses/{response_id}/cancel",
"/v1/responses/{response_id}/cancel",
# vector stores
"/vector_stores",
"/v1/vector_stores",

View file

@ -46,6 +46,7 @@ from litellm.proxy._types import (
RoleBasedPermissions,
SpecialModelNames,
UserAPIKeyAuth,
NewTeamRequest,
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.route_llm_request import route_request
@ -889,10 +890,17 @@ async def _get_team_db_check(
)
if response is None and team_id_upsert:
response = await prisma_client.db.litellm_teamtable.create(
data={"team_id": team_id}
)
from litellm.proxy.management_endpoints.team_endpoints import new_team
new_team_data = NewTeamRequest(team_id=team_id)
mock_request = Request(scope={"type": "http"})
system_admin_user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
created_team_dict = await new_team(
data=new_team_data, http_request=mock_request, user_api_key_dict=system_admin_user
)
response = LiteLLM_TeamTable(**created_team_dict)
return response

View file

@ -259,6 +259,7 @@ class ProxyBaseLLMRequestProcessing:
"_arealtime",
"aget_responses",
"adelete_responses",
"acancel_responses",
"acreate_batch",
"aretrieve_batch",
"afile_content",
@ -355,6 +356,7 @@ class ProxyBaseLLMRequestProcessing:
"_arealtime",
"aget_responses",
"adelete_responses",
"acancel_responses",
"atext_completion",
"aimage_edit",
"alist_input_items",

View file

@ -18,6 +18,7 @@ from fastapi import (
Response,
)
from typing_extensions import TypedDict
from pydantic import BaseModel
import litellm
from litellm._logging import verbose_proxy_logger
@ -29,6 +30,7 @@ from litellm.proxy._types import (
Member,
NewTeamRequest,
NewUserRequest,
NewUserResponse,
TeamMemberAddRequest,
TeamMemberDeleteRequest,
UserAPIKeyAuth,
@ -101,6 +103,13 @@ class ScimUserData(TypedDict):
active: Optional[bool]
class GroupMemberExtractionResult(BaseModel):
"""Result of extracting and processing group members."""
existing_member_ids: List[str]
created_users: List[NewUserResponse]
all_member_ids: List[str] # existing + newly created
scim_router = APIRouter(
prefix="/scim/v2",
tags=["✨ SCIM v2 (Enterprise Only)"],
@ -190,21 +199,47 @@ def _build_scim_metadata(given_name: Optional[str], family_name: Optional[str],
return metadata
async def _extract_group_member_ids(group: SCIMGroup) -> List[str]:
"""Extract valid member IDs from SCIMGroup, verifying users exist."""
async def _extract_group_member_ids(group: SCIMGroup) -> GroupMemberExtractionResult:
"""
Extract member IDs from SCIMGroup, creating users that don't exist.
Returns:
GroupMemberExtractionResult with existing members, created users, and all member IDs
"""
prisma_client = await _get_prisma_client_or_raise_exception()
member_ids = []
existing_member_ids = []
created_users = []
all_member_ids = []
if group.members:
for member in group.members:
user_id = member.value
# Check if user exists
user = await prisma_client.db.litellm_usertable.find_unique(
where={"user_id": member.value}
where={"user_id": user_id}
)
if user:
member_ids.append(member.value)
existing_member_ids.append(user_id)
all_member_ids.append(user_id)
else:
# Create the user if they don't exist using our helper
created_user = await _create_user_if_not_exists(
user_id=user_id,
created_via="scim_group_membership"
)
if created_user:
created_users.append(created_user)
all_member_ids.append(user_id)
# If creation failed, user is skipped (logged in helper)
return member_ids
return GroupMemberExtractionResult(
existing_member_ids=existing_member_ids,
created_users=created_users,
all_member_ids=all_member_ids
)
async def _get_team_members_display(member_ids: List[str]) -> List[SCIMMember]:
@ -239,6 +274,51 @@ async def _handle_team_membership_changes(user_id: str, existing_teams: List[str
)
async def _create_user_if_not_exists(user_id: str, created_via: str = "scim_group") -> Optional[NewUserResponse]:
"""
Helper function to create a user if they don't exist.
Args:
user_id: The user ID to create
created_via: Context for where the user was created from
Returns:
LiteLLM_UserTable if user was created, None if creation failed
"""
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
try:
# Get default role for new internal users
default_role: Optional[
Literal[
LitellmUserRoles.PROXY_ADMIN,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
LitellmUserRoles.INTERNAL_USER,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
]
] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
if litellm.default_internal_user_params:
default_role = litellm.default_internal_user_params.get("user_role")
new_user_request = NewUserRequest(
user_id=user_id,
user_email=user_id, # We don't have email from group membership
user_alias=None,
teams=[], # Teams will be added separately
metadata={"created_via": created_via},
auto_create_key=False,
user_role=default_role,
)
created_user = await new_user(data=new_user_request)
verbose_proxy_logger.info(f"Created user {user_id} via {created_via}")
return created_user
except Exception as e:
verbose_proxy_logger.exception(f"Failed to create user {user_id}: {e}")
return None
async def _get_team_member_user_ids_from_team(team: LiteLLM_TeamTable) -> List[str]:
"""
Get the IDs of the members from a team.
@ -256,6 +336,8 @@ async def _get_team_member_user_ids_from_team(team: LiteLLM_TeamTable) -> List[s
member_user_ids.append(user_id)
return member_user_ids
# Dependency to set the correct SCIM Content-Type
async def set_scim_content_type(response: Response):
"""Sets the Content-Type header to application/scim+json"""
@ -914,9 +996,9 @@ async def create_group(
detail={"error": f"Group already exists with ID: {team_id}"},
)
# Extract valid member IDs
member_ids = await _extract_group_member_ids(group)
members_with_roles = [Member(user_id=member_id, role="user") for member_id in member_ids]
# Extract and process group members (creating users that don't exist)
member_result = await _extract_group_member_ids(group)
members_with_roles = [Member(user_id=member_id, role="user") for member_id in member_result.all_member_ids]
# Create team in database
created_team = await new_team(
@ -959,9 +1041,10 @@ async def update_group(
prisma_client = await _get_prisma_client_or_raise_exception()
existing_team = await _check_team_exists(group_id)
# Extract valid member IDs
member_ids = await _extract_group_member_ids(group)
verbose_proxy_logger.debug(f"SCIM PUT GROUP member_ids: {member_ids}")
# Extract and process group members (creating users that don't exist)
member_result = await _extract_group_member_ids(group)
verbose_proxy_logger.debug(f"SCIM PUT GROUP all_member_ids: {member_result.all_member_ids}")
verbose_proxy_logger.debug(f"SCIM PUT GROUP created_users: {len(member_result.created_users)}")
# Prepare update data
existing_metadata = existing_team.metadata if existing_team.metadata else {}
@ -978,10 +1061,10 @@ async def update_group(
data=update_data,
)
# Handle user-team relationship changes using the same approach as patch_group
# Handle user-team relationship changes
current_members = set(await _get_team_member_user_ids_from_team(existing_team))
verbose_proxy_logger.debug(f"SCIM PUT GROUP current_members: {current_members}")
final_members = set(member_ids)
final_members = set(member_result.all_member_ids)
verbose_proxy_logger.debug(f"SCIM PUT GROUP final_members: {final_members}")
await _handle_group_membership_changes(
@ -1075,7 +1158,7 @@ async def _process_group_patch_operations(
elif path.startswith("members"):
# Handle member operations
member_values = _extract_group_values(value)
# Validate that users exist
# Create users that don't exist and get all valid member IDs
valid_members = []
for member_id in member_values:
user = await prisma_client.db.litellm_usertable.find_unique(
@ -1083,6 +1166,16 @@ async def _process_group_patch_operations(
)
if user:
valid_members.append(member_id)
else:
# Create the user if they don't exist using our helper
created_user = await _create_user_if_not_exists(
user_id=member_id,
created_via="scim_group_patch"
)
if created_user:
valid_members.append(member_id)
# If creation failed, user is skipped (logged in helper)
if op_type == "replace":
final_members = set(valid_members)

View file

@ -383,6 +383,19 @@ async def new_team( # noqa: PLR0915
"error": f"Team id = {data.team_id} already exists. Please use a different team id."
},
)
# If max_budget is not explicitly provided in the request,
# check for a default value in the proxy configuration.
if data.max_budget is None:
if (
isinstance(litellm.default_team_settings, list)
and len(litellm.default_team_settings) > 0
and isinstance(litellm.default_team_settings[0], dict)
):
default_settings = litellm.default_team_settings[0]
default_budget = default_settings.get("max_budget")
if default_budget is not None:
data.max_budget = default_budget
if (
user_api_key_dict.user_role is None

View file

@ -285,3 +285,75 @@ async def get_response_input_items(
proxy_logging_obj=proxy_logging_obj,
version=version,
)
@router.post(
"/v1/responses/{response_id}/cancel",
dependencies=[Depends(user_api_key_auth)],
tags=["responses"],
)
@router.post(
"/responses/{response_id}/cancel",
dependencies=[Depends(user_api_key_auth)],
tags=["responses"],
)
async def cancel_response(
response_id: str,
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Cancel a response by ID.
Follows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses/cancel
```bash
curl -X POST http://localhost:4000/v1/responses/resp_abc123/cancel \
-H "Authorization: Bearer sk-1234"
```
"""
from litellm.proxy.proxy_server import (
_read_request_body,
general_settings,
llm_router,
proxy_config,
proxy_logging_obj,
select_data_generator,
user_api_base,
user_max_tokens,
user_model,
user_request_timeout,
user_temperature,
version,
)
data = await _read_request_body(request=request)
data["response_id"] = response_id
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type="acancel_responses",
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
model=None,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)

View file

@ -24,6 +24,7 @@ ROUTE_ENDPOINT_MAPPING = {
"aresponses": "/responses",
"alist_input_items": "/responses/{response_id}/input_items",
"aimage_edit": "/images/edits",
"acancel_responses": "/responses/{response_id}/cancel",
}
@ -70,6 +71,8 @@ async def route_request(
"aresponses",
"aget_responses",
"adelete_responses",
"acancel_responses",
"acreate_response_reply",
"alist_input_items",
"_arealtime", # private function for realtime API
"aimage_edit",
@ -86,6 +89,11 @@ async def route_request(
team_id = get_team_id_from_data(data)
router_model_names = llm_router.model_names if llm_router is not None else []
# Preprocess Google GenAI generate content requests
if route_type in ["agenerate_content", "agenerate_content_stream"]:
# Map generationConfig to config parameter for Google GenAI compatibility
if "generationConfig" in data and "config" not in data:
data["config"] = data.pop("generationConfig")
if "api_key" in data or "api_base" in data:
if llm_router is not None:
return getattr(llm_router, f"{route_type}")(**data)
@ -149,6 +157,7 @@ async def route_request(
"amoderation",
"aget_responses",
"adelete_responses",
"acancel_responses",
"alist_input_items",
"avector_store_create",
"avector_store_search",

View file

@ -167,13 +167,17 @@ async def aresponses_api_with_mcp(
# Process MCP tools through the complete pipeline (fetch + filter + deduplicate + transform)
user_api_key_auth = kwargs.get("user_api_key_auth")
# Get original MCP tools (for events) and OpenAI tools (for LLM) by reusing existing methods
original_mcp_tools = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
user_api_key_auth=user_api_key_auth,
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy
original_mcp_tools = (
await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
user_api_key_auth=user_api_key_auth,
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
)
)
openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(
original_mcp_tools
)
openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(original_mcp_tools)
# Combine with other tools
all_tools = openai_tools + other_tools if (openai_tools or other_tools) else None
@ -212,15 +216,15 @@ async def aresponses_api_with_mcp(
from litellm.responses.mcp.mcp_streaming_iterator import (
create_mcp_list_tools_events,
)
base_item_id = f"mcp_{uuid.uuid4().hex[:8]}"
mcp_discovery_events = await create_mcp_list_tools_events(
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
user_api_key_auth=user_api_key_auth,
base_item_id=base_item_id,
pre_processed_mcp_tools=original_mcp_tools
pre_processed_mcp_tools=original_mcp_tools,
)
return LiteLLM_Proxy_MCP_Handler._create_mcp_streaming_response(
input=input,
model=model,
@ -229,23 +233,21 @@ async def aresponses_api_with_mcp(
mcp_discovery_events=mcp_discovery_events,
call_params=call_params,
previous_response_id=previous_response_id,
**kwargs
**kwargs,
)
# Determine if we should auto-execute tools
should_auto_execute = (
bool(mcp_tools_with_litellm_proxy)
and LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy
)
should_auto_execute = bool(
mcp_tools_with_litellm_proxy
) and LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy
)
# Prepare parameters for the initial call
initial_call_params = LiteLLM_Proxy_MCP_Handler._prepare_initial_call_params(
call_params=call_params,
should_auto_execute=should_auto_execute
call_params=call_params, should_auto_execute=should_auto_execute
)
#########################################################
# Make initial response API call
#########################################################
@ -263,9 +265,8 @@ async def aresponses_api_with_mcp(
# Auto-Execute Tools Handling
# If auto-execute tools is True, then we need to execute the tool calls
#########################################################
if (
should_auto_execute
and isinstance(response, ResponsesAPIResponse)
if should_auto_execute and isinstance(
response, ResponsesAPIResponse
): # type: ignore
tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_response(
response=response
@ -285,19 +286,21 @@ async def aresponses_api_with_mcp(
)
# Prepare parameters for follow-up call (restores original stream setting)
follow_up_call_params = LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params(
call_params=call_params,
original_stream_setting=stream or False
follow_up_call_params = (
LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params(
call_params=call_params, original_stream_setting=stream or False
)
)
# Create tool execution events for streaming if needed
tool_execution_events = []
if stream:
tool_execution_events = LiteLLM_Proxy_MCP_Handler._create_tool_execution_events(
tool_calls=tool_calls,
tool_results=tool_results
tool_execution_events = (
LiteLLM_Proxy_MCP_Handler._create_tool_execution_events(
tool_calls=tool_calls, tool_results=tool_results
)
)
final_response = await LiteLLM_Proxy_MCP_Handler._make_follow_up_call(
follow_up_input=follow_up_input,
model=model,
@ -307,13 +310,20 @@ async def aresponses_api_with_mcp(
)
# If streaming and we have tool execution events, wrap the response
if stream and tool_execution_events and (hasattr(final_response, '__aiter__') or hasattr(final_response, '__iter__')):
if (
stream
and tool_execution_events
and (
hasattr(final_response, "__aiter__")
or hasattr(final_response, "__iter__")
)
):
from litellm.responses.mcp.mcp_streaming_iterator import (
MCPEnhancedStreamingIterator,
)
final_response = MCPEnhancedStreamingIterator(
base_iterator=final_response,
mcp_events=tool_execution_events
base_iterator=final_response, mcp_events=tool_execution_events
)
# Add custom output elements to the final response (for non-streaming)
@ -321,7 +331,7 @@ async def aresponses_api_with_mcp(
# Fetch MCP tools again for output elements (without OpenAI transformation)
mcp_tools_for_output = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
user_api_key_auth=user_api_key_auth,
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
)
final_response = (
LiteLLM_Proxy_MCP_Handler._add_mcp_output_elements_to_response(
@ -1150,4 +1160,163 @@ def list_input_items(
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
)
@client
async def acancel_responses(
response_id: str,
# 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,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
# LiteLLM specific params,
custom_llm_provider: Optional[str] = None,
**kwargs,
) -> ResponsesAPIResponse:
"""
Async version of the POST Cancel Responses API
POST /v1/responses/{response_id}/cancel endpoint in the responses API
"""
local_vars = locals()
try:
loop = asyncio.get_event_loop()
kwargs["acancel_responses"] = True
# get custom llm provider from response_id
decoded_response_id: DecodedResponseId = (
ResponsesAPIRequestUtils._decode_responses_api_response_id(
response_id=response_id,
)
)
response_id = decoded_response_id.get("response_id") or response_id
custom_llm_provider = (
decoded_response_id.get("custom_llm_provider") or custom_llm_provider
)
func = partial(
cancel_responses,
response_id=response_id,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
extra_query=extra_query,
extra_body=extra_body,
timeout=timeout,
**kwargs,
)
ctx = contextvars.copy_context()
func_with_context = partial(ctx.run, func)
init_response = await loop.run_in_executor(None, func_with_context)
if asyncio.iscoroutine(init_response):
response = await init_response
else:
response = init_response
return response
except Exception as e:
raise litellm.exception_type(
model=None,
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
@client
def cancel_responses(
response_id: str,
# 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,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
# LiteLLM specific params,
custom_llm_provider: Optional[str] = None,
**kwargs,
) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]:
"""
Synchronous version of the POST Responses API
POST /v1/responses/{response_id}/cancel endpoint in the responses API
"""
local_vars = locals()
try:
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
_is_async = kwargs.pop("acancel_responses", False) is True
# get llm provider logic
litellm_params = GenericLiteLLMParams(**kwargs)
# get custom llm provider from response_id
decoded_response_id: DecodedResponseId = (
ResponsesAPIRequestUtils._decode_responses_api_response_id(
response_id=response_id,
)
)
response_id = decoded_response_id.get("response_id") or response_id
custom_llm_provider = (
decoded_response_id.get("custom_llm_provider") or custom_llm_provider
)
if custom_llm_provider is None:
raise ValueError("custom_llm_provider is required but passed as None")
# get provider config
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=litellm.LlmProviders(custom_llm_provider),
)
if responses_api_provider_config is None:
raise ValueError(
f"CANCEL responses is not supported for {custom_llm_provider}"
)
local_vars.update(kwargs)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
model=None,
optional_params={
"response_id": response_id,
},
litellm_params={
"litellm_call_id": litellm_call_id,
},
custom_llm_provider=custom_llm_provider,
)
# Call the handler with _is_async flag instead of directly calling the async handler
response = base_llm_http_handler.cancel_response_api_handler(
response_id=response_id,
custom_llm_provider=custom_llm_provider,
responses_api_provider_config=responses_api_provider_config,
litellm_params=litellm_params,
logging_obj=litellm_logging_obj,
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout or request_timeout,
_is_async=_is_async,
client=kwargs.get("client"),
)
return response
except Exception as e:
raise litellm.exception_type(
model=None,
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)

View file

@ -359,9 +359,9 @@ class Router:
) # names of models under litellm_params. ex. azure/chatgpt-v-2
self.deployment_latency_map = {}
### CACHING ###
cache_type: Literal["local", "redis", "redis-semantic", "s3", "disk"] = (
"local" # default to an in-memory cache
)
cache_type: Literal[
"local", "redis", "redis-semantic", "s3", "disk"
] = "local" # default to an in-memory cache
redis_cache = None
cache_config: Dict[str, Any] = {}
@ -403,9 +403,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.team_pattern_routers: Dict[
str, PatternMatchRouter
] = {} # {"TEAM_ID": PatternMatchRouter}
self.auto_routers: Dict[str, "AutoRouter"] = {}
if model_list is not None:
@ -587,9 +587,9 @@ class Router:
)
)
self.model_group_retry_policy: Optional[Dict[str, RetryPolicy]] = (
model_group_retry_policy
)
self.model_group_retry_policy: Optional[
Dict[str, RetryPolicy]
] = model_group_retry_policy
self.allowed_fails_policy: Optional[AllowedFailsPolicy] = None
if allowed_fails_policy is not None:
@ -782,6 +782,9 @@ class Router:
self.aget_responses = self.factory_function(
litellm.aget_responses, call_type="aget_responses"
)
self.acancel_responses = self.factory_function(
litellm.acancel_responses, call_type="acancel_responses"
)
self.adelete_responses = self.factory_function(
litellm.adelete_responses, call_type="adelete_responses"
)
@ -873,7 +876,6 @@ class Router:
def add_optional_pre_call_checks(
self, optional_pre_call_checks: Optional[OptionalPreCallChecks]
):
if optional_pre_call_checks is not None:
for pre_call_check in optional_pre_call_checks:
_callback: Optional[CustomLogger] = None
@ -1209,10 +1211,7 @@ class Router:
async def _acompletion(
self, model: str, messages: List[Dict[str, str]], **kwargs
) -> Union[
ModelResponse,
CustomStreamWrapper,
]:
) -> Union[ModelResponse, CustomStreamWrapper,]:
"""
- Get an available deployment
- call it with a semaphore over the call
@ -2713,7 +2712,6 @@ class Router:
passthrough_on_no_deployment = kwargs.pop("passthrough_on_no_deployment", False)
function_name = "_ageneric_api_call_with_fallbacks"
try:
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs)
try:
deployment = await self.async_get_available_deployment(
@ -3046,7 +3044,7 @@ class Router:
from litellm.router_utils.common_utils import add_model_file_id_mappings
verbose_router_logger.debug(
f"Inside _acreate_file()- model: {model}; kwargs: {kwargs}"
f"Inside _atext_completion()- model: {model}; kwargs: {kwargs}"
)
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs)
healthy_deployments = await self.async_get_healthy_deployments(
@ -3157,9 +3155,9 @@ class Router:
healthy_deployments=healthy_deployments, responses=responses
)
returned_response = cast(OpenAIFileObject, responses[0])
returned_response._hidden_params["model_file_id_mapping"] = (
model_file_id_mapping
)
returned_response._hidden_params[
"model_file_id_mapping"
] = model_file_id_mapping
return returned_response
except Exception as e:
verbose_router_logger.exception(
@ -3485,6 +3483,7 @@ class Router:
"moderation",
"anthropic_messages",
"aresponses",
"acancel_responses",
"responses",
"aget_responses",
"adelete_responses",
@ -3578,6 +3577,7 @@ class Router:
)
elif call_type in (
"aget_responses",
"acancel_responses",
"adelete_responses",
"alist_input_items",
):
@ -3625,7 +3625,7 @@ class Router:
"""
Initialize the Responses API endpoints on the router.
GET, DELETE Responses API Requests encode the model_id in the response_id, this function decodes the response_id and sets the model to the model_id.
GET, DELETE, CANCEL Responses API Requests encode the model_id in the response_id, this function decodes the response_id and sets the model to the model_id.
"""
from litellm.responses.utils import ResponsesAPIRequestUtils
@ -3720,11 +3720,11 @@ class Router:
if isinstance(e, litellm.ContextWindowExceededError):
if context_window_fallbacks is not None:
context_window_fallback_model_group: Optional[List[str]] = (
self._get_fallback_model_group_from_fallbacks(
fallbacks=context_window_fallbacks,
model_group=model_group,
)
context_window_fallback_model_group: Optional[
List[str]
] = self._get_fallback_model_group_from_fallbacks(
fallbacks=context_window_fallbacks,
model_group=model_group,
)
if context_window_fallback_model_group is None:
raise original_exception
@ -3756,11 +3756,11 @@ class Router:
e.message += "\n{}".format(error_message)
elif isinstance(e, litellm.ContentPolicyViolationError):
if content_policy_fallbacks is not None:
content_policy_fallback_model_group: Optional[List[str]] = (
self._get_fallback_model_group_from_fallbacks(
fallbacks=content_policy_fallbacks,
model_group=model_group,
)
content_policy_fallback_model_group: Optional[
List[str]
] = self._get_fallback_model_group_from_fallbacks(
fallbacks=content_policy_fallbacks,
model_group=model_group,
)
if content_policy_fallback_model_group is None:
raise original_exception
@ -4414,7 +4414,7 @@ class Router:
return tpm_key
except Exception as e:
verbose_router_logger.debug(
verbose_router_logger.exception(
"litellm.router.Router::deployment_callback_on_success(): Exception occured - {}".format(
str(e)
)
@ -4992,26 +4992,26 @@ class Router:
"""
from litellm.router_strategy.auto_router.auto_router import AutoRouter
auto_router_config_path: Optional[str] = (
deployment.litellm_params.auto_router_config_path
)
auto_router_config_path: Optional[
str
] = deployment.litellm_params.auto_router_config_path
auto_router_config: Optional[str] = deployment.litellm_params.auto_router_config
if auto_router_config_path is None and auto_router_config is None:
raise ValueError(
"auto_router_config_path or auto_router_config is required for auto-router deployments. Please set it in the litellm_params"
)
default_model: Optional[str] = (
deployment.litellm_params.auto_router_default_model
)
default_model: Optional[
str
] = deployment.litellm_params.auto_router_default_model
if default_model is None:
raise ValueError(
"auto_router_default_model is required for auto-router deployments. Please set it in the litellm_params"
)
embedding_model: Optional[str] = (
deployment.litellm_params.auto_router_embedding_model
)
embedding_model: Optional[
str
] = deployment.litellm_params.auto_router_embedding_model
if embedding_model is None:
raise ValueError(
"auto_router_embedding_model is required for auto-router deployments. Please set it in the litellm_params"

View file

@ -2327,6 +2327,7 @@ class LlmProviders(str, Enum):
DATABRICKS = "databricks"
EMPOWER = "empower"
GITHUB = "github"
COMPACTIFAI = "compactifai"
CUSTOM = "custom"
LITELLM_PROXY = "litellm_proxy"
HOSTED_VLLM = "hosted_vllm"

View file

@ -6979,6 +6979,8 @@ class ProviderConfigManager:
return litellm.EmpowerChatConfig()
elif litellm.LlmProviders.GITHUB == provider:
return litellm.GithubChatConfig()
elif litellm.LlmProviders.COMPACTIFAI == provider:
return litellm.CompactifAIChatConfig()
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
return litellm.GithubCopilotConfig()
elif (

View file

@ -13477,6 +13477,23 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"bedrock/us.anthropic.claude-3-5-haiku-20241022-v1:0": {
"max_tokens": 8192,
"max_input_tokens": 200000,
"max_output_tokens": 8192,
"input_cost_per_token": 8e-07,
"output_cost_per_token": 4e-06,
"cache_creation_input_token_cost": 1e-06,
"cache_read_input_token_cost": 8e-08,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_assistant_prefill": true,
"supports_pdf_input": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"us.anthropic.claude-3-opus-20240229-v1:0": {
"max_tokens": 4096,
"max_input_tokens": 200000,

View file

@ -590,3 +590,66 @@ class BaseResponsesAPITest(ABC):
assert function_call_item["status"] == "completed", "status value should be preserved"
print("✅ OpenAI Responses API dict input filtering test passed")
@pytest.mark.parametrize("sync_mode", [False, True])
@pytest.mark.flaky(retries=3, delay=2)
@pytest.mark.asyncio
async def test_basic_openai_responses_cancel_endpoint(self, sync_mode):
try:
litellm._turn_on_debug()
litellm.set_verbose = True
base_completion_call_args = self.get_base_completion_call_args()
if sync_mode:
response = litellm.responses(
input="Basic ping", max_output_tokens=20, background=True, **base_completion_call_args
)
# cancel the response
if isinstance(response, ResponsesAPIResponse):
cancel_result = litellm.cancel_responses(
response_id=response.id, **base_completion_call_args
)
assert cancel_result is not None
assert hasattr(cancel_result, "id")
# The actual response structure depends on the provider implementation
assert isinstance(cancel_result, ResponsesAPIResponse)
else:
raise ValueError("response is not a ResponsesAPIResponse")
else:
response = await litellm.aresponses(
input="Basic ping", max_output_tokens=20, background=True, **base_completion_call_args
)
# async cancel the response
if isinstance(response, ResponsesAPIResponse):
cancel_result = await litellm.acancel_responses(
response_id=response.id, **base_completion_call_args
)
assert cancel_result is not None
assert hasattr(cancel_result, "id")
# The actual response structure depends on the provider implementation
assert isinstance(cancel_result, ResponsesAPIResponse)
else:
raise ValueError("response is not a ResponsesAPIResponse")
except Exception as e:
if "Cannot cancel a completed response" in str(e):
pass
else:
raise e
@pytest.mark.parametrize("sync_mode", [False, True])
@pytest.mark.asyncio
async def test_cancel_responses_invalid_response_id(self, sync_mode):
"""Test cancel_responses with invalid response ID should raise appropriate error"""
base_completion_call_args = self.get_base_completion_call_args()
if sync_mode:
with pytest.raises(Exception):
litellm.cancel_responses(
response_id="invalid_response_id_12345", **base_completion_call_args
)
else:
with pytest.raises(Exception):
await litellm.acancel_responses(
response_id="invalid_response_id_12345", **base_completion_call_args
)

View file

@ -34,14 +34,19 @@ class TestAnthropicResponsesAPITest(BaseResponsesAPITest):
}
async def test_basic_openai_responses_delete_endpoint(self, sync_mode=False):
pass
pytest.skip("DELETE responses is not supported for anthropic")
async def test_basic_openai_responses_streaming_delete_endpoint(self, sync_mode=False):
pass
pytest.skip("DELETE responses is not supported for anthropic")
async def test_basic_openai_responses_get_endpoint(self, sync_mode=False):
pass
pytest.skip("GET responses is not supported for anthropic")
async def test_basic_openai_responses_cancel_endpoint(self, sync_mode=False):
pytest.skip("CANCEL responses is not supported for anthropic")
async def test_cancel_responses_invalid_response_id(self, sync_mode=False):
pytest.skip("CANCEL responses is not supported for anthropic")

View file

@ -93,13 +93,20 @@ class TestGoogleAIStudioResponsesAPITest(BaseResponsesAPITest):
}
async def test_basic_openai_responses_delete_endpoint(self, sync_mode=False):
pass
pytest.skip("DELETE responses is not supported for Google AI Studio")
async def test_basic_openai_responses_streaming_delete_endpoint(self, sync_mode=False):
pass
pytest.skip("DELETE responses is not supported for Google AI Studio")
async def test_basic_openai_responses_get_endpoint(self, sync_mode=False):
pass
pytest.skip("GET responses is not supported for Google AI Studio")
async def test_basic_openai_responses_cancel_endpoint(self, sync_mode=False):
pytest.skip("CANCEL responses is not supported for Google AI Studio")
async def test_cancel_responses_invalid_response_id(self, sync_mode=False):
pytest.skip("CANCEL responses is not supported for Google AI Studio")

View file

@ -207,6 +207,7 @@ class DummyCredentials:
("aws_role_name", "dummy_role_name"),
("aws_web_identity_token", "dummy_web_identity_token"),
("aws_sts_endpoint", "dummy_sts_endpoint"),
("aws_external_id", "dummy_external_id"),
],
)
def test_dynamic_aws_params_propagation(model, param_name, param_value):

View file

@ -128,3 +128,69 @@ def test_anthropic_with_responses_api():
previous_response_id="hi",
)
print("anthropic response=", response)
def test_cancel_response():
try:
client = get_test_client()
from litellm.types.llms.openai import ResponsesAPIResponse
response = client.responses.create(
model="gpt-4o", input="just respond with the word 'ping'", background=True
)
print("basic response=", response)
# cancel the response
cancel_response = client.responses.cancel(response.id)
print("CANCEL response=", cancel_response)
# verify cancel response structure
assert hasattr(cancel_response, "id")
# Note: Cancel response returns ResponsesAPIResponse, not DeleteResponseResult
# The actual response structure depends on the provider implementation
assert isinstance(cancel_response, ResponsesAPIResponse)
except Exception as e:
if "Cannot cancel a completed response" in str(e):
pass
else:
raise e
def test_cancel_streaming_response():
try:
client = get_test_client()
from litellm.types.llms.openai import ResponsesAPIResponse
stream = client.responses.create(
model="gpt-4o", input="just respond with the word 'ping'", stream=True, background=True
)
collected_chunks = []
response_id = None
for chunk in stream:
print("stream chunk=", chunk)
collected_chunks.append(chunk)
# Extract response ID from the first chunk that has it
if response_id is None and hasattr(chunk, 'response') and hasattr(chunk.response, 'id'):
response_id = chunk.response.id
assert len(collected_chunks) > 0
# cancel the response if we got a response ID
if response_id:
cancel_response = client.responses.cancel(response_id)
print("CANCEL streaming response=", cancel_response)
assert hasattr(cancel_response, "id")
# Note: Cancel response returns ResponsesAPIResponse, not DeleteResponseResult
# The actual response structure depends on the provider implementation
assert isinstance(cancel_response, ResponsesAPIResponse)
except Exception as e:
if "Cannot cancel a completed response" in str(e):
pass
else:
raise e
def test_cancel_invalid_response_id():
client = get_test_client()
with pytest.raises(Exception):
# Try to cancel a non-existent response ID
client.responses.cancel("invalid_response_id_12345")

View file

@ -10,6 +10,7 @@ from litellm.types.utils import StandardLoggingPayload
class TestS3V2UnitTests:
"""Test that S3 v2 integration only uses safe_dumps and not json.dumps"""
def test_s3_v2_source_code_analysis(self):
"""Test that S3 v2 source code only imports and uses safe_dumps"""
import inspect
@ -18,7 +19,139 @@ class TestS3V2UnitTests:
# Get the source code of the s3_v2 module
source_code = inspect.getsource(s3_v2)
# Verify that json.dumps is not used directly in the code
assert "json.dumps(" not in source_code, \
"S3 v2 should not use json.dumps directly"
assert (
"json.dumps(" not in source_code
), "S3 v2 should not use json.dumps directly"
@patch('asyncio.create_task')
@patch('litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush')
def test_s3_v2_endpoint_url(self, mock_periodic_flush, mock_create_task):
"""testing s3 endpoint url"""
from unittest.mock import AsyncMock, MagicMock
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
# Mock periodic_flush and create_task to prevent async task creation during init
mock_periodic_flush.return_value = None
mock_create_task.return_value = None
# Mock response for all tests
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.raise_for_status = MagicMock()
# Create a test batch logging element
test_element = s3BatchLoggingElement(
s3_object_key="2025-09-14/test-key.json",
payload={"test": "data"},
s3_object_download_filename="test-file.json"
)
# Test 1: Custom endpoint URL with bucket name
s3_logger = S3Logger(
s3_bucket_name="test-bucket",
s3_endpoint_url="https://s3.amazonaws.com",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1"
)
s3_logger.async_httpx_client = AsyncMock()
s3_logger.async_httpx_client.put.return_value = mock_response
asyncio.run(s3_logger.async_upload_data_to_s3(test_element))
call_args = s3_logger.async_httpx_client.put.call_args
assert call_args is not None
url = call_args[0][0]
expected_url = "https://s3.amazonaws.com/test-bucket/2025-09-14/test-key.json"
assert url == expected_url, f"Expected URL {expected_url}, got {url}"
# Test 2: MinIO-compatible endpoint
s3_logger_minio = S3Logger(
s3_bucket_name="litellm-logs",
s3_endpoint_url="https://minio.example.com:9000",
s3_aws_access_key_id="minio-key",
s3_aws_secret_access_key="minio-secret",
s3_region_name="us-east-1"
)
s3_logger_minio.async_httpx_client = AsyncMock()
s3_logger_minio.async_httpx_client.put.return_value = mock_response
asyncio.run(s3_logger_minio.async_upload_data_to_s3(test_element))
call_args_minio = s3_logger_minio.async_httpx_client.put.call_args
assert call_args_minio is not None
url_minio = call_args_minio[0][0]
expected_minio_url = "https://minio.example.com:9000/litellm-logs/2025-09-14/test-key.json"
assert url_minio == expected_minio_url, f"Expected MinIO URL {expected_minio_url}, got {url_minio}"
# Test 3: Custom endpoint without bucket name (should fall back to default)
s3_logger_no_bucket = S3Logger(
s3_endpoint_url="https://s3.amazonaws.com",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1"
)
s3_logger_no_bucket.async_httpx_client = AsyncMock()
s3_logger_no_bucket.async_httpx_client.put.return_value = mock_response
asyncio.run(s3_logger_no_bucket.async_upload_data_to_s3(test_element))
call_args_no_bucket = s3_logger_no_bucket.async_httpx_client.put.call_args
assert call_args_no_bucket is not None
url_no_bucket = call_args_no_bucket[0][0]
# Should use default S3 URL format when bucket is missing (bucket becomes None in URL)
assert "s3.us-east-1.amazonaws.com" in url_no_bucket
assert "https://" in url_no_bucket
# Should not include the custom endpoint since bucket is missing
assert "https://s3.amazonaws.com/" not in url_no_bucket
# Test 4: Sync upload method with custom endpoint
s3_logger_sync = S3Logger(
s3_bucket_name="sync-bucket",
s3_endpoint_url="https://custom.s3.endpoint.com",
s3_aws_access_key_id="sync-key",
s3_aws_secret_access_key="sync-secret",
s3_region_name="us-east-1"
)
mock_sync_client = MagicMock()
mock_sync_client.put.return_value = mock_response
with patch('litellm.integrations.s3_v2._get_httpx_client', return_value=mock_sync_client):
s3_logger_sync.upload_data_to_s3(test_element)
call_args_sync = mock_sync_client.put.call_args
assert call_args_sync is not None
url_sync = call_args_sync[0][0]
expected_sync_url = "https://custom.s3.endpoint.com/sync-bucket/2025-09-14/test-key.json"
assert url_sync == expected_sync_url, f"Expected sync URL {expected_sync_url}, got {url_sync}"
# Test 5: Download method with custom endpoint
s3_logger_download = S3Logger(
s3_bucket_name="download-bucket",
s3_endpoint_url="https://download.s3.endpoint.com",
s3_aws_access_key_id="download-key",
s3_aws_secret_access_key="download-secret",
s3_region_name="us-east-1"
)
mock_download_response = MagicMock()
mock_download_response.status_code = 200
mock_download_response.json = MagicMock(return_value={"downloaded": "data"})
s3_logger_download.async_httpx_client = AsyncMock()
s3_logger_download.async_httpx_client.get.return_value = mock_download_response
result = asyncio.run(s3_logger_download._download_object_from_s3("2025-09-14/download-test-key.json"))
call_args_download = s3_logger_download.async_httpx_client.get.call_args
assert call_args_download is not None
url_download = call_args_download[0][0]
expected_download_url = "https://download.s3.endpoint.com/download-bucket/2025-09-14/download-test-key.json"
assert url_download == expected_download_url, f"Expected download URL {expected_download_url}, got {url_download}"
assert result == {"downloaded": "data"}

View file

@ -293,3 +293,55 @@ class TestAzureResponsesAPIConfig:
litellm_params={"api_version": None},
)
assert result_none_version == expected_url
def test_azure_cancel_response_api_request(self):
"""Test Azure cancel response API request transformation"""
from litellm.types.router import GenericLiteLLMParams
response_id = "resp_test123"
api_base = "https://test.openai.azure.com/openai/responses?api-version=2024-05-01-preview"
litellm_params = GenericLiteLLMParams(api_version="2024-05-01-preview")
headers = {"Authorization": "Bearer test-key"}
url, data = self.config.transform_cancel_response_api_request(
response_id=response_id,
api_base=api_base,
litellm_params=litellm_params,
headers=headers,
)
expected_url = "https://test.openai.azure.com/openai/responses/resp_test123/cancel?api-version=2024-05-01-preview"
assert url == expected_url
assert data == {}
def test_azure_cancel_response_api_response(self):
"""Test Azure cancel response API response transformation"""
from unittest.mock import Mock
from litellm.types.llms.openai import ResponsesAPIResponse
# Mock response
mock_response = Mock()
mock_response.json.return_value = {
"id": "resp_test123",
"object": "response",
"created_at": 1234567890,
"output": [],
"parallel_tool_calls": True,
"tool_choice": "auto",
"tools": [],
"top_p": 1.0,
"status": "cancelled"
}
mock_response.text = "test response"
mock_response.status_code = 200
# Mock logging object
mock_logging_obj = Mock()
result = self.config.transform_cancel_response_api_response(
raw_response=mock_response,
logging_obj=mock_logging_obj,
)
assert isinstance(result, ResponsesAPIResponse)
assert result.id == "resp_test123"

View file

@ -1026,7 +1026,7 @@ def test_auth_with_aws_role_irsa_environment():
def test_auth_with_aws_role_same_role_irsa():
"""Test that when IRSA role matches the requested role, we skip assumption"""
base_llm = BaseAWSLLM()
# Set IRSA environment variables
with patch.dict(os.environ, {
'AWS_ROLE_ARN': 'arn:aws:iam::111111111111:role/LitellmRole',
@ -1037,7 +1037,7 @@ def test_auth_with_aws_role_same_role_irsa():
mock_creds.access_key = 'irsa-access-key'
mock_creds.secret_key = 'irsa-secret-key'
mock_creds.token = 'irsa-session-token'
with patch.object(base_llm, '_auth_with_env_vars', return_value=(mock_creds, None)) as mock_env_auth:
# Call get_credentials instead of _auth_with_aws_role directly
# This tests the full flow
@ -1048,9 +1048,146 @@ def test_auth_with_aws_role_same_role_irsa():
aws_session_name='test-session',
aws_region_name='us-east-1'
)
# Verify it used the env vars auth (no role assumption)
mock_env_auth.assert_called_once()
# Verify the returned credentials
assert creds.access_key == 'irsa-access-key'
def test_assume_role_with_external_id():
"""Test that assume_role STS call includes ExternalId parameter when provided"""
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client
mock_sts_client = MagicMock()
mock_expiry = datetime.now(timezone.utc) + timedelta(hours=1)
mock_sts_response = {
"Credentials": {
"AccessKeyId": "test-access-key",
"SecretAccessKey": "test-secret-key",
"SessionToken": "test-session-token",
"Expiration": mock_expiry,
}
}
mock_sts_client.assume_role.return_value = mock_sts_response
with patch("boto3.client", return_value=mock_sts_client):
# Call _auth_with_aws_role with external ID
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::123456789012:role/ExampleRole",
aws_session_name="test-session",
aws_external_id="UniqueExternalID123"
)
# Verify assume_role was called with ExternalId
mock_sts_client.assume_role.assert_called_once_with(
RoleArn="arn:aws:iam::123456789012:role/ExampleRole",
RoleSessionName="test-session",
ExternalId="UniqueExternalID123"
)
def test_assume_role_without_external_id():
"""Test that assume_role STS call excludes ExternalId parameter when not provided"""
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client
mock_sts_client = MagicMock()
mock_expiry = datetime.now(timezone.utc) + timedelta(hours=1)
mock_sts_response = {
"Credentials": {
"AccessKeyId": "test-access-key",
"SecretAccessKey": "test-secret-key",
"SessionToken": "test-session-token",
"Expiration": mock_expiry,
}
}
mock_sts_client.assume_role.return_value = mock_sts_response
with patch("boto3.client", return_value=mock_sts_client):
# Call _auth_with_aws_role without external ID
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::123456789012:role/ExampleRole",
aws_session_name="test-session"
)
# Verify assume_role was called without ExternalId
mock_sts_client.assume_role.assert_called_once_with(
RoleArn="arn:aws:iam::123456789012:role/ExampleRole",
RoleSessionName="test-session"
)
def test_converse_handler_external_id_extraction():
"""Test that BedrockConverseLLM properly extracts and passes aws_external_id parameter"""
from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM
converse_llm = BedrockConverseLLM()
# Mock get_credentials to capture parameters
def mock_get_credentials(**kwargs):
mock_get_credentials.called_kwargs = kwargs
mock_credentials = MagicMock()
mock_credentials.access_key = "test-access-key"
mock_credentials.secret_key = "test-secret-key"
mock_credentials.token = "test-session-token"
return mock_credentials
with patch.object(converse_llm, 'get_credentials', side_effect=mock_get_credentials):
with patch.object(converse_llm, '_get_aws_region_name', return_value="us-west-2"):
with patch.object(converse_llm, 'get_runtime_endpoint', return_value=("https://test", "https://test")):
with patch('litellm.AmazonConverseConfig') as mock_config:
mock_config.return_value._transform_request.return_value = {"test": "data"}
with patch.object(converse_llm, 'get_request_headers') as mock_headers:
mock_headers.return_value = MagicMock()
mock_headers.return_value.headers = {"Authorization": "test"}
with patch('litellm.llms.custom_httpx.http_handler._get_httpx_client') as mock_client:
mock_http_client = MagicMock()
mock_response = MagicMock()
mock_response.raise_for_status.return_value = None
mock_http_client.post.return_value = mock_response
mock_client.return_value = mock_http_client
# Mock the transform_response method
mock_config.return_value._transform_response.return_value = MagicMock()
# Call completion with aws_external_id in optional_params
optional_params = {
"aws_role_name": "arn:aws:iam::123456789012:role/ExampleRole",
"aws_session_name": "test-session",
"aws_external_id": "TestExternalID123"
}
try:
converse_llm.completion(
model="anthropic.claude-3-sonnet-20240229-v1:0",
messages=[{"role": "user", "content": "Hello"}],
api_base=None,
custom_prompt_dict={},
model_response=MagicMock(),
encoding="utf-8",
logging_obj=MagicMock(),
optional_params=optional_params,
acompletion=False,
timeout=None,
litellm_params={}
)
except Exception:
# We expect this to fail due to mocking, but that's OK
# We just want to verify the parameter extraction
pass
# Verify aws_external_id was extracted and passed to get_credentials
assert hasattr(mock_get_credentials, 'called_kwargs')
assert "aws_external_id" in mock_get_credentials.called_kwargs
assert mock_get_credentials.called_kwargs["aws_external_id"] == "TestExternalID123"

View file

@ -0,0 +1,43 @@
"""Test Bedrock cross-region inference profile model mapping"""
import os
import sys
sys.path.insert(0, os.path.abspath("../../../.."))
from litellm.utils import _get_model_info_helper
from litellm.cost_calculator import completion_cost
from litellm.types.utils import ModelResponse, Usage, Choices, Message
def test_bedrock_cross_region_inference_profile_mapping():
"""Test that bedrock cross-region inference profile model is mapped"""
model = "bedrock/us.anthropic.claude-3-5-haiku-20241022-v1:0"
model_info = _get_model_info_helper(model=model, custom_llm_provider="bedrock")
assert model_info is not None
assert model_info["litellm_provider"] == "bedrock"
assert model_info["input_cost_per_token"] == 8e-07
def test_proxy_cost_calculation_scenario():
"""Test exact GitHub issue scenario: proxy cost calculation"""
model = "litellm_proxy/bedrock/us.anthropic.claude-3-5-haiku-20241022-v1:0"
# Test model info lookup works
model_info = _get_model_info_helper(model=model, custom_llm_provider="litellm_proxy")
assert model_info is not None
# Test cost calculation works
response = ModelResponse(
id="test",
created=1234567890,
model=model,
object="chat.completion",
choices=[Choices(finish_reason="stop", index=0, message=Message(content="Test", role="assistant"))],
usage=Usage(total_tokens=150, prompt_tokens=100, completion_tokens=50),
)
cost = completion_cost(completion_response=response, model=model, custom_llm_provider="litellm_proxy")
expected_cost = (100 * 8e-07) + (50 * 4e-06)
assert cost == expected_cost

View file

@ -0,0 +1,344 @@
import json
import os
import sys
from unittest.mock import AsyncMock, patch
from typing import Optional
import httpx
import pytest
import respx
from respx import MockRouter
import litellm
from litellm import Choices, Message, ModelResponse
@pytest.mark.respx()
def test_compactifai_completion_basic(respx_mock):
"""Test basic CompactifAI completion functionality"""
litellm.disable_aiohttp_transport = True
mock_response = {
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "cai-llama-3-1-8b-slim",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "Hello! How can I help you today?"
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": 9,
"completion_tokens": 12,
"total_tokens": 21
}
}
respx_mock.post("https://api.compactif.ai/v1/chat/completions").respond(
json=mock_response, status_code=200
)
response = litellm.completion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[{"role": "user", "content": "Hello"}],
api_key="test-key"
)
assert response.choices[0].message.content == "Hello! How can I help you today?"
assert response.model == "compactifai/cai-llama-3-1-8b-slim"
assert response.usage.total_tokens == 21
@pytest.mark.respx()
def test_compactifai_completion_streaming(respx_mock):
"""Test CompactifAI streaming completion"""
litellm.disable_aiohttp_transport = True
mock_chunks = [
"data: " + json.dumps({
"id": "chatcmpl-123",
"object": "chat.completion.chunk",
"created": 1677652288,
"model": "cai-llama-3-1-8b-slim",
"choices": [
{
"index": 0,
"delta": {"content": "Hello"},
"finish_reason": None
}
]
}) + "\n\n",
"data: " + json.dumps({
"id": "chatcmpl-123",
"object": "chat.completion.chunk",
"created": 1677652288,
"model": "cai-llama-3-1-8b-slim",
"choices": [
{
"index": 0,
"delta": {"content": "!"},
"finish_reason": "stop"
}
]
}) + "\n\n",
"data: [DONE]\n\n"
]
respx_mock.post("https://api.compactif.ai/v1/chat/completions").respond(
status_code=200,
headers={"content-type": "text/plain"},
content="".join(mock_chunks)
)
response = litellm.completion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[{"role": "user", "content": "Hello"}],
api_key="test-key",
stream=True
)
chunks = list(response)
assert len(chunks) >= 2
assert chunks[0].choices[0].delta.content == "Hello"
@pytest.mark.respx()
def test_compactifai_models_endpoint(respx_mock):
"""Test CompactifAI models listing"""
litellm.disable_aiohttp_transport = True
mock_response = {
"object": "list",
"data": [
{
"id": "cai-llama-3-1-8b-slim",
"object": "model",
"created": 1677610602,
"owned_by": "compactifai"
},
{
"id": "mistral-7b-compressed",
"object": "model",
"created": 1677610602,
"owned_by": "compactifai"
}
]
}
respx_mock.post("https://api.compactif.ai/v1/chat/completions").respond(
json={
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "cai-llama-3-1-8b-slim",
"choices": [{
"index": 0,
"message": {
"role": "assistant",
"content": "Test response"
},
"finish_reason": "stop"
}],
"usage": {
"prompt_tokens": 5,
"completion_tokens": 10,
"total_tokens": 15
}
},
status_code=200
)
# This would be tested if litellm had a models() function
# For now, we'll test that the provider is properly configured
response = litellm.completion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[{"role": "user", "content": "test"}],
api_key="test-key"
)
@pytest.mark.respx()
def test_compactifai_authentication_error(respx_mock):
"""Test CompactifAI authentication error handling"""
litellm.disable_aiohttp_transport = True
mock_error = {
"error": {
"message": "Invalid API key provided",
"type": "invalid_request_error",
"param": None,
"code": "invalid_api_key"
}
}
respx_mock.post("https://api.compactif.ai/v1/chat/completions").respond(
json=mock_error, status_code=401
)
with pytest.raises(litellm.APIConnectionError) as exc_info:
litellm.completion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[{"role": "user", "content": "test"}],
api_key="invalid-key"
)
# Verify the error contains the expected authentication error message
assert "Invalid API key provided" in str(exc_info.value)
@pytest.mark.respx()
def test_compactifai_provider_detection(respx_mock):
"""Test that CompactifAI provider is properly detected from model name"""
from litellm.utils import get_llm_provider
model, provider, dynamic_api_key, api_base = get_llm_provider(
model="compactifai/cai-llama-3-1-8b-slim"
)
assert provider == "compactifai"
assert model == "cai-llama-3-1-8b-slim"
@pytest.mark.respx()
def test_compactifai_with_optional_params(respx_mock):
"""Test CompactifAI with optional parameters like temperature, max_tokens"""
litellm.disable_aiohttp_transport = True
mock_response = {
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "cai-llama-3-1-8b-slim",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "This is a test response with custom parameters."
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": 15,
"completion_tokens": 20,
"total_tokens": 35
}
}
request_mock = respx_mock.post("https://api.compactif.ai/v1/chat/completions").respond(
json=mock_response, status_code=200
)
response = litellm.completion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[{"role": "user", "content": "Hello with params"}],
api_key="test-key",
temperature=0.7,
max_tokens=100,
top_p=0.9
)
assert response.choices[0].message.content == "This is a test response with custom parameters."
# Verify the request was made with correct parameters
assert request_mock.called
request_data = request_mock.calls[0].request.content
parsed_data = json.loads(request_data)
assert parsed_data["temperature"] == 0.7
assert parsed_data["max_tokens"] == 100
assert parsed_data["top_p"] == 0.9
@pytest.mark.respx()
def test_compactifai_headers_authentication(respx_mock):
"""Test that CompactifAI request includes proper authorization headers"""
litellm.disable_aiohttp_transport = True
mock_response = {
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "cai-llama-3-1-8b-slim",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "Test response"
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": 5,
"completion_tokens": 10,
"total_tokens": 15
}
}
request_mock = respx_mock.post("https://api.compactif.ai/v1/chat/completions").respond(
json=mock_response, status_code=200
)
response = litellm.completion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[{"role": "user", "content": "Test auth"}],
api_key="test-api-key-123"
)
assert response.choices[0].message.content == "Test response"
# Verify authorization header was set correctly
assert request_mock.called
request_headers = request_mock.calls[0].request.headers
assert "authorization" in request_headers
assert request_headers["authorization"] == "Bearer test-api-key-123"
@pytest.mark.asyncio
@pytest.mark.respx()
async def test_compactifai_async_completion(respx_mock):
"""Test CompactifAI async completion"""
litellm.disable_aiohttp_transport = True
mock_response = {
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "cai-llama-3-1-8b-slim",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "Async response from CompactifAI"
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": 8,
"completion_tokens": 15,
"total_tokens": 23
}
}
respx_mock.post("https://api.compactif.ai/v1/chat/completions").respond(
json=mock_response, status_code=200
)
response = await litellm.acompletion(
model="compactifai/cai-llama-3-1-8b-slim",
messages=[{"role": "user", "content": "Async test"}],
api_key="test-key"
)
assert response.choices[0].message.content == "Async response from CompactifAI"
assert response.usage.total_tokens == 23

View file

@ -1,4 +1,7 @@
from litellm.llms.vertex_ai.gemini.transformation import check_if_part_exists_in_parts
from litellm.llms.vertex_ai.gemini.transformation import (
check_if_part_exists_in_parts,
_transform_request_body,
)
def test_check_if_part_exists_in_parts():
@ -73,3 +76,82 @@ def test_check_if_part_exists_in_parts_camel_case_snake_case():
}
assert check_if_part_exists_in_parts(parts_mixed, part_mixed_casing)
# Tests for issue #14556: Labels field provider-aware filtering
def test_google_genai_excludes_labels():
"""Test that Google GenAI/AI Studio endpoints exclude labels when custom_llm_provider='gemini'"""
messages = [{"role": "user", "content": "test"}]
optional_params = {"labels": {"project": "test", "team": "ai"}}
litellm_params = {}
result = _transform_request_body(
messages=messages,
model="gemini-2.5-pro",
optional_params=optional_params,
custom_llm_provider="gemini",
litellm_params=litellm_params,
cached_content=None,
)
# Google GenAI/AI Studio should NOT include labels
assert "labels" not in result
assert "contents" in result
def test_vertex_ai_includes_labels():
"""Test that Vertex AI endpoints include labels when custom_llm_provider='vertex_ai'"""
messages = [{"role": "user", "content": "test"}]
optional_params = {"labels": {"project": "test", "team": "ai"}}
litellm_params = {}
result = _transform_request_body(
messages=messages,
model="gemini-2.5-pro",
optional_params=optional_params,
custom_llm_provider="vertex_ai",
litellm_params=litellm_params,
cached_content=None,
)
# Vertex AI SHOULD include labels
assert "labels" in result
assert result["labels"] == {"project": "test", "team": "ai"}
def test_metadata_to_labels_vertex_only():
"""Test that metadata->labels conversion only happens for Vertex AI"""
messages = [{"role": "user", "content": "test"}]
optional_params = {}
litellm_params = {
"metadata": {
"requester_metadata": {
"user": "john_doe",
"project": "test-project"
}
}
}
# Google GenAI/AI Studio should not include labels from metadata
result = _transform_request_body(
messages=messages,
model="gemini-2.5-pro",
optional_params=optional_params.copy(),
custom_llm_provider="gemini",
litellm_params=litellm_params.copy(),
cached_content=None,
)
assert "labels" not in result
# Vertex AI should include labels from metadata
result = _transform_request_body(
messages=messages,
model="gemini-2.5-pro",
optional_params=optional_params.copy(),
custom_llm_provider="vertex_ai",
litellm_params=litellm_params.copy(),
cached_content=None,
)
assert "labels" in result
assert result["labels"] == {"user": "john_doe", "project": "test-project"}

View file

@ -14,7 +14,7 @@ class TestVolcEngineConfig:
supported_params = config.get_supported_openai_params(model="doubao-seed-1.6")
assert "thinking" in supported_params
# Test thinking disabled - should NOT appear in extra_body
# Test thinking disabled - should appear in extra_body
mapped_params = config.map_openai_params(
non_default_params={
"thinking": {"type": "disabled"},
@ -24,8 +24,10 @@ class TestVolcEngineConfig:
drop_params=False,
)
# Fixed: thinking disabled should be omitted from extra_body
assert mapped_params == {}
# Fixed: thinking disabled should appear in extra_body
assert mapped_params == {
"extra_body": {"thinking": {"type": "disabled"}}
}
e2e_mapped_params = get_optional_params(
model="doubao-seed-1.6",
@ -43,7 +45,7 @@ class TestVolcEngineConfig:
def test_thinking_parameter_handling(self):
"""Test comprehensive thinking parameter handling scenarios"""
config = VolcEngineConfig()
# Test 1: thinking enabled - should appear in extra_body
result_enabled = config.map_openai_params(
non_default_params={"thinking": {"type": "enabled"}},
@ -54,38 +56,36 @@ class TestVolcEngineConfig:
assert result_enabled == {
"extra_body": {"thinking": {"type": "enabled"}}
}
# Test 2: thinking None - should appear in extra_body as None
# Test 2: thinking None - should NOT appear in extra_body
result_none = config.map_openai_params(
non_default_params={"thinking": None},
optional_params={},
model="doubao-seed-1.6",
model="doubao-seed-1.6",
drop_params=False,
)
assert result_none == {
"extra_body": {"thinking": None}
}
# Test 3: thinking with custom value - should appear in extra_body
assert result_none == {}
# Test 3: thinking with custom value - should NOT appear in extra_body (invalid value)
result_custom = config.map_openai_params(
non_default_params={"thinking": "custom_mode"},
optional_params={},
model="doubao-seed-1.6",
drop_params=False,
)
assert result_custom == {
"extra_body": {"thinking": "custom_mode"}
}
# Test 4: thinking disabled - should NOT appear in extra_body
assert result_custom == {}
# Test 4: thinking disabled - should appear in extra_body with original structure
result_disabled = config.map_openai_params(
non_default_params={"thinking": {"type": "disabled"}},
optional_params={},
model="doubao-seed-1.6",
drop_params=False,
)
assert result_disabled == {}
assert result_disabled == {
"extra_body": {"thinking": {"type": "disabled"}}
}
# Test 5: No thinking parameter - should return empty dict
result_no_thinking = config.map_openai_params(
non_default_params={},
@ -95,6 +95,24 @@ class TestVolcEngineConfig:
)
assert result_no_thinking == {}
# Test 6: invalid thinking type - should NOT appear in extra_body (invalid type)
result_no_thinking = config.map_openai_params(
non_default_params={"thinking": {"type": "invalid_type"}},
optional_params={},
model="doubao-seed-1.6",
drop_params=False,
)
assert result_no_thinking == {}
# Test 7: invalid thinking type - should NOT appear in extra_body (value is None)
result_no_thinking = config.map_openai_params(
non_default_params={"thinking": {"type": None}},
optional_params={},
model="doubao-seed-1.6",
drop_params=False,
)
assert result_no_thinking == {}
def test_e2e_completion(self):
from openai import OpenAI
@ -131,5 +149,5 @@ class TestVolcEngineConfig:
mock_create.assert_called_once()
print(mock_create.call_args.kwargs)
# Fixed: thinking disabled should NOT appear in extra_body
assert "extra_body" not in mock_create.call_args.kwargs or "thinking" not in mock_create.call_args.kwargs.get("extra_body", {})
# Fixed: thinking disabled should appear in extra_body with original structure
assert "extra_body" in mock_create.call_args.kwargs and "thinking" in mock_create.call_args.kwargs.get("extra_body", {}) and mock_create.call_args.kwargs.get("extra_body", {})["thinking"] == {"type": "disabled"}

View file

@ -342,3 +342,82 @@ async def test_concurrent_initialize_session_managers():
mcp_server._SESSION_MANAGERS_INITIALIZED = original_initialized
mcp_server._session_manager_cm = original_session_cm
mcp_server._sse_session_manager_cm = original_sse_session_cm
@pytest.mark.asyncio
async def test_mcp_routing_with_conflicting_alias_and_group_name():
"""
Tests (GH #14536) where an MCP server alias (e.g., "group/id")
conflicts with an access group name (e.g., "group").
"""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_mcp_servers_in_path,
_get_tools_from_mcp_servers,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.proxy._types import MCPTransport, MCPSpecVersion
except ImportError:
pytest.skip("MCP server not available")
global_mcp_server_manager.registry.clear()
# Create two in-memory servers
specific_server = MCPServer(
server_id="specific_server_id",
name="custom_solutions/user_123",
alias="custom_solutions/user_123",
transport=MCPTransport.http,
spec_version=MCPSpecVersion.jun_2025,
)
other_server = MCPServer(
server_id="other_server_in_group_id",
name="custom_solutions/another_user_456",
alias="custom_solutions/another_user_456",
transport=MCPTransport.http,
spec_version=MCPSpecVersion.jun_2025,
)
global_mcp_server_manager.registry[specific_server.server_id] = specific_server
global_mcp_server_manager.registry[other_server.server_id] = other_server
user_key = UserAPIKeyAuth(api_key="sk-test", team_id="team_custom_solutions")
# Define the request path that triggers the bug
test_path = "/mcp/custom_solutions/user_123/chat/completions"
# This mock will be our "spy" to see which servers are ultimately contacted
mock_get_tools_spy = AsyncMock(return_value=[])
# Mock the function that checks DB for an access group named "custom_solutions"
mock_db_lookup = AsyncMock(return_value=[specific_server.server_id, other_server.server_id])
mock_get_allowed = AsyncMock(return_value=[specific_server.server_id, other_server.server_id])
with patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers",
mock_get_allowed,
), patch(
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler._get_mcp_servers_from_access_groups",
mock_db_lookup,
), patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server",
mock_get_tools_spy,
):
mcp_servers_from_path = _get_mcp_servers_in_path(test_path)
await _get_tools_from_mcp_servers(
user_api_key_auth=user_key,
mcp_servers=mcp_servers_from_path,
mcp_auth_header=None,
)
# Get the list of actual server objects that the orchestrator tried to contact
called_servers = [call.kwargs["server"] for call in mock_get_tools_spy.call_args_list]
assert len(called_servers) == 1, "Should have resolved to exactly one server."
assert (
called_servers[0].server_id == specific_server.server_id
), "Should have contacted the specific server alias, not the group."

View file

@ -28,6 +28,7 @@ from litellm.proxy.auth.auth_checks import (
_can_object_call_vector_stores,
get_user_object,
vector_store_access_check,
_get_team_db_check,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.utils import get_utc_datetime
@ -192,6 +193,64 @@ async def test_default_internal_user_params_with_get_user_object(monkeypatch):
assert creation_args["user_role"] == "internal_user"
@pytest.mark.asyncio
@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock)
async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeypatch):
"""
Test that _get_team_db_check correctly calls the `new_team` function
when a team does not exist and upsert is enabled.
"""
mock_prisma_client = MagicMock()
mock_db = AsyncMock()
mock_prisma_client.db = mock_db
mock_prisma_client.db.litellm_teamtable.find_unique.return_value = None
# Define what our mocked `new_team` function should return
team_id_to_create = "new-jwt-team"
mock_new_team.return_value = {"team_id": team_id_to_create, "max_budget": 123.45}
await _get_team_db_check(
team_id=team_id_to_create,
prisma_client=mock_prisma_client,
team_id_upsert=True,
)
# Verify that our mocked `new_team` function was called exactly once
mock_new_team.assert_called_once()
call_args = mock_new_team.call_args[1]
data_arg = call_args["data"]
# Verify that `new_team` was called with the correct team_id and that
# `max_budget` was None, as our function's job is to delegate, not to set defaults.
assert data_arg.team_id == team_id_to_create
assert data_arg.max_budget is None
@pytest.mark.asyncio
@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock)
async def test_get_team_db_check_does_not_call_new_team_if_exists(mock_new_team, monkeypatch):
"""
Test that _get_team_db_check does NOT call the `new_team` function
if the team already exists in the database.
"""
mock_prisma_client = MagicMock()
mock_db = AsyncMock()
mock_prisma_client.db = mock_db
mock_prisma_client.db.litellm_teamtable.find_unique.return_value = MagicMock()
team_id_to_find = "existing-jwt-team"
await _get_team_db_check(
team_id=team_id_to_find,
prisma_client=mock_prisma_client,
team_id_upsert=True,
)
# Verify that `new_team` was NEVER called, because the team was found.
mock_new_team.assert_not_called()
# Vector Store Auth Check Tests

View file

@ -7,6 +7,7 @@ from litellm.proxy._types import LitellmUserRoles, NewUserRequest, ProxyExceptio
from litellm.proxy.management_endpoints.scim.scim_v2 import (
UserProvisionerHelpers,
_handle_team_membership_changes,
create_group,
create_user,
get_service_provider_config,
patch_user,
@ -910,4 +911,282 @@ async def test_update_group_e2e(mocker):
assert len(result.members) == 3
# Verify SCIM transformation was called with updated team
ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with(updated_team)
ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with(updated_team)
@pytest.mark.asyncio
async def test_create_group_with_nonexistent_users_creates_users(mocker):
"""
Test that creating a group with non-existent users creates those users.
This tests the scenario: Group Push ['new user', existing users...]
"""
# Test data
group_id = "test-group-123"
scim_group = SCIMGroup(
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
id=group_id,
displayName="Test Group",
members=[
SCIMMember(value="existing-user", display="Existing User"), # This user exists
SCIMMember(value="new-user-1", display="New User 1"), # This user doesn't exist
SCIMMember(value="new-user-2", display="New User 2"), # This user doesn't exist
]
)
#########################################################
# We expect new-user-1 and new-user-2 to be created
#########################################################
# Mock prisma client
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_teamtable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
# Mock team operations - team doesn't exist yet
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
# Mock user lookup - only existing-user exists
def mock_user_lookup(where):
user_id = where["user_id"]
if user_id == "existing-user":
mock_user = mocker.MagicMock()
mock_user.user_id = user_id
return mock_user
return None # new-user-1 and new-user-2 don't exist
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
# Mock dependencies
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client)
)
# Mock new_user function to track user creation
mock_new_user = mocker.patch(
"litellm.proxy.management_endpoints.internal_user_endpoints.new_user",
AsyncMock()
)
# Mock created users return values
def mock_new_user_side_effect(data):
from litellm.proxy._types import NewUserResponse
return NewUserResponse(
key="sk-test-key-" + data.user_id, # Required field from GenerateKeyResponse
user_id=data.user_id,
user_email=data.user_email,
metadata=data.metadata,
teams=data.teams,
user_role=data.user_role
)
mock_new_user.side_effect = mock_new_user_side_effect
# Mock new_team function
mock_created_team = mocker.MagicMock()
mock_created_team.team_id = group_id
mock_created_team.team_alias = "Test Group"
mock_new_team = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.new_team",
AsyncMock(return_value=mock_created_team)
)
# Mock SCIM transformation
expected_scim_response = SCIMGroup(
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
id=group_id,
displayName="Test Group",
members=[
SCIMMember(value="existing-user", display="existing-user"),
SCIMMember(value="new-user-1", display="new-user-1"),
SCIMMember(value="new-user-2", display="new-user-2")
]
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group",
AsyncMock(return_value=expected_scim_response)
)
# Execute the create_group function
result = await create_group(group=scim_group)
#########################################################
# Assert that new-user-1 and new-user-2 were created
#########################################################
# Verify that new_user was called exactly twice (for new-user-1 and new-user-2)
assert mock_new_user.call_count == 2
# Check the user creation calls
created_user_ids = set()
for call in mock_new_user.call_args_list:
user_request = call.kwargs["data"]
created_user_ids.add(user_request.user_id)
assert user_request.metadata["created_via"] == "scim_group_membership"
assert user_request.user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
assert user_request.auto_create_key is False
assert user_request.teams == [] # Teams added separately
assert created_user_ids == {"new-user-1", "new-user-2"}
# Verify team creation was called with all members (existing + created)
mock_new_team.assert_called_once()
team_request = mock_new_team.call_args.kwargs["data"]
assert team_request.team_id == group_id
assert team_request.team_alias == "Test Group"
# Verify all members are in the team (existing + newly created)
member_user_ids = {member.user_id for member in team_request.members_with_roles}
assert member_user_ids == {"existing-user", "new-user-1", "new-user-2"}
# Verify response
assert result.id == group_id
assert result.displayName == "Test Group"
assert len(result.members) == 3
@pytest.mark.asyncio
async def test_update_group_with_nonexistent_users_creates_users(mocker):
"""
Test that updating a group with non-existent users creates those users.
This tests the scenario where a group is updated with members that don't exist in user table.
"""
# Test data
group_id = "existing-group-456"
# Mock existing team
mock_existing_team = mocker.MagicMock()
mock_existing_team.team_id = group_id
mock_existing_team.team_alias = "Old Group Name"
mock_existing_team.members = ["old-user"]
mock_existing_team.members_with_roles = [{"user_id": "old-user", "role": "user"}]
mock_existing_team.metadata = {"existing": "data"}
# SCIM group update request
scim_group_update = SCIMGroup(
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
id=group_id,
displayName="Updated Group Name",
members=[
SCIMMember(value="existing-user", display="Existing User"), # This user exists
SCIMMember(value="new-user-3", display="New User 3"), # This user doesn't exist
SCIMMember(value="new-user-4", display="New User 4"), # This user doesn't exist
]
)
# Mock prisma client
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_teamtable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
# Mock team operations
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team)
# Mock updated team response
mock_updated_team = mocker.MagicMock()
mock_updated_team.team_id = group_id
mock_updated_team.team_alias = "Updated Group Name"
mock_updated_team.members = ["existing-user", "new-user-3", "new-user-4"]
mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team)
# Mock user lookup - only existing-user exists
def mock_user_lookup(where):
user_id = where["user_id"]
if user_id == "existing-user":
mock_user = mocker.MagicMock()
mock_user.user_id = user_id
return mock_user
return None # new-user-3 and new-user-4 don't exist
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
# Mock dependencies
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client)
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._check_team_exists",
AsyncMock(return_value=mock_existing_team)
)
# Mock new_user function to track user creation
mock_new_user = mocker.patch(
"litellm.proxy.management_endpoints.internal_user_endpoints.new_user",
AsyncMock()
)
# Mock created users return values
def mock_new_user_side_effect(data):
from litellm.proxy._types import NewUserResponse
return NewUserResponse(
key="sk-test-key-" + data.user_id, # Required field from GenerateKeyResponse
user_id=data.user_id,
user_email=data.user_email,
metadata=data.metadata,
teams=data.teams,
user_role=data.user_role
)
mock_new_user.side_effect = mock_new_user_side_effect
# Mock group membership changes
mock_handle_group_membership_changes = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._handle_group_membership_changes",
AsyncMock()
)
# Mock SCIM transformation
expected_scim_response = SCIMGroup(
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
id=group_id,
displayName="Updated Group Name",
members=[
SCIMMember(value="existing-user", display="existing-user"),
SCIMMember(value="new-user-3", display="new-user-3"),
SCIMMember(value="new-user-4", display="new-user-4")
]
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group",
AsyncMock(return_value=expected_scim_response)
)
# Execute the update_group function
result = await update_group(group_id=group_id, group=scim_group_update)
# Verify that new_user was called exactly twice (for new-user-3 and new-user-4)
assert mock_new_user.call_count == 2
# Check the user creation calls
created_user_ids = set()
for call in mock_new_user.call_args_list:
user_request = call.kwargs["data"]
created_user_ids.add(user_request.user_id)
assert user_request.metadata["created_via"] == "scim_group_membership"
assert user_request.user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
assert user_request.auto_create_key is False
assert user_request.teams == [] # Teams added separately
assert created_user_ids == {"new-user-3", "new-user-4"}
# Verify team update was called
mock_prisma_client.db.litellm_teamtable.update.assert_called_once()
update_call = mock_prisma_client.db.litellm_teamtable.update.call_args
assert update_call[1]["where"]["team_id"] == group_id
assert update_call[1]["data"]["team_alias"] == "Updated Group Name"
# Verify group membership changes were handled with all members (existing + created)
mock_handle_group_membership_changes.assert_called_once()
membership_call = mock_handle_group_membership_changes.call_args
assert membership_call[1]["group_id"] == group_id
assert membership_call[1]["final_members"] == {"existing-user", "new-user-3", "new-user-4"}
# Verify response
assert result.id == group_id
assert result.displayName == "Updated Group Name"
assert len(result.members) == 3

View file

@ -40,6 +40,7 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
data: mcpServers,
isLoading: isLoadingServers,
refetch,
dataUpdatedAt,
} = useQuery({
queryKey: ["mcpServers"],
queryFn: () => {
@ -47,7 +48,7 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
return fetchMCPServers(accessToken)
},
enabled: !!accessToken,
}) as { data: MCPServer[]; isLoading: boolean; refetch: () => void }
}) as { data: MCPServer[]; isLoading: boolean; refetch: () => void; dataUpdatedAt: number }
// state
const [serverIdToDelete, setServerToDelete] = useState<string | null>(null)
@ -117,11 +118,10 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
setFilteredServers(filtered)
}
// Initial and effect-based filtering
// Initial and effect-based filtering (trigger on query data updates)
useEffect(() => {
filterServers(selectedTeam, selectedMcpAccessGroup)
// eslint-disable-next-line
}, [mcpServers])
}, [dataUpdatedAt])
const columns = React.useMemo(
() =>