mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge branch 'BerriAI:main' into wandb-inference
This commit is contained in:
commit
a9667e5930
52 changed files with 2762 additions and 248 deletions
|
|
@ -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) | ✅ | ✅ | ✅ | ✅ |
|
||||
|
|
|
|||
223
docs/my-website/docs/providers/compactifai.md
Normal file
223
docs/my-website/docs/providers/compactifai.md
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -453,6 +453,7 @@ const sidebars = {
|
|||
"providers/elevenlabs",
|
||||
"providers/fireworks_ai",
|
||||
"providers/clarifai",
|
||||
"providers/compactifai",
|
||||
"providers/vllm",
|
||||
"providers/llamafile",
|
||||
"providers/infinity",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 #######
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 ###
|
||||
|
|
|
|||
1
litellm/llms/compactifai/__init__.py
Normal file
1
litellm/llms/compactifai/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
# CompactifAI provider for LiteLLM
|
||||
1
litellm/llms/compactifai/chat/__init__.py
Normal file
1
litellm/llms/compactifai/chat/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
# CompactifAI chat completions
|
||||
100
litellm/llms/compactifai/chat/transformation.py
Normal file
100
litellm/llms/compactifai/chat/transformation.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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")
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
@ -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"}
|
||||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
344
tests/test_litellm/llms/compactifai/test_compactifai.py
Normal file
344
tests/test_litellm/llms/compactifai/test_compactifai.py
Normal 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
|
||||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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(
|
||||
() =>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue