Merge main into fix/add-pytest-postgresql-dependency

Resolved poetry.lock conflict by regenerating with Poetry 2.3.2.

Co-Authored-By: Claude Sonnet 4.5 <noreply@anthropic.com>
This commit is contained in:
Julio Quinteros Pro 2026-02-16 12:03:20 -03:00
commit af9b6f6e0d
62 changed files with 2977 additions and 444 deletions

View file

@ -1,22 +1,121 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# OpenAI Agents SDK
The [OpenAI Agents SDK](https://github.com/openai/openai-agents-python) is a lightweight framework for building multi-agent workflows.
It includes an official LiteLLM extension that lets you use any of the 100+ supported providers (Anthropic, Gemini, Mistral, Bedrock, etc.)
Use OpenAI Agents SDK with any LLM provider through LiteLLM Proxy.
The [OpenAI Agents SDK](https://github.com/openai/openai-agents-python) is a lightweight framework for building multi-agent workflows. It includes an official LiteLLM extension that lets you use any of the 100+ supported providers.
## Quick Start
### 1. Install Dependencies
```bash
pip install "openai-agents[litellm]"
```
### 2. Add Model to Config
```yaml title="config.yaml"
model_list:
- model_name: gpt-4o
litellm_params:
model: "openai/gpt-4o"
api_key: "os.environ/OPENAI_API_KEY"
- model_name: claude-sonnet
litellm_params:
model: "anthropic/claude-3-5-sonnet-20241022"
api_key: "os.environ/ANTHROPIC_API_KEY"
- model_name: gemini-pro
litellm_params:
model: "gemini/gemini-2.0-flash-exp"
api_key: "os.environ/GEMINI_API_KEY"
```
### 3. Start LiteLLM Proxy
```bash
litellm --config config.yaml
```
### 4. Use with Proxy
<Tabs>
<TabItem value="proxy" label="Via Proxy">
```python
from agents import Agent, Runner
from agents.extensions.models.litellm_model import LitellmModel
# Point to LiteLLM proxy
agent = Agent(
name="Assistant",
instructions="You are a helpful assistant.",
model=LitellmModel(model="provider/model-name")
model=LitellmModel(
model="claude-sonnet", # Model from config.yaml
api_key="sk-1234", # LiteLLM API key
base_url="http://localhost:4000"
)
)
result = Runner.run_sync(agent, "your_prompt_here")
print("Result:", result.final_output)
result = await Runner.run(agent, "What is LiteLLM?")
print(result.final_output)
```
- [GitHub](https://github.com/openai/openai-agents-python)
- [LiteLLM Extension Docs](https://openai.github.io/openai-agents-python/ref/extensions/litellm/)
</TabItem>
<TabItem value="direct" label="Direct (No Proxy)">
```python
from agents import Agent, Runner
from agents.extensions.models.litellm_model import LitellmModel
# Use any provider directly
agent = Agent(
name="Assistant",
instructions="You are a helpful assistant.",
model=LitellmModel(
model="anthropic/claude-3-5-sonnet-20241022",
api_key="your-anthropic-key"
)
)
result = await Runner.run(agent, "What is LiteLLM?")
print(result.final_output)
```
</TabItem>
</Tabs>
## Track Usage
Enable usage tracking to monitor token consumption:
```python
from agents import Agent, ModelSettings
from agents.extensions.models.litellm_model import LitellmModel
agent = Agent(
name="Assistant",
model=LitellmModel(model="claude-sonnet", api_key="sk-1234"),
model_settings=ModelSettings(include_usage=True)
)
result = await Runner.run(agent, "Hello")
print(result.context_wrapper.usage) # Token counts
```
## Environment Variables
| Variable | Value | Description |
|----------|-------|-------------|
| `LITELLM_BASE_URL` | `http://localhost:4000` | LiteLLM proxy URL |
| `LITELLM_API_KEY` | `sk-1234` | Your LiteLLM API key |
## Related Resources
- [OpenAI Agents SDK Documentation](https://openai.github.io/openai-agents-python/)
- [LiteLLM Extension Docs](https://openai.github.io/openai-agents-python/models/litellm/)
- [LiteLLM Proxy Quick Start](../proxy/quick_start)

View file

@ -769,6 +769,7 @@ router_settings:
| LITELM_ENVIRONMENT | Environment of LiteLLM Instance, used by logging services. Currently only used by DeepEval.
| LITELLM_KEY_ROTATION_ENABLED | Enable auto-key rotation for LiteLLM (boolean). Default is false.
| LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS | Interval in seconds for how often to run job that auto-rotates keys. Default is 86400 (24 hours).
| LITELLM_KEY_ROTATION_GRACE_PERIOD | Duration to keep old key valid after rotation (e.g. "24h", "2d"). Default is empty (immediate revoke). Used for scheduled rotations and as fallback when not specified in regenerate request.
| LITELLM_LICENSE | License key for LiteLLM usage
| LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS | Set to `True` to use the local bundled Anthropic beta headers config only, disabling remote fetching. Default is `False`
| LITELLM_LOCAL_MODEL_COST_MAP | Local configuration for model cost mapping in LiteLLM

View file

@ -1338,6 +1338,7 @@ litellm_settings:
s3_aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY # AWS Secret Access Key for S3
s3_path: my-test-path # [OPTIONAL] set path in bucket you want to write logs to
s3_endpoint_url: https://s3.amazonaws.com # [OPTIONAL] S3 endpoint URL, if you want to use Backblaze/cloudflare s3 buckets
s3_use_virtual_hosted_style: false # [OPTIONAL] use virtual-hosted-style URLs (bucket.endpoint/key) instead of path-style (endpoint/bucket/key). Useful for S3-compatible services like MinIO
s3_strip_base64_files: false # [OPTIONAL] remove base64 files before storing in s3
```

View file

@ -549,11 +549,14 @@ curl 'http://localhost:4000/key/sk-1234/regenerate' \
"models": [
"gpt-4",
"gpt-3.5-turbo"
]
],
"grace_period": "48h"
}'
```
**Grace period (optional)**: Set `grace_period` (e.g. `"24h"`, `"2d"`, `"1w"`) to keep the old key valid for a transitional period. Both old and new keys work until the grace period elapses, enabling seamless cutover without production downtime. Omitted or empty = immediate revoke. Can also be set via `LITELLM_KEY_ROTATION_GRACE_PERIOD` env var for scheduled rotations.
**Read More**
- [Write rotated keys to secrets manager](https://docs.litellm.ai/docs/secret#aws-secret-manager)
@ -640,11 +643,13 @@ Set these environment variables when starting the proxy:
|----------|-------------|---------|
| `LITELLM_KEY_ROTATION_ENABLED` | Enable the rotation worker | `false` |
| `LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS` | How often to scan for keys to rotate (in seconds) | `86400` (24 hours) |
| `LITELLM_KEY_ROTATION_GRACE_PERIOD` | Duration to keep old key valid after rotation (e.g. `24h`, `2d`) | `""` (immediate revoke) |
**Example:**
```bash
export LITELLM_KEY_ROTATION_ENABLED=true
export LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS=3600 # Check every hour
export LITELLM_KEY_ROTATION_GRACE_PERIOD=48h # Keep old key valid for 48h during cutover
litellm --config config.yaml
```

View file

@ -176,6 +176,7 @@ const sidebars = {
"tutorials/copilotkit_sdk",
"tutorials/google_adk",
"tutorials/livekit_xai_realtime",
"projects/openai-agents"
]
},

View file

@ -4,7 +4,7 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if t
from litellm._uuid import uuid
from datetime import datetime
from typing import TYPE_CHECKING, Optional, cast
from typing import TYPE_CHECKING, Optional
from litellm._logging import verbose_proxy_logger
@ -35,14 +35,11 @@ class CheckBatchCost:
- if not, return False
- if so, return True
"""
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
from litellm.batches.batch_utils import (
_get_file_content_as_dictionary,
calculate_batch_cost_and_usage,
)
from litellm.files.main import afile_content
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.proxy.openai_files_endpoints.common_utils import (
@ -102,27 +99,29 @@ class CheckBatchCost:
continue
## RETRIEVE THE BATCH JOB OUTPUT FILE
managed_files_obj = cast(
Optional[_PROXY_LiteLLMManagedFiles],
self.proxy_logging_obj.get_proxy_hook("managed_files"),
)
if (
response.status == "completed"
and response.output_file_id is not None
and managed_files_obj is not None
):
verbose_proxy_logger.info(
f"Batch ID: {batch_id} is complete, tracking cost and usage"
)
# track cost
model_file_id_mapping = {
response.output_file_id: {model_id: response.output_file_id}
}
_file_content = await managed_files_obj.afile_content(
file_id=response.output_file_id,
litellm_parent_otel_span=None,
llm_router=self.llm_router,
model_file_id_mapping=model_file_id_mapping,
# This background job runs as default_user_id, so going through the HTTP endpoint
# would trigger check_managed_file_id_access and get 403. Instead, extract the raw
# provider file ID and call afile_content directly with deployment credentials.
raw_output_file_id = response.output_file_id
decoded = _is_base64_encoded_unified_file_id(raw_output_file_id)
if decoded:
try:
raw_output_file_id = decoded.split("llm_output_file_id,")[1].split(";")[0]
except (IndexError, AttributeError):
pass
credentials = self.llm_router.get_deployment_credentials_with_provider(model_id) or {}
_file_content = await afile_content(
file_id=raw_output_file_id,
**credentials,
)
file_content_as_dict = _get_file_content_as_dictionary(
@ -143,11 +142,15 @@ class CheckBatchCost:
custom_llm_provider=custom_llm_provider,
)
# Pass deployment model_info so custom batch pricing
# (input_cost_per_token_batches etc.) is used for cost calc
deployment_model_info = deployment_info.model_info.model_dump() if deployment_info.model_info else {}
batch_cost, batch_usage, batch_models = (
await calculate_batch_cost_and_usage(
file_content_dictionary=file_content_as_dict,
custom_llm_provider=llm_provider, # type: ignore
model_name=model_name,
model_info=deployment_model_info,
)
)
logging_obj = LiteLLMLogging(

View file

@ -230,12 +230,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
if managed_file:
return managed_file.created_by == user_id
return False
raise HTTPException(
status_code=404,
detail=f"File not found: {unified_file_id}",
)
async def can_user_call_unified_object_id(
self, unified_object_id: str, user_api_key_dict: UserAPIKeyAuth
) -> bool:
## check if the user has access to the unified object id
## check if the user has access to the unified object id
user_id = user_api_key_dict.user_id
managed_object = (
@ -246,7 +248,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
if managed_object:
return managed_object.created_by == user_id
return True # don't raise error if managed object is not found
raise HTTPException(
status_code=404,
detail=f"Object not found: {unified_object_id}",
)
async def list_user_batches(
self,
@ -911,15 +916,22 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
setattr(response, file_attr, unified_file_id)
# Fetch the actual file object from the provider
# Use llm_router credentials when available. Without credentials,
# Azure and other auth-required providers return 500/401.
file_object = None
try:
# Use litellm to retrieve the file object from the provider
from litellm import afile_retrieve
file_object = await afile_retrieve(
custom_llm_provider=model_name.split("/")[0] if model_name and "/" in model_name else "openai",
file_id=original_file_id
)
from litellm.proxy.proxy_server import llm_router as _llm_router
if _llm_router is not None and model_id:
_creds = _llm_router.get_deployment_credentials_with_provider(model_id) or {}
file_object = await litellm.afile_retrieve(
file_id=original_file_id,
**_creds,
)
else:
file_object = await litellm.afile_retrieve(
custom_llm_provider=model_name.split("/")[0] if model_name and "/" in model_name else "openai",
file_id=original_file_id,
)
verbose_logger.debug(
f"Successfully retrieved file object for {file_attr}={original_file_id}"
)
@ -1004,7 +1016,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
# Case 2: Managed file and the file object exists in the database
# The stored file_object has the raw provider ID. Replace with the unified ID
# so callers see a consistent ID (matching Case 3 which does response.id = file_id).
if stored_file_object and stored_file_object.file_object:
stored_file_object.file_object.id = file_id
return stored_file_object.file_object
# Case 3: Managed file exists in the database but not the file object (for. e.g the batch task might not have run)

View file

@ -0,0 +1,19 @@
-- CreateTable
CREATE TABLE "LiteLLM_DeprecatedVerificationToken" (
"id" TEXT NOT NULL,
"token" TEXT NOT NULL,
"active_token_id" TEXT NOT NULL,
"revoke_at" TIMESTAMP(3) NOT NULL,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
CONSTRAINT "LiteLLM_DeprecatedVerificationToken_pkey" PRIMARY KEY ("id")
);
-- CreateIndex
CREATE UNIQUE INDEX "LiteLLM_DeprecatedVerificationToken_token_key" ON "LiteLLM_DeprecatedVerificationToken"("token");
-- CreateIndex
CREATE INDEX "LiteLLM_DeprecatedVerificationToken_token_revoke_at_idx" ON "LiteLLM_DeprecatedVerificationToken"("token", "revoke_at");
-- CreateIndex
CREATE INDEX "LiteLLM_DeprecatedVerificationToken_revoke_at_idx" ON "LiteLLM_DeprecatedVerificationToken"("revoke_at");

View file

@ -325,6 +325,19 @@ model LiteLLM_VerificationToken {
@@index([budget_reset_at, expires])
}
// Deprecated keys during grace period - allows old key to work until revoke_at
model LiteLLM_DeprecatedVerificationToken {
id String @id @default(uuid())
token String // Hashed old key
active_token_id String // Current token hash in LiteLLM_VerificationToken
revoke_at DateTime // When the old key stops working
created_at DateTime @default(now()) @map("created_at")
@@unique([token])
@@index([token, revoke_at])
@@index([revoke_at])
}
// Audit table for deleted keys - preserves spend and key information for historical tracking
model LiteLLM_DeletedVerificationToken {
id String @id @default(uuid())

View file

@ -312,10 +312,12 @@ class ServiceLogging(CustomLogger):
_duration, type(_duration)
)
) # invalid _duration value
# Batch polling callbacks (check_batch_cost) don't include call_type in kwargs.
# Use .get() to avoid KeyError.
await self.async_service_success_hook(
service=ServiceTypes.LITELLM,
duration=_duration,
call_type=kwargs["call_type"],
call_type=kwargs.get("call_type", "unknown")
)
except Exception as e:
raise e

View file

@ -8,7 +8,7 @@ import litellm
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.types.llms.openai import Batch
from litellm.types.utils import CallTypes, ModelResponse, Usage
from litellm.types.utils import CallTypes, ModelInfo, ModelResponse, Usage
from litellm.utils import token_counter
@ -16,14 +16,22 @@ async def calculate_batch_cost_and_usage(
file_content_dictionary: List[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],
model_name: Optional[str] = None,
model_info: Optional[ModelInfo] = None,
) -> Tuple[float, Usage, List[str]]:
"""
Calculate the cost and usage of a batch
Calculate the cost and usage of a batch.
Args:
model_info: Optional deployment-level model info with custom batch
pricing. Threaded through to batch_cost_calculator so that
deployment-specific pricing (e.g. input_cost_per_token_batches)
is used instead of the global cost map.
"""
batch_cost = _batch_cost_calculator(
custom_llm_provider=custom_llm_provider,
file_content_dictionary=file_content_dictionary,
model_name=model_name,
model_info=model_info,
)
batch_usage = _get_batch_job_total_usage_from_file_content(
file_content_dictionary=file_content_dictionary,
@ -94,6 +102,7 @@ def _batch_cost_calculator(
file_content_dictionary: List[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
model_name: Optional[str] = None,
model_info: Optional[ModelInfo] = None,
) -> float:
"""
Calculate the cost of a batch based on the output file id
@ -108,6 +117,7 @@ def _batch_cost_calculator(
total_cost = _get_batch_job_cost_from_file_content(
file_content_dictionary=file_content_dictionary,
custom_llm_provider=custom_llm_provider,
model_info=model_info,
)
verbose_logger.debug("total_cost=%s", total_cost)
return total_cost
@ -290,10 +300,13 @@ def _get_file_content_as_dictionary(file_content: bytes) -> List[dict]:
def _get_batch_job_cost_from_file_content(
file_content_dictionary: List[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
model_info: Optional[ModelInfo] = None,
) -> float:
"""
Get the cost of a batch job from the file content
"""
from litellm.cost_calculator import batch_cost_calculator
try:
total_cost: float = 0.0
# parse the file content as json
@ -303,11 +316,22 @@ def _get_batch_job_cost_from_file_content(
for _item in file_content_dictionary:
if _batch_response_was_successful(_item):
_response_body = _get_response_from_batch_job_output_file(_item)
total_cost += litellm.completion_cost(
completion_response=_response_body,
custom_llm_provider=custom_llm_provider,
call_type=CallTypes.aretrieve_batch.value,
)
if model_info is not None:
usage = _get_batch_job_usage_from_response_body(_response_body)
model = _response_body.get("model", "")
prompt_cost, completion_cost = batch_cost_calculator(
usage=usage,
model=model,
custom_llm_provider=custom_llm_provider,
model_info=model_info,
)
total_cost += prompt_cost + completion_cost
else:
total_cost += litellm.completion_cost(
completion_response=_response_body,
custom_llm_provider=custom_llm_provider,
call_type=CallTypes.aretrieve_batch.value,
)
verbose_logger.debug("total_cost=%s", total_cost)
return total_cost
except Exception as e:

View file

@ -319,6 +319,9 @@ NON_LLM_CONNECTION_TIMEOUT = int(
MAX_EXCEPTION_MESSAGE_LENGTH = int(os.getenv("MAX_EXCEPTION_MESSAGE_LENGTH", 2000))
MAX_STRING_LENGTH_PROMPT_IN_DB = int(os.getenv("MAX_STRING_LENGTH_PROMPT_IN_DB", 2048))
BEDROCK_MAX_POLICY_SIZE = int(os.getenv("BEDROCK_MAX_POLICY_SIZE", 75))
BEDROCK_MIN_THINKING_BUDGET_TOKENS = int(
os.getenv("BEDROCK_MIN_THINKING_BUDGET_TOKENS", 1024)
)
REPLICATE_POLLING_DELAY_SECONDS = float(
os.getenv("REPLICATE_POLLING_DELAY_SECONDS", 0.5)
)
@ -1258,6 +1261,9 @@ LITELLM_KEY_ROTATION_ENABLED = os.getenv("LITELLM_KEY_ROTATION_ENABLED", "false"
LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS = int(
os.getenv("LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS", 86400)
) # 24 hours default
LITELLM_KEY_ROTATION_GRACE_PERIOD: str = os.getenv(
"LITELLM_KEY_ROTATION_GRACE_PERIOD", ""
) # Duration to keep old key valid after rotation (e.g. "24h", "2d"); empty = immediate revoke (default)
UI_SESSION_TOKEN_TEAM_ID = "litellm-dashboard"
LITELLM_PROXY_ADMIN_NAME = "default_user_id"

View file

@ -1896,9 +1896,16 @@ def batch_cost_calculator(
usage: Usage,
model: str,
custom_llm_provider: Optional[str] = None,
model_info: Optional[ModelInfo] = None,
) -> Tuple[float, float]:
"""
Calculate the cost of a batch job
Calculate the cost of a batch job.
Args:
model_info: Optional deployment-level model info containing custom
batch pricing (e.g. input_cost_per_token_batches). When provided,
skips the global litellm.get_model_info() lookup so that
deployment-specific pricing is used.
"""
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
@ -1911,12 +1918,13 @@ def batch_cost_calculator(
custom_llm_provider,
)
try:
model_info: Optional[ModelInfo] = litellm.get_model_info(
model=model, custom_llm_provider=custom_llm_provider
)
except Exception:
model_info = None
if model_info is None:
try:
model_info = litellm.get_model_info(
model=model, custom_llm_provider=custom_llm_provider
)
except Exception:
model_info = None
if not model_info:
return 0.0, 0.0

View file

@ -51,6 +51,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_use_team_prefix: bool = False,
s3_strip_base64_files: bool = False,
s3_use_key_prefix: bool = False,
s3_use_virtual_hosted_style: bool = False,
**kwargs,
):
try:
@ -78,7 +79,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_path=s3_path,
s3_use_team_prefix=s3_use_team_prefix,
s3_strip_base64_files=s3_strip_base64_files,
s3_use_key_prefix=s3_use_key_prefix
s3_use_key_prefix=s3_use_key_prefix,
s3_use_virtual_hosted_style=s3_use_virtual_hosted_style
)
verbose_logger.debug(f"s3 logger using endpoint url {s3_endpoint_url}")
@ -135,6 +137,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_use_team_prefix: bool = False,
s3_strip_base64_files: bool = False,
s3_use_key_prefix: bool = False,
s3_use_virtual_hosted_style: bool = False,
):
"""
Initialize the s3 params for this logging callback
@ -217,6 +220,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
or s3_strip_base64_files
)
self.s3_use_virtual_hosted_style = (
bool(litellm.s3_callback_params.get("s3_use_virtual_hosted_style", False))
or s3_use_virtual_hosted_style
)
return
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
@ -247,8 +255,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
standard_logging_payload=kwargs.get("standard_logging_object", None),
)
# afile_delete and other non-model call types never produce a standard_logging_object,
# so s3_batch_logging_element is None. Skip gracefully instead of raising ValueError.
if s3_batch_logging_element is None:
raise ValueError("s3_batch_logging_element is None")
verbose_logger.debug(
"s3 Logging - skipping event, no standard_logging_object for call_type=%s",
kwargs.get("call_type", "unknown"),
)
return
verbose_logger.debug(
"\ns3 Logger - Logging payload = %s", s3_batch_logging_element
@ -302,13 +316,20 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
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 and self.s3_bucket_name:
url = (
self.s3_endpoint_url
+ "/"
+ self.s3_bucket_name
+ "/"
+ batch_logging_element.s3_object_key
)
if self.s3_use_virtual_hosted_style:
# Virtual-hosted-style: bucket.endpoint/key
endpoint_host = self.s3_endpoint_url.replace("https://", "").replace("http://", "")
protocol = "https://" if self.s3_endpoint_url.startswith("https://") else "http://"
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{batch_logging_element.s3_object_key}"
else:
# Path-style: endpoint/bucket/key
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)
@ -456,13 +477,20 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
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 and self.s3_bucket_name:
url = (
self.s3_endpoint_url
+ "/"
+ self.s3_bucket_name
+ "/"
+ batch_logging_element.s3_object_key
)
if self.s3_use_virtual_hosted_style:
# Virtual-hosted-style: bucket.endpoint/key
endpoint_host = self.s3_endpoint_url.replace("https://", "").replace("http://", "")
protocol = "https://" if self.s3_endpoint_url.startswith("https://") else "http://"
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{batch_logging_element.s3_object_key}"
else:
# Path-style: endpoint/bucket/key
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)
@ -550,13 +578,20 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{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
)
if self.s3_use_virtual_hosted_style:
# Virtual-hosted-style: bucket.endpoint/key
endpoint_host = self.s3_endpoint_url.replace("https://", "").replace("http://", "")
protocol = "https://" if self.s3_endpoint_url.startswith("https://") else "http://"
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{s3_object_key}"
else:
# Path-style: endpoint/bucket/key
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
@ -618,4 +653,4 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
verbose_logger.exception(
f"Error retrieving object {object_key} from cold storage: {str(e)}"
)
return None
return None

View file

@ -1282,9 +1282,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
output_config = optional_params.get("output_config")
if output_config and isinstance(output_config, dict):
effort = output_config.get("effort")
if effort and effort not in ["high", "medium", "low"]:
if effort and effort not in ["high", "medium", "low", "max"]:
raise ValueError(
f"Invalid effort value: {effort}. Must be one of: 'high', 'medium', 'low'"
f"Invalid effort value: {effort}. Must be one of: 'high', 'medium', 'low', 'max'"
)
if effort == "max" and not self._is_claude_opus_4_6(model):
raise ValueError(
f"effort='max' is only supported by Claude Opus 4.6. Got model: {model}"
)
data["output_config"] = output_config

View file

@ -19,6 +19,7 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
from litellm.types.utils import ModelResponse
from litellm.utils import get_model_info
if TYPE_CHECKING:
pass
@ -63,6 +64,14 @@ class LiteLLMMessagesToCompletionTransformationHandler:
return
model = completion_kwargs.get("model")
try:
model_info = get_model_info(model=cast(str, model), custom_llm_provider=custom_llm_provider)
if model_info and model_info.get("supports_reasoning") is False:
# Model doesn't support reasoning/responses API, don't route
return
except Exception:
pass
if isinstance(model, str) and model and not model.startswith("responses/"):
# Prefix model with "responses/" to route to OpenAI Responses API
completion_kwargs["model"] = f"responses/{model}"

View file

@ -239,8 +239,13 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
merged_chunk["delta"] = {}
# Add usage to the held chunk
uncached_input_tokens = chunk.usage.prompt_tokens or 0
if hasattr(chunk.usage, "prompt_tokens_details") and chunk.usage.prompt_tokens_details:
cached_tokens = getattr(chunk.usage.prompt_tokens_details, "cached_tokens", 0) or 0
uncached_input_tokens -= cached_tokens
usage_dict: UsageDelta = {
"input_tokens": chunk.usage.prompt_tokens or 0,
"input_tokens": uncached_input_tokens,
"output_tokens": chunk.usage.completion_tokens or 0,
}
# Add cache tokens if available (for prompt caching support)
@ -412,6 +417,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
if block_type == "tool_use":
# Type narrowing: content_block_start is ToolUseBlock when block_type is "tool_use"
from typing import cast
from litellm.types.llms.anthropic import ToolUseBlock
tool_block = cast(ToolUseBlock, content_block_start)
@ -430,6 +436,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
# if we get a function name since it signals a new tool call
if block_type == "tool_use":
from typing import cast
from litellm.types.llms.anthropic import ToolUseBlock
tool_block = cast(ToolUseBlock, content_block_start)

View file

@ -1070,8 +1070,13 @@ class LiteLLMAnthropicMessagesAdapter:
)
# extract usage
usage: Usage = getattr(response, "usage")
uncached_input_tokens = usage.prompt_tokens or 0
if hasattr(usage, "prompt_tokens_details") and usage.prompt_tokens_details:
cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
uncached_input_tokens -= cached_tokens
anthropic_usage = AnthropicUsage(
input_tokens=usage.prompt_tokens or 0,
input_tokens=uncached_input_tokens,
output_tokens=usage.completion_tokens or 0,
)
# Add cache tokens if available (for prompt caching support)
@ -1230,8 +1235,13 @@ class LiteLLMAnthropicMessagesAdapter:
else:
litellm_usage_chunk = None
if litellm_usage_chunk is not None:
uncached_input_tokens = litellm_usage_chunk.prompt_tokens or 0
if hasattr(litellm_usage_chunk, "prompt_tokens_details") and litellm_usage_chunk.prompt_tokens_details:
cached_tokens = getattr(litellm_usage_chunk.prompt_tokens_details, "cached_tokens", 0) or 0
uncached_input_tokens -= cached_tokens
usage_delta = UsageDelta(
input_tokens=litellm_usage_chunk.prompt_tokens or 0,
input_tokens=uncached_input_tokens,
output_tokens=litellm_usage_chunk.completion_tokens or 0,
)
# Add cache tokens if available (for prompt caching support)

View file

@ -11,7 +11,10 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.constants import (
BEDROCK_MIN_THINKING_BUDGET_TOKENS,
RESPONSE_FORMAT_TOOL_NAME,
)
from litellm.litellm_core_utils.core_helpers import (
filter_exceptions_from_params,
filter_internal_params,
@ -434,6 +437,25 @@ class AmazonConverseConfig(BaseConfig):
reasoning_effort=reasoning_effort, model=model
)
@staticmethod
def _clamp_thinking_budget_tokens(optional_params: dict) -> None:
"""
Clamp thinking.budget_tokens to the Bedrock minimum (1024).
Bedrock returns a 400 error if budget_tokens < 1024.
"""
thinking = optional_params.get("thinking")
if isinstance(thinking, dict):
budget = thinking.get("budget_tokens")
if isinstance(budget, int) and budget < BEDROCK_MIN_THINKING_BUDGET_TOKENS:
verbose_logger.debug(
"Bedrock requires thinking.budget_tokens >= %d, got %d. "
"Clamping to minimum.",
BEDROCK_MIN_THINKING_BUDGET_TOKENS,
budget,
)
thinking["budget_tokens"] = BEDROCK_MIN_THINKING_BUDGET_TOKENS
def get_supported_openai_params(self, model: str) -> List[str]:
from litellm.utils import supports_function_calling
@ -871,9 +893,14 @@ class AmazonConverseConfig(BaseConfig):
Checks 'non_default_params' for 'thinking' and 'max_tokens'
if 'thinking' is enabled and 'max_tokens' is not specified, set 'max_tokens' to the thinking token budget + DEFAULT_MAX_TOKENS
Also clamps thinking.budget_tokens to the Bedrock minimum (1024) to
prevent 400 errors from the Bedrock API.
"""
from litellm.constants import DEFAULT_MAX_TOKENS
self._clamp_thinking_budget_tokens(optional_params)
is_thinking_enabled = self.is_thinking_enabled(optional_params)
is_max_tokens_in_request = self.is_max_tokens_in_request(non_default_params)
if is_thinking_enabled and not is_max_tokens_in_request:

View file

@ -73,10 +73,6 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
litellm_params,
headers,
)
request.pop("max_output_tokens", None)
request.pop("max_tokens", None)
request.pop("max_completion_tokens", None)
request.pop("metadata", None)
base_instructions = get_chatgpt_default_instructions()
existing_instructions = request.get("instructions")
if existing_instructions:
@ -92,7 +88,22 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
if "reasoning.encrypted_content" not in include:
include.append("reasoning.encrypted_content")
request["include"] = include
return request
allowed_keys = {
"model",
"input",
"instructions",
"stream",
"store",
"include",
"tools",
"tool_choice",
"reasoning",
"previous_response_id",
"truncation",
}
return {k: v for k, v in request.items() if k in allowed_keys}
def transform_response_api_response(
self,

View file

@ -119,8 +119,13 @@ class AiohttpResponseStream(httpx.AsyncByteStream):
class AiohttpTransport(httpx.AsyncBaseTransport):
def __init__(self, client: Union[ClientSession, Callable[[], ClientSession]]) -> None:
def __init__(
self,
client: Union[ClientSession, Callable[[], ClientSession]],
owns_session: bool = True,
) -> None:
self.client = client
self._owns_session = owns_session
#########################################################
# Class variables for proxy settings
@ -128,7 +133,7 @@ class AiohttpTransport(httpx.AsyncBaseTransport):
self.proxy_cache: Dict[str, Optional[str]] = {}
async def aclose(self) -> None:
if isinstance(self.client, ClientSession):
if self._owns_session and isinstance(self.client, ClientSession):
await self.client.close()
@ -144,10 +149,11 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
self,
client: Union[ClientSession, Callable[[], ClientSession]],
ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None,
owns_session: bool = True,
):
self.client = client
self._ssl_verify = ssl_verify # Store for per-request SSL override
super().__init__(client=client)
super().__init__(client=client, owns_session=owns_session)
# Store the client factory for recreating sessions when needed
if callable(client):
self._client_factory = client

View file

@ -866,6 +866,7 @@ class AsyncHTTPHandler:
return LiteLLMAiohttpTransport(
client=shared_session,
ssl_verify=ssl_for_transport,
owns_session=False,
)
# Create new session only if none provided or existing one is invalid

View file

@ -12456,6 +12456,19 @@
"supports_tool_choice": true,
"supports_web_search": true
},
"fireworks_ai/accounts/fireworks/models/kimi-k2p5": {
"input_cost_per_token": 6e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 3e-06,
"source": "https://fireworks.ai/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"fireworks_ai/accounts/fireworks/models/llama-v3p1-405b-instruct": {
"input_cost_per_token": 3e-06,
"litellm_provider": "fireworks_ai",

View file

@ -854,9 +854,9 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
allowed_cache_controls: Optional[list] = []
config: Optional[dict] = {}
permissions: Optional[dict] = {}
model_max_budget: Optional[dict] = (
{}
) # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
model_max_budget: Optional[
dict
] = {} # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
model_config = ConfigDict(protected_namespaces=())
model_rpm_limit: Optional[dict] = None
@ -995,6 +995,9 @@ class RegenerateKeyRequest(GenerateKeyRequest):
spend: Optional[float] = None
metadata: Optional[dict] = None
new_master_key: Optional[str] = None
grace_period: Optional[
str
] = None # Duration to keep old key valid (e.g. "24h", "2d"); None = immediate revoke
class ResetSpendRequest(LiteLLMPydanticObjectBase):
@ -1406,12 +1409,12 @@ class NewCustomerRequest(BudgetNewRequest):
blocked: bool = False # allow/disallow requests for this end-user
budget_id: Optional[str] = None # give either a budget_id or max_budget
spend: Optional[float] = None
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
@model_validator(mode="before")
@classmethod
@ -1433,12 +1436,12 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase):
blocked: bool = False # allow/disallow requests for this end-user
max_budget: Optional[float] = None
budget_id: Optional[str] = None # give either a budget_id or max_budget
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
class DeleteCustomerRequest(LiteLLMPydanticObjectBase):
@ -1527,15 +1530,15 @@ class NewTeamRequest(TeamBase):
] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm
model_tpm_limit: Optional[Dict[str, int]] = None
team_member_budget: Optional[float] = (
None # allow user to set a budget for all team members
)
team_member_rpm_limit: Optional[int] = (
None # allow user to set RPM limit for all team members
)
team_member_tpm_limit: Optional[int] = (
None # allow user to set TPM limit for all team members
)
team_member_budget: Optional[
float
] = None # allow user to set a budget for all team members
team_member_rpm_limit: Optional[
int
] = None # allow user to set RPM limit for all team members
team_member_tpm_limit: Optional[
int
] = None # allow user to set TPM limit for all team members
team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m"
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
@ -1627,9 +1630,9 @@ class BlockKeyRequest(LiteLLMPydanticObjectBase):
class AddTeamCallback(LiteLLMPydanticObjectBase):
callback_name: str
callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = (
"success_and_failure"
)
callback_type: Optional[
Literal["success", "failure", "success_and_failure"]
] = "success_and_failure"
callback_vars: Dict[str, str]
@model_validator(mode="before")
@ -1961,9 +1964,9 @@ class ConfigList(LiteLLMPydanticObjectBase):
stored_in_db: Optional[bool]
field_default_value: Any
premium_field: bool = False
nested_fields: Optional[List[FieldDetail]] = (
None # For nested dictionary or Pydantic fields
)
nested_fields: Optional[
List[FieldDetail]
] = None # For nested dictionary or Pydantic fields
class UserHeaderMapping(LiteLLMPydanticObjectBase):
@ -2403,9 +2406,9 @@ class LiteLLM_OrganizationMembershipTable(LiteLLMPydanticObjectBase):
budget_id: Optional[str] = None
created_at: datetime
updated_at: datetime
user: Optional[Any] = (
None # You might want to replace 'Any' with a more specific type if available
)
user: Optional[
Any
] = None # You might want to replace 'Any' with a more specific type if available
litellm_budget_table: Optional[LiteLLM_BudgetTable] = None
model_config = ConfigDict(protected_namespaces=())
@ -3396,9 +3399,9 @@ class TeamModelDeleteRequest(BaseModel):
# Organization Member Requests
class OrganizationMemberAddRequest(OrgMemberAddRequest):
organization_id: str
max_budget_in_organization: Optional[float] = (
None # Users max budget within the organization
)
max_budget_in_organization: Optional[
float
] = None # Users max budget within the organization
class OrganizationMemberDeleteRequest(MemberDeleteRequest):
@ -3616,9 +3619,9 @@ class ProviderBudgetResponse(LiteLLMPydanticObjectBase):
Maps provider names to their budget configs.
"""
providers: Dict[str, ProviderBudgetResponseObject] = (
{}
) # Dictionary mapping provider names to their budget configurations
providers: Dict[
str, ProviderBudgetResponseObject
] = {} # Dictionary mapping provider names to their budget configurations
class ProxyStateVariables(TypedDict):
@ -3761,9 +3764,9 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
enforce_rbac: bool = False
roles_jwt_field: Optional[str] = None # v2 on role mappings
role_mappings: Optional[List[RoleMapping]] = None
object_id_jwt_field: Optional[str] = (
None # can be either user / team, inferred from the role mapping
)
object_id_jwt_field: Optional[
str
] = None # can be either user / team, inferred from the role mapping
scope_mappings: Optional[List[ScopeMapping]] = None
enforce_scope_based_access: bool = False
enforce_team_based_model_access: bool = False

View file

@ -29,6 +29,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
get_models_from_unified_file_id,
get_original_file_id,
prepare_data_with_credentials,
resolve_input_file_id_to_unified,
update_batch_in_database,
)
from litellm.proxy.utils import handle_exception_on_proxy, is_known_model
@ -305,7 +306,7 @@ async def create_batch( # noqa: PLR0915
dependencies=[Depends(user_api_key_auth)],
tags=["batch"],
)
async def retrieve_batch(
async def retrieve_batch( # noqa: PLR0915
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
@ -377,6 +378,11 @@ async def retrieve_batch(
response = await proxy_logging_obj.post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
# async_post_call_success_hook replaces batch.id and output_file_id with unified IDs
# but not input_file_id. Resolve raw provider ID to unified ID.
if unified_batch_id:
await resolve_input_file_id_to_unified(response, prisma_client)
asyncio.create_task(
proxy_logging_obj.update_request_status(
@ -479,6 +485,11 @@ async def retrieve_batch(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
# Fix: bug_feb14_batch_retrieve_returns_raw_input_file_id
# Resolve raw provider input_file_id to unified ID.
if unified_batch_id:
await resolve_input_file_id_to_unified(response, prisma_client)
### ALERTING ###
asyncio.create_task(
proxy_logging_obj.update_request_status(

View file

@ -8,7 +8,10 @@ from datetime import datetime, timezone
from typing import List
from litellm._logging import verbose_proxy_logger
from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
from litellm.constants import (
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
LITELLM_KEY_ROTATION_GRACE_PERIOD,
)
from litellm.proxy._types import (
GenerateKeyResponse,
LiteLLM_VerificationToken,
@ -37,6 +40,9 @@ class KeyRotationManager:
try:
verbose_proxy_logger.info("Starting scheduled key rotation check...")
# Clean up expired deprecated keys first
await self._cleanup_expired_deprecated_keys()
# Find keys that are due for rotation
keys_to_rotate = await self._find_keys_needing_rotation()
@ -97,6 +103,24 @@ class KeyRotationManager:
return keys_with_rotation
async def _cleanup_expired_deprecated_keys(self) -> None:
"""
Remove deprecated key entries whose revoke_at has passed.
"""
try:
now = datetime.now(timezone.utc)
result = await self.prisma_client.db.litellm_deprecatedverificationtoken.delete_many(
where={"revoke_at": {"lt": now}}
)
if result > 0:
verbose_proxy_logger.debug(
"Cleaned up %s expired deprecated key(s)", result
)
except Exception as e:
verbose_proxy_logger.debug(
"Deprecated key cleanup skipped (table may not exist): %s", e
)
def _should_rotate_key(self, key: LiteLLM_VerificationToken, now: datetime) -> bool:
"""
Determine if a key should be rotated based on key_rotation_at timestamp.
@ -115,10 +139,11 @@ class KeyRotationManager:
"""
Rotate a single key using existing regenerate_key_fn and call the rotation hook
"""
# Create regenerate request
# Create regenerate request with grace period for seamless cutover
regenerate_request = RegenerateKeyRequest(
key=key.token or "",
key_alias=key.key_alias, # Pass key alias to ensure correct secret is updated in AWS Secrets Manager
grace_period=LITELLM_KEY_ROTATION_GRACE_PERIOD or None,
)
# Create a system user for key rotation

View file

@ -1725,13 +1725,6 @@ class DBSpendUpdateWriter:
"prisma_client is None. Skipping writing spend logs to db."
)
return
base_daily_transaction = (
await self._common_add_spend_log_transaction_to_daily_transaction(
payload, prisma_client, "agent"
)
)
if base_daily_transaction is None:
return
if payload["agent_id"] is None:
verbose_proxy_logger.debug(
"agent_id is None for request. Skipping incrementing agent spend."

View file

@ -259,9 +259,10 @@ class _PROXY_BatchRateLimiter(CustomLogger):
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
# Managed files require bypassing the HTTP endpoint (which runs access-check hooks)
# and calling the managed files hook directly with the user's credentials.
is_managed_file = _is_base64_encoded_unified_file_id(file_id)
if is_managed_file and user_api_key_dict is not None:
# For managed files, use the managed files hook directly
file_content = await self._fetch_managed_file_content(
file_id=file_id,
user_api_key_dict=user_api_key_dict,

View file

@ -202,6 +202,14 @@ class _ProxyDBLogger(CustomLogger):
max_budget=end_user_max_budget,
)
else:
# Non-model call types (health checks, afile_delete) have no model or standard_logging_object.
# Use .get() for "stream" to avoid KeyError on health checks.
if sl_object is None and not kwargs.get("model"):
verbose_proxy_logger.warning(
"Cost tracking - skipping, no standard_logging_object and no model for call_type=%s",
kwargs.get("call_type", "unknown"),
)
return
if kwargs.get("stream") is not True or (
kwargs.get("stream") is True and "complete_streaming_response" in kwargs
):

View file

@ -12,6 +12,7 @@ All /key management endpoints
import asyncio
import copy
import json
import os
import secrets
import traceback
from datetime import datetime, timedelta, timezone
@ -629,7 +630,11 @@ async def _common_key_generation_helper( # noqa: PLR0915
# Validate user-provided key format
if data.key is not None and not data.key.startswith("sk-"):
_masked = "{}****{}".format(data.key[:4], data.key[-4:]) if len(data.key) > 8 else "****"
_masked = (
"{}****{}".format(data.key[:4], data.key[-4:])
if len(data.key) > 8
else "****"
)
raise HTTPException(
status_code=400,
detail={
@ -1343,6 +1348,7 @@ async def prepare_key_update_data(
data_json: dict = data.model_dump(exclude_unset=True)
data_json.pop("key", None)
data_json.pop("new_key", None)
data_json.pop("grace_period", None) # Request-only param, not a DB column
if (
data.metadata is not None
and data.metadata.get("service_account_id") is not None
@ -3181,6 +3187,67 @@ def get_new_token(data: Optional[RegenerateKeyRequest]) -> str:
return new_token
async def _insert_deprecated_key(
prisma_client: "PrismaClient",
old_token_hash: str,
new_token_hash: str,
grace_period: Optional[str],
) -> None:
"""
Insert old key into deprecated table so it remains valid during grace period.
Uses upsert to handle concurrent rotations gracefully.
Parameters:
prisma_client: DB client
old_token_hash: Hash of the old key being rotated out
new_token_hash: Hash of the new replacement key
grace_period: Duration string (e.g. "24h", "2d") or None/empty for immediate revoke
"""
grace_period_value = grace_period or os.getenv(
"LITELLM_KEY_ROTATION_GRACE_PERIOD", ""
)
if not grace_period_value:
return
try:
grace_seconds = duration_in_seconds(grace_period_value)
except ValueError:
verbose_proxy_logger.warning(
"Invalid grace_period format: %s. Expected format like '24h', '2d'.",
grace_period_value,
)
return
if grace_seconds <= 0:
return
try:
revoke_at = datetime.now(timezone.utc) + timedelta(seconds=grace_seconds)
await prisma_client.db.litellm_deprecatedverificationtoken.upsert(
where={"token": old_token_hash},
data={
"create": {
"token": old_token_hash,
"active_token_id": new_token_hash,
"revoke_at": revoke_at,
},
"update": {
"active_token_id": new_token_hash,
"revoke_at": revoke_at,
},
},
)
verbose_proxy_logger.debug(
"Deprecated key retained for %s (revoke_at: %s)",
grace_period_value,
revoke_at,
)
except Exception as deprecated_err:
verbose_proxy_logger.warning(
"Failed to insert deprecated key for grace period: %s",
deprecated_err,
)
async def _execute_virtual_key_regeneration(
*,
prisma_client: PrismaClient,
@ -3288,6 +3355,7 @@ async def regenerate_key_fn( # noqa: PLR0915
- permissions: Optional[dict] - Key-specific permissions
- guardrails: Optional[List[str]] - List of active guardrails for the key
- blocked: Optional[bool] - Whether the key is blocked
- grace_period: Optional[str] - Duration to keep old key valid after rotation (e.g. "24h", "2d"). Omitted = immediate revoke. Env: LITELLM_KEY_ROTATION_GRACE_PERIOD
Returns:
@ -3406,6 +3474,58 @@ async def regenerate_key_fn( # noqa: PLR0915
)
verbose_proxy_logger.debug("key_in_db: %s", _key_in_db)
new_token = get_new_token(data=data)
new_token_hash = hash_token(new_token)
new_token_key_name = f"sk-...{new_token[-4:]}"
# Prepare the update data
update_data = {
"token": new_token_hash,
"key_name": new_token_key_name,
}
non_default_values = {}
if data is not None:
# Update with any provided parameters from GenerateKeyRequest
non_default_values = await prepare_key_update_data(
data=data, existing_key_row=_key_in_db
)
verbose_proxy_logger.debug("non_default_values: %s", non_default_values)
update_data.update(non_default_values)
update_data = prisma_client.jsonify_object(data=update_data)
# If grace period set, insert deprecated key so old key remains valid
await _insert_deprecated_key(
prisma_client=prisma_client,
old_token_hash=hashed_api_key,
new_token_hash=new_token_hash,
grace_period=data.grace_period if data else None,
)
# Update the token in the database
updated_token = await prisma_client.db.litellm_verificationtoken.update(
where={"token": hashed_api_key},
data=update_data, # type: ignore
)
updated_token_dict = {}
if updated_token is not None:
updated_token_dict = dict(updated_token)
updated_token_dict["key"] = new_token
updated_token_dict["token_id"] = updated_token_dict.pop("token")
### 3. remove existing key entry from cache
######################################################################
if hashed_api_key or key:
await _delete_cache_key_object(
hashed_token=hash_token(key),
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# Normalize litellm_changed_by: if it's a Header object or not a string, convert to None
if litellm_changed_by is not None and not isinstance(litellm_changed_by, str):
litellm_changed_by = None

View file

@ -14,7 +14,7 @@ import hashlib
import os
import secrets
from copy import deepcopy
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast
from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi.responses import RedirectResponse
@ -82,7 +82,15 @@ from litellm.proxy.utils import (
get_server_root_path,
)
from litellm.secret_managers.main import get_secret_bool, str_to_bool
from litellm.types.proxy.management_endpoints.ui_sso import *
from litellm.types.proxy.management_endpoints.ui_sso import (
DefaultTeamSSOParams,
MicrosoftGraphAPIUserGroupDirectoryObject,
MicrosoftGraphAPIUserGroupResponse,
MicrosoftServicePrincipalTeam,
RoleMappings,
TeamMappings,
)
from litellm.types.proxy.management_endpoints.ui_sso import * # noqa: F403, F401
from litellm.types.proxy.ui_sso import ParsedOpenIDResult
if TYPE_CHECKING:
@ -96,15 +104,15 @@ router = APIRouter()
def normalize_email(email: Optional[str]) -> Optional[str]:
"""
Normalize email address to lowercase for consistent storage and comparison.
Email addresses should be treated as case-insensitive for SSO purposes,
even though RFC 5321 technically allows case-sensitive local parts.
This prevents issues where SSO providers return emails with different casing
than what's stored in the database.
Args:
email: Email address to normalize, can be None
Returns:
Lowercased email address, or None if input is None
"""
@ -336,7 +344,7 @@ async def google_login(
# check if user defined a custom auth sso sign in handler, if yes, use it
if user_custom_ui_sso_sign_in_handler is not None:
try:
from litellm_enterprise.proxy.auth.custom_sso_handler import (
from litellm_enterprise.proxy.auth.custom_sso_handler import ( # type: ignore[import-untyped]
EnterpriseCustomSSOHandler,
)
@ -494,7 +502,9 @@ def generic_response_convertor(
display_name=get_nested_value(
response, generic_user_display_name_attribute_name
),
email=normalize_email(get_nested_value(response, generic_user_email_attribute_name)),
email=normalize_email(
get_nested_value(response, generic_user_email_attribute_name)
),
first_name=get_nested_value(response, generic_user_first_name_attribute_name),
last_name=get_nested_value(response, generic_user_last_name_attribute_name),
provider=get_nested_value(response, generic_provider_attribute_name),
@ -584,6 +594,7 @@ async def _setup_team_mappings() -> Optional["TeamMappings"]:
if team_mappings_data:
from litellm.types.proxy.management_endpoints.ui_sso import TeamMappings
if isinstance(team_mappings_data, dict):
team_mappings = TeamMappings(**team_mappings_data)
elif isinstance(team_mappings_data, TeamMappings):
@ -621,6 +632,7 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]:
if role_mappings_data:
from litellm.types.proxy.management_endpoints.ui_sso import RoleMappings
if isinstance(role_mappings_data, dict):
role_mappings = RoleMappings(**role_mappings_data)
elif isinstance(role_mappings_data, RoleMappings):
@ -634,7 +646,7 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]:
verbose_proxy_logger.debug(
f"Could not load role_mappings from database: {e}. Continuing with existing role logic."
)
generic_role_mappings = os.getenv("GENERIC_ROLE_MAPPINGS_ROLES", None)
generic_role_mappings_group_claim = os.getenv(
"GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", None
@ -644,8 +656,8 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]:
)
if generic_role_mappings is not None:
verbose_proxy_logger.debug(
"Found role_mappings for generic provider in environment variables"
)
"Found role_mappings for generic provider in environment variables"
)
import ast
try:
@ -670,7 +682,9 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]:
)
return role_mappings
except TypeError as e:
verbose_proxy_logger.warning(f"Error decoding role mappings from environment variables: {e}. Continuing with existing role logic.")
verbose_proxy_logger.warning(
f"Error decoding role mappings from environment variables: {e}. Continuing with existing role logic."
)
return role_mappings
@ -747,7 +761,7 @@ async def get_generic_sso_response(
try:
result = await generic_sso.verify_and_process(
request,
params=SSOAuthenticationHandler.prepare_token_exchange_parameters(
params=await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=request,
generic_include_client_id=generic_include_client_id,
),
@ -942,7 +956,7 @@ def _build_sso_user_update_data(
Returns:
dict: Update data containing user_email and optionally user_role if valid
"""
"""
update_data: dict = {"user_email": normalize_email(user_email)}
# Get SSO role from result and include if valid
@ -1740,7 +1754,7 @@ class SSOAuthenticationHandler:
"""
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
from litellm.proxy.proxy_server import user_api_key_cache
from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache
with generic_sso:
# TODO: state should be a random string and added to the user session with cookie
@ -1769,13 +1783,21 @@ class SSOAuthenticationHandler:
# If PKCE is enabled, add PKCE parameters to the redirect URL
if code_verifier and "state" in redirect_params:
# Store code_verifier in cache (10 min TTL)
# Store code_verifier in cache (10 min TTL). Use Redis when available
# so callbacks landing on another pod can retrieve it (multi-pod SSO).
cache_key = f"pkce_verifier:{redirect_params['state']}"
user_api_key_cache.set_cache(
key=cache_key,
value=code_verifier,
ttl=600,
)
if redis_usage_cache is not None:
await redis_usage_cache.async_set_cache(
key=cache_key,
value=code_verifier,
ttl=600,
)
else:
await user_api_key_cache.async_set_cache(
key=cache_key,
value=code_verifier,
ttl=600,
)
# Add PKCE parameters to the authorization URL
if pkce_params:
@ -2372,7 +2394,7 @@ class SSOAuthenticationHandler:
return redirect_response
@staticmethod
def prepare_token_exchange_parameters(
async def prepare_token_exchange_parameters(
request: Request,
generic_include_client_id: bool,
) -> dict:
@ -2386,27 +2408,38 @@ class SSOAuthenticationHandler:
Returns:
dict: Token exchange parameters
"""
# Prepare token exchange parameters
token_params = {"include_client_id": generic_include_client_id}
# Prepare token exchange parameters (may add code_verifier: str later)
token_params: Dict[str, Any] = {"include_client_id": generic_include_client_id}
# Retrieve PKCE code_verifier if PKCE was used in authorization
# Retrieve PKCE code_verifier if PKCE was used in authorization.
# Use same cache as store: Redis when available (multi-pod), else in-memory.
query_params = dict(request.query_params)
state = query_params.get("state")
if state:
from litellm.proxy.proxy_server import user_api_key_cache
from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache
cache_key = f"pkce_verifier:{state}"
code_verifier = user_api_key_cache.get_cache(key=cache_key)
if redis_usage_cache is not None:
code_verifier = await redis_usage_cache.async_get_cache(key=cache_key)
else:
code_verifier = await user_api_key_cache.async_get_cache(key=cache_key)
if code_verifier:
# Add code_verifier to token exchange parameters
token_params["code_verifier"] = code_verifier
# Add code_verifier to token exchange parameters (Redis returns decoded string)
token_params["code_verifier"] = (
code_verifier
if isinstance(code_verifier, str)
else str(code_verifier)
)
verbose_proxy_logger.debug(
"PKCE code_verifier retrieved and will be included in token exchange"
)
# Clean up the cache entry (single-use verifier)
user_api_key_cache.delete_cache(key=cache_key)
if redis_usage_cache is not None:
await redis_usage_cache.async_delete_cache(key=cache_key)
else:
await user_api_key_cache.async_delete_cache(key=cache_key)
return token_params
@staticmethod
@ -2549,7 +2582,9 @@ class MicrosoftSSOHandler:
response = response or {}
verbose_proxy_logger.debug(f"Microsoft SSO Callback Response: {response}")
openid_response = CustomOpenID(
email=normalize_email(response.get(MICROSOFT_USER_EMAIL_ATTRIBUTE) or response.get("mail")),
email=normalize_email(
response.get(MICROSOFT_USER_EMAIL_ATTRIBUTE) or response.get("mail")
),
display_name=response.get(MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE),
provider="microsoft",
id=response.get(MICROSOFT_USER_ID_ATTRIBUTE),

View file

@ -644,6 +644,28 @@ def _extract_model_param(request: "Request", request_body: dict) -> Optional[str
# ============================================================================
async def resolve_input_file_id_to_unified(response, prisma_client) -> None:
"""
If the batch response contains a raw provider input_file_id (not already a
unified ID), look up the corresponding unified file ID from the managed file
table and replace it in-place.
"""
if (
hasattr(response, "input_file_id")
and response.input_file_id
and not _is_base64_encoded_unified_file_id(response.input_file_id)
and prisma_client
):
try:
managed_file = await prisma_client.db.litellm_managedfiletable.find_first(
where={"flat_model_file_ids": {"has": response.input_file_id}}
)
if managed_file:
response.input_file_id = managed_file.unified_file_id
except Exception:
pass
async def get_batch_from_database(
batch_id: str,
unified_batch_id: Union[str, Literal[False]],
@ -687,6 +709,9 @@ async def get_batch_from_database(
batch_data = json.loads(db_batch_object.file_object) if isinstance(db_batch_object.file_object, str) else db_batch_object.file_object
response = LiteLLMBatch(**batch_data)
response.id = batch_id
# The stored batch object has the raw provider input_file_id. Resolve to unified ID.
await resolve_input_file_id_to_unified(response, prisma_client)
verbose_proxy_logger.debug(
f"Retrieved batch {batch_id} from ManagedObjectTable with status={response.status}"

View file

@ -325,6 +325,19 @@ model LiteLLM_VerificationToken {
@@index([budget_reset_at, expires])
}
// Deprecated keys during grace period - allows old key to work until revoke_at
model LiteLLM_DeprecatedVerificationToken {
id String @id @default(uuid())
token String // Hashed old key
active_token_id String // Current token hash in LiteLLM_VerificationToken
revoke_at DateTime // When the old key stops working
created_at DateTime @default(now()) @map("created_at")
@@unique([token])
@@index([token, revoke_at])
@@index([revoke_at])
}
// Audit table for deleted keys - preserves spend and key information for historical tracking
model LiteLLM_DeletedVerificationToken {
id String @id @default(uuid())

View file

@ -7,7 +7,7 @@ import smtplib
import threading
import time
import traceback
from datetime import date, datetime, timedelta
from datetime import date, datetime, timedelta, timezone
from email.mime.multipart import MIMEMultipart
from email.mime.text import MIMEText
from typing import (
@ -76,6 +76,7 @@ from litellm import (
from litellm._logging import verbose_proxy_logger
from litellm._service_logger import ServiceLogging, ServiceTypes
from litellm.caching.caching import DualCache, RedisCache
from litellm.caching.dual_cache import LimitedSizeOrderedDict
from litellm.exceptions import RejectedRequestError
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
@ -2154,6 +2155,58 @@ def jsonify_object(data: dict) -> dict:
return db_data
# In-memory cache for deprecated key lookups: maps old_token_hash -> (active_token_id, expires_at_ts)
# Avoids a DB query on every auth request for non-deprecated keys.
# Bounded to prevent memory leaks from accumulated rotations.
_deprecated_key_cache: LimitedSizeOrderedDict = LimitedSizeOrderedDict(max_size=1000)
_DEPRECATED_KEY_CACHE_TTL_SECONDS = 60
async def _lookup_deprecated_key(
db: Any,
hashed_token: str,
) -> Optional[str]:
"""
Check if a token exists in the deprecated keys table and is still within its grace period.
Returns the active_token_id if found and valid, otherwise None.
Uses an in-memory cache to avoid DB queries on every auth request.
"""
now = datetime.now(timezone.utc)
now_ts = now.timestamp()
# Check cache first
cached = _deprecated_key_cache.get(hashed_token)
cached = _deprecated_key_cache.get(hashed_token)
if cached is not None:
active_token_id, cache_expires_at_ts, revoke_at_ts = cached
if now_ts < cache_expires_at_ts and now_ts < revoke_at_ts:
return active_token_id
else:
_deprecated_key_cache.pop(hashed_token, None)
try:
deprecated_row = await db.litellm_deprecatedverificationtoken.find_first(
where={
"token": hashed_token,
"revoke_at": {"gt": now},
},
select={"active_token_id": True},
)
if deprecated_row and deprecated_row.active_token_id:
_deprecated_key_cache[hashed_token] = (
deprecated_row.active_token_id,
now_ts + _DEPRECATED_KEY_CACHE_TTL_SECONDS,
)
return deprecated_row.active_token_id
# Only cache positive results; negative lookups are fast on indexed columns
# and caching them risks evicting real deprecated key entries.
except Exception as e:
verbose_proxy_logger.debug("Deprecated key lookup skipped: %s", e)
return None
class PrismaClient:
spend_log_transactions: List = []
_spend_log_transactions_lock = asyncio.Lock()
@ -2489,6 +2542,7 @@ class PrismaClient:
parent_otel_span: Optional[Span] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
budget_id_list: Optional[List[str]] = None,
check_deprecated: bool = True,
):
args_passed_in = locals()
start_time = time.time()
@ -2786,6 +2840,30 @@ class PrismaClient:
sql_query
)
# If not found in main table, check deprecated keys (grace period)
# check_deprecated=False on the recursive call prevents unbounded chaining
if (
response is None
and hashed_token is not None
and check_deprecated
):
active_token_id = await _lookup_deprecated_key(
db=self.db, hashed_token=hashed_token
)
if active_token_id:
response = await self.get_data(
token=active_token_id,
table_name="combined_view",
query_type="find_unique",
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_deprecated=False,
)
if response is not None:
verbose_proxy_logger.debug(
"Deprecated key used during grace period"
)
if response is not None:
if response["team_models"] is None:
response["team_models"] = []

View file

@ -12456,6 +12456,19 @@
"supports_tool_choice": true,
"supports_web_search": true
},
"fireworks_ai/accounts/fireworks/models/kimi-k2p5": {
"input_cost_per_token": 6e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 3e-06,
"source": "https://fireworks.ai/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"fireworks_ai/accounts/fireworks/models/llama-v3p1-405b-instruct": {
"input_cost_per_token": 3e-06,
"litellm_provider": "fireworks_ai",

56
poetry.lock generated
View file

@ -330,14 +330,14 @@ zookeeper = ["kazoo"]
name = "async-timeout"
version = "5.0.1"
description = "Timeout context manager for asyncio programs"
optional = true
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "python_full_version < \"3.11.3\" and (extra == \"extra-proxy\" or extra == \"proxy\") or python_version < \"3.11\""
groups = ["main", "dev"]
files = [
{file = "async_timeout-5.0.1-py3-none-any.whl", hash = "sha256:39e3809566ff85354557ec2398b55e096c8364bacac9405a7a1fa429e77fe76c"},
{file = "async_timeout-5.0.1.tar.gz", hash = "sha256:d9321a7a3d5a6a5e187e824d2fa0793ce379a202935782d555d6e9d2735677d3"},
]
markers = {main = "python_full_version < \"3.11.3\" and (extra == \"extra-proxy\" or extra == \"proxy\") or python_version < \"3.11\"", dev = "python_full_version < \"3.11.3\""}
[[package]]
name = "attrs"
@ -1271,6 +1271,34 @@ typing-extensions = {version = ">=4.6.0", markers = "python_version < \"3.13\""}
[package.extras]
test = ["pytest (>=6)"]
[[package]]
name = "fakeredis"
version = "2.33.0"
description = "Python implementation of redis API, can be used for testing purposes."
optional = false
python-versions = ">=3.7"
groups = ["dev"]
files = [
{file = "fakeredis-2.33.0-py3-none-any.whl", hash = "sha256:de535f3f9ccde1c56672ab2fdd6a8efbc4f2619fc2f1acc87b8737177d71c965"},
{file = "fakeredis-2.33.0.tar.gz", hash = "sha256:d7bc9a69d21df108a6451bbffee23b3eba432c21a654afc7ff2d295428ec5770"},
]
[package.dependencies]
redis = [
{version = ">=4.3", markers = "python_version > \"3.8\""},
{version = ">=4.3,<7.1.0", markers = "python_version < \"3.10\" and python_version > \"3.8\""},
]
sortedcontainers = ">=2"
typing-extensions = {version = ">=4.7,<5.0", markers = "python_version < \"3.11\""}
[package.extras]
bf = ["pyprobables (>=0.6)"]
cf = ["pyprobables (>=0.6)"]
json = ["jsonpath-ng (>=1.6)"]
lua = ["lupa (>=2.1)"]
probabilistic = ["pyprobables (>=0.6)"]
valkey = ["valkey (>=6) ; python_version >= \"3.8\""]
[[package]]
name = "fastapi"
version = "0.121.3"
@ -5306,7 +5334,7 @@ version = "2.10.1"
description = "JSON Web Token implementation in Python"
optional = false
python-versions = ">=3.9"
groups = ["main", "proxy-dev"]
groups = ["main", "dev", "proxy-dev"]
files = [
{file = "PyJWT-2.10.1-py3-none-any.whl", hash = "sha256:dcdd193e30abefd5debf142f9adfcdd2b58004e644f25406ffaebd50bd98dacb"},
{file = "pyjwt-2.10.1.tar.gz", hash = "sha256:3cc5772eb20009233caf06e9d8a0577824723b44e6648ee0a2aedb6cf9381953"},
@ -5705,14 +5733,14 @@ files = [
name = "redis"
version = "5.3.1"
description = "Python client for Redis database and key-value store"
optional = true
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "(python_version < \"3.14\" or extra == \"proxy\") and (python_version <= \"3.13\" or extra == \"proxy\" or extra == \"extra-proxy\") and (extra == \"extra-proxy\" or extra == \"proxy\")"
groups = ["main", "dev"]
files = [
{file = "redis-5.3.1-py3-none-any.whl", hash = "sha256:dc1909bd24669cc31b5f67a039700b16ec30571096c5f1f0d9d2324bff31af97"},
{file = "redis-5.3.1.tar.gz", hash = "sha256:ca49577a531ea64039b5a36db3d6cd1a0c7a60c34124d46924a45b956e8cf14c"},
]
markers = {main = "(python_version < \"3.14\" or extra == \"proxy\") and (python_version <= \"3.13\" or extra == \"proxy\" or extra == \"extra-proxy\") and (extra == \"extra-proxy\" or extra == \"proxy\")"}
[package.dependencies]
async-timeout = {version = ">=4.0.3", markers = "python_full_version < \"3.11.3\""}
@ -6583,6 +6611,18 @@ files = [
{file = "snowballstemmer-3.0.1.tar.gz", hash = "sha256:6d5eeeec8e9f84d4d56b847692bacf79bc2c8e90c7f80ca4444ff8b6f2e52895"},
]
[[package]]
name = "sortedcontainers"
version = "2.4.0"
description = "Sorted Containers -- Sorted List, Sorted Dict, Sorted Set"
optional = false
python-versions = "*"
groups = ["dev"]
files = [
{file = "sortedcontainers-2.4.0-py2.py3-none-any.whl", hash = "sha256:a163dcaede0f1c021485e957a39245190e74249897e2ae4b2aa38595db237ee0"},
{file = "sortedcontainers-2.4.0.tar.gz", hash = "sha256:25caa5a06cc30b6b83d11423433f65d1f9d76c4c6a0c90e3379eaa43b9bfdb88"},
]
[[package]]
name = "soundfile"
version = "0.12.1"
@ -7894,4 +7934,4 @@ utils = ["numpydoc"]
[metadata]
lock-version = "2.1"
python-versions = ">=3.9,<4.0"
content-hash = "6110d55c765f8f0fd050c0cbc09808f16e27de3bc2ae69f607b8fc5f16215de4"
content-hash = "20ca098d83da3b9364b05930a74e9ff8512e31d626018fc9f056b6fbd50a69af"

View file

@ -165,6 +165,7 @@ opentelemetry-sdk = "^1.28.0"
opentelemetry-exporter-otlp = "^1.28.0"
langfuse = "^2.45.0"
fastapi-offline = "^1.7.3"
fakeredis = "^2.27.1"
[tool.poetry.group.proxy-dev.dependencies]
prisma = "0.11.0"

View file

@ -325,6 +325,19 @@ model LiteLLM_VerificationToken {
@@index([budget_reset_at, expires])
}
// Deprecated keys during grace period - allows old key to work until revoke_at
model LiteLLM_DeprecatedVerificationToken {
id String @id @default(uuid())
token String // Hashed old key
active_token_id String // Current token hash in LiteLLM_VerificationToken
revoke_at DateTime // When the old key stops working
created_at DateTime @default(now()) @map("created_at")
@@unique([token])
@@index([token, revoke_at])
@@index([revoke_at])
}
// Audit table for deleted keys - preserves spend and key information for historical tracking
model LiteLLM_DeletedVerificationToken {
id String @id @default(uuid())

View file

@ -0,0 +1,131 @@
"""
Test that batch cost calculation uses custom deployment-level pricing
when model_info is provided.
Reproduces the bug where `input_cost_per_token_batches` /
`output_cost_per_token_batches` set on a proxy deployment's model_info
are ignored by the batch cost pipeline because they are never threaded
through to `batch_cost_calculator`.
"""
import pytest
from litellm.batches.batch_utils import (
_batch_cost_calculator,
_get_batch_job_cost_from_file_content,
calculate_batch_cost_and_usage,
)
from litellm.cost_calculator import batch_cost_calculator
from litellm.types.utils import Usage
# --- helpers ---
def _make_batch_output_line(prompt_tokens: int = 10, completion_tokens: int = 5):
"""Return a single successful batch output line (OpenAI JSONL format)."""
return {
"id": "batch_req_1",
"custom_id": "req-1",
"response": {
"status_code": 200,
"body": {
"id": "chatcmpl-test",
"object": "chat.completion",
"model": "fake-batch-model",
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
},
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "Hello"},
"finish_reason": "stop",
}
],
},
},
"error": None,
}
CUSTOM_MODEL_INFO = {
"input_cost_per_token_batches": 0.00125,
"output_cost_per_token_batches": 0.005,
}
# --- tests ---
def test_batch_cost_calculator_uses_custom_model_info():
"""batch_cost_calculator should use model_info override when provided."""
usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
prompt_cost, completion_cost = batch_cost_calculator(
usage=usage,
model="fake-batch-model",
custom_llm_provider="openai",
model_info=CUSTOM_MODEL_INFO,
)
expected_prompt = 10 * 0.00125
expected_completion = 5 * 0.005
assert prompt_cost == pytest.approx(expected_prompt), (
f"Expected prompt cost {expected_prompt}, got {prompt_cost}"
)
assert completion_cost == pytest.approx(expected_completion), (
f"Expected completion cost {expected_completion}, got {completion_cost}"
)
def test_get_batch_job_cost_from_file_content_uses_custom_model_info():
"""_get_batch_job_cost_from_file_content should thread model_info to completion_cost."""
file_content = [_make_batch_output_line(prompt_tokens=10, completion_tokens=5)]
cost = _get_batch_job_cost_from_file_content(
file_content_dictionary=file_content,
custom_llm_provider="openai",
model_info=CUSTOM_MODEL_INFO,
)
expected = (10 * 0.00125) + (5 * 0.005)
assert cost == pytest.approx(expected), (
f"Expected total cost {expected}, got {cost}"
)
def test_batch_cost_calculator_func_uses_custom_model_info():
"""_batch_cost_calculator should thread model_info."""
file_content = [_make_batch_output_line(prompt_tokens=10, completion_tokens=5)]
cost = _batch_cost_calculator(
file_content_dictionary=file_content,
custom_llm_provider="openai",
model_info=CUSTOM_MODEL_INFO,
)
expected = (10 * 0.00125) + (5 * 0.005)
assert cost == pytest.approx(expected), (
f"Expected total cost {expected}, got {cost}"
)
@pytest.mark.asyncio
async def test_calculate_batch_cost_and_usage_uses_custom_model_info():
"""calculate_batch_cost_and_usage should thread model_info."""
file_content = [_make_batch_output_line(prompt_tokens=10, completion_tokens=5)]
batch_cost, batch_usage, batch_models = await calculate_batch_cost_and_usage(
file_content_dictionary=file_content,
custom_llm_provider="openai",
model_info=CUSTOM_MODEL_INFO,
)
expected = (10 * 0.00125) + (5 * 0.005)
assert batch_cost == pytest.approx(expected), (
f"Expected total cost {expected}, got {batch_cost}"
)
assert batch_usage.prompt_tokens == 10
assert batch_usage.completion_tokens == 5

View file

@ -1170,17 +1170,25 @@ def test_team_callback_metadata_none_values(none_key):
assert none_key not in resp
def test_proxy_config_state_post_init_callback_call():
def test_proxy_config_state_post_init_callback_call(monkeypatch):
"""
Ensures team_id is still in config, after callback is called
Addresses issue: https://github.com/BerriAI/litellm/issues/6787
Where team_id was being popped from config, after callback was called
Note: Environment variables are mocked to avoid validation errors
in parallel execution where env vars may not be set.
"""
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.proxy_server import ProxyConfig
# Mock environment variables to avoid Pydantic validation errors
# when env vars are resolved to None in parallel execution
monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "test_public_key")
monkeypatch.setenv("LANGFUSE_SECRET_KEY", "test_secret_key")
pc = ProxyConfig()
pc.update_config_state(

View file

@ -0,0 +1,67 @@
"""
Test that managed_files.afile_retrieve returns the unified file ID, not the
raw provider file ID, when file_object is already stored in the database.
Bug: managed_files.py Case 2 returns stored_file_object.file_object directly
without replacing .id with the unified ID. Case 3 (fetch from provider) does
it correctly at line 1028.
"""
import pytest
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LiteLLM_ManagedFileTable
from litellm.types.llms.openai import OpenAIFileObject
def _make_managed_files_instance():
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
instance = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=MagicMock(),
prisma_client=MagicMock(),
)
return instance
@pytest.mark.asyncio
async def test_should_return_unified_id_when_file_object_exists_in_db():
"""
When get_unified_file_id returns a stored file_object (Case 2),
afile_retrieve must set .id to the unified file ID before returning.
"""
unified_id = "bGl0ZWxsbV9wcm94eTp1bmlmaWVkX291dHB1dF9maWxl"
raw_provider_id = "batch_20260214-output-file-1"
stored = LiteLLM_ManagedFileTable(
unified_file_id=unified_id,
file_object=OpenAIFileObject(
id=raw_provider_id,
bytes=489,
created_at=1700000000,
filename="batch_output.jsonl",
object="file",
purpose="batch_output",
status="processed",
),
model_mappings={"model-abc": raw_provider_id},
flat_model_file_ids=[raw_provider_id],
created_by="test-user",
updated_by="test-user",
)
managed_files = _make_managed_files_instance()
managed_files.get_unified_file_id = AsyncMock(return_value=stored)
result = await managed_files.afile_retrieve(
file_id=unified_id,
litellm_parent_otel_span=None,
llm_router=None,
)
assert result.id == unified_id, (
f"afile_retrieve should return the unified ID '{unified_id}', "
f"but got raw provider ID '{result.id}'"
)

View file

@ -0,0 +1,75 @@
"""
Test that batch retrieve endpoint resolves raw input_file_id to the
unified managed file ID before returning.
Bug: After batch completion, batches.retrieve returns the raw provider
input_file_id instead of the LiteLLM unified ID.
"""
import base64
import json
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
DECODED_UNIFIED_INPUT_FILE_ID = "litellm_proxy:application/octet-stream;unified_id,test-uuid;target_model_names,azure-gpt-4"
B64_UNIFIED_INPUT_FILE_ID = base64.urlsafe_b64encode(DECODED_UNIFIED_INPUT_FILE_ID.encode()).decode().rstrip("=")
RAW_INPUT_FILE_ID = "file-raw-provider-abc123"
DECODED_UNIFIED_BATCH_ID = "litellm_proxy;model_id:model-xyz;llm_batch_id:batch-123"
B64_UNIFIED_BATCH_ID = base64.urlsafe_b64encode(DECODED_UNIFIED_BATCH_ID.encode()).decode().rstrip("=")
@pytest.mark.asyncio
async def test_should_resolve_raw_input_file_id_to_unified():
"""
When a completed batch has a raw input_file_id and the managed file table
contains a record for that raw ID, the retrieve endpoint should resolve
it to the unified file ID.
"""
unified_batch_id = _is_base64_encoded_unified_file_id(B64_UNIFIED_BATCH_ID)
assert unified_batch_id, "Test setup: batch_id should decode as unified"
from litellm.types.utils import LiteLLMBatch
batch_data = {
"id": B64_UNIFIED_BATCH_ID,
"completion_window": "24h",
"created_at": 1700000000,
"endpoint": "/v1/chat/completions",
"input_file_id": RAW_INPUT_FILE_ID,
"object": "batch",
"status": "completed",
"output_file_id": "file-output-xyz",
}
mock_db_object = MagicMock()
mock_db_object.file_object = json.dumps(batch_data)
mock_managed_file = MagicMock()
mock_managed_file.unified_file_id = B64_UNIFIED_INPUT_FILE_ID
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=mock_db_object)
mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=mock_managed_file)
from litellm.proxy.openai_files_endpoints.common_utils import get_batch_from_database
_, response = await get_batch_from_database(
batch_id=B64_UNIFIED_BATCH_ID,
unified_batch_id=unified_batch_id,
managed_files_obj=MagicMock(),
prisma_client=mock_prisma,
verbose_proxy_logger=MagicMock(),
)
assert response is not None, "Batch should be found in DB"
assert response.input_file_id == B64_UNIFIED_INPUT_FILE_ID, (
f"input_file_id should be unified '{B64_UNIFIED_INPUT_FILE_ID}', "
f"got raw '{response.input_file_id}'"
)

View file

@ -0,0 +1,124 @@
"""
Test that get_batch_from_database resolves raw input_file_id to the
unified/managed file ID when reading a batch from the database.
Bug: The batch retrieve path stores the raw provider input_file_id in the
DB (via async_post_call_success_hook on the retrieve endpoint). When the
batch is later read from DB, get_batch_from_database returns the raw ID
without resolving it to the unified ID.
"""
import json
import pytest
from typing import Optional
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy.openai_files_endpoints.common_utils import get_batch_from_database
def _mock_prisma(batch_json: str, managed_file_record=None):
"""Create a mock prisma client with canned responses."""
prisma = MagicMock()
batch_db_record = MagicMock()
batch_db_record.file_object = batch_json
prisma.db.litellm_managedobjecttable.find_first = AsyncMock(
return_value=batch_db_record
)
prisma.db.litellm_managedfiletable.find_first = AsyncMock(
return_value=managed_file_record
)
return prisma
@pytest.mark.asyncio
async def test_should_resolve_raw_input_file_id_to_unified_id():
"""
When input_file_id in the stored batch is a raw provider ID,
get_batch_from_database must look up the unified ID from the
managed files table.
"""
unified_batch_id = "bGl0ZWxsbV9wcm94eTpiYXRjaF9pZA"
unified_input_file_id = "bGl0ZWxsbV9wcm94eTp1bmlmaWVkX2lucHV0"
raw_input_file_id = "file-abc123-raw"
batch_data = {
"id": "batch-raw-123",
"completion_window": "24h",
"created_at": 1700000000,
"endpoint": "/v1/chat/completions",
"input_file_id": raw_input_file_id,
"object": "batch",
"status": "completed",
"output_file_id": "file-output-raw",
}
managed_file_record = MagicMock()
managed_file_record.unified_file_id = unified_input_file_id
prisma = _mock_prisma(
batch_json=json.dumps(batch_data),
managed_file_record=managed_file_record,
)
_, response = await get_batch_from_database(
batch_id=unified_batch_id,
unified_batch_id="decoded_unified_batch_id",
managed_files_obj=MagicMock(),
prisma_client=prisma,
verbose_proxy_logger=MagicMock(),
)
assert response is not None
assert response.input_file_id == unified_input_file_id, (
f"input_file_id should be resolved to '{unified_input_file_id}', "
f"got raw: '{response.input_file_id}'"
)
prisma.db.litellm_managedfiletable.find_first.assert_called_once_with(
where={"flat_model_file_ids": {"has": raw_input_file_id}}
)
@pytest.mark.asyncio
async def test_should_preserve_already_managed_input_file_id():
"""
When input_file_id is already a managed/unified ID, it should
not be modified.
"""
import base64
unified_batch_id = "bGl0ZWxsbV9wcm94eTpiYXRjaF9pZA"
decoded_unified = "litellm_proxy:application/octet-stream;unified_id,test-123"
base64_input_file_id = base64.urlsafe_b64encode(decoded_unified.encode()).decode().rstrip("=")
batch_data = {
"id": "batch-raw-123",
"completion_window": "24h",
"created_at": 1700000000,
"endpoint": "/v1/chat/completions",
"input_file_id": base64_input_file_id,
"object": "batch",
"status": "completed",
}
prisma = _mock_prisma(batch_json=json.dumps(batch_data))
_, response = await get_batch_from_database(
batch_id=unified_batch_id,
unified_batch_id="decoded_unified_batch_id",
managed_files_obj=MagicMock(),
prisma_client=prisma,
verbose_proxy_logger=MagicMock(),
)
assert response is not None
assert response.input_file_id == base64_input_file_id, (
f"input_file_id was already managed, should be preserved as '{base64_input_file_id}', "
f"got: '{response.input_file_id}'"
)
prisma.db.litellm_managedfiletable.find_first.assert_not_called()

View file

@ -0,0 +1,119 @@
"""
Regression test: deleted managed files should return 404, not 403.
When a managed file's DB record has been deleted, can_user_call_unified_file_id()
raises HTTPException(404) directly — rather than returning True (which would
weaken access control) or False (which would cause a misleading 403).
"""
import base64
import pytest
from unittest.mock import AsyncMock, MagicMock
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
def _make_user_api_key_dict(user_id: str) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="sk-test",
user_id=user_id,
parent_otel_span=None,
)
def _make_unified_file_id() -> str:
raw = "litellm_proxy:application/octet-stream;unified_id,test-deleted-file;target_model_names,azure-gpt-4"
return base64.b64encode(raw.encode()).decode()
def _make_managed_files_with_no_db_record():
"""Create a _PROXY_LiteLLMManagedFiles where the DB returns None (file was deleted)."""
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
return _PROXY_LiteLLMManagedFiles(
internal_usage_cache=MagicMock(),
prisma_client=mock_prisma,
)
@pytest.mark.asyncio
async def test_should_raise_404_for_deleted_file():
"""
When a managed file record has been deleted from the DB,
check_managed_file_id_access should raise 404 (not 403).
"""
unified_file_id = _make_unified_file_id()
managed_files = _make_managed_files_with_no_db_record()
user = _make_user_api_key_dict("any-user")
data = {"file_id": unified_file_id}
with pytest.raises(HTTPException) as exc_info:
await managed_files.check_managed_file_id_access(data, user)
assert exc_info.value.status_code == 404
@pytest.mark.asyncio
async def test_should_allow_owner_access_when_record_exists():
"""Baseline: file owner can access their own file."""
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
unified_file_id = _make_unified_file_id()
mock_db_record = MagicMock()
mock_db_record.created_by = "user-A"
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(
return_value=mock_db_record
)
managed_files = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=MagicMock(),
prisma_client=mock_prisma,
)
user = _make_user_api_key_dict("user-A")
data = {"file_id": unified_file_id}
result = await managed_files.check_managed_file_id_access(data, user)
assert result is True
@pytest.mark.asyncio
async def test_should_block_different_user_when_record_exists():
"""Baseline: different user cannot access another user's file."""
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
unified_file_id = _make_unified_file_id()
mock_db_record = MagicMock()
mock_db_record.created_by = "user-A"
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(
return_value=mock_db_record
)
managed_files = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=MagicMock(),
prisma_client=mock_prisma,
)
user = _make_user_api_key_dict("user-B")
data = {"file_id": unified_file_id}
with pytest.raises(HTTPException) as exc_info:
await managed_files.check_managed_file_id_access(data, user)
assert exc_info.value.status_code == 403

View file

@ -0,0 +1,200 @@
"""
Tests for managed files access control in batch polling context.
Regression test for: batch polling job running as default_user_id gets 403
when trying to access managed files created by a real user.
The fix (Option C) makes check_batch_cost call litellm.afile_content directly
with deployment credentials, bypassing the managed files access-control hooks.
"""
import base64
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
def _make_user_api_key_dict(user_id: str) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="sk-test",
user_id=user_id,
parent_otel_span=None,
)
def _make_unified_file_id() -> str:
"""Create a base64-encoded unified file ID that passes _is_base64_encoded_unified_file_id."""
raw = "litellm_proxy:application/octet-stream;unified_id,test-123;target_model_names,azure-gpt-4"
return base64.b64encode(raw.encode()).decode()
def _make_managed_files_instance(file_created_by: str, unified_file_id: str):
"""Create a _PROXY_LiteLLMManagedFiles with a mocked DB that returns a file owned by file_created_by."""
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
mock_db_record = MagicMock()
mock_db_record.created_by = file_created_by
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(
return_value=mock_db_record
)
instance = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=MagicMock(),
prisma_client=mock_prisma,
)
return instance
# --- Access control unit tests (document existing behavior) ---
@pytest.mark.asyncio
async def test_should_allow_file_owner_access():
"""File owner can access their own file — baseline sanity check."""
unified_file_id = _make_unified_file_id()
managed_files = _make_managed_files_instance(
file_created_by="user-A",
unified_file_id=unified_file_id,
)
user = _make_user_api_key_dict("user-A")
data = {"file_id": unified_file_id}
result = await managed_files.check_managed_file_id_access(data, user)
assert result is True
@pytest.mark.asyncio
async def test_should_block_different_user_access():
"""A different regular user cannot access another user's file — correct behavior."""
unified_file_id = _make_unified_file_id()
managed_files = _make_managed_files_instance(
file_created_by="user-A",
unified_file_id=unified_file_id,
)
user = _make_user_api_key_dict("user-B")
data = {"file_id": unified_file_id}
with pytest.raises(HTTPException) as exc_info:
await managed_files.check_managed_file_id_access(data, user)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_should_block_default_user_id_access():
"""
default_user_id is correctly blocked by the access check.
This documents the existing behavior that the Option C fix works around.
"""
unified_file_id = _make_unified_file_id()
managed_files = _make_managed_files_instance(
file_created_by="user-A",
unified_file_id=unified_file_id,
)
system_user = _make_user_api_key_dict("default_user_id")
data = {"file_id": unified_file_id}
with pytest.raises(HTTPException) as exc_info:
await managed_files.check_managed_file_id_access(data, system_user)
assert exc_info.value.status_code == 403
# --- Option C fix test: check_batch_cost bypasses managed files hook ---
@pytest.mark.asyncio
async def test_check_batch_cost_should_call_afile_content_directly_with_credentials():
"""
check_batch_cost should call litellm.afile_content directly with deployment
credentials, bypassing managed_files_obj.afile_content and its access-control
hooks. This avoids the 403 that occurs when the background job runs as
default_user_id.
"""
from litellm_enterprise.proxy.common_utils.check_batch_cost import CheckBatchCost
# Build a unified object ID in the expected format:
# litellm_proxy;model_id:{};llm_batch_id:{};llm_output_file_id:{}
unified_raw = "litellm_proxy;model_id:model-deploy-xyz;llm_batch_id:batch-123;llm_output_file_id:file-raw-output"
unified_object_id = base64.b64encode(unified_raw.encode()).decode()
# Mock a pending job from the DB
mock_job = MagicMock()
mock_job.unified_object_id = unified_object_id
mock_job.created_by = "user-A"
mock_job.id = "job-1"
# Mock prisma
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_prisma.db.litellm_managedobjecttable.update_many = AsyncMock()
# Mock proxy_logging_obj — should NOT be called for file content
mock_proxy_logging = MagicMock()
mock_managed_files_hook = MagicMock()
mock_managed_files_hook.afile_content = AsyncMock()
mock_proxy_logging.get_proxy_hook = MagicMock(return_value=mock_managed_files_hook)
# Mock the batch response (completed, with output file)
from litellm.types.utils import LiteLLMBatch
batch_response = LiteLLMBatch(
id="batch-123",
completion_window="24h",
created_at=1700000000,
endpoint="/v1/chat/completions",
input_file_id="file-input",
object="batch",
status="completed",
output_file_id="file-raw-output",
)
# Mock router
mock_router = MagicMock()
mock_router.aretrieve_batch = AsyncMock(return_value=batch_response)
mock_router.get_deployment_credentials_with_provider = MagicMock(
return_value={
"api_key": "test-key",
"api_base": "https://test.azure.com/",
"custom_llm_provider": "azure",
}
)
mock_deployment = MagicMock()
mock_deployment.litellm_params.custom_llm_provider = "azure"
mock_deployment.litellm_params.model = "azure/gpt-4"
mock_router.get_deployment = MagicMock(return_value=mock_deployment)
checker = CheckBatchCost(
proxy_logging_obj=mock_proxy_logging,
prisma_client=mock_prisma,
llm_router=mock_router,
)
mock_file_content = MagicMock()
mock_file_content.content = b'{"id":"req-1","response":{"status_code":200,"body":{"id":"cmpl-1","object":"chat.completion","created":1700000000,"model":"gpt-4","choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}],"usage":{"prompt_tokens":10,"completion_tokens":5,"total_tokens":15}}}}\n'
with patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
) as mock_direct_afile_content:
await checker.check_batch_cost()
# afile_content should be called directly (not through managed_files_obj)
mock_direct_afile_content.assert_called_once()
call_kwargs = mock_direct_afile_content.call_args.kwargs
assert call_kwargs.get("api_key") == "test-key", (
f"afile_content should receive api_key from deployment credentials. "
f"Got: {call_kwargs}"
)
# managed_files_obj.afile_content should NOT have been called
mock_managed_files_hook.afile_content.assert_not_called()

View file

@ -0,0 +1,167 @@
"""
Tests for enterprise/litellm_enterprise/proxy/hooks/managed_files.py
Regression test for afile_retrieve called without credentials in
async_post_call_success_hook when processing completed batch responses.
"""
import pytest
from typing import Optional
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import OpenAIFileObject
from litellm.types.utils import LiteLLMBatch
def _make_file_object(file_id: str = "file-output-abc") -> OpenAIFileObject:
return OpenAIFileObject(
id=file_id,
bytes=100,
created_at=1700000000,
filename="output.jsonl",
object="file",
purpose="batch_output",
status="processed",
)
def _make_batch_response(
batch_id: str = "batch-123",
output_file_id: Optional[str] = "file-output-abc",
status: str = "completed",
model_id: str = "model-deploy-xyz",
model_name: str = "azure/gpt-4",
) -> LiteLLMBatch:
"""Create a LiteLLMBatch response with hidden params set as the router would."""
batch = LiteLLMBatch(
id=batch_id,
completion_window="24h",
created_at=1700000000,
endpoint="/v1/chat/completions",
input_file_id="file-input-abc",
object="batch",
status=status,
output_file_id=output_file_id,
)
batch._hidden_params = {
"unified_file_id": "some-unified-id",
"unified_batch_id": "some-unified-batch-id",
"model_id": model_id,
"model_name": model_name,
}
return batch
def _make_user_api_key_dict() -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="sk-test",
user_id="test-user",
parent_otel_span=None,
)
def _make_managed_files_instance():
"""Create a _PROXY_LiteLLMManagedFiles with storage methods mocked out."""
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
mock_cache = MagicMock()
mock_prisma = MagicMock()
instance = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=mock_cache,
prisma_client=mock_prisma,
)
instance.store_unified_file_id = AsyncMock()
instance.store_unified_object_id = AsyncMock()
return instance
@pytest.mark.asyncio
async def test_should_pass_credentials_to_afile_retrieve():
"""
When async_post_call_success_hook processes a completed batch with an output_file_id,
it calls afile_retrieve to fetch file metadata. It must pass credentials from the
router deployment, not just custom_llm_provider and file_id.
Regression test for: managed_files.py:919 calling afile_retrieve without api_key/api_base.
"""
managed_files = _make_managed_files_instance()
batch_response = _make_batch_response(
model_id="model-deploy-xyz",
model_name="azure/gpt-4",
output_file_id="file-output-abc",
)
user_api_key_dict = _make_user_api_key_dict()
mock_credentials = {
"api_key": "test-azure-key",
"api_base": "https://my-azure.openai.azure.com/",
"api_version": "2025-03-01-preview",
"custom_llm_provider": "azure",
}
mock_router = MagicMock()
mock_router.get_deployment_credentials_with_provider = MagicMock(
return_value=mock_credentials
)
mock_afile_retrieve = AsyncMock(return_value=_make_file_object("file-output-abc"))
with patch(
"litellm.afile_retrieve", mock_afile_retrieve
), patch(
"litellm.proxy.proxy_server.llm_router", mock_router
):
await managed_files.async_post_call_success_hook(
data={},
user_api_key_dict=user_api_key_dict,
response=batch_response,
)
mock_afile_retrieve.assert_called()
call_kwargs = mock_afile_retrieve.call_args
assert call_kwargs.kwargs.get("api_key") == "test-azure-key", (
f"afile_retrieve must receive api_key from router credentials. "
f"Got kwargs: {call_kwargs.kwargs}"
)
assert call_kwargs.kwargs.get("api_base") == "https://my-azure.openai.azure.com/", (
f"afile_retrieve must receive api_base from router credentials. "
f"Got kwargs: {call_kwargs.kwargs}"
)
@pytest.mark.asyncio
async def test_should_fallback_when_no_router():
"""
When llm_router is not available, afile_retrieve should still be called
with the fallback behavior (custom_llm_provider extracted from model_name).
"""
managed_files = _make_managed_files_instance()
batch_response = _make_batch_response(
model_id="model-deploy-xyz",
model_name="azure/gpt-4",
output_file_id="file-output-abc",
)
user_api_key_dict = _make_user_api_key_dict()
mock_afile_retrieve = AsyncMock(return_value=_make_file_object("file-output-abc"))
with patch(
"litellm.afile_retrieve", mock_afile_retrieve
), patch(
"litellm.proxy.proxy_server.llm_router", None
):
await managed_files.async_post_call_success_hook(
data={},
user_api_key_dict=user_api_key_dict,
response=batch_response,
)
mock_afile_retrieve.assert_called()
call_kwargs = mock_afile_retrieve.call_args
assert call_kwargs.kwargs.get("custom_llm_provider") == "azure"
assert call_kwargs.kwargs.get("file_id") == "file-output-abc"

View file

@ -268,22 +268,21 @@ class TestLangfuseUsageDetails(unittest.TestCase):
Test that _log_langfuse_v2 correctly handles None values in the usage object
by converting them to 0, preventing validation errors.
"""
# Create fresh mocks for this test to avoid state pollution from setUp's side_effect
# The setUp configures trace.side_effect which can interfere with return_value
mock_trace = MagicMock()
mock_generation = MagicMock()
mock_generation.trace_id = "test-trace-id"
# Reset the mock to ensure clean state
self.mock_langfuse_client.reset_mock()
self.mock_langfuse_trace.reset_mock()
self.mock_langfuse_generation.reset_mock()
# Re-setup the trace and generation chain with clean state
self.mock_langfuse_generation.trace_id = "test-trace-id"
mock_span = MagicMock()
mock_span.end = MagicMock()
mock_trace.generation.return_value = mock_generation
mock_trace.span.return_value = mock_span
mock_client = MagicMock()
mock_client.trace.return_value = mock_trace
# Use our fresh mock client
self.logger.Langfuse = mock_client
self.mock_langfuse_trace.span.return_value = mock_span
self.mock_langfuse_trace.generation.return_value = self.mock_langfuse_generation
# Ensure trace returns our mock
self.mock_langfuse_client.trace.return_value = self.mock_langfuse_trace
self.logger.Langfuse = self.mock_langfuse_client
with patch(
"litellm.integrations.langfuse.langfuse._add_prompt_to_generation_params",
@ -337,13 +336,13 @@ class TestLangfuseUsageDetails(unittest.TestCase):
)
except Exception as e:
self.fail(f"_log_langfuse_v2 raised an exception: {e}")
# Verify that trace was called first
mock_client.trace.assert_called()
self.mock_langfuse_client.trace.assert_called()
# Check the arguments passed to the mocked langfuse generation call
mock_trace.generation.assert_called_once()
call_args, call_kwargs = mock_trace.generation.call_args
self.mock_langfuse_trace.generation.assert_called_once()
call_args, call_kwargs = self.mock_langfuse_trace.generation.call_args
# Inspect the usage and usage_details dictionaries
usage_arg = call_kwargs.get("usage")

View file

@ -157,6 +157,186 @@ class TestS3V2UnitTests:
assert result == {"downloaded": "data"}
@patch('asyncio.create_task')
@patch('litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush')
def test_s3_v2_virtual_hosted_style(self, mock_periodic_flush, mock_create_task):
"""Test s3_use_virtual_hosted_style parameter for virtual-hosted-style URLs"""
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: Virtual-hosted-style with custom endpoint
s3_logger_virtual = S3Logger(
s3_bucket_name="test-bucket",
s3_endpoint_url="https://s3.custom-endpoint.com",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_use_virtual_hosted_style=True
)
s3_logger_virtual.async_httpx_client = AsyncMock()
s3_logger_virtual.async_httpx_client.put.return_value = mock_response
asyncio.run(s3_logger_virtual.async_upload_data_to_s3(test_element))
call_args = s3_logger_virtual.async_httpx_client.put.call_args
assert call_args is not None
url = call_args[0][0]
expected_url = "https://test-bucket.s3.custom-endpoint.com/2025-09-14/test-key.json"
assert url == expected_url, f"Expected virtual-hosted-style URL {expected_url}, got {url}"
# Test 2: Path-style (default behavior with s3_use_virtual_hosted_style=False)
s3_logger_path = S3Logger(
s3_bucket_name="test-bucket",
s3_endpoint_url="https://s3.custom-endpoint.com",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_use_virtual_hosted_style=False
)
s3_logger_path.async_httpx_client = AsyncMock()
s3_logger_path.async_httpx_client.put.return_value = mock_response
asyncio.run(s3_logger_path.async_upload_data_to_s3(test_element))
call_args_path = s3_logger_path.async_httpx_client.put.call_args
assert call_args_path is not None
url_path = call_args_path[0][0]
expected_path_url = "https://s3.custom-endpoint.com/test-bucket/2025-09-14/test-key.json"
assert url_path == expected_path_url, f"Expected path-style URL {expected_path_url}, got {url_path}"
# Test 3: Virtual-hosted-style with http protocol
s3_logger_http = S3Logger(
s3_bucket_name="http-bucket",
s3_endpoint_url="http://minio.local:9000",
s3_aws_access_key_id="minio-key",
s3_aws_secret_access_key="minio-secret",
s3_region_name="us-east-1",
s3_use_virtual_hosted_style=True
)
s3_logger_http.async_httpx_client = AsyncMock()
s3_logger_http.async_httpx_client.put.return_value = mock_response
asyncio.run(s3_logger_http.async_upload_data_to_s3(test_element))
call_args_http = s3_logger_http.async_httpx_client.put.call_args
assert call_args_http is not None
url_http = call_args_http[0][0]
expected_http_url = "http://http-bucket.minio.local:9000/2025-09-14/test-key.json"
assert url_http == expected_http_url, f"Expected virtual-hosted-style URL with http {expected_http_url}, got {url_http}"
# Test 4: Sync upload method with virtual-hosted-style
s3_logger_sync_virtual = S3Logger(
s3_bucket_name="sync-bucket",
s3_endpoint_url="https://storage.example.com",
s3_aws_access_key_id="sync-key",
s3_aws_secret_access_key="sync-secret",
s3_region_name="us-east-1",
s3_use_virtual_hosted_style=True
)
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_virtual.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://sync-bucket.storage.example.com/2025-09-14/test-key.json"
assert url_sync == expected_sync_url, f"Expected virtual-hosted-style sync URL {expected_sync_url}, got {url_sync}"
# Test 5: Download method with virtual-hosted-style
s3_logger_download_virtual = S3Logger(
s3_bucket_name="download-bucket",
s3_endpoint_url="https://download.endpoint.com",
s3_aws_access_key_id="download-key",
s3_aws_secret_access_key="download-secret",
s3_region_name="us-east-1",
s3_use_virtual_hosted_style=True
)
mock_download_response = MagicMock()
mock_download_response.status_code = 200
mock_download_response.json = MagicMock(return_value={"downloaded": "data"})
s3_logger_download_virtual.async_httpx_client = AsyncMock()
s3_logger_download_virtual.async_httpx_client.get.return_value = mock_download_response
result = asyncio.run(s3_logger_download_virtual._download_object_from_s3("2025-09-14/download-test-key.json"))
call_args_download = s3_logger_download_virtual.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-bucket.download.endpoint.com/2025-09-14/download-test-key.json"
assert url_download == expected_download_url, f"Expected virtual-hosted-style download URL {expected_download_url}, got {url_download}"
assert result == {"downloaded": "data"}
@pytest.mark.asyncio
async def test_async_log_event_skips_when_standard_logging_object_missing():
"""
Reproduces the bug where _async_log_event_base raises ValueError when
kwargs has no standard_logging_object (e.g. call_type=afile_delete).
The S3 logger should skip gracefully, not raise.
"""
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_region_name="us-east-1",
s3_aws_access_key_id="fake",
s3_aws_secret_access_key="fake",
)
kwargs_without_slo = {
"call_type": "afile_delete",
"model": None,
"litellm_call_id": "test-call-id",
}
start_time = datetime.utcnow()
end_time = datetime.utcnow()
# Spy on handle_callback_failure — should NOT be called if we skip gracefully.
# Without the fix, the ValueError is caught by the except block which calls
# handle_callback_failure. With the fix, we return early and never hit except.
with patch.object(logger, "handle_callback_failure") as mock_failure:
await logger._async_log_event_base(
kwargs=kwargs_without_slo,
response_obj=None,
start_time=start_time,
end_time=end_time,
)
assert not mock_failure.called, (
"handle_callback_failure should not be called — "
"missing standard_logging_object should be a graceful skip, not an error"
)
# Nothing should have been queued (catches the case where code falls
# through without returning and appends None to the queue)
assert len(logger.log_queue) == 0, "log_queue should be empty when standard_logging_object is missing"
@pytest.mark.asyncio
async def test_strip_base64_removes_file_and_nontext_entries():
logger = S3Logger(s3_strip_base64_files=True)

View file

@ -491,7 +491,20 @@ from unittest.mock import MagicMock, patch
from litellm.utils import _select_tokenizer_helper, claude_json_str, encoding
# Clear the cache at module load to ensure clean state
_select_tokenizer_helper.cache_clear()
class TestTokenizerSelection(unittest.TestCase):
def setUp(self):
"""Clear the LRU cache before each test method.
The _select_tokenizer_helper function is decorated with @lru_cache,
which can cause cache hits from previous tests when running with
--dist=loadscope (tests from same file run on same worker).
"""
_select_tokenizer_helper.cache_clear()
@patch("litellm.utils.Tokenizer.from_pretrained")
def test_llama3_tokenizer_api_failure(self, mock_from_pretrained):
# Setup mock to raise an error

View file

@ -1638,6 +1638,41 @@ def test_effort_with_claude_opus_45():
assert result["model"] == "claude-opus-4-5-20251101"
def test_effort_validation_with_opus_46():
"""Test that all four effort levels are accepted for Claude Opus 4.6."""
config = AnthropicConfig()
messages = [{"role": "user", "content": "Test"}]
for effort in ["high", "medium", "low", "max"]:
optional_params = {"output_config": {"effort": effort}}
result = config.transform_request(
model="claude-opus-4-6-20260205",
messages=messages,
optional_params=optional_params,
litellm_params={},
headers={}
)
assert result["output_config"]["effort"] == effort
def test_max_effort_rejected_for_opus_45():
"""Test that effort='max' is rejected when using Claude Opus 4.5."""
config = AnthropicConfig()
messages = [{"role": "user", "content": "Test"}]
with pytest.raises(ValueError, match="effort='max' is only supported by Claude Opus 4.6"):
optional_params = {"output_config": {"effort": "max"}}
config.transform_request(
model="claude-opus-4-5-20251101",
messages=messages,
optional_params=optional_params,
litellm_params={},
headers={}
)
def test_effort_with_other_features():
"""Test effort works alongside other features (thinking, tools)."""
config = AnthropicConfig()

View file

@ -1706,3 +1706,108 @@ def test_translate_openai_response_restores_tool_names():
assert len(tool_use_blocks) == 1
# Name should be restored to original
assert tool_use_blocks[0]["name"] == original_name
def test_translate_openai_response_to_anthropic_input_tokens_excludes_cached_tokens():
"""
Regression test: input_tokens in Anthropic format should NOT include cached tokens.
Issue: v1/messages API was returning incorrect input_token count when using prompt caching.
The OpenAI format includes cached tokens in prompt_tokens, but Anthropic format should not.
According to Anthropic's spec:
- input_tokens = uncached input tokens only
- cache_read_input_tokens = tokens read from cache
In OpenAI format:
- prompt_tokens = all input tokens (including cached)
- prompt_tokens_details.cached_tokens = cached tokens
Expected: anthropic.input_tokens = openai.prompt_tokens - openai.prompt_tokens_details.cached_tokens
"""
from litellm.types.utils import PromptTokensDetailsWrapper
# Create OpenAI format response with cached tokens
# Scenario: 100 total prompt tokens, 30 of which are cached
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=30
),
cache_read_input_tokens=30, # Anthropic format cache info
)
response = ModelResponse(
id="test-id",
choices=[
Choices(
index=0,
finish_reason="stop",
message=Message(
role="assistant",
content="Test response",
),
)
],
model="claude-3-sonnet-20240229",
usage=usage,
)
# Convert to Anthropic format
adapter = LiteLLMAnthropicMessagesAdapter()
anthropic_response = adapter.translate_openai_response_to_anthropic(
response=response,
tool_name_mapping=None,
)
# Validate: input_tokens should be 70 (100 - 30 cached), not 100
assert anthropic_response["usage"]["input_tokens"] == 70, (
f"Expected input_tokens=70 (100 total - 30 cached), "
f"but got {anthropic_response['usage']['input_tokens']}. "
f"input_tokens should NOT include cached tokens per Anthropic spec."
)
assert anthropic_response["usage"]["output_tokens"] == 50
assert anthropic_response["usage"]["cache_read_input_tokens"] == 30
def test_translate_openai_response_to_anthropic_input_tokens_no_cache():
"""
Regression test: input_tokens should equal prompt_tokens when there are no cached tokens.
"""
from litellm.types.utils import PromptTokensDetailsWrapper
# Create OpenAI format response without cached tokens
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
)
response = ModelResponse(
id="test-id",
choices=[
Choices(
index=0,
finish_reason="stop",
message=Message(
role="assistant",
content="Test response",
),
)
],
model="claude-3-sonnet-20240229",
usage=usage,
)
# Convert to Anthropic format
adapter = LiteLLMAnthropicMessagesAdapter()
anthropic_response = adapter.translate_openai_response_to_anthropic(
response=response,
tool_name_mapping=None,
)
# Validate: input_tokens should equal prompt_tokens when no caching
assert anthropic_response["usage"]["input_tokens"] == 100
assert anthropic_response["usage"]["output_tokens"] == 50

View file

@ -5,10 +5,9 @@ This test file verifies that Pydantic models with various constraints
are properly converted to Anthropic-compatible JSON schemas.
"""
from typing import List
import pytest
from pydantic import BaseModel, Field
from typing import List
class TestAnthropicStructuredOutput:
@ -24,7 +23,7 @@ class TestAnthropicStructuredOutput:
Related issue: https://github.com/BerriAI/litellm/issues/19444
"""
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
# Define a Pydantic model with max_length on a List field
class ResponseModel(BaseModel):
items: List[str] = Field(max_length=5, description="List of items")
@ -131,12 +130,10 @@ class TestAnthropicStructuredOutput:
def test_other_constraints_preserved(self):
"""
Test that string and numeric constraints are moved to description.
Test that constraints are properly handled (removed from schema, added to description).
Anthropic's output_format API doesn't support minLength/maxLength for
strings or minimum/maximum for numbers. Per Anthropic's SDK approach,
these constraints are removed from the schema and added to the
description text instead.
Per Anthropic API requirements, constraints like minLength/maxLength and
minimum/maximum must be removed from the schema but documented in descriptions.
"""
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
@ -150,7 +147,7 @@ class TestAnthropicStructuredOutput:
response_format = {
"type": "json_schema",
"json_schema": json_schema["json_schema"],
"json_schema": json_schema["json_schema"]
}
output_format = config.map_response_format_to_anthropic_output_format(
@ -160,16 +157,20 @@ class TestAnthropicStructuredOutput:
assert output_format is not None
transformed_schema = output_format["schema"]
# String constraints moved to description (not preserved in schema)
# String constraints should be REMOVED from schema (Anthropic doesn't support them)
name_schema = transformed_schema["properties"]["name"]
assert "maxLength" not in name_schema
assert "minLength" not in name_schema
assert "maximum length: 100" in name_schema["description"]
# But constraint info should be added to description
assert "description" in name_schema
assert "minimum length: 1" in name_schema["description"]
assert "maximum length: 100" in name_schema["description"]
# Number constraints moved to description (not preserved in schema)
# Number constraints should be REMOVED from schema (Anthropic doesn't support them)
age_schema = transformed_schema["properties"]["age"]
assert "minimum" not in age_schema
assert "maximum" not in age_schema
# But constraint info should be added to description
assert "description" in age_schema
assert "minimum value: 0" in age_schema["description"]
assert "maximum value: 150" in age_schema["description"]

View file

@ -2934,3 +2934,47 @@ def test_drop_thinking_param_when_thinking_blocks_missing():
finally:
# Restore original modify_params setting
litellm.modify_params = original_modify_params
class TestBedrockMinThinkingBudgetTokens:
"""Test that thinking.budget_tokens is clamped to the Bedrock minimum (1024)."""
def _map_params(
self, thinking_value, model="anthropic.claude-3-7-sonnet-20250219-v1:0"
):
"""Helper to call map_openai_params with the given thinking value."""
config = AmazonConverseConfig()
non_default_params = {"thinking": thinking_value}
optional_params = {"thinking": thinking_value}
return config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=False,
)
def test_budget_tokens_below_minimum_is_clamped(self):
"""budget_tokens < 1024 should be clamped to 1024."""
result = self._map_params({"type": "enabled", "budget_tokens": 499})
assert result["thinking"]["budget_tokens"] == 1024
def test_budget_tokens_at_minimum_is_unchanged(self):
"""budget_tokens == 1024 should remain 1024."""
result = self._map_params({"type": "enabled", "budget_tokens": 1024})
assert result["thinking"]["budget_tokens"] == 1024
def test_budget_tokens_above_minimum_is_unchanged(self):
"""budget_tokens > 1024 should remain unchanged."""
result = self._map_params({"type": "enabled", "budget_tokens": 2048})
assert result["thinking"]["budget_tokens"] == 2048
def test_no_thinking_param_does_not_error(self):
"""When thinking is not provided, map_openai_params should not raise."""
config = AmazonConverseConfig()
result = config.map_openai_params(
non_default_params={},
optional_params={},
model="anthropic.claude-3-7-sonnet-20250219-v1:0",
drop_params=False,
)
assert "thinking" not in result or result.get("thinking") is None

View file

@ -88,6 +88,45 @@ class TestChatGPTResponsesAPITransformation:
"You are Codex, based on GPT-5."
)
def test_chatgpt_drops_unsupported_responses_params(self):
config = ChatGPTResponsesAPIConfig()
request = config.transform_responses_api_request(
model="chatgpt/gpt-5.2-codex",
input="hi",
response_api_optional_request_params={
# unsupported by ChatGPT Codex
"user": "user_123",
"temperature": 0.2,
"top_p": 0.9,
"context_management": [{"type": "compaction", "compact_threshold": 200000}],
"metadata": {"foo": "bar"},
"max_output_tokens": 123,
"stream_options": {"include_usage": True},
# supported and should be preserved
"truncation": "auto",
"previous_response_id": "resp_123",
"reasoning": {"effort": "medium"},
"tools": [{"type": "function", "function": {"name": "hello"}}],
"tool_choice": {"type": "function", "function": {"name": "hello"}},
},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert "user" not in request
assert "temperature" not in request
assert "top_p" not in request
assert "context_management" not in request
assert "metadata" not in request
assert "max_output_tokens" not in request
assert "stream_options" not in request
assert request["truncation"] == "auto"
assert request["previous_response_id"] == "resp_123"
assert request["reasoning"] == {"effort": "medium"}
assert request["tools"] == [{"type": "function", "function": {"name": "hello"}}]
assert request["tool_choice"] == {"type": "function", "function": {"name": "hello"}}
def test_chatgpt_non_stream_sse_response_parsing(self):
config = ChatGPTResponsesAPIConfig()
response_payload = {

View file

@ -12,10 +12,42 @@ sys.path.insert(0, os.path.abspath("../../../..")) # Adds the parent directory
from litellm.llms.custom_httpx.aiohttp_transport import (
AiohttpResponseStream,
AiohttpTransport,
LiteLLMAiohttpTransport,
)
@pytest.mark.asyncio
async def test_aclose_does_not_close_shared_session():
"""Test that aclose() does not close a session it does not own (shared session)."""
session = aiohttp.ClientSession()
try:
transport = LiteLLMAiohttpTransport(client=session, owns_session=False)
await transport.aclose()
assert not session.closed, "Shared session should not be closed by transport"
finally:
await session.close()
@pytest.mark.asyncio
async def test_aclose_closes_owned_session():
"""Test that aclose() closes a session it owns."""
session = aiohttp.ClientSession()
transport = LiteLLMAiohttpTransport(client=session, owns_session=True)
await transport.aclose()
assert session.closed, "Owned session should be closed by transport"
@pytest.mark.asyncio
async def test_owns_session_defaults_to_true():
"""Test that owns_session defaults to True for backwards compatibility."""
session = aiohttp.ClientSession()
transport = AiohttpTransport(client=session)
assert transport._owns_session is True
await transport.aclose()
assert session.closed
class MockAiohttpResponse:
"""Mock aiohttp ClientResponse for testing"""

View file

@ -73,7 +73,7 @@ class TestVertexAIGPTOSSTransformation:
@pytest.mark.asyncio
async def test_vertex_ai_gpt_oss_simple_request():
"""
Test that a simple request to vertex_ai/openai/gpt-oss-20b-maas lands at the correct URL
Test that a simple request to vertex_ai/openai/gpt-oss-20b-maas lands at the correct URL
with the correct request body.
"""
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
@ -106,14 +106,20 @@ async def test_vertex_ai_gpt_oss_simple_request():
"total_tokens": 70
}
}
client = AsyncHTTPHandler()
async def mock_post_func(*args, **kwargs):
return mock_response
# Mock vertexai module to prevent import from triggering authentication
mock_vertexai = MagicMock()
mock_vertexai.preview = MagicMock()
mock_vertexai.preview.language_models = MagicMock()
with patch.object(client, "post", side_effect=mock_post_func) as mock_post, \
patch.object(VertexLLM, "_ensure_access_token", return_value=("fake-token", "pathrise-convert-1606954137718")):
patch.object(VertexLLM, "_ensure_access_token", return_value=("fake-token", "pathrise-convert-1606954137718")), \
patch.dict('sys.modules', {'vertexai': mock_vertexai, 'vertexai.preview': mock_vertexai.preview}):
response = await litellm.acompletion(
model="vertex_ai/openai/gpt-oss-20b-maas",
messages=[
@ -122,7 +128,7 @@ async def test_vertex_ai_gpt_oss_simple_request():
"content": "Your name is Litellm Bot, you are a helpful assistant"
},
{
"role": "user",
"role": "user",
"content": "Hello, what is your name and can you tell me the weather?"
}
],
@ -170,7 +176,7 @@ async def test_vertex_ai_gpt_oss_simple_request():
@pytest.mark.asyncio
async def test_vertex_ai_gpt_oss_reasoning_effort():
"""
Test that reasoning_effort parameter is correctly passed in the request body
Test that reasoning_effort parameter is correctly passed in the request body
for GPT-OSS models.
"""
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
@ -184,7 +190,7 @@ async def test_vertex_ai_gpt_oss_reasoning_effort():
mock_response.headers = {}
mock_response.json.return_value = {
"id": "chatcmpl-test456",
"object": "chat.completion",
"object": "chat.completion",
"created": 1234567890,
"model": "openai/gpt-oss-20b-maas",
"choices": [
@ -203,14 +209,20 @@ async def test_vertex_ai_gpt_oss_reasoning_effort():
"total_tokens": 67
}
}
client = AsyncHTTPHandler()
async def mock_post_func(*args, **kwargs):
return mock_response
# Mock vertexai module to prevent import from triggering authentication
mock_vertexai = MagicMock()
mock_vertexai.preview = MagicMock()
mock_vertexai.preview.language_models = MagicMock()
with patch.object(client, "post", side_effect=mock_post_func) as mock_post, \
patch.object(VertexLLM, "_ensure_access_token", return_value=("fake-token", "pathrise-convert-1606954137718")):
patch.object(VertexLLM, "_ensure_access_token", return_value=("fake-token", "pathrise-convert-1606954137718")), \
patch.dict('sys.modules', {'vertexai': mock_vertexai, 'vertexai.preview': mock_vertexai.preview}):
response = await litellm.acompletion(
model="vertex_ai/openai/gpt-oss-20b-maas",
messages=[

View file

@ -4,7 +4,7 @@ Test key rotation manager functionality
import os
import sys
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock
import pytest
@ -24,7 +24,7 @@ class TestKeyRotationManager:
async def test_should_rotate_key_logic(self):
"""
Test the core logic for determining when a key should be rotated.
This tests:
- Keys with null key_rotation_at should rotate immediately
- Keys with future key_rotation_at should not rotate
@ -33,69 +33,69 @@ class TestKeyRotationManager:
# Setup
mock_prisma_client = AsyncMock()
manager = KeyRotationManager(mock_prisma_client)
now = datetime.now(timezone.utc)
# Test Case 1: No rotation time set (key_rotation_at = None) - should rotate
key_no_rotation_time = LiteLLM_VerificationToken(
token="test-token-1",
auto_rotate=True,
rotation_interval="30s",
key_rotation_at=None,
rotation_count=0
rotation_count=0,
)
assert manager._should_rotate_key(key_no_rotation_time, now) == True
assert manager._should_rotate_key(key_no_rotation_time, now) is True
# Test Case 2: Future rotation time - should NOT rotate
key_future_rotation = LiteLLM_VerificationToken(
token="test-token-2",
auto_rotate=True,
rotation_interval="30s",
key_rotation_at=now + timedelta(seconds=10),
rotation_count=1
rotation_count=1,
)
assert manager._should_rotate_key(key_future_rotation, now) == False
assert manager._should_rotate_key(key_future_rotation, now) is False
# Test Case 3: Past rotation time - should rotate
key_past_rotation = LiteLLM_VerificationToken(
token="test-token-3",
auto_rotate=True,
rotation_interval="30s",
key_rotation_at=now - timedelta(seconds=10),
rotation_count=2
rotation_count=2,
)
assert manager._should_rotate_key(key_past_rotation, now) == True
assert manager._should_rotate_key(key_past_rotation, now) is True
# Test Case 4: Exact rotation time - should rotate
key_exact_rotation = LiteLLM_VerificationToken(
token="test-token-4",
auto_rotate=True,
rotation_interval="30s",
key_rotation_at=now,
rotation_count=1
rotation_count=1,
)
assert manager._should_rotate_key(key_exact_rotation, now) == True
assert manager._should_rotate_key(key_exact_rotation, now) is True
# Test Case 5: No rotation interval - should NOT rotate
key_no_interval = LiteLLM_VerificationToken(
token="test-token-5",
auto_rotate=True,
rotation_interval=None,
key_rotation_at=None,
rotation_count=0
rotation_count=0,
)
assert manager._should_rotate_key(key_no_interval, now) == False
assert manager._should_rotate_key(key_no_interval, now) is False
@pytest.mark.asyncio
async def test_find_keys_needing_rotation(self):
"""
Test finding keys that need rotation from database.
This tests:
- Only keys with auto_rotate=True are considered
- Database query filters by key_rotation_at properly
@ -104,10 +104,10 @@ class TestKeyRotationManager:
# Setup
mock_prisma_client = AsyncMock()
manager = KeyRotationManager(mock_prisma_client)
# Use a fixed timestamp to avoid timing issues in tests
now = datetime(2024, 1, 1, 12, 0, 0, tzinfo=timezone.utc)
# Mock database response - these are the keys the database query would return
mock_keys = [
LiteLLM_VerificationToken(
@ -115,42 +115,47 @@ class TestKeyRotationManager:
auto_rotate=True,
rotation_interval="30s",
key_rotation_at=None, # Should rotate (null key_rotation_at)
rotation_count=0
rotation_count=0,
),
LiteLLM_VerificationToken(
token="token-2",
auto_rotate=True,
rotation_interval="60s",
key_rotation_at=now - timedelta(seconds=10), # Should rotate (past time)
rotation_count=1
)
key_rotation_at=now
- timedelta(seconds=10), # Should rotate (past time)
rotation_count=1,
),
]
mock_prisma_client.db.litellm_verificationtoken.find_many.return_value = mock_keys
mock_prisma_client.db.litellm_verificationtoken.find_many.return_value = (
mock_keys
)
# Mock datetime.now to return our fixed timestamp
from unittest.mock import patch
with patch('litellm.proxy.common_utils.key_rotation_manager.datetime') as mock_datetime:
with patch(
"litellm.proxy.common_utils.key_rotation_manager.datetime"
) as mock_datetime:
mock_datetime.now.return_value = now
mock_datetime.side_effect = lambda *args, **kwargs: datetime(*args, **kwargs)
mock_datetime.side_effect = lambda *args, **kwargs: datetime(
*args, **kwargs
)
# Execute
keys_needing_rotation = await manager._find_keys_needing_rotation()
# Verify database query - should use OR condition for key_rotation_at
mock_prisma_client.db.litellm_verificationtoken.find_many.assert_called_once_with(
where={
"auto_rotate": True,
"OR": [
{"key_rotation_at": None},
{"key_rotation_at": {"lte": now}}
]
"OR": [{"key_rotation_at": None}, {"key_rotation_at": {"lte": now}}],
}
)
# Verify all keys returned by database query are included (no additional filtering)
assert len(keys_needing_rotation) == 2
tokens_needing_rotation = [key.token for key in keys_needing_rotation]
assert "token-1" in tokens_needing_rotation # Null key_rotation_at
assert "token-2" in tokens_needing_rotation # Past key_rotation_at
@ -159,7 +164,7 @@ class TestKeyRotationManager:
async def test_rotate_key_updates_database(self):
"""
Test that key rotation properly updates the database with new rotation info.
This tests:
- Rotation count is incremented
- last_rotation_at is set to current time
@ -169,7 +174,7 @@ class TestKeyRotationManager:
# Setup
mock_prisma_client = AsyncMock()
manager = KeyRotationManager(mock_prisma_client)
# Mock key to rotate
key_to_rotate = LiteLLM_VerificationToken(
token="old-token",
@ -177,31 +182,35 @@ class TestKeyRotationManager:
rotation_interval="30s",
last_rotation_at=None,
key_rotation_at=None,
rotation_count=0
rotation_count=0,
)
# Mock regenerate_key_fn response
mock_response = GenerateKeyResponse(
key="new-api-key",
token_id="new-token-id",
user_id="test-user"
key="new-api-key", token_id="new-token-id", user_id="test-user"
)
# Mock the regenerate function
from unittest.mock import patch
with patch('litellm.proxy.common_utils.key_rotation_manager.regenerate_key_fn', return_value=mock_response):
with patch('litellm.proxy.common_utils.key_rotation_manager.KeyManagementEventHooks.async_key_rotated_hook'):
with patch(
"litellm.proxy.common_utils.key_rotation_manager.regenerate_key_fn",
return_value=mock_response,
):
with patch(
"litellm.proxy.common_utils.key_rotation_manager.KeyManagementEventHooks.async_key_rotated_hook"
):
# Execute
await manager._rotate_key(key_to_rotate)
# Verify database update was called with correct data
mock_prisma_client.db.litellm_verificationtoken.update.assert_called_once()
call_args = mock_prisma_client.db.litellm_verificationtoken.update.call_args
# Check the WHERE clause targets the new token
assert call_args[1]["where"]["token"] == "new-token-id"
# Check the data being updated
update_data = call_args[1]["data"]
assert update_data["rotation_count"] == 1 # Incremented from 0
@ -209,9 +218,75 @@ class TestKeyRotationManager:
assert isinstance(update_data["last_rotation_at"], datetime)
assert "key_rotation_at" in update_data
assert isinstance(update_data["key_rotation_at"], datetime)
# Verify key_rotation_at is set to future time (30s from now)
now = datetime.now(timezone.utc)
next_rotation = update_data["key_rotation_at"]
time_diff = (next_rotation - now).total_seconds()
assert 25 <= time_diff <= 35 # Should be around 30 seconds, allow some tolerance
assert (
25 <= time_diff <= 35
) # Should be around 30 seconds, allow some tolerance
@pytest.mark.asyncio
async def test_cleanup_expired_deprecated_keys(self):
"""
Test that _cleanup_expired_deprecated_keys deletes expired deprecated keys.
"""
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_deprecatedverificationtoken.delete_many.return_value = (
3
)
manager = KeyRotationManager(mock_prisma_client)
await manager._cleanup_expired_deprecated_keys()
mock_prisma_client.db.litellm_deprecatedverificationtoken.delete_many.assert_called_once()
call_args = (
mock_prisma_client.db.litellm_deprecatedverificationtoken.delete_many.call_args
)
assert "revoke_at" in call_args[1]["where"]
assert call_args[1]["where"]["revoke_at"]["lt"] is not None
@pytest.mark.asyncio
async def test_rotate_key_passes_grace_period(self):
"""
Test that _rotate_key passes grace_period in RegenerateKeyRequest.
"""
mock_prisma_client = AsyncMock()
manager = KeyRotationManager(mock_prisma_client)
key_to_rotate = LiteLLM_VerificationToken(
token="old-token",
auto_rotate=True,
rotation_interval="30s",
key_rotation_at=None,
rotation_count=0,
)
mock_response = GenerateKeyResponse(
key="new-api-key",
token_id="new-token-id",
user_id="test-user",
)
from unittest.mock import patch
with patch(
"litellm.proxy.common_utils.key_rotation_manager.regenerate_key_fn",
new_callable=AsyncMock,
) as mock_regenerate:
mock_regenerate.return_value = mock_response
with patch(
"litellm.proxy.common_utils.key_rotation_manager.KeyManagementEventHooks.async_key_rotated_hook",
new_callable=AsyncMock,
):
with patch(
"litellm.proxy.common_utils.key_rotation_manager.LITELLM_KEY_ROTATION_GRACE_PERIOD",
"48h",
):
await manager._rotate_key(key_to_rotate)
mock_regenerate.assert_called_once()
call_args = mock_regenerate.call_args
regenerate_request = call_args[1]["data"]
assert regenerate_request.grace_period == "48h"

View file

@ -756,6 +756,45 @@ async def test_add_spend_log_transaction_to_daily_agent_transaction_injects_agen
assert transaction["custom_llm_provider"] == "openai"
@pytest.mark.asyncio
async def test_add_spend_log_transaction_to_daily_agent_transaction_calls_common_helper_once():
writer = DBSpendUpdateWriter()
mock_prisma = MagicMock()
mock_prisma.get_request_status = MagicMock(return_value="success")
payload = {
"request_id": "req-common-helper",
"agent_id": "agent-abc",
"user": "test-user",
"startTime": "2024-01-01T12:00:00",
"api_key": "test-key",
"model": "gpt-4",
"custom_llm_provider": "openai",
"model_group": "gpt-4-group",
"prompt_tokens": 12,
"completion_tokens": 6,
"spend": 0.25,
"metadata": '{"usage_object": {}}',
}
writer.daily_agent_spend_update_queue.add_update = AsyncMock()
original_common_helper = (
writer._common_add_spend_log_transaction_to_daily_transaction
)
writer._common_add_spend_log_transaction_to_daily_transaction = AsyncMock(
wraps=original_common_helper
)
await writer.add_spend_log_transaction_to_daily_agent_transaction(
payload=payload,
prisma_client=mock_prisma,
)
assert (
writer._common_add_spend_log_transaction_to_daily_transaction.await_count == 1
)
@pytest.mark.asyncio
async def test_add_spend_log_transaction_to_daily_agent_transaction_skips_when_agent_id_missing():
"""
@ -960,4 +999,4 @@ async def test_update_daily_spend_re_raises_exception_after_logging():
entity_id_field="user_id",
table_name="litellm_dailyuserspend",
unique_constraint_name="user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint",
)
)

View file

@ -126,3 +126,77 @@ async def test_async_post_call_failure_hook_non_llm_route():
# Assert that update_database was NOT called for non-LLM routes
mock_update_database.assert_not_called()
@pytest.mark.asyncio
async def test_track_cost_callback_skips_when_no_standard_logging_object():
"""
Reproduces the bug where _PROXY_track_cost_callback raises
'Cost tracking failed for model=None' when kwargs has no
standard_logging_object (e.g. call_type=afile_delete).
File operations have no model and no standard_logging_object.
The callback should skip gracefully instead of raising.
"""
logger = _ProxyDBLogger()
kwargs = {
"call_type": "afile_delete",
"model": None,
"litellm_call_id": "test-call-id",
"litellm_params": {},
"stream": False,
}
with patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
) as mock_proxy_logging:
mock_proxy_logging.failed_tracking_alert = AsyncMock()
mock_proxy_logging.db_spend_update_writer = MagicMock()
mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock()
await logger._PROXY_track_cost_callback(
kwargs=kwargs,
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
# update_database should NOT be called — nothing to track
mock_proxy_logging.db_spend_update_writer.update_database.assert_not_called()
# failed_tracking_alert should NOT be called — this is not an error
mock_proxy_logging.failed_tracking_alert.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("model_value", [None, ""])
async def test_track_cost_callback_skips_for_falsy_model_and_no_slo(model_value):
"""
Same bug as above but model can also be empty string (e.g. health check callbacks).
The guard should catch all falsy model values when sl_object is missing.
"""
logger = _ProxyDBLogger()
kwargs = {
"call_type": "acompletion",
"model": model_value,
"litellm_params": {},
"stream": False,
}
with patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
) as mock_proxy_logging:
mock_proxy_logging.failed_tracking_alert = AsyncMock()
mock_proxy_logging.db_spend_update_writer = MagicMock()
mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock()
await logger._PROXY_track_cost_callback(
kwargs=kwargs,
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
mock_proxy_logging.failed_tracking_alert.assert_not_called()

View file

@ -2,12 +2,10 @@ import asyncio
import json
import os
import sys
from typing import Optional, cast
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import Request
from fastapi.testclient import TestClient
from litellm._uuid import uuid
@ -16,7 +14,7 @@ sys.path.insert(
) # Adds the parent directory to the system path
import litellm
from litellm.proxy._types import LiteLLM_UserTable, NewTeamRequest, NewUserResponse
from litellm.proxy._types import LiteLLM_UserTable, NewUserResponse
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.management_endpoints.sso import CustomMicrosoftSSO
from litellm.proxy.management_endpoints.types import CustomOpenID
@ -136,16 +134,32 @@ def test_microsoft_sso_handler_openid_from_response_with_custom_attributes():
expected_team_ids = ["team1"]
# Act
with patch("litellm.constants.MICROSOFT_USER_EMAIL_ATTRIBUTE", "custom_email_field"), \
patch("litellm.constants.MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE", "custom_display_name"), \
patch("litellm.constants.MICROSOFT_USER_ID_ATTRIBUTE", "custom_id_field"), \
patch("litellm.constants.MICROSOFT_USER_FIRST_NAME_ATTRIBUTE", "custom_first_name"), \
patch("litellm.constants.MICROSOFT_USER_LAST_NAME_ATTRIBUTE", "custom_last_name"), \
patch("litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_EMAIL_ATTRIBUTE", "custom_email_field"), \
patch("litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE", "custom_display_name"), \
patch("litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_ID_ATTRIBUTE", "custom_id_field"), \
patch("litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_FIRST_NAME_ATTRIBUTE", "custom_first_name"), \
patch("litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_LAST_NAME_ATTRIBUTE", "custom_last_name"):
with patch(
"litellm.constants.MICROSOFT_USER_EMAIL_ATTRIBUTE", "custom_email_field"
), patch(
"litellm.constants.MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE", "custom_display_name"
), patch(
"litellm.constants.MICROSOFT_USER_ID_ATTRIBUTE", "custom_id_field"
), patch(
"litellm.constants.MICROSOFT_USER_FIRST_NAME_ATTRIBUTE", "custom_first_name"
), patch(
"litellm.constants.MICROSOFT_USER_LAST_NAME_ATTRIBUTE", "custom_last_name"
), patch(
"litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_EMAIL_ATTRIBUTE",
"custom_email_field",
), patch(
"litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE",
"custom_display_name",
), patch(
"litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_ID_ATTRIBUTE",
"custom_id_field",
), patch(
"litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_FIRST_NAME_ATTRIBUTE",
"custom_first_name",
), patch(
"litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_LAST_NAME_ATTRIBUTE",
"custom_last_name",
):
result = MicrosoftSSOHandler.openid_from_response(
response=mock_response, team_ids=expected_team_ids, user_role=None
)
@ -231,7 +245,6 @@ def test_get_microsoft_callback_response_raw_sso_response():
)
# Assert
print("result from verify_and_process", result)
assert isinstance(result, dict)
assert result["mail"] == "microsoft_user@example.com"
assert result["displayName"] == "Microsoft User"
@ -455,10 +468,6 @@ async def test_default_team_params(team_params):
# Assert
# Verify team was created with correct parameters
mock_prisma.db.litellm_teamtable.create.assert_called_once()
print(
"mock_prisma.db.litellm_teamtable.create.call_args",
mock_prisma.db.litellm_teamtable.create.call_args,
)
create_call_args = mock_prisma.db.litellm_teamtable.create.call_args.kwargs[
"data"
]
@ -583,7 +592,7 @@ def test_apply_user_info_values_to_sso_user_defined_values_with_models():
def test_apply_user_info_values_sso_role_takes_precedence():
"""
Test that SSO role takes precedence over DB role.
When Microsoft SSO returns a user_role, it should be used instead of the role stored in the database.
This ensures SSO is the authoritative source for user roles.
"""
@ -678,16 +687,16 @@ def test_normalize_email():
"""
# Test with lowercase email
assert normalize_email("test@example.com") == "test@example.com"
# Test with uppercase email
assert normalize_email("TEST@EXAMPLE.COM") == "test@example.com"
# Test with mixed case email
assert normalize_email("Test.User@Example.COM") == "test.user@example.com"
# Test with None
assert normalize_email(None) is None
# Test with empty string
assert normalize_email("") == ""
@ -900,7 +909,7 @@ async def test_upsert_sso_user_no_role_in_sso_response():
def test_get_user_email_and_id_extracts_microsoft_role():
"""
Test that _get_user_email_and_id_from_result extracts user_role from Microsoft SSO.
This ensures Microsoft SSO roles (from app_roles in id_token) are properly
extracted and converted from enum to string.
"""
@ -966,7 +975,7 @@ async def test_get_user_info_from_db_user_exists():
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_object"
) as mock_get_user_object:
user_info = await get_user_info_from_db(**args)
await get_user_info_from_db(**args)
mock_get_user_object.assert_called_once()
assert mock_get_user_object.call_args.kwargs["user_id"] == "krrishd"
@ -1008,7 +1017,7 @@ async def test_get_user_info_from_db_user_exists_alternate_user_id():
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_object"
) as mock_get_user_object:
user_info = await get_user_info_from_db(**args)
await get_user_info_from_db(**args)
mock_get_user_object.assert_called_once()
assert mock_get_user_object.call_args.kwargs["user_id"] == "krrishd-email1234"
@ -1017,7 +1026,7 @@ async def test_get_user_info_from_db_user_exists_alternate_user_id():
async def test_get_user_info_from_db_user_not_exists_creates_user():
"""
Test that get_user_info_from_db creates a new user when user doesn't exist in DB.
When get_existing_user_info_from_db returns None, get_user_info_from_db should:
1. Call upsert_sso_user with user_info=None
2. upsert_sso_user should call insert_sso_user to create the user
@ -1105,7 +1114,7 @@ async def test_get_user_info_from_db_user_not_exists_creates_user():
async def test_get_user_info_from_db_user_exists_updates_user():
"""
Test that get_user_info_from_db updates existing user when user exists in DB.
When get_existing_user_info_from_db returns a user, get_user_info_from_db should:
1. Call upsert_sso_user with the existing user_info
2. upsert_sso_user should update the user in the database
@ -1197,6 +1206,7 @@ async def test_get_user_info_from_db_user_exists_updates_user():
# Should return the updated user
assert user_info == updated_user
@pytest.mark.asyncio
async def test_check_and_update_if_proxy_admin_id():
"""
@ -1305,10 +1315,10 @@ async def test_get_generic_sso_response_with_additional_headers():
mock_sso_class = MagicMock(return_value=mock_sso_instance)
with patch.dict(os.environ, test_env_vars):
with patch("fastapi_sso.sso.base.DiscoveryDocument") as mock_discovery:
with patch("fastapi_sso.sso.base.DiscoveryDocument"):
with patch(
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
) as mock_create_provider:
):
# Act
result, received_response = await get_generic_sso_response(
request=mock_request,
@ -1367,10 +1377,10 @@ async def test_get_generic_sso_response_with_empty_headers():
mock_sso_class = MagicMock(return_value=mock_sso_instance)
with patch.dict(os.environ, test_env_vars):
with patch("fastapi_sso.sso.base.DiscoveryDocument") as mock_discovery:
with patch("fastapi_sso.sso.base.DiscoveryDocument"):
with patch(
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
) as mock_create_provider:
):
# Act
result, received_response = await get_generic_sso_response(
request=mock_request,
@ -1755,8 +1765,6 @@ class TestCustomUISSO:
"""Test that proper error is raised when enterprise module is not available"""
from unittest.mock import MagicMock, patch
from litellm.proxy.management_endpoints.ui_sso import google_login
# Mock request
mock_request = MagicMock()
mock_request.base_url = "https://test.example.com/"
@ -1778,7 +1786,7 @@ class TestCustomUISSO:
# This mimics the relevant part of google_login that would trigger the import error
try:
from enterprise.litellm_enterprise.proxy.auth.custom_sso_handler import (
EnterpriseCustomSSOHandler,
EnterpriseCustomSSOHandler, # noqa: F401
)
return "success"
@ -1982,59 +1990,56 @@ class TestCLIKeyRegenerationFlow:
# Test data
session_key = "sk-session-456"
# Mock user info
mock_user_info = LiteLLM_UserTable(
user_id="test-user-123",
user_role="internal_user",
teams=["team1", "team2"],
models=["gpt-4"]
models=["gpt-4"],
)
# Mock SSO result
mock_sso_result = {
"user_email": "test@example.com",
"user_id": "test-user-123"
}
mock_sso_result = {"user_email": "test@example.com", "user_id": "test-user-123"}
# Mock cache
mock_cache = MagicMock()
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
return_value=mock_user_info
), patch(
"litellm.proxy.proxy_server.prisma_client", MagicMock()
), patch(
return_value=mock_user_info,
), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
), patch(
"litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page",
return_value="<html>Success</html>",
):
# Act
result = await cli_sso_callback(
request=mock_request, key=session_key, existing_key=None, result=mock_sso_result
request=mock_request,
key=session_key,
existing_key=None,
result=mock_sso_result,
)
# Assert - verify session was stored in cache
mock_cache.set_cache.assert_called_once()
call_args = mock_cache.set_cache.call_args
# Verify cache key format
assert "cli_sso_session:" in call_args.kwargs["key"]
assert session_key in call_args.kwargs["key"]
# Verify session data structure
session_data = call_args.kwargs["value"]
assert session_data["user_id"] == "test-user-123"
assert session_data["user_role"] == "internal_user"
assert session_data["teams"] == ["team1", "team2"]
assert session_data["models"] == ["gpt-4"]
# Verify TTL
assert call_args.kwargs["ttl"] == 600 # 10 minutes
assert result.status_code == 200
# Verify response contains success message (response is HTML)
assert result.body is not None
@ -2050,17 +2055,14 @@ class TestCLIKeyRegenerationFlow:
"user_id": "test-user-456",
"user_role": "internal_user",
"teams": ["team-a", "team-b", "team-c"],
"models": ["gpt-4"]
"models": ["gpt-4"],
}
# Mock cache
mock_cache = MagicMock()
mock_cache.get_cache.return_value = session_data
with patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
):
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
# Act - First poll without team_id
result = await cli_poll_key(key_id=session_key, team_id=None)
@ -2070,7 +2072,7 @@ class TestCLIKeyRegenerationFlow:
assert result["user_id"] == "test-user-456"
assert result["teams"] == ["team-a", "team-b", "team-c"]
assert "key" not in result # JWT should not be generated yet
# Verify session was NOT deleted
mock_cache.delete_cache.assert_not_called()
@ -2174,34 +2176,33 @@ class TestCLIKeyRegenerationFlow:
"user_role": "internal_user",
"teams": ["team-a", "team-b", "team-c"],
"models": ["gpt-4"],
"user_email": "test@example.com"
"user_email": "test@example.com",
}
# Mock user info
mock_user_info = LiteLLM_UserTable(
user_id="test-user-789",
user_role="internal_user",
teams=["team-a", "team-b", "team-c"],
models=["gpt-4"]
models=["gpt-4"],
)
# Mock cache
mock_cache = MagicMock()
mock_cache.get_cache.return_value = session_data
mock_jwt_token = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.token"
with patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
), patch(
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), patch(
"litellm.proxy.proxy_server.prisma_client"
) as mock_prisma, patch(
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
return_value=mock_jwt_token
return_value=mock_jwt_token,
) as mock_get_jwt:
# Mock the user lookup
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user_info)
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(
return_value=mock_user_info
)
# Act - Second poll with team_id
result = await cli_poll_key(key_id=session_key, team_id=selected_team)
@ -2212,12 +2213,12 @@ class TestCLIKeyRegenerationFlow:
assert result["user_id"] == "test-user-789"
assert result["team_id"] == selected_team
assert result["teams"] == ["team-a", "team-b", "team-c"]
# Verify JWT was generated with correct team
mock_get_jwt.assert_called_once()
jwt_call_args = mock_get_jwt.call_args
assert jwt_call_args.kwargs["team_id"] == selected_team
# Verify session was deleted after JWT generation
mock_cache.delete_cache.assert_called_once()
@ -2227,7 +2228,6 @@ class TestGetAppRolesFromIdToken:
def test_roles_picked_when_app_roles_not_exists(self):
"""Test that 'roles' is picked when 'app_roles' doesn't exist"""
import jwt
# Create a token with only 'roles' claim
token_payload = {
@ -2251,7 +2251,6 @@ class TestGetAppRolesFromIdToken:
def test_app_roles_picked_when_both_exist(self):
"""Test that 'app_roles' takes precedence when both 'app_roles' and 'roles' exist"""
import jwt
# Create a token with both 'app_roles' and 'roles' claims
token_payload = {
@ -2272,7 +2271,6 @@ class TestGetAppRolesFromIdToken:
def test_roles_picked_when_app_roles_is_empty(self):
"""Test that 'roles' is picked when 'app_roles' exists but is empty"""
import jwt
# Create a token with empty 'app_roles' and populated 'roles'
token_payload = {
@ -2293,7 +2291,6 @@ class TestGetAppRolesFromIdToken:
def test_empty_list_when_neither_exists(self):
"""Test that empty list is returned when neither 'app_roles' nor 'roles' exist"""
import jwt
# Create a token without roles claims
token_payload = {"sub": "user123", "email": "test@example.com"}
@ -2317,7 +2314,6 @@ class TestGetAppRolesFromIdToken:
def test_empty_list_when_roles_not_a_list(self):
"""Test that empty list is returned when roles is not a list"""
import jwt
# Create a token with non-list roles
token_payload = {
@ -2337,7 +2333,6 @@ class TestGetAppRolesFromIdToken:
def test_error_handling_on_jwt_decode_exception(self):
"""Test that exceptions during JWT decode are handled gracefully"""
import jwt
mock_token = "invalid.jwt.token"
@ -2788,12 +2783,6 @@ class TestGenericResponseConvertorNestedAttributes:
# to handle dotted paths like "attributes.userId"
# Current behavior: returns None for nested paths
print(f"User ID result: {result.id}")
print(f"Email result: {result.email}")
print(f"First name result: {result.first_name}")
print(f"Last name result: {result.last_name}")
print(f"Display name result: {result.display_name}")
# Expected behavior with current implementation (no nested path support):
assert result.id == "nested-user-456"
assert (
@ -2883,14 +2872,15 @@ class TestGetGenericSSORedirectParams:
# Arrange
cli_state = "litellm-session-token:sk-test123"
with patch.dict(os.environ, {"GENERIC_CLIENT_STATE": "env_state_value"}):
# Act
redirect_params, code_verifier = (
SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=cli_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=cli_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
@ -2905,14 +2895,15 @@ class TestGetGenericSSORedirectParams:
# Arrange
env_state = "custom_env_state_value"
with patch.dict(os.environ, {"GENERIC_CLIENT_STATE": env_state}):
# Act
redirect_params, code_verifier = (
SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=None,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=None,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
@ -2929,13 +2920,14 @@ class TestGetGenericSSORedirectParams:
with patch.dict(os.environ, {}, clear=False):
# Remove GENERIC_CLIENT_STATE if it exists
os.environ.pop("GENERIC_CLIENT_STATE", None)
# Act
redirect_params, code_verifier = (
SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=None,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=None,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
@ -2955,26 +2947,27 @@ class TestGetGenericSSORedirectParams:
# Arrange
test_state = "test_state_123"
with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}):
# Act
redirect_params, code_verifier = (
SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=test_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=test_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert state
assert redirect_params["state"] == test_state
# Assert PKCE parameters
assert code_verifier is not None
assert len(code_verifier) == 43 # Standard PKCE verifier length
assert "code_challenge" in redirect_params
assert "code_challenge_method" in redirect_params
assert redirect_params["code_challenge_method"] == "S256"
# Verify code_challenge is correctly derived from code_verifier
expected_challenge_bytes = hashlib.sha256(
code_verifier.encode("utf-8")
@ -2994,14 +2987,15 @@ class TestGetGenericSSORedirectParams:
# Arrange
test_state = "test_state_456"
with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "false"}):
# Act
redirect_params, code_verifier = (
SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=test_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=test_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
@ -3019,7 +3013,7 @@ class TestGetGenericSSORedirectParams:
# Arrange
cli_state = "cli_state_priority"
env_state = "env_state_should_not_be_used"
with patch.dict(
os.environ,
{
@ -3028,17 +3022,18 @@ class TestGetGenericSSORedirectParams:
},
):
# Act
redirect_params, code_verifier = (
SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=cli_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=cli_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
assert redirect_params["state"] == cli_state # CLI state takes priority
assert redirect_params["state"] != env_state
# PKCE should still be generated
assert code_verifier is not None
assert "code_challenge" in redirect_params
@ -3052,14 +3047,15 @@ class TestGetGenericSSORedirectParams:
# Arrange
env_state = "env_state_for_empty_cli"
with patch.dict(os.environ, {"GENERIC_CLIENT_STATE": env_state}):
# Act
redirect_params, code_verifier = (
SSOAuthenticationHandler._get_generic_sso_redirect_params(
state="", # Empty string
generic_authorization_endpoint="https://auth.example.com/authorize",
)
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state="", # Empty string
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert - empty string is falsy, so env variable should be used
@ -3076,7 +3072,7 @@ class TestGetGenericSSORedirectParams:
# Arrange - no state provided
with patch.dict(os.environ, {}, clear=False):
os.environ.pop("GENERIC_CLIENT_STATE", None)
# Act
params1, _ = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=None,
@ -3139,15 +3135,18 @@ class TestPKCEFunctionality:
test_state = "test_oauth_state_123"
mock_request.query_params = {"state": test_state}
# Mock cache
# Mock cache with async methods
mock_cache = MagicMock()
test_code_verifier = "test_code_verifier_abc123xyz"
mock_cache.get_cache.return_value = test_code_verifier
mock_cache.async_get_cache = AsyncMock(return_value=test_code_verifier)
mock_cache.async_delete_cache = AsyncMock()
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
with patch("litellm.proxy.proxy_server.redis_usage_cache", None), patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
# Act
token_params = SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
token_params = (
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
)
# Assert
@ -3155,10 +3154,10 @@ class TestPKCEFunctionality:
assert token_params["code_verifier"] == test_code_verifier
# Verify cache was accessed and deleted
mock_cache.get_cache.assert_called_once_with(
mock_cache.async_get_cache.assert_called_once_with(
key=f"pkce_verifier:{test_state}"
)
mock_cache.delete_cache.assert_called_once_with(
mock_cache.async_delete_cache.assert_called_once_with(
key=f"pkce_verifier:{test_state}"
)
@ -3183,6 +3182,8 @@ class TestPKCEFunctionality:
test_state = "test456"
mock_cache = MagicMock()
mock_cache.async_set_cache = AsyncMock()
with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}):
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
# Act
@ -3193,9 +3194,9 @@ class TestPKCEFunctionality:
)
# Assert
# Verify cache was called to store code_verifier
mock_cache.set_cache.assert_called_once()
cache_call = mock_cache.set_cache.call_args
# Verify async cache was called to store code_verifier
mock_cache.async_set_cache.assert_called_once()
cache_call = mock_cache.async_set_cache.call_args
assert cache_call.kwargs["key"] == f"pkce_verifier:{test_state}"
assert cache_call.kwargs["ttl"] == 600
assert len(cache_call.kwargs["value"]) == 43
@ -3207,6 +3208,178 @@ class TestPKCEFunctionality:
assert "code_challenge_method=S256" in updated_location
assert f"state={test_state}" in updated_location
@pytest.mark.asyncio
async def test_pkce_redis_multi_pod_verifier_roundtrip(self):
"""
Mock Redis to verify PKCE code_verifier round-trip across "pods":
Pod A stores verifier in Redis; Pod B retrieves it (no real IdP).
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# In-memory mock of Redis (shared between "pods")
class MockRedisCache:
def __init__(self):
self._store = {}
async def async_set_cache(self, key, value, **kwargs):
self._store[key] = json.dumps(value)
async def async_get_cache(self, key, **kwargs):
val = self._store.get(key)
if val is None:
return None
# Simulate RedisCache._get_cache_logic: stored as JSON string, return decoded
if isinstance(val, str):
try:
return json.loads(val)
except (ValueError, TypeError):
return val
return val
async def async_delete_cache(self, key):
self._store.pop(key, None)
mock_redis = MockRedisCache()
mock_in_memory = MagicMock()
mock_sso = MagicMock()
mock_redirect_response = MagicMock()
mock_redirect_response.headers = {
"location": "https://auth.example.com/authorize?state=multi_pod_state_xyz&client_id=abc"
}
mock_sso.get_login_redirect = AsyncMock(return_value=mock_redirect_response)
mock_sso.__enter__ = MagicMock(return_value=mock_sso)
mock_sso.__exit__ = MagicMock(return_value=False)
with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}):
with patch("litellm.proxy.proxy_server.redis_usage_cache", mock_redis):
with patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_in_memory
):
# Pod A: start login, store code_verifier in "Redis"
await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_sso,
state="multi_pod_state_xyz",
generic_authorization_endpoint="https://auth.example.com/authorize",
)
mock_in_memory.async_set_cache.assert_not_called()
# MockRedisCache is a real class; assert on state, not .assert_called_*
stored_key = "pkce_verifier:multi_pod_state_xyz"
assert stored_key in mock_redis._store
stored_value = mock_redis._store[stored_key]
assert isinstance(stored_value, str) and len(json.loads(stored_value)) == 43
# Pod B: callback with same state, retrieve from "Redis"
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "multi_pod_state_xyz"}
token_params = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
assert "code_verifier" in token_params
assert token_params["code_verifier"] == json.loads(stored_value)
mock_in_memory.async_get_cache.assert_not_called()
# delete_cache called; key removed (asserted below)
# Verifier consumed (single-use); key removed from "Redis"
assert "pkce_verifier:multi_pod_state_xyz" not in mock_redis._store
@pytest.mark.asyncio
async def test_pkce_fallback_in_memory_roundtrip_when_redis_none(self):
"""
Regression: When redis_usage_cache is None (no Redis configured),
code_verifier is stored and retrieved via user_api_key_cache.
Roundtrip works when callback hits same pod (same in-memory cache).
Single-pod or no-Redis deployments must continue to work.
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# In-memory store (simulates user_api_key_cache on one pod)
in_memory_store = {}
async def async_set_cache(key, value, **kwargs):
in_memory_store[key] = value
async def async_get_cache(key, **kwargs):
return in_memory_store.get(key)
async def async_delete_cache(key):
in_memory_store.pop(key, None)
mock_in_memory = MagicMock()
mock_in_memory.async_set_cache = AsyncMock(side_effect=async_set_cache)
mock_in_memory.async_get_cache = AsyncMock(side_effect=async_get_cache)
mock_in_memory.async_delete_cache = AsyncMock(side_effect=async_delete_cache)
mock_sso = MagicMock()
mock_redirect_response = MagicMock()
mock_redirect_response.headers = {
"location": "https://auth.example.com/authorize?state=fallback_state_xyz&client_id=abc"
}
mock_sso.get_login_redirect = AsyncMock(return_value=mock_redirect_response)
mock_sso.__enter__ = MagicMock(return_value=mock_sso)
mock_sso.__exit__ = MagicMock(return_value=False)
with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}):
with patch("litellm.proxy.proxy_server.redis_usage_cache", None):
with patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_in_memory
):
# Pod A: start login, store code_verifier in in-memory cache
await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_sso,
state="fallback_state_xyz",
generic_authorization_endpoint="https://auth.example.com/authorize",
)
mock_in_memory.async_set_cache.assert_called_once()
stored_key = mock_in_memory.async_set_cache.call_args.kwargs["key"]
stored_value = mock_in_memory.async_set_cache.call_args.kwargs[
"value"
]
assert stored_key == "pkce_verifier:fallback_state_xyz"
assert isinstance(stored_value, str) and len(stored_value) == 43
# Same pod: callback retrieves from in-memory cache
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "fallback_state_xyz"}
token_params = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
assert "code_verifier" in token_params
assert token_params["code_verifier"] == stored_value
mock_in_memory.async_get_cache.assert_called_once_with(
key=stored_key
)
mock_in_memory.async_delete_cache.assert_called_once_with(
key=stored_key
)
# Verifier consumed; key removed from in-memory
assert "pkce_verifier:fallback_state_xyz" not in in_memory_store
@pytest.mark.asyncio
async def test_pkce_prepare_token_exchange_returns_nothing_when_no_state(self):
"""
Regression: prepare_token_exchange_parameters with no state in request
does not call cache and does not add code_verifier.
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
mock_redis = MagicMock()
mock_in_memory = MagicMock()
with patch("litellm.proxy.proxy_server.redis_usage_cache", mock_redis):
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_in_memory):
mock_request = MagicMock(spec=Request)
mock_request.query_params = {}
token_params = (
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
)
assert "code_verifier" not in token_params
mock_redis.async_get_cache.assert_not_called()
mock_in_memory.async_get_cache.assert_not_called()
# Tests for SSO user team assignment bug (Issue: SSO Users Not Added to Entra-Synced Teams on First Login)
class TestAddMissingTeamMember:
@ -3330,9 +3503,7 @@ class TestAddMissingTeamMember:
team_member_calls = []
async def track_team_member_add(team_id, user_info):
team_member_calls.append(
{"team_id": team_id, "user_id": user_info.user_id}
)
team_member_calls.append({"team_id": team_id, "user_id": user_info.user_id})
# New SSO user with Entra groups
new_user = NewUserResponse(
@ -3393,7 +3564,6 @@ class TestAddMissingTeamMember:
"""
Parametrized test ensuring add_missing_team_member works for all user types.
"""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.ui_sso import add_missing_team_member
user_info = user_info_factory("test-user-id")
@ -3483,7 +3653,7 @@ async def test_role_mappings_override_default_internal_user_params():
return_value=mock_new_user_response,
) as mock_new_user:
# Act
result = await insert_sso_user(
_ = await insert_sso_user(
result_openid=mock_result_openid,
user_defined_values=user_defined_values,
)
@ -3505,7 +3675,7 @@ async def test_role_mappings_override_default_internal_user_params():
assert (
new_user_request.budget_duration == "30d"
), "budget_duration from default_internal_user_params should be applied"
# Note: models are applied via _update_internal_new_user_params inside new_user,
# not in insert_sso_user, so we verify user_defined_values was updated correctly
# by checking that the function completed successfully and other defaults were applied
@ -3620,7 +3790,10 @@ class TestSSOReadinessEndpoint:
assert data["sso_configured"] is True
assert data["provider"] == "google"
assert "GOOGLE_CLIENT_SECRET" in data["missing_environment_variables"]
assert "Google SSO is configured but missing required environment variables" in data["message"]
assert (
"Google SSO is configured but missing required environment variables"
in data["message"]
)
finally:
app.dependency_overrides.clear()
@ -3669,7 +3842,7 @@ class TestSSOReadinessEndpoint:
response = client.get("/sso/readiness")
assert response.status_code == expected_status
if expected_status == 200:
data = response.json()
assert data["sso_configured"] is True
@ -3739,7 +3912,7 @@ class TestSSOReadinessEndpoint:
response = client.get("/sso/readiness")
assert response.status_code == expected_status
if expected_status == 200:
data = response.json()
assert data["sso_configured"] is True
@ -3784,8 +3957,14 @@ class TestCustomMicrosoftSSO:
discovery = await sso.get_discovery_document()
assert discovery["authorization_endpoint"] == "https://login.microsoftonline.com/test-tenant/oauth2/v2.0/authorize"
assert discovery["token_endpoint"] == "https://login.microsoftonline.com/test-tenant/oauth2/v2.0/token"
assert (
discovery["authorization_endpoint"]
== "https://login.microsoftonline.com/test-tenant/oauth2/v2.0/authorize"
)
assert (
discovery["token_endpoint"]
== "https://login.microsoftonline.com/test-tenant/oauth2/v2.0/token"
)
assert discovery["userinfo_endpoint"] == "https://graph.microsoft.com/v1.0/me"
@pytest.mark.asyncio
@ -3849,8 +4028,13 @@ class TestCustomMicrosoftSSO:
# Custom auth endpoint
assert discovery["authorization_endpoint"] == custom_auth_endpoint
# Default token and userinfo endpoints
assert discovery["token_endpoint"] == "https://login.microsoftonline.com/test-tenant/oauth2/v2.0/token"
assert discovery["userinfo_endpoint"] == "https://graph.microsoft.com/v1.0/me"
assert (
discovery["token_endpoint"]
== "https://login.microsoftonline.com/test-tenant/oauth2/v2.0/token"
)
assert (
discovery["userinfo_endpoint"] == "https://graph.microsoft.com/v1.0/me"
)
def test_custom_microsoft_sso_uses_common_tenant_when_none(self):
"""
@ -3887,11 +4071,7 @@ async def test_setup_team_mappings():
# Arrange
mock_prisma = MagicMock()
mock_sso_config = MagicMock()
mock_sso_config.sso_settings = {
"team_mappings": {
"team_ids_jwt_field": "groups"
}
}
mock_sso_config.sso_settings = {"team_mappings": {"team_ids_jwt_field": "groups"}}
mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(
return_value=mock_sso_config
)

View file

@ -0,0 +1,97 @@
"""
Tests for litellm/_service_logger.py
Regression test for KeyError: 'call_type' when async_log_success_event
is called without call_type in kwargs (e.g. from batch polling callbacks).
"""
import pytest
from datetime import datetime, timedelta
from unittest.mock import AsyncMock, patch
from litellm._service_logger import ServiceLogging
@pytest.mark.asyncio
async def test_async_log_success_event_should_not_raise_when_call_type_missing():
"""
When async_log_success_event is called with kwargs that omit 'call_type',
it should not raise a KeyError. This happens in the batch polling flow
where check_batch_cost.py creates a Logging object whose model_call_details
don't include call_type.
"""
service_logger = ServiceLogging(mock_testing=True)
start_time = datetime(2026, 2, 13, 22, 35, 0)
end_time = datetime(2026, 2, 13, 22, 35, 1)
kwargs_without_call_type = {"model": "gpt-4", "stream": False}
with patch.object(
service_logger, "async_service_success_hook", new_callable=AsyncMock
) as mock_hook:
await service_logger.async_log_success_event(
kwargs=kwargs_without_call_type,
response_obj=None,
start_time=start_time,
end_time=end_time,
)
mock_hook.assert_called_once()
call_kwargs = mock_hook.call_args
assert call_kwargs.kwargs["call_type"] == "unknown"
@pytest.mark.asyncio
async def test_async_log_success_event_should_pass_call_type_when_present():
"""
When call_type IS present in kwargs, it should be forwarded correctly.
"""
service_logger = ServiceLogging(mock_testing=True)
start_time = datetime(2026, 2, 13, 22, 35, 0)
end_time = datetime(2026, 2, 13, 22, 35, 1)
kwargs_with_call_type = {
"model": "gpt-4",
"stream": False,
"call_type": "aretrieve_batch",
}
with patch.object(
service_logger, "async_service_success_hook", new_callable=AsyncMock
) as mock_hook:
await service_logger.async_log_success_event(
kwargs=kwargs_with_call_type,
response_obj=None,
start_time=start_time,
end_time=end_time,
)
mock_hook.assert_called_once()
call_kwargs = mock_hook.call_args
assert call_kwargs.kwargs["call_type"] == "aretrieve_batch"
@pytest.mark.asyncio
async def test_async_log_success_event_should_handle_float_duration():
"""
When start_time and end_time produce a float duration (not timedelta),
it should still work correctly.
"""
service_logger = ServiceLogging(mock_testing=True)
start_time = 1000.0
end_time = 1001.5
with patch.object(
service_logger, "async_service_success_hook", new_callable=AsyncMock
) as mock_hook:
await service_logger.async_log_success_event(
kwargs={"call_type": "completion"},
response_obj=None,
start_time=start_time,
end_time=end_time,
)
mock_hook.assert_called_once()
call_kwargs = mock_hook.call_args
assert call_kwargs.kwargs["duration"] == 1.5

View file

@ -37,6 +37,7 @@ export function RegenerateKeyModal({ selectedToken, visible, onClose, onKeyUpdat
tpm_limit: selectedToken.tpm_limit,
rpm_limit: selectedToken.rpm_limit,
duration: selectedToken.duration || "",
grace_period: "",
});
// Initialize the current access token
@ -223,6 +224,23 @@ export function RegenerateKeyModal({ selectedToken, visible, onClose, onKeyUpdat
Current expiry: {selectedToken?.expires ? new Date(selectedToken.expires).toLocaleString() : "Never"}
</div>
{newExpiryTime && <div className="mt-2 text-sm text-green-600">New expiry: {newExpiryTime}</div>}
<Form.Item
name="grace_period"
label="Grace Period (eg: 24h, 2d)"
tooltip="Keep the old key valid for this duration after rotation. Both keys work during this period for seamless cutover. Empty = immediate revoke."
className="mt-8"
rules={[
{
pattern: /^(\d+(s|m|h|d|w|mo))?$/,
message: "Must be a duration like 30s, 30m, 24h, 2d, 1w, or 1mo",
},
]}
>
<TextInput placeholder="e.g. 24h, 2d (empty = immediate revoke)" />
</Form.Item>
<div className="mt-2 text-sm text-gray-500">
Recommended: 24h to 72h for production keys to allow seamless client migration.
</div>
</Form>
)}
</Modal>