diff --git a/README.md b/README.md
index f1e13c00d1..46acf6cef0 100644
--- a/README.md
+++ b/README.md
@@ -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) | ✅ | ✅ | ✅ | ✅ |
diff --git a/docs/my-website/docs/providers/compactifai.md b/docs/my-website/docs/providers/compactifai.md
new file mode 100644
index 0000000000..0e6e8f4ed3
--- /dev/null
+++ b/docs/my-website/docs/providers/compactifai.md
@@ -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
+
+
+
+
+```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)
+```
+
+
+
+
+```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
+```
+
+
+
+
+## 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)
\ No newline at end of file
diff --git a/docs/my-website/docs/proxy/team_budgets.md b/docs/my-website/docs/proxy/team_budgets.md
index 66ba679c65..3847406641 100644
--- a/docs/my-website/docs/proxy/team_budgets.md
+++ b/docs/my-website/docs/proxy/team_budgets.md
@@ -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
diff --git a/docs/my-website/release_notes/v1.77.2-stable/index.md b/docs/my-website/release_notes/v1.77.2-stable/index.md
index cdbe6595fe..6d54db84df 100644
--- a/docs/my-website/release_notes/v1.77.2-stable/index.md
+++ b/docs/my-website/release_notes/v1.77.2-stable/index.md
@@ -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.
+
+:::
+
``` showLineNumbers title="docker run litellm"
-docker run \
--e STORE_MODEL_IN_DB=True \
--p 4000:4000 \
-ghcr.io/berriai/litellm:v1.77.2
```
``` showLineNumbers title="pip install litellm"
-pip install litellm==1.77.2
```
diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js
index 64de5c2927..f3bab0219f 100644
--- a/docs/my-website/sidebars.js
+++ b/docs/my-website/sidebars.js
@@ -453,6 +453,7 @@ const sidebars = {
"providers/elevenlabs",
"providers/fireworks_ai",
"providers/clarifai",
+ "providers/compactifai",
"providers/vllm",
"providers/llamafile",
"providers/infinity",
diff --git a/litellm/__init__.py b/litellm/__init__.py
index 736cd2c2fe..f5628d5486 100644
--- a/litellm/__init__.py
+++ b/litellm/__init__.py
@@ -1023,6 +1023,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
diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py
index 7cbb2afcb2..414ccb7ab8 100644
--- a/litellm/litellm_core_utils/get_llm_provider_logic.py
+++ b/litellm/litellm_core_utils/get_llm_provider_logic.py
@@ -372,6 +372,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:
diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py
index 2adddd52e7..65f49cf08b 100644
--- a/litellm/litellm_core_utils/prompt_templates/factory.py
+++ b/litellm/litellm_core_utils/prompt_templates/factory.py
@@ -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
)
diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py
index ce196757f9..0ddf8896fd 100644
--- a/litellm/llms/bedrock/base_aws_llm.py
+++ b/litellm/llms/bedrock/base_aws_llm.py
@@ -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)
diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py
index 15a5002f0e..54c603e596 100644
--- a/litellm/llms/bedrock/chat/converse_handler.py
+++ b/litellm/llms/bedrock/chat/converse_handler.py
@@ -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 ###
diff --git a/litellm/llms/compactifai/__init__.py b/litellm/llms/compactifai/__init__.py
new file mode 100644
index 0000000000..16b0c04cda
--- /dev/null
+++ b/litellm/llms/compactifai/__init__.py
@@ -0,0 +1 @@
+# CompactifAI provider for LiteLLM
\ No newline at end of file
diff --git a/litellm/llms/compactifai/chat/__init__.py b/litellm/llms/compactifai/chat/__init__.py
new file mode 100644
index 0000000000..d1a4463166
--- /dev/null
+++ b/litellm/llms/compactifai/chat/__init__.py
@@ -0,0 +1 @@
+# CompactifAI chat completions
\ No newline at end of file
diff --git a/litellm/llms/compactifai/chat/transformation.py b/litellm/llms/compactifai/chat/transformation.py
new file mode 100644
index 0000000000..5cb8cd9a4a
--- /dev/null
+++ b/litellm/llms/compactifai/chat/transformation.py
@@ -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,
+ )
\ No newline at end of file
diff --git a/litellm/llms/volcengine/chat/transformation.py b/litellm/llms/volcengine/chat/transformation.py
index 216570a1ab..6df1cd3826 100644
--- a/litellm/llms/volcengine/chat/transformation.py
+++ b/litellm/llms/volcengine/chat/transformation.py
@@ -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
diff --git a/litellm/main.py b/litellm/main.py
index bfae912cca..2f24f8b3be 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -2549,6 +2549,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(
diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py
index d0461f91e9..51c19beb78 100644
--- a/litellm/proxy/_experimental/mcp_server/server.py
+++ b/litellm/proxy/_experimental/mcp_server/server.py
@@ -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:
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index 4bd539ede4..2ef67c507b 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -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",
diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py
index f1ee2ce43f..be9e494042 100644
--- a/litellm/proxy/auth/auth_checks.py
+++ b/litellm/proxy/auth/auth_checks.py
@@ -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
diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py
index 7b0df5a562..9766386344 100644
--- a/litellm/proxy/management_endpoints/team_endpoints.py
+++ b/litellm/proxy/management_endpoints/team_endpoints.py
@@ -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
diff --git a/litellm/types/utils.py b/litellm/types/utils.py
index 289291bde0..9ed66003e7 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -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"
diff --git a/litellm/utils.py b/litellm/utils.py
index d279ce5c14..e3dba0c7e4 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -6954,6 +6954,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 (
diff --git a/tests/llm_responses_api_testing/base_responses_api.py b/tests/llm_responses_api_testing/base_responses_api.py
index 5ed4fbbb7b..8436f130e1 100644
--- a/tests/llm_responses_api_testing/base_responses_api.py
+++ b/tests/llm_responses_api_testing/base_responses_api.py
@@ -595,41 +595,47 @@ class BaseResponsesAPITest(ABC):
@pytest.mark.flaky(retries=3, delay=2)
@pytest.mark.asyncio
async def test_basic_openai_responses_cancel_endpoint(self, sync_mode):
- 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
+ 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
)
- 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)
+ # 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:
- raise ValueError("response is not a ResponsesAPIResponse")
+ 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
diff --git a/tests/llm_responses_api_testing/test_anthropic_responses_api.py b/tests/llm_responses_api_testing/test_anthropic_responses_api.py
index 8f7a96a016..d633cd0f1d 100644
--- a/tests/llm_responses_api_testing/test_anthropic_responses_api.py
+++ b/tests/llm_responses_api_testing/test_anthropic_responses_api.py
@@ -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")
diff --git a/tests/llm_responses_api_testing/test_google_ai_studio_responses_api.py b/tests/llm_responses_api_testing/test_google_ai_studio_responses_api.py
index 81daaea238..203ee252b3 100644
--- a/tests/llm_responses_api_testing/test_google_ai_studio_responses_api.py
+++ b/tests/llm_responses_api_testing/test_google_ai_studio_responses_api.py
@@ -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")
+
diff --git a/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py b/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py
index 7220ffbb2c..06a3086857 100644
--- a/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py
+++ b/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py
@@ -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):
diff --git a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py b/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py
index 7e7def0ee0..de60881820 100644
--- a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py
+++ b/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py
@@ -131,50 +131,62 @@ def test_anthropic_with_responses_api():
def test_cancel_response():
- 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)
+ 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)
-
-
-def test_cancel_streaming_response():
- 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)
+ # 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():
diff --git a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py
index 5effa6fa01..f5856cd12d 100644
--- a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py
+++ b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py
@@ -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"
diff --git a/tests/test_litellm/llms/compactifai/test_compactifai.py b/tests/test_litellm/llms/compactifai/test_compactifai.py
new file mode 100644
index 0000000000..99b8acc3dc
--- /dev/null
+++ b/tests/test_litellm/llms/compactifai/test_compactifai.py
@@ -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
\ No newline at end of file
diff --git a/tests/test_litellm/llms/volcengine/test_volcengine.py b/tests/test_litellm/llms/volcengine/test_volcengine.py
index 5931791419..f43167efa3 100644
--- a/tests/test_litellm/llms/volcengine/test_volcengine.py
+++ b/tests/test_litellm/llms/volcengine/test_volcengine.py
@@ -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"}
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py
index 42c64c1581..088438556f 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py
@@ -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."
diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py
index eb26eb776f..9a50986a1b 100644
--- a/tests/test_litellm/proxy/auth/test_auth_checks.py
+++ b/tests/test_litellm/proxy/auth/test_auth_checks.py
@@ -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
diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py
index 5cbd602268..230e251a5d 100644
--- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py
@@ -971,8 +971,9 @@ async def test_create_group_with_nonexistent_users_creates_users(mocker):
# Mock created users return values
def mock_new_user_side_effect(data):
- from litellm.proxy._types import LiteLLM_UserTable
- return LiteLLM_UserTable(
+ 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,
@@ -1121,8 +1122,9 @@ async def test_update_group_with_nonexistent_users_creates_users(mocker):
# Mock created users return values
def mock_new_user_side_effect(data):
- from litellm.proxy._types import LiteLLM_UserTable
- return LiteLLM_UserTable(
+ 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,
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx
index 2b95d27a6f..4253d13837 100644
--- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx
+++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx
@@ -40,6 +40,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID })
data: mcpServers,
isLoading: isLoadingServers,
refetch,
+ dataUpdatedAt,
} = useQuery({
queryKey: ["mcpServers"],
queryFn: () => {
@@ -47,7 +48,7 @@ const MCPServers: React.FC = ({ 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(null)
@@ -117,11 +118,10 @@ const MCPServers: React.FC = ({ 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(
() =>