mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
commit
af9b6f6e0d
62 changed files with 2977 additions and 444 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
```
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
```
|
||||
|
|
|
|||
|
|
@ -176,6 +176,7 @@ const sidebars = {
|
|||
"tutorials/copilotkit_sdk",
|
||||
"tutorials/google_adk",
|
||||
"tutorials/livekit_xai_realtime",
|
||||
"projects/openai-agents"
|
||||
]
|
||||
},
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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"] = []
|
||||
|
|
|
|||
|
|
@ -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
56
poetry.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
131
tests/batches_tests/test_batch_custom_pricing.py
Normal file
131
tests/batches_tests/test_batch_custom_pricing.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}'"
|
||||
)
|
||||
|
|
@ -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}'"
|
||||
)
|
||||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
|
|
@ -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()
|
||||
167
tests/test_litellm/enterprise/proxy/test_managed_files_hook.py
Normal file
167
tests/test_litellm/enterprise/proxy/test_managed_files_hook.py
Normal 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"
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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=[
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
97
tests/test_litellm/test_service_logger.py
Normal file
97
tests/test_litellm/test_service_logger.py
Normal 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
|
||||
|
|
@ -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>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue