Merge pull request #19839 from BerriAI/litellm_oss_staging_01_27_2026

Litellm oss staging 01 27 2026
This commit is contained in:
Sameer Kankute 2026-01-28 17:33:27 +05:30 • committed by GitHub
commit 7386621d04
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
29 changed files with 1215 additions and 274 deletions

View file

@ -121,8 +121,8 @@ Use this to track overall LiteLLM Proxy usage.
| Metric Name | Description |
|----------------------|--------------------------------------|
| `litellm_proxy_failed_requests_metric` | Total number of failed responses from proxy - the client did not get a success response from litellm proxy. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "exception_status", "exception_class", "route"` |
| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code", "user_email", "route"` |
| `litellm_proxy_failed_requests_metric` | Total number of failed responses from proxy - the client did not get a success response from litellm proxy. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "user_email", "exception_status", "exception_class", "route", "model_id"` |
| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code", "user_email", "route", "model_id"` |
### Callback Logging Metrics
@ -191,10 +191,10 @@ Use this for LLM API Error monitoring and tracking remaining rate limits and tok
| Metric Name | Description |
|----------------------|--------------------------------------|
| `litellm_request_total_latency_metric` | Total latency (seconds) for a request to LiteLLM Proxy Server - tracked for labels "end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model" |
| `litellm_request_total_latency_metric` | Total latency (seconds) for a request to LiteLLM Proxy Server - tracked for labels "end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model", "model_id" |
| `litellm_overhead_latency_metric` | Latency overhead (seconds) added by LiteLLM processing - tracked for labels "model_group", "api_provider", "api_base", "litellm_model_name", "hashed_api_key", "api_key_alias" |
| `litellm_llm_api_latency_metric` | Latency (seconds) for just the LLM API call - tracked for labels "model", "hashed_api_key", "api_key_alias", "team", "team_alias", "requested_model", "end_user", "user" |
| `litellm_llm_api_time_to_first_token_metric` | Time to first token for LLM API call - tracked for labels `model`, `hashed_api_key`, `api_key_alias`, `team`, `team_alias` [Note: only emitted for streaming requests] |
| `litellm_llm_api_time_to_first_token_metric` | Time to first token for LLM API call - tracked for labels `model`, `hashed_api_key`, `api_key_alias`, `team`, `team_alias`, `requested_model`, `end_user`, `user`, `model_id` [Note: only emitted for streaming requests] |
## Tracking `end_user` on Prometheus

View file

@ -244,6 +244,78 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
return managed_object.created_by == user_id
return True # don't raise error if managed object is not found
async def list_user_batches(
self,
user_api_key_dict: UserAPIKeyAuth,
limit: Optional[int] = None,
after: Optional[str] = None,
provider: Optional[str] = None,
target_model_names: Optional[str] = None,
llm_router: Optional[Router] = None,
) -> Dict[str, Any]:
# Provider filtering is not supported for managed batches
# This is because the encoded object ids stored in the managed objects table do not contain the provider information
# To support provider filtering, we would need to store the provider information in the encoded object ids
if provider:
raise Exception(
"Filtering by 'provider' is not supported when using managed batches."
)
# Model name filtering is not supported for managed batches
# This is because the encoded object ids stored in the managed objects table do not contain the model name
# A hash of the model name + litellm_params for the model name is encoded as the model id. This is not sufficient to reliably map the target model names to the model ids.
if target_model_names:
raise Exception(
"Filtering by 'target_model_names' is not supported when using managed batches."
)
where_clause: Dict[str, Any] = {"file_purpose": "batch"}
# Filter by user who created the batch
if user_api_key_dict.user_id:
where_clause["created_by"] = user_api_key_dict.user_id
if after:
where_clause["id"] = {"gt": after}
# Fetch more than needed to allow for post-fetch filtering
fetch_limit = limit or 20
if target_model_names:
# Fetch extra to account for filtering
fetch_limit = max(fetch_limit * 3, 100)
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
where=where_clause,
take=fetch_limit,
order={"created_at": "desc"},
)
batch_objects: List[LiteLLMBatch] = []
for batch in batches:
try:
# Stop once we have enough after filtering
if len(batch_objects) >= (limit or 20):
break
batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object
batch_obj = LiteLLMBatch(**batch_data)
batch_obj.id = batch.unified_object_id
batch_objects.append(batch_obj)
except Exception as e:
verbose_logger.warning(
f"Failed to parse batch object {batch.unified_object_id}: {e}"
)
continue
return {
"object": "list",
"data": batch_objects,
"first_id": batch_objects[0].id if batch_objects else None,
"last_id": batch_objects[-1].id if batch_objects else None,
"has_more": len(batch_objects) == (limit or 20),
}
async def get_user_created_file_ids(
self, user_api_key_dict: UserAPIKeyAuth, model_object_ids: List[str]
) -> List[OpenAIFileObject]:
@ -673,6 +745,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
bytes=file_objects[0].bytes,
filename=file_objects[0].filename,
status="uploaded",
expires_at=file_objects[0].expires_at,
)
return response

View file

@ -18,14 +18,15 @@ def str_to_bool(value: Optional[str]) -> bool:
return value.lower() in ("true", "1", "t", "y", "yes")
def _get_prisma_env() -> dict:
"""Get environment variables for Prisma, handling offline mode if configured."""
prisma_env = os.environ.copy()
if str_to_bool(os.getenv("PRISMA_OFFLINE_MODE")):
# These env vars prevent Prisma from attempting downloads
prisma_env["NPM_CONFIG_PREFER_OFFLINE"] = "true"
prisma_env["NPM_CONFIG_CACHE"] = os.getenv("NPM_CONFIG_CACHE", "/app/.cache/npm")
prisma_env["NPM_CONFIG_CACHE"] = os.getenv(
"NPM_CONFIG_CACHE", "/app/.cache/npm"
)
return prisma_env
@ -34,29 +35,28 @@ def _get_prisma_command() -> str:
if str_to_bool(os.getenv("PRISMA_OFFLINE_MODE")):
# Primary location where Prisma Python package installs the CLI
default_cli_path = "/app/.cache/prisma-python/binaries/node_modules/.bin/prisma"
# Check if custom path is provided (for flexibility)
custom_cli_path = os.getenv("PRISMA_CLI_PATH")
if custom_cli_path and os.path.exists(custom_cli_path):
logger.info(f"Using custom Prisma CLI at {custom_cli_path}")
return custom_cli_path
# Check the default location
if os.path.exists(default_cli_path):
logger.info(f"Using cached Prisma CLI at {default_cli_path}")
return default_cli_path
# If not found, log warning and fall back
logger.warning(
f"Prisma CLI not found at {default_cli_path}. "
"Falling back to Python wrapper (may attempt downloads)"
)
# Fall back to the Python wrapper (will work in online mode)
return "prisma"
class ProxyExtrasDBManager:
@staticmethod
def _get_prisma_dir() -> str:
@ -119,7 +119,7 @@ class ProxyExtrasDBManager:
stdout=open(migration_file, "w"),
check=True,
timeout=30,
env=prisma_env
env=prisma_env,
)
# 3. Mark the migration as applied since it represents current state
@ -134,7 +134,7 @@ class ProxyExtrasDBManager:
],
check=True,
timeout=30,
env=prisma_env
env=prisma_env,
)
return True
@ -159,14 +159,20 @@ class ProxyExtrasDBManager:
@staticmethod
def _roll_back_migration(migration_name: str):
"""Mark a specific migration as rolled back"""
# Set up environment for offline mode if configured
# Set up environment for offline mode if configured
prisma_env = _get_prisma_env()
subprocess.run(
[_get_prisma_command(), "migrate", "resolve", "--rolled-back", migration_name],
[
_get_prisma_command(),
"migrate",
"resolve",
"--rolled-back",
migration_name,
],
timeout=60,
check=True,
capture_output=True,
env=prisma_env
env=prisma_env,
)
@staticmethod
@ -178,7 +184,7 @@ class ProxyExtrasDBManager:
timeout=60,
check=True,
capture_output=True,
env=prisma_env
env=prisma_env,
)
@staticmethod
@ -228,6 +234,8 @@ class ProxyExtrasDBManager:
r"duplicate key value violates",
r"relation .* already exists",
r"constraint .* already exists",
r"does not exist",
r"Can't drop database.* because it doesn't exist",
]
for pattern in idempotent_patterns:
@ -248,7 +256,7 @@ class ProxyExtrasDBManager:
if not database_url:
logger.error("DATABASE_URL not set")
return
diff_dir = (
Path(migrations_dir)
/ "migrations"
@ -283,7 +291,7 @@ class ProxyExtrasDBManager:
check=True,
timeout=60,
stdout=f,
env=_get_prisma_env()
env=_get_prisma_env(),
)
except subprocess.CalledProcessError as e:
logger.warning(f"Failed to generate migration diff: {e.stderr}")
@ -313,7 +321,7 @@ class ProxyExtrasDBManager:
check=True,
capture_output=True,
text=True,
env=_get_prisma_env()
env=_get_prisma_env(),
)
logger.info(f"prisma db execute stdout: {result.stdout}")
logger.info("✅ Migration diff applied successfully")
@ -331,12 +339,18 @@ class ProxyExtrasDBManager:
try:
logger.info(f"Resolving migration: {migration_name}")
subprocess.run(
[_get_prisma_command(), "migrate", "resolve", "--applied", migration_name],
[
_get_prisma_command(),
"migrate",
"resolve",
"--applied",
migration_name,
],
timeout=60,
check=True,
capture_output=True,
text=True,
env=_get_prisma_env()
env=_get_prisma_env(),
)
logger.debug(f"Resolved migration: {migration_name}")
except subprocess.CalledProcessError as e:
@ -375,7 +389,7 @@ class ProxyExtrasDBManager:
check=True,
capture_output=True,
text=True,
env=_get_prisma_env()
env=_get_prisma_env(),
)
logger.info(f"prisma migrate deploy stdout: {result.stdout}")
@ -397,27 +411,42 @@ class ProxyExtrasDBManager:
)
if migration_match:
failed_migration = migration_match.group(1)
logger.info(
f"Found failed migration: {failed_migration}, marking as rolled back"
)
# Mark the failed migration as rolled back
subprocess.run(
[
_get_prisma_command(),
"migrate",
"resolve",
"--rolled-back",
failed_migration,
],
timeout=60,
check=True,
capture_output=True,
text=True,
env=_get_prisma_env()
)
logger.info(
f"✅ Migration {failed_migration} marked as rolled back... retrying"
)
if ProxyExtrasDBManager._is_idempotent_error(e.stderr):
logger.info(
f"Migration {failed_migration} failed due to idempotent error (e.g., column already exists), resolving as applied"
)
ProxyExtrasDBManager._roll_back_migration(
failed_migration
)
ProxyExtrasDBManager._resolve_specific_migration(
failed_migration
)
logger.info(
f"✅ Migration {failed_migration} resolved."
)
return True
else:
logger.info(
f"Found failed migration: {failed_migration}, marking as rolled back"
)
# Mark the failed migration as rolled back
subprocess.run(
[
_get_prisma_command(),
"migrate",
"resolve",
"--rolled-back",
failed_migration,
],
timeout=60,
check=True,
capture_output=True,
text=True,
env=_get_prisma_env(),
)
logger.info(
f"✅ Migration {failed_migration} marked as rolled back... retrying"
)
elif (
"P3005" in e.stderr
and "database schema is not empty" in e.stderr

View file

@ -80,6 +80,10 @@ import dotenv
litellm_mode = os.getenv("LITELLM_MODE", "DEV") # "PRODUCTION", "DEV"
if litellm_mode == "DEV":
dotenv.load_dotenv()
# Import default_encoding to ensure environment variables are initialized at import time
from litellm.litellm_core_utils import default_encoding # noqa: F401
####################################################
if set_verbose:
_turn_on_debug()

View file

@ -23,7 +23,11 @@ from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import
from litellm.litellm_core_utils.llm_cost_calc.utils import (
CostCalculatorUtils,
_generic_cost_per_character,
_get_service_tier_cost_key,
_parse_prompt_tokens_details,
calculate_cost_component,
generic_cost_per_token,
get_billable_input_tokens,
select_cost_metric_for_model,
)
from litellm.llms.anthropic.cost_calculation import (
@ -431,12 +435,18 @@ def cost_per_token( # noqa: PLR0915
model=model, custom_llm_provider=custom_llm_provider
)
if model_info["input_cost_per_token"] > 0:
## COST PER TOKEN ##
prompt_tokens_cost_usd_dollar = (
model_info["input_cost_per_token"] * prompt_tokens
if (
model_info.get("input_cost_per_token", 0) > 0
or model_info.get("output_cost_per_token", 0) > 0
):
return generic_cost_per_token(
model=model,
usage=usage_block,
custom_llm_provider=custom_llm_provider,
service_tier=service_tier,
)
elif (
if (
model_info.get("input_cost_per_second", None) is not None
and response_time_ms is not None
):
@ -451,11 +461,7 @@ def cost_per_token( # noqa: PLR0915
model_info["input_cost_per_second"] * response_time_ms / 1000 # type: ignore
)
if model_info["output_cost_per_token"] > 0:
completion_tokens_cost_usd_dollar = (
model_info["output_cost_per_token"] * completion_tokens
)
elif (
if (
model_info.get("output_cost_per_second", None) is not None
and response_time_ms is not None
):
@ -955,7 +961,10 @@ def completion_cost( # noqa: PLR0915
router_model_id=router_model_id,
)
potential_model_names = [selected_model, _get_response_model(completion_response)]
potential_model_names = [
selected_model,
_get_response_model(completion_response),
]
if model is not None:
potential_model_names.append(model)
@ -1710,10 +1719,16 @@ def default_image_cost_calculator(
)
# Priority 1: Use per-image pricing if available (for gpt-image-1 and similar models)
if "input_cost_per_image" in cost_info and cost_info["input_cost_per_image"] is not None:
if (
"input_cost_per_image" in cost_info
and cost_info["input_cost_per_image"] is not None
):
return cost_info["input_cost_per_image"] * n
# Priority 2: Fall back to per-pixel pricing for backward compatibility
elif "input_cost_per_pixel" in cost_info and cost_info["input_cost_per_pixel"] is not None:
elif (
"input_cost_per_pixel" in cost_info
and cost_info["input_cost_per_pixel"] is not None
):
return cost_info["input_cost_per_pixel"] * height * width * n
else:
raise Exception(
@ -1833,9 +1848,22 @@ def batch_cost_calculator(
if input_cost_per_token_batches:
total_prompt_cost = usage.prompt_tokens * input_cost_per_token_batches
elif input_cost_per_token:
# Subtract cached tokens from prompt_tokens before calculating cost
# Fixes issue where cached tokens are being charged again
total_prompt_cost = (
usage.prompt_tokens * (input_cost_per_token) / 2
get_billable_input_tokens(usage) * (input_cost_per_token) / 2
) # batch cost is usually half of the regular token cost
# Add cache read cost if applicable
details = _parse_prompt_tokens_details(usage)
cache_read_tokens = details["cache_hit_tokens"]
cache_read_cost_key = _get_service_tier_cost_key(
"cache_read_input_token_cost", None
)
total_prompt_cost += (
calculate_cost_component(model_info, cache_read_cost_key, cache_read_tokens)
/ 2
)
if output_cost_per_token_batches:
total_completion_cost = usage.completion_tokens * output_cost_per_token_batches
elif output_cost_per_token:

View file

@ -901,7 +901,7 @@ class PrometheusLogger(CustomLogger):
model = kwargs.get("model", "")
litellm_params = kwargs.get("litellm_params", {}) or {}
_metadata = litellm_params.get("metadata", {})
_metadata = litellm_params.get("metadata") or {}
get_end_user_id_for_cost_tracking = _get_cached_end_user_id_for_cost_tracking()
end_user_id = get_end_user_id_for_cost_tracking(
@ -1178,26 +1178,15 @@ class PrometheusLogger(CustomLogger):
response_cost: float,
user_id: Optional[str] = None,
):
_team_spend = litellm_params.get("metadata", {}).get(
"user_api_key_team_spend", None
)
_team_max_budget = litellm_params.get("metadata", {}).get(
"user_api_key_team_max_budget", None
)
_metadata = litellm_params.get("metadata") or {}
_team_spend = _metadata.get("user_api_key_team_spend", None)
_team_max_budget = _metadata.get("user_api_key_team_max_budget", None)
_api_key_spend = litellm_params.get("metadata", {}).get(
"user_api_key_spend", None
)
_api_key_max_budget = litellm_params.get("metadata", {}).get(
"user_api_key_max_budget", None
)
_api_key_spend = _metadata.get("user_api_key_spend", None)
_api_key_max_budget = _metadata.get("user_api_key_max_budget", None)
_user_spend = litellm_params.get("metadata", {}).get(
"user_api_key_user_spend", None
)
_user_max_budget = litellm_params.get("metadata", {}).get(
"user_api_key_user_max_budget", None
)
_user_spend = _metadata.get("user_api_key_user_spend", None)
_user_max_budget = _metadata.get("user_api_key_user_max_budget", None)
await self._set_api_key_budget_metrics_after_api_request(
user_api_key=user_api_key,
@ -1310,12 +1299,14 @@ class PrometheusLogger(CustomLogger):
time_to_first_token_seconds is not None
and kwargs.get("stream", False) is True # only emit for streaming requests
):
_ttft_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_llm_api_time_to_first_token_metric"
),
enum_values=enum_values,
)
self.litellm_llm_api_time_to_first_token_metric.labels(
model,
user_api_key,
user_api_key_alias,
user_api_team,
user_api_team_alias,
**_ttft_labels
).observe(time_to_first_token_seconds)
else:
verbose_logger.debug(
@ -1355,7 +1346,7 @@ class PrometheusLogger(CustomLogger):
# request queue time (time from arrival to processing start)
_litellm_params = kwargs.get("litellm_params", {}) or {}
queue_time_seconds = _litellm_params.get("metadata", {}).get(
queue_time_seconds = (_litellm_params.get("metadata") or {}).get(
"queue_time_seconds"
)
if queue_time_seconds is not None and queue_time_seconds >= 0:
@ -2509,8 +2500,8 @@ class PrometheusLogger(CustomLogger):
self,
user_api_team: Optional[str],
user_api_team_alias: Optional[str],
team_spend: float,
team_max_budget: float,
team_spend: Optional[float],
team_max_budget: Optional[float],
response_cost: float,
):
"""
@ -2672,7 +2663,7 @@ class PrometheusLogger(CustomLogger):
user_api_key: Optional[str],
user_api_key_alias: Optional[str],
response_cost: float,
key_max_budget: float,
key_max_budget: Optional[float],
key_spend: Optional[float],
):
if user_api_key:
@ -2689,7 +2680,7 @@ class PrometheusLogger(CustomLogger):
self,
user_api_key: str,
user_api_key_alias: str,
key_max_budget: float,
key_max_budget: Optional[float],
key_spend: Optional[float],
response_cost: float,
) -> UserAPIKeyAuth:

View file

@ -23,6 +23,15 @@ def _is_above_128k(tokens: float) -> bool:
return False
def get_billable_input_tokens(usage: Usage) -> int:
"""
Returns the number of billable input tokens.
Subtracts cached tokens from prompt tokens if applicable.
"""
details = _parse_prompt_tokens_details(usage)
return usage.prompt_tokens - details["cache_hit_tokens"]
def select_cost_metric_for_model(
model_info: ModelInfo,
) -> Literal["cost_per_character", "cost_per_token"]:
@ -190,7 +199,6 @@ def _get_token_base_cost(
1000 if "k" in threshold_str else 1
)
if usage.prompt_tokens > threshold:
prompt_base_cost = cast(
float, _get_cost_per_unit(model_info, key, prompt_base_cost)
)
@ -566,14 +574,28 @@ def generic_cost_per_token( # noqa: PLR0915
if usage.prompt_tokens_details:
prompt_tokens_details = _parse_prompt_tokens_details(usage)
## EDGE CASE - text tokens not set inside PromptTokensDetails
## EDGE CASE - text tokens not set or includes cached tokens (double-counting)
## Some providers (like xAI) report text_tokens = prompt_tokens (including cached)
## We detect this when: text_tokens + cached_tokens + other > prompt_tokens
## Ref: https://github.com/BerriAI/litellm/issues/19680, #14874, #14875
if prompt_tokens_details["text_tokens"] == 0:
cache_hit = prompt_tokens_details["cache_hit_tokens"]
text_tokens = prompt_tokens_details["text_tokens"]
audio_tokens = prompt_tokens_details["audio_tokens"]
cache_creation = prompt_tokens_details["cache_creation_tokens"]
image_tokens = prompt_tokens_details["image_tokens"]
# Check for double-counting: sum of details > prompt_tokens means overlap
total_details = text_tokens + cache_hit + audio_tokens + cache_creation + image_tokens
has_double_counting = cache_hit > 0 and total_details > usage.prompt_tokens
if text_tokens == 0 or has_double_counting:
text_tokens = (
usage.prompt_tokens
- prompt_tokens_details["cache_hit_tokens"]
- prompt_tokens_details["audio_tokens"]
- prompt_tokens_details["cache_creation_tokens"]
- cache_hit
- audio_tokens
- cache_creation
- image_tokens
)
prompt_tokens_details["text_tokens"] = text_tokens
@ -619,7 +641,11 @@ def generic_cost_per_token( # noqa: PLR0915
# Calculate text tokens as remainder when we have a breakdown
# This handles cases like OpenAI's reasoning models where text_tokens isn't provided
text_tokens = max(
0, usage.completion_tokens - reasoning_tokens - audio_tokens - image_tokens
0,
usage.completion_tokens
- reasoning_tokens
- audio_tokens
- image_tokens,
)
else:
# No breakdown at all, all tokens are text tokens

View file

@ -1677,13 +1677,16 @@ def convert_to_anthropic_tool_result(
] = []
for content in content_list:
if content["type"] == "text":
anthropic_content_list.append(
AnthropicMessagesToolResultContent(
type="text",
text=content["text"],
cache_control=content.get("cache_control", None),
)
)
# Only include cache_control if explicitly set and not None
# to avoid sending "cache_control": null which breaks some API channels
text_content: AnthropicMessagesToolResultContent = {
"type": "text",
"text": content["text"],
}
cache_control_value = content.get("cache_control")
if cache_control_value is not None:
text_content["cache_control"] = cache_control_value
anthropic_content_list.append(text_content)
elif content["type"] == "image_url":
format = (
content["image_url"].get("format")

View file

@ -1,11 +1,12 @@
"""
Helper util for handling azure openai-specific cost calculation
- e.g.: prompt caching
- e.g.: prompt caching, audio tokens
"""
from typing import Optional, Tuple
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
from litellm.types.utils import Usage
from litellm.utils import get_model_info
@ -18,34 +19,15 @@ def cost_per_token(
Input:
- model: str, the model name without provider prefix
- usage: LiteLLM Usage block, containing anthropic caching information
- usage: LiteLLM Usage block, containing caching and audio token information
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
"""
## GET MODEL INFO
model_info = get_model_info(model=model, custom_llm_provider="azure")
cached_tokens: Optional[int] = None
## CALCULATE INPUT COST
non_cached_text_tokens = usage.prompt_tokens
if usage.prompt_tokens_details and usage.prompt_tokens_details.cached_tokens:
cached_tokens = usage.prompt_tokens_details.cached_tokens
non_cached_text_tokens = non_cached_text_tokens - cached_tokens
prompt_cost: float = non_cached_text_tokens * model_info["input_cost_per_token"]
## CALCULATE OUTPUT COST
completion_cost: float = (
usage["completion_tokens"] * model_info["output_cost_per_token"]
)
## Prompt Caching cost calculation
if model_info.get("cache_read_input_token_cost") is not None and cached_tokens:
# Note: We read ._cache_read_input_tokens from the Usage - since cost_calculator.py standardizes the cache read tokens on usage._cache_read_input_tokens
prompt_cost += cached_tokens * (
model_info.get("cache_read_input_token_cost", 0) or 0
)
## Speech / Audio cost calculation
## Speech / Audio cost calculation (cost per second for TTS models)
if (
"output_cost_per_second" in model_info
and model_info["output_cost_per_second"] is not None
@ -55,7 +37,14 @@ def cost_per_token(
f"For model={model} - output_cost_per_second: {model_info.get('output_cost_per_second')}; response time: {response_time_ms}"
)
## COST PER SECOND ##
prompt_cost = 0
prompt_cost = 0.0
completion_cost = model_info["output_cost_per_second"] * response_time_ms / 1000
return prompt_cost, completion_cost
return prompt_cost, completion_cost
## Use generic cost calculator for all other cases
## This properly handles: text tokens, audio tokens, cached tokens, reasoning tokens, etc.
return generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider="azure",
)

View file

@ -31,6 +31,16 @@ else:
GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1"
def is_valid_json(value: str) -> bool:
"""Checks whether the value passed is a valid serialized JSON string"""
try:
json.loads(value)
except json.JSONDecodeError:
return False
else:
return True
class GigaChatError(BaseLLMException):
"""GigaChat API error."""
@ -101,7 +111,11 @@ class GigaChatConfig(BaseConfig):
Set up headers with OAuth token.
"""
# Get access token
credentials = api_key or get_secret_str("GIGACHAT_CREDENTIALS") or get_secret_str("GIGACHAT_API_KEY")
credentials = (
api_key
or get_secret_str("GIGACHAT_CREDENTIALS")
or get_secret_str("GIGACHAT_API_KEY")
)
access_token = get_access_token(credentials=credentials)
# Store credentials for image uploads
@ -193,11 +207,13 @@ class GigaChatConfig(BaseConfig):
for tool in tools:
if tool.get("type") == "function":
func = tool.get("function", {})
functions.append({
"name": func.get("name", ""),
"description": func.get("description", ""),
"parameters": func.get("parameters", {}),
})
functions.append(
{
"name": func.get("name", ""),
"description": func.get("description", ""),
"parameters": func.get("parameters", {}),
}
)
return functions
def _map_tool_choice(
@ -281,8 +297,14 @@ class GigaChatConfig(BaseConfig):
}
# Add optional params
for key in ["temperature", "top_p", "max_tokens", "stream",
"repetition_penalty", "profanity_check"]:
for key in [
"temperature",
"top_p",
"max_tokens",
"stream",
"repetition_penalty",
"profanity_check",
]:
if key in optional_params:
request_data[key] = optional_params[key]
@ -314,7 +336,7 @@ class GigaChatConfig(BaseConfig):
elif role == "tool":
message["role"] = "function"
content = message.get("content", "")
if not isinstance(content, str):
if not isinstance(content, str) or not is_valid_json(content):
message["content"] = json.dumps(content, ensure_ascii=False)
# Handle None content
@ -441,14 +463,16 @@ class GigaChatConfig(BaseConfig):
# Convert to tool_calls format
if isinstance(args, dict):
args = json.dumps(args, ensure_ascii=False)
message_data["tool_calls"] = [{
"id": f"call_{uuid.uuid4().hex[:24]}",
"type": "function",
"function": {
"name": func_call.get("name", ""),
"arguments": args,
message_data["tool_calls"] = [
{
"id": f"call_{uuid.uuid4().hex[:24]}",
"type": "function",
"function": {
"name": func_call.get("name", ""),
"arguments": args,
},
}
}]
]
message_data.pop("function_call", None)
finish_reason = "tool_calls"

View file

@ -32,6 +32,7 @@ from litellm.types.llms.oci import (
OCICompletionResponse,
OCIContentPartUnion,
OCIImageContentPart,
OCIImageUrl,
OCIMessage,
OCIRoles,
OCIServingMode,
@ -1129,7 +1130,7 @@ def adapt_messages_to_generic_oci_standard_content_message(
image_url = image_url.get("url")
if not isinstance(image_url, str):
raise Exception("Prop `image_url` must be a string or an object with a `url` property")
new_content.append(OCIImageContentPart(imageUrl=image_url))
new_content.append(OCIImageContentPart(imageUrl=OCIImageUrl(url=image_url)))
return OCIMessage(
role=open_ai_to_generic_oci_role_map[role],

View file

@ -30461,6 +30461,7 @@
"supports_web_search": true
},
"xai/grok-3": {
"cache_read_input_token_cost": 7.5e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30475,6 +30476,7 @@
"supports_web_search": true
},
"xai/grok-3-beta": {
"cache_read_input_token_cost": 7.5e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30489,6 +30491,7 @@
"supports_web_search": true
},
"xai/grok-3-fast-beta": {
"cache_read_input_token_cost": 1.25e-06,
"input_cost_per_token": 5e-06,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30503,6 +30506,7 @@
"supports_web_search": true
},
"xai/grok-3-fast-latest": {
"cache_read_input_token_cost": 1.25e-06,
"input_cost_per_token": 5e-06,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30517,6 +30521,7 @@
"supports_web_search": true
},
"xai/grok-3-latest": {
"cache_read_input_token_cost": 7.5e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30531,6 +30536,7 @@
"supports_web_search": true
},
"xai/grok-3-mini": {
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 3e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30546,6 +30552,7 @@
"supports_web_search": true
},
"xai/grok-3-mini-beta": {
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 3e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30561,6 +30568,7 @@
"supports_web_search": true
},
"xai/grok-3-mini-fast": {
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 6e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30576,6 +30584,7 @@
"supports_web_search": true
},
"xai/grok-3-mini-fast-beta": {
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 6e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30591,6 +30600,7 @@
"supports_web_search": true
},
"xai/grok-3-mini-fast-latest": {
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 6e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30606,6 +30616,7 @@
"supports_web_search": true
},
"xai/grok-3-mini-latest": {
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 3e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,

View file

@ -542,14 +542,26 @@ async def list_batches(
route_type="alist_batches",
)
model_param = (
# Try to use managed objects table for listing batches (returns encoded IDs)
managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files")
if managed_files_obj is not None and hasattr(managed_files_obj, "list_user_batches"):
verbose_proxy_logger.debug(
"Using managed objects table for batch listing"
)
response = await managed_files_obj.list_user_batches(
user_api_key_dict=user_api_key_dict,
limit=limit,
after=after,
provider=provider,
target_model_names=target_model_names,
llm_router=llm_router,
)
elif (model_param := (
data.get("model")
or request.query_params.get("model")
or request.headers.get("x-litellm-model")
)
# SCENARIO 2: Use model-based routing from header/query/body
if model_param:
)):
# SCENARIO 2: Use model-based routing from header/query/body
credentials = get_credentials_for_model(
llm_router=llm_router,
model_id=model_param,

View file

@ -1224,6 +1224,7 @@ redis_usage_cache: Optional[
RedisCache
] = None # redis cache used for tracking spend, tpm/rpm limits
polling_via_cache_enabled: Union[Literal["all"], List[str], bool] = False
native_background_mode: List[str] = [] # Models that should use native provider background mode instead of polling
polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache
user_custom_auth = None
user_custom_key_generate = None
@ -2456,14 +2457,17 @@ class ProxyConfig:
pass
elif key == "responses":
# Initialize global polling via cache settings
global polling_via_cache_enabled, polling_cache_ttl
global polling_via_cache_enabled, native_background_mode, polling_cache_ttl
background_mode = value.get("background_mode", {})
polling_via_cache_enabled = background_mode.get(
"polling_via_cache", False
)
native_background_mode = background_mode.get(
"native_background_mode", []
)
polling_cache_ttl = background_mode.get("ttl", 3600)
verbose_proxy_logger.debug(
f"{blue_color_code} Initialized polling via cache: enabled={polling_via_cache_enabled}, ttl={polling_cache_ttl}{reset_color_code}"
f"{blue_color_code} Initialized polling via cache: enabled={polling_via_cache_enabled}, native_background_mode={native_background_mode}, ttl={polling_cache_ttl}{reset_color_code}"
)
elif key == "default_team_settings":
for idx, team_setting in enumerate(

View file

@ -68,6 +68,7 @@ async def responses_api(
_read_request_body,
general_settings,
llm_router,
native_background_mode,
polling_cache_ttl,
polling_via_cache_enabled,
proxy_config,
@ -95,6 +96,7 @@ async def responses_api(
redis_cache=redis_usage_cache,
model=data.get("model", ""),
llm_router=llm_router,
native_background_mode=native_background_mode,
)
# If polling is enabled, use polling mode

View file

@ -3,7 +3,7 @@ Response Polling Handler for Background Responses with Cache
"""
import json
from datetime import datetime, timezone
from typing import Any, Dict, Optional
from typing import Any, Dict, List, Optional
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid4
@ -257,6 +257,7 @@ def should_use_polling_for_request(
redis_cache, # RedisCache or None
model: str,
llm_router, # Router instance or None
native_background_mode: Optional[List[str]] = None, # List of models that should use native background mode
) -> bool:
"""
Determine if polling via cache should be used for a request.
@ -267,6 +268,8 @@ def should_use_polling_for_request(
redis_cache: Redis cache instance (required for polling)
model: Model name from the request (e.g., "gpt-5" or "openai/gpt-4o")
llm_router: LiteLLM router instance for looking up model deployments
native_background_mode: List of model names that should use native provider
background mode instead of polling via cache
Returns:
True if polling should be used, False otherwise
@ -275,6 +278,13 @@ def should_use_polling_for_request(
if not (background_mode and polling_via_cache_enabled and redis_cache):
return False
# Check if model is in native_background_mode list - these use native provider background mode
if native_background_mode and model in native_background_mode:
verbose_proxy_logger.debug(
f"Model {model} is in native_background_mode list, skipping polling via cache"
)
return False
# "all" enables polling for all providers
if polling_via_cache_enabled == "all":
return True

View file

@ -434,6 +434,8 @@ async def aresponses(
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model, api_base=local_vars.get("base_url", None)
)
# Update local_vars with detected provider (fixes #19782)
local_vars["custom_llm_provider"] = custom_llm_provider
func = partial(
responses,
@ -583,6 +585,9 @@ def responses(
api_key=litellm_params.api_key,
)
# Update local_vars with detected provider (fixes #19782)
local_vars["custom_llm_provider"] = custom_llm_provider
# Use dynamic credentials from get_llm_provider (e.g., when use_litellm_proxy=True)
if dynamic_api_key is not None:
litellm_params.api_key = dynamic_api_key
@ -1411,6 +1416,8 @@ async def acompact_responses(
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model, api_base=local_vars.get("base_url", None)
)
# Update local_vars with detected provider (fixes #19782)
local_vars["custom_llm_provider"] = custom_llm_provider
func = partial(
compact_responses,
@ -1498,6 +1505,9 @@ def compact_responses(
api_key=litellm_params.api_key,
)
# Update local_vars with detected provider (fixes #19782)
local_vars["custom_llm_provider"] = custom_llm_provider
# Use dynamic credentials from get_llm_provider (e.g., when use_litellm_proxy=True)
if dynamic_api_key is not None:
litellm_params.api_key = dynamic_api_key

View file

@ -227,6 +227,9 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.TEAM.value,
UserAPIKeyLabelNames.TEAM_ALIAS.value,
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
UserAPIKeyLabelNames.END_USER.value,
UserAPIKeyLabelNames.USER.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]

View file

@ -35,11 +35,18 @@ class OCITextContentPart(OCIContentPart):
text: str
class OCIImageUrl(BaseModel):
"""ImageUrl object for OCI API. See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/generative_ai_inference/models/oci.generative_ai_inference.models.ImageUrl.html"""
url: str
detail: Optional[Literal["AUTO", "HIGH", "LOW"]] = None
class OCIImageContentPart(OCIContentPart):
"""Image content part for the OCI API."""
type: Literal["IMAGE"] = "IMAGE"
imageUrl: str
imageUrl: OCIImageUrl
OCIContentPartUnion = Union[OCITextContentPart, OCIImageContentPart]

View file

@ -3653,10 +3653,9 @@
"max_input_tokens": 128000,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"mode": "responses",
"output_cost_per_token": 1.4e-05,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
@ -23892,8 +23891,11 @@
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"mode": "responses",
"output_cost_per_token": 1.4e-05,
"supported_endpoints": [
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
@ -30461,6 +30463,7 @@
"supports_web_search": true
},
"xai/grok-3": {
"cache_read_input_token_cost": 7.5e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30475,6 +30478,7 @@
"supports_web_search": true
},
"xai/grok-3-beta": {
"cache_read_input_token_cost": 7.5e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30489,6 +30493,7 @@
"supports_web_search": true
},
"xai/grok-3-fast-beta": {
"cache_read_input_token_cost": 1.25e-06,
"input_cost_per_token": 5e-06,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30503,6 +30508,7 @@
"supports_web_search": true
},
"xai/grok-3-fast-latest": {
"cache_read_input_token_cost": 1.25e-06,
"input_cost_per_token": 5e-06,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30517,6 +30523,7 @@
"supports_web_search": true
},
"xai/grok-3-latest": {
"cache_read_input_token_cost": 7.5e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30531,6 +30538,7 @@
"supports_web_search": true
},
"xai/grok-3-mini": {
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 3e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30546,6 +30554,7 @@
"supports_web_search": true
},
"xai/grok-3-mini-beta": {
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 3e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30561,6 +30570,7 @@
"supports_web_search": true
},
"xai/grok-3-mini-fast": {
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 6e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30576,6 +30586,7 @@
"supports_web_search": true
},
"xai/grok-3-mini-fast-beta": {
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 6e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30591,6 +30602,7 @@
"supports_web_search": true
},
"xai/grok-3-mini-fast-latest": {
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 6e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,
@ -30606,6 +30618,7 @@
"supports_web_search": true
},
"xai/grok-3-mini-latest": {
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 3e-07,
"litellm_provider": "xai",
"max_input_tokens": 131072,

View file

@ -411,7 +411,15 @@ def test_set_latency_metrics(prometheus_logger):
# completion_start_time - api_call_start_time
prometheus_logger.litellm_llm_api_time_to_first_token_metric.labels.assert_called_once_with(
"gpt-3.5-turbo", "key1", "alias1", "team1", "team_alias1"
end_user=None,
user="test_user",
hashed_api_key="test_hash",
api_key_alias="test_alias",
team="test_team",
team_alias="test_team_alias",
requested_model="openai-gpt",
model="gpt-3.5-turbo",
model_id="model-123",
)
prometheus_logger.litellm_llm_api_time_to_first_token_metric.labels().observe.assert_called_once_with(
0.5
@ -442,6 +450,7 @@ def test_set_latency_metrics(prometheus_logger):
team_alias="test_team_alias",
requested_model="openai-gpt",
model="gpt-3.5-turbo",
model_id="model-123",
)
prometheus_logger.litellm_request_total_latency_metric.labels().observe.assert_called_once_with(
2.0
@ -737,6 +746,7 @@ async def test_async_post_call_failure_hook(prometheus_logger):
exception_status="429",
exception_class="Openai.RateLimitError",
route=user_api_key_dict.request_route,
model_id=None,
)
prometheus_logger.litellm_proxy_failed_requests_metric.labels().inc.assert_called_once()
@ -752,6 +762,7 @@ async def test_async_post_call_failure_hook(prometheus_logger):
status_code="429",
user_email=None,
route=user_api_key_dict.request_route,
model_id=None,
)
prometheus_logger.litellm_proxy_total_requests_metric.labels().inc.assert_called_once()
@ -798,6 +809,7 @@ async def test_async_post_call_success_hook(prometheus_logger):
status_code="200",
user_email=None,
route=user_api_key_dict.request_route,
model_id=None,
)
prometheus_logger.litellm_proxy_total_requests_metric.labels().inc.assert_called_once()

View file

@ -827,3 +827,262 @@ async def test_afile_retrieve_raises_error_for_non_managed_file():
)
assert "not found" in str(exc_info.value)
@pytest.mark.asyncio
async def test_list_batches_from_managed_objects_table():
from litellm.proxy._types import UserAPIKeyAuth
from openai.types.batch import BatchRequestCounts
prisma_client = AsyncMock()
batch_record_1 = MagicMock()
batch_record_1.unified_object_id = "unified-batch-id-1"
batch_record_1.file_object = json.dumps({
"id": "batch_abc123",
"object": "batch",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"status": "completed",
"created_at": 1234567890,
"input_file_id": "file-input-1",
"request_counts": {"total": 1, "completed": 1, "failed": 0},
})
batch_record_2 = MagicMock()
batch_record_2.unified_object_id = "unified-batch-id-2"
batch_record_2.file_object = json.dumps({
"id": "batch_xyz789",
"object": "batch",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"status": "in_progress",
"created_at": 1234567891,
"input_file_id": "file-input-2",
"request_counts": {"total": 5, "completed": 2, "failed": 0},
})
prisma_client.db.litellm_managedobjecttable.find_many.return_value = [
batch_record_1,
batch_record_2,
]
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
result = await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
limit=10,
)
assert result["object"] == "list"
assert len(result["data"]) == 2
assert result["data"][0].id == "unified-batch-id-1"
assert result["data"][1].id == "unified-batch-id-2"
assert result["first_id"] == "unified-batch-id-1"
assert result["last_id"] == "unified-batch-id-2"
# Should filter by user_id (created_by)
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with(
where={"file_purpose": "batch", "created_by": "test-user"},
take=10,
order={"created_at": "desc"},
)
@pytest.mark.asyncio
async def test_list_batches_from_managed_objects_table_empty_list():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_many.return_value = []
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
result = await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
)
assert result["object"] == "list"
assert len(result["data"]) == 0
assert result["first_id"] is None
assert result["last_id"] is None
assert result["has_more"] is False
# Verify where clause includes created_by filter
# Default take is 20 when no limit is provided
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with(
where={"file_purpose": "batch", "created_by": "test-user"},
take=20,
order={"created_at": "desc"},
)
def _create_unified_batch_id(model_id: str, batch_id: str) -> str:
import base64
unified_str = f"litellm_proxy;model_id:{model_id};llm_batch_id:{batch_id}"
return base64.urlsafe_b64encode(unified_str.encode()).decode().rstrip("=")
@pytest.mark.asyncio
async def test_list_batches_from_managed_objects_table_provider_filter_raises_exception():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
# Filtering by provider should raise Exception
with pytest.raises(Exception) as exc_info:
await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
limit=10,
provider="openai",
)
assert str(exc_info.value) == (
"Filtering by 'provider' is not supported when using managed batches."
)
# Verify find_many was NOT called since exception is raised before database query
prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called()
@pytest.mark.asyncio
async def test_list_batches_from_managed_objects_table_target_model_name_filter_raises_exception():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
# Filtering by provider should raise Exception
with pytest.raises(Exception) as exc_info:
await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
limit=10,
target_model_names="gpt-4o,gpt-3.5",
)
assert str(exc_info.value) == (
"Filtering by 'target_model_names' is not supported when using managed batches."
)
# Verify find_many was NOT called since exception is raised before database query
prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called()
@pytest.mark.asyncio
async def test_list_batches_from_managed_objects_table_filters_by_created_by():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
# Create batch for user1
batch_user1 = MagicMock()
batch_user1.unified_object_id = "unified-batch-user1"
batch_user1.file_object = json.dumps({
"id": "batch_user1_abc",
"object": "batch",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"status": "completed",
"created_at": 1234567890,
"input_file_id": "file-input-user1",
"request_counts": {"total": 1, "completed": 1, "failed": 0},
})
# Create batch for user2
batch_user2 = MagicMock()
batch_user2.unified_object_id = "unified-batch-user2"
batch_user2.file_object = json.dumps({
"id": "batch_user2_xyz",
"object": "batch",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"status": "completed",
"created_at": 1234567891,
"input_file_id": "file-input-user2",
"request_counts": {"total": 2, "completed": 2, "failed": 0},
})
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
# Query with user1's API key - should only return user1's batch
prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user1]
result_user1 = await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="user1"),
limit=10,
)
assert len(result_user1["data"]) == 1
assert result_user1["data"][0].id == "unified-batch-user1"
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_with(
where={"file_purpose": "batch", "created_by": "user1"},
take=10,
order={"created_at": "desc"},
)
# Query with user2's API key - should only return user2's batch
prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user2]
result_user2 = await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="user2"),
limit=10,
)
assert len(result_user2["data"]) == 1
assert result_user2["data"][0].id == "unified-batch-user2"
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_with(
where={"file_purpose": "batch", "created_by": "user2"},
take=10,
order={"created_at": "desc"},
)
@pytest.mark.asyncio
async def test_return_unified_file_id_includes_expires_at():
from litellm.types.llms.openai import OpenAIFileObject
# Create a mock file object with expires_at set
file_object = OpenAIFileObject(
id="file-abc123",
object="file",
bytes=1234,
created_at=1234567890,
filename="test.jsonl",
purpose="batch",
status="uploaded",
expires_at=1234657890,
)
file_object._hidden_params = {"model_id": "test-model-id"}
create_file_request = {
"file": ("test.jsonl", b"test content", "application/jsonl"),
"purpose": "batch",
}
internal_usage_cache = MagicMock()
result = await _PROXY_LiteLLMManagedFiles.return_unified_file_id(
file_objects=[file_object],
create_file_request=create_file_request,
internal_usage_cache=internal_usage_cache,
litellm_parent_otel_span=None,
target_model_names_list=["gpt-4o"],
)
# Verify expires_at is passed through
assert result.expires_at == 1234657890
assert result.purpose == "batch"
assert result.filename == "test.jsonl"
assert result.bytes == 1234
assert result.created_at == 1234567890
assert _is_base64_encoded_unified_file_id(result.id)

View file

@ -97,6 +97,11 @@ class TestIdempotentErrorDetection:
error_message = "constraint 'fk_user_id' already exists"
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
def test_is_idempotent_error_does_not_exist(self):
"""Test detection of 'does not exist' error"""
error_message = "ERROR: index 'idx' does not exist"
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
def test_is_idempotent_error_case_insensitive(self):
"""Test that idempotent error detection is case insensitive"""
error_message = "COLUMN 'ID' ALREADY EXISTS"

View file

@ -203,6 +203,7 @@ class TestOCIImageUrlTransformation:
"""Tests for OCI image_url format handling in multimodal messages.
Fixes: https://github.com/BerriAI/litellm/issues/18270
Fixes: https://github.com/BerriAI/litellm/issues/19589
"""
def test_image_url_as_string(self):
@ -224,7 +225,8 @@ class TestOCIImageUrlTransformation:
assert len(result) == 1
assert result[0].role == "USER"
assert len(result[0].content) == 2
assert result[0].content[1].imageUrl == "https://example.com/image.png"
# imageUrl is now an OCIImageUrl object with a 'url' property
assert result[0].content[1].imageUrl.url == "https://example.com/image.png"
def test_image_url_as_openai_object(self):
"""Test that image_url as OpenAI-style object {"url": "..."} works."""
@ -245,7 +247,38 @@ class TestOCIImageUrlTransformation:
assert len(result) == 1
assert result[0].role == "USER"
assert len(result[0].content) == 2
assert result[0].content[1].imageUrl == "https://example.com/image.png"
# imageUrl is now an OCIImageUrl object with a 'url' property
assert result[0].content[1].imageUrl.url == "https://example.com/image.png"
def test_image_url_serializes_as_object(self):
"""Test that imageUrl serializes as {"url": "..."} for OCI API.
Fixes: https://github.com/BerriAI/litellm/issues/19589
OCI expects imageUrl to be an object with a 'url' property, not a plain string.
"""
from litellm.llms.oci.chat.transformation import adapt_messages_to_generic_oci_standard
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "Describe this image."},
{"type": "image_url", "image_url": {"url": "data:image/png;base64,ABC123"}},
],
}
]
result = adapt_messages_to_generic_oci_standard(messages)
image_part = result[0].content[1]
# Serialize as OCI would receive it (with exclude_none=True)
serialized = image_part.model_dump(exclude_none=True)
# Verify the structure matches OCI's expected format
assert serialized == {
"type": "IMAGE",
"imageUrl": {"url": "data:image/png;base64,ABC123"}
}
def test_image_url_invalid_type_raises_error(self):
"""Test that invalid image_url type raises an error."""

View file

@ -5,9 +5,7 @@ Tests message transformation, parameter handling, and response transformation.
Run with: pytest tests/llm_translation/test_gigachat.py -v
"""
import json
import pytest
from unittest.mock import Mock, MagicMock
class TestGigaChatMessageTransformation:
@ -16,6 +14,7 @@ class TestGigaChatMessageTransformation:
@pytest.fixture
def config(self):
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
return GigaChatConfig()
def test_simple_user_message(self, config):
@ -52,20 +51,46 @@ class TestGigaChatMessageTransformation:
assert result[0]["role"] == "function"
def test_tool_content_convertation_non_string_value(self, config):
"""Non string tool content should be serialized"""
messages = [{"role": "tool", "content": {"output": 42}}]
result = config._transform_messages(messages)
assert result[0]["content"] == '{"output": 42}'
def test_tool_content_convertation_json_string_value(self, config):
"""JSON string tool content left unchanged"""
valid_json = '{"output": "red car"}'
messages = [{"role": "tool", "content": valid_json}]
result = config._transform_messages(messages)
assert result[0]["content"] == valid_json
def test_tool_content_convertation_random_string_value(self, config):
"""Non JSON tool content should be serialized"""
messages = [{"role": "tool", "content": "random string"}]
result = config._transform_messages(messages)
assert result[0]["content"] == '"random string"'
def test_tool_calls_to_function_call(self, config):
"""tool_calls should be converted to function_call"""
messages = [{
"role": "assistant",
"content": "",
"tool_calls": [{
"id": "call_123",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"city": "Moscow"}'
}
}]
}]
messages = [
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_123",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"city": "Moscow"}',
},
}
],
}
]
result = config._transform_messages(messages)
assert "function_call" in result[0]
@ -94,6 +119,7 @@ class TestGigaChatCollapseUserMessages:
@pytest.fixture
def config(self):
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
return GigaChatConfig()
def test_no_collapse_single_message(self, config):
@ -136,23 +162,24 @@ class TestGigaChatToolsTransformation:
@pytest.fixture
def config(self):
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
return GigaChatConfig()
def test_single_tool_conversion(self, config):
"""Single tool should be converted correctly"""
tools = [{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather for a city",
"parameters": {
"type": "object",
"properties": {
"city": {"type": "string"}
}
}
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather for a city",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
},
},
}
}]
]
result = config._convert_tools_to_functions(tools)
assert len(result) == 1
@ -162,8 +189,22 @@ class TestGigaChatToolsTransformation:
def test_multiple_tools_conversion(self, config):
"""Multiple tools should all be converted"""
tools = [
{"type": "function", "function": {"name": "func1", "description": "First", "parameters": {"type": "object", "properties": {}}}},
{"type": "function", "function": {"name": "func2", "description": "Second", "parameters": {"type": "object", "properties": {}}}},
{
"type": "function",
"function": {
"name": "func1",
"description": "First",
"parameters": {"type": "object", "properties": {}},
},
},
{
"type": "function",
"function": {
"name": "func2",
"description": "Second",
"parameters": {"type": "object", "properties": {}},
},
},
]
result = config._convert_tools_to_functions(tools)
@ -178,6 +219,7 @@ class TestGigaChatParamsTransformation:
@pytest.fixture
def config(self):
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
return GigaChatConfig()
def test_temperature_zero_becomes_top_p_zero(self, config):
@ -229,10 +271,10 @@ class TestGigaChatParamsTransformation:
"type": "object",
"properties": {
"name": {"type": "string"},
"age": {"type": "integer"}
}
}
}
"age": {"type": "integer"},
},
},
},
}
}
result = config.map_openai_params(
@ -283,6 +325,7 @@ class TestGigaChatTransformRequest:
@pytest.fixture
def config(self):
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
return GigaChatConfig()
def test_basic_request(self, config):
@ -335,6 +378,7 @@ class TestGigaChatSupportedParams:
@pytest.fixture
def config(self):
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
return GigaChatConfig()
def test_supported_params(self, config):

View file

@ -1005,6 +1005,142 @@ class TestPollingConditionChecks:
assert result is False
# ==================== Native Background Mode Tests ====================
def test_polling_disabled_when_model_in_native_background_mode(self):
"""Test that polling is disabled when model is in native_background_mode list"""
from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request
result = should_use_polling_for_request(
background_mode=True,
polling_via_cache_enabled="all",
redis_cache=Mock(),
model="o4-mini-deep-research",
llm_router=None,
native_background_mode=["o4-mini-deep-research", "o3-deep-research"],
)
assert result is False
def test_polling_disabled_for_native_background_mode_with_provider_list(self):
"""Test that native_background_mode takes precedence even when provider matches"""
from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request
result = should_use_polling_for_request(
background_mode=True,
polling_via_cache_enabled=["openai"],
redis_cache=Mock(),
model="openai/o4-mini-deep-research",
llm_router=None,
native_background_mode=["openai/o4-mini-deep-research"],
)
assert result is False
def test_polling_enabled_when_model_not_in_native_background_mode(self):
"""Test that polling is enabled when model is not in native_background_mode list"""
from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request
result = should_use_polling_for_request(
background_mode=True,
polling_via_cache_enabled="all",
redis_cache=Mock(),
model="gpt-4o",
llm_router=None,
native_background_mode=["o4-mini-deep-research"],
)
assert result is True
def test_polling_enabled_when_native_background_mode_is_none(self):
"""Test that polling works normally when native_background_mode is None"""
from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request
result = should_use_polling_for_request(
background_mode=True,
polling_via_cache_enabled="all",
redis_cache=Mock(),
model="gpt-4o",
llm_router=None,
native_background_mode=None,
)
assert result is True
def test_polling_enabled_when_native_background_mode_is_empty_list(self):
"""Test that polling works normally when native_background_mode is empty list"""
from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request
result = should_use_polling_for_request(
background_mode=True,
polling_via_cache_enabled="all",
redis_cache=Mock(),
model="gpt-4o",
llm_router=None,
native_background_mode=[],
)
assert result is True
def test_native_background_mode_exact_match_required(self):
"""Test that native_background_mode uses exact model name matching"""
from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request
# "o4-mini" should not match "o4-mini-deep-research"
result = should_use_polling_for_request(
background_mode=True,
polling_via_cache_enabled="all",
redis_cache=Mock(),
model="o4-mini",
llm_router=None,
native_background_mode=["o4-mini-deep-research"],
)
assert result is True
def test_native_background_mode_with_provider_prefix_in_request(self):
"""Test native_background_mode matching when request model has provider prefix"""
from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request
# Model in native_background_mode without provider prefix
# Request comes in with provider prefix - should not match
result = should_use_polling_for_request(
background_mode=True,
polling_via_cache_enabled=["openai"],
redis_cache=Mock(),
model="openai/o4-mini-deep-research",
llm_router=None,
native_background_mode=["o4-mini-deep-research"], # Without prefix
)
# Should return True because "openai/o4-mini-deep-research" != "o4-mini-deep-research"
assert result is True
def test_native_background_mode_with_router_lookup(self):
"""Test that native_background_mode works with router-resolved models"""
from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request
mock_router = Mock()
mock_router.model_name_to_deployment_indices = {"deep-research": [0]}
mock_router.model_list = [
{
"model_name": "deep-research",
"litellm_params": {"model": "openai/o4-mini-deep-research"}
}
]
# Model alias "deep-research" is in native_background_mode
result = should_use_polling_for_request(
background_mode=True,
polling_via_cache_enabled=["openai"],
redis_cache=Mock(),
model="deep-research",
llm_router=mock_router,
native_background_mode=["deep-research"],
)
assert result is False
class TestStreamingEventParsing:
"""

View file

@ -106,6 +106,30 @@ def test_prometheus_metric_labels_structure():
print(f"✅ {metric_name} has proper label structure with user_email")
def test_model_id_in_required_metrics():
"""
Test that model_id label is present in all the metrics that should have it:
- litellm_proxy_total_requests_metric
- litellm_proxy_failed_requests_metric
- litellm_request_total_latency_metric
- litellm_llm_api_time_to_first_token_metric
"""
model_id_label = UserAPIKeyLabelNames.MODEL_ID.value
# Metrics that should have model_id
metrics_with_model_id = [
"litellm_proxy_total_requests_metric",
"litellm_proxy_failed_requests_metric",
"litellm_request_total_latency_metric",
"litellm_llm_api_time_to_first_token_metric"
]
for metric_name in metrics_with_model_id:
labels = PrometheusMetricLabels.get_labels(metric_name)
assert model_id_label in labels, f"Metric {metric_name} should contain model_id label"
print(f"✅ {metric_name} contains model_id label")
def test_route_normalization_for_responses_api():
"""
Test that route normalization prevents high cardinality in Prometheus metrics

View file

@ -309,3 +309,46 @@ class TestResponseAPILoggingUtils:
assert result.completion_tokens_details.reasoning_tokens == 30
assert result.completion_tokens_details.image_tokens == 100
assert result.completion_tokens_details.text_tokens == 70
class TestResponsesAPIProviderSpecificParams:
"""
Tests for fix #19782: provider-specific params (aws_*, vertex_*) should work
without explicitly passing custom_llm_provider.
"""
def test_provider_specific_params_no_crash_with_bedrock(self):
"""Test that processing aws_* params with bedrock provider doesn't crash."""
params = {
"temperature": 0.7,
"custom_llm_provider": "bedrock",
"kwargs": {"aws_region_name": "eu-central-1"},
}
# Should not raise any exception
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params)
assert "temperature" in result
def test_provider_specific_params_no_crash_with_openai(self):
"""Test that processing aws_* params with openai provider doesn't crash."""
params = {
"temperature": 0.7,
"custom_llm_provider": "openai",
"kwargs": {"aws_region_name": "eu-central-1"},
}
# Should not raise any exception
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params)
assert "temperature" in result
def test_provider_specific_params_no_crash_with_vertex_ai(self):
"""Test that processing vertex_* params with vertex_ai provider doesn't crash."""
params = {
"temperature": 0.7,
"custom_llm_provider": "vertex_ai",
"kwargs": {"vertex_project": "my-project"},
}
# Should not raise any exception
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params)
assert "temperature" in result

View file

@ -1,4 +1,3 @@
import json
import os
import sys
@ -8,7 +7,6 @@ sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
from unittest.mock import MagicMock, patch
from pydantic import BaseModel
@ -77,7 +75,9 @@ def test_cost_calculator_with_usage(monkeypatch):
prompt_tokens=120,
completion_tokens=100,
prompt_tokens_details=PromptTokensDetailsWrapper(
text_tokens=10, audio_tokens=90, image_tokens=20,
text_tokens=10,
audio_tokens=90,
image_tokens=20,
),
)
mr = ModelResponse(usage=usage, model="gemini-2.0-flash-001")
@ -96,7 +96,9 @@ def test_cost_calculator_with_usage(monkeypatch):
# Step 1: Test a model where input_cost_per_image_token is not set.
# In this case the calculation should use input_cost_per_token as fallback.
assert model_info.get("input_cost_per_image_token") is None, "Test case expects that input_cost_per_image_token is not set"
assert (
model_info.get("input_cost_per_image_token") is None
), "Test case expects that input_cost_per_image_token is not set"
expected_cost = (
usage.prompt_tokens_details.audio_tokens
@ -116,9 +118,7 @@ def test_cost_calculator_with_usage(monkeypatch):
monkeypatch.setattr(
litellm,
"model_cost",
{
"gemini-2.0-flash-001": temp_model_info_object
},
{"gemini-2.0-flash-001": temp_model_info_object},
)
# Invalidate caches after modifying litellm.model_cost
@ -138,8 +138,10 @@ def test_cost_calculator_with_usage(monkeypatch):
expected_cost = (
usage.prompt_tokens_details.audio_tokens
* temp_model_info_object["input_cost_per_audio_token"]
+ usage.prompt_tokens_details.text_tokens * temp_model_info_object["input_cost_per_token"]
+ usage.prompt_tokens_details.image_tokens * temp_model_info_object["input_cost_per_image_token"]
+ usage.prompt_tokens_details.text_tokens
* temp_model_info_object["input_cost_per_token"]
+ usage.prompt_tokens_details.image_tokens
* temp_model_info_object["input_cost_per_image_token"]
+ usage.completion_tokens * temp_model_info_object["output_cost_per_token"]
)
@ -331,8 +333,6 @@ def test_custom_pricing_with_router_model_id():
def test_azure_realtime_cost_calculator():
from litellm import get_model_info
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
@ -357,6 +357,90 @@ def test_azure_realtime_cost_calculator():
assert cost > 0
def test_azure_audio_output_cost_calculation():
"""
Test that Azure audio models correctly calculate costs for audio output tokens.
Reproduces issue: https://github.com/BerriAI/litellm/issues/19764
Audio tokens should be charged at output_cost_per_audio_token rate,
not at the text token rate (output_cost_per_token).
"""
from litellm.types.utils import (
Choices,
CompletionTokensDetailsWrapper,
Message,
)
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
# Scenario from issue #19764:
# Input: 17 text tokens, 0 audio tokens
# Output: 110 text tokens, 482 audio tokens
usage_object = Usage(
prompt_tokens=17,
completion_tokens=592, # 110 text + 482 audio
total_tokens=609,
prompt_tokens_details=PromptTokensDetailsWrapper(
audio_tokens=0,
cached_tokens=0,
text_tokens=17,
image_tokens=0,
),
completion_tokens_details=CompletionTokensDetailsWrapper(
audio_tokens=482,
reasoning_tokens=0,
text_tokens=110,
),
)
completion = ModelResponse(
id="test-azure-audio-cost",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(
content="Test response",
role="assistant",
),
)
],
created=1729282652,
model="azure/gpt-audio-2025-08-28",
object="chat.completion",
usage=usage_object,
)
cost = completion_cost(completion, model="azure/gpt-audio-2025-08-28")
model_info = litellm.get_model_info("azure/gpt-audio-2025-08-28")
# Calculate expected cost
expected_input_cost = (
model_info["input_cost_per_token"] * 17 # text tokens
)
expected_output_cost = (
model_info["output_cost_per_token"] * 110 # text tokens
+ model_info["output_cost_per_audio_token"] * 482 # audio tokens
)
expected_total_cost = expected_input_cost + expected_output_cost
# The bug was: all output tokens charged at text rate
wrong_output_cost = model_info["output_cost_per_token"] * 592
wrong_total_cost = expected_input_cost + wrong_output_cost
# Verify audio tokens are NOT charged at text rate (the bug)
assert abs(cost - wrong_total_cost) > 0.001, (
"Bug: Audio tokens are being charged at text token rate"
)
# Verify cost matches
assert abs(cost - expected_total_cost) < 0.0000001, (
f"Expected cost {expected_total_cost}, got {cost}"
)
def test_default_image_cost_calculator(monkeypatch):
from litellm.cost_calculator import default_image_cost_calculator
@ -389,9 +473,7 @@ def test_cost_calculator_with_cache_creation():
from litellm import completion_cost
from litellm.types.utils import (
Choices,
CompletionTokensDetailsWrapper,
Message,
PromptTokensDetailsWrapper,
Usage,
)
@ -447,7 +529,7 @@ def test_cost_calculator_with_cache_creation():
def test_bedrock_cost_calculator_comparison_with_without_cache():
"""Test that Bedrock caching reduces costs compared to non-cached requests"""
from litellm import completion_cost
from litellm.types.utils import Choices, Message, PromptTokensDetailsWrapper, Usage
from litellm.types.utils import Choices, Message, Usage
# Response WITHOUT caching
response_no_cache = ModelResponse(
@ -698,7 +780,7 @@ def test_log_context_cost_calculation():
f"DEBUG: Tiered input cost per token (>200k): ${input_cost_above_200k:.2e}"
)
else:
print(f"DEBUG: No tiered input pricing available, using base pricing")
print("DEBUG: No tiered input pricing available, using base pricing")
input_cost_above_200k = input_cost_per_token
if output_cost_above_200k is not None:
@ -706,7 +788,7 @@ def test_log_context_cost_calculation():
f"DEBUG: Tiered output cost per token (>200k): ${output_cost_above_200k:.2e}"
)
else:
print(f"DEBUG: No tiered output pricing available, using base pricing")
print("DEBUG: No tiered output pricing available, using base pricing")
output_cost_above_200k = output_cost_per_token
if cache_creation_above_200k is not None:
@ -714,7 +796,7 @@ def test_log_context_cost_calculation():
f"DEBUG: Tiered cache creation cost per token (>200k): ${cache_creation_above_200k:.2e}"
)
else:
print(f"DEBUG: No tiered cache creation pricing available, using base pricing")
print("DEBUG: No tiered cache creation pricing available, using base pricing")
cache_creation_above_200k = cache_creation_cost_per_token
# Since we're above 200k tokens, we should use tiered pricing if available
@ -923,7 +1005,7 @@ def test_cost_discount_vertex_ai():
expected_cost = cost_without_discount * 0.95
assert cost_with_discount == pytest.approx(expected_cost, rel=1e-9)
print(f"✓ Cost discount test passed:")
print("✓ Cost discount test passed:")
print(f" - Original cost: ${cost_without_discount:.6f}")
print(f" - Discounted cost (5% off): ${cost_with_discount:.6f}")
print(f" - Savings: ${cost_without_discount - cost_with_discount:.6f}")
@ -973,7 +1055,7 @@ def test_cost_discount_not_applied_to_other_providers():
# Costs should be the same (no discount applied to OpenAI)
assert cost_with_selective_discount == cost_without_discount
print(f"✓ Selective discount test passed:")
print("✓ Selective discount test passed:")
print(f" - OpenAI cost (no discount configured): ${cost_without_discount:.6f}")
print(f" - Cost remains unchanged: ${cost_with_selective_discount:.6f}")
@ -1023,7 +1105,7 @@ def test_cost_margin_percentage():
expected_cost = cost_without_margin * 1.10
assert cost_with_margin == pytest.approx(expected_cost, rel=1e-9)
print(f"✓ Cost margin percentage test passed:")
print("✓ Cost margin percentage test passed:")
print(f" - Original cost: ${cost_without_margin:.6f}")
print(f" - Cost with margin (10%): ${cost_with_margin:.6f}")
print(f" - Margin added: ${cost_with_margin - cost_without_margin:.6f}")
@ -1074,7 +1156,7 @@ def test_cost_margin_fixed_amount():
expected_cost = cost_without_margin + 0.001
assert cost_with_margin == pytest.approx(expected_cost, rel=1e-9)
print(f"✓ Cost margin fixed amount test passed:")
print("✓ Cost margin fixed amount test passed:")
print(f" - Original cost: ${cost_without_margin:.6f}")
print(f" - Cost with margin ($0.001): ${cost_with_margin:.6f}")
print(f" - Margin added: ${cost_with_margin - cost_without_margin:.6f}")
@ -1109,7 +1191,9 @@ def test_cost_margin_combined():
)
# Set 8% margin + $0.0005 fixed for openai
litellm.cost_margin_config = {"openai": {"percentage": 0.08, "fixed_amount": 0.0005}}
litellm.cost_margin_config = {
"openai": {"percentage": 0.08, "fixed_amount": 0.0005}
}
# Calculate cost with margin
cost_with_margin = completion_cost(
@ -1125,7 +1209,7 @@ def test_cost_margin_combined():
expected_cost = cost_without_margin * 1.08 + 0.0005
assert cost_with_margin == pytest.approx(expected_cost, rel=1e-9)
print(f"✓ Cost margin combined test passed:")
print("✓ Cost margin combined test passed:")
print(f" - Original cost: ${cost_without_margin:.6f}")
print(f" - Cost with margin (8% + $0.0005): ${cost_with_margin:.6f}")
print(f" - Margin added: ${cost_with_margin - cost_without_margin:.6f}")
@ -1176,7 +1260,7 @@ def test_cost_margin_global():
expected_cost = cost_without_margin * 1.05
assert cost_with_global_margin == pytest.approx(expected_cost, rel=1e-9)
print(f"✓ Cost margin global test passed:")
print("✓ Cost margin global test passed:")
print(f" - Original cost: ${cost_without_margin:.6f}")
print(f" - Cost with global margin (5%): ${cost_with_global_margin:.6f}")
print(f" - Margin added: ${cost_with_global_margin - cost_without_margin:.6f}")
@ -1227,9 +1311,11 @@ def test_cost_margin_provider_overrides_global():
expected_cost = cost_without_margin * 1.10 # 10% from provider, not 5% from global
assert cost_with_provider_margin == pytest.approx(expected_cost, rel=1e-9)
print(f"✓ Cost margin provider override test passed:")
print("✓ Cost margin provider override test passed:")
print(f" - Original cost: ${cost_without_margin:.6f}")
print(f" - Cost with provider margin (10%, overrides 5% global): ${cost_with_provider_margin:.6f}")
print(
f" - Cost with provider margin (10%, overrides 5% global): ${cost_with_provider_margin:.6f}"
)
print(f" - Margin added: ${cost_with_provider_margin - cost_without_margin:.6f}")
@ -1283,7 +1369,7 @@ def test_cost_margin_with_discount():
expected_cost = base_cost * 0.95 * 1.10
assert cost_with_both == pytest.approx(expected_cost, rel=1e-9)
print(f"✓ Cost margin with discount test passed:")
print("✓ Cost margin with discount test passed:")
print(f" - Base cost: ${base_cost:.6f}")
print(f" - Cost with 5% discount + 10% margin: ${cost_with_both:.6f}")
print(f" - Expected: ${expected_cost:.6f}")
@ -1352,14 +1438,10 @@ def test_completion_cost_extracts_service_tier_from_response():
# Test with gpt-5-nano which has flex pricing
model = "gpt-5-nano"
# Create usage object
usage = Usage(
prompt_tokens=1000,
completion_tokens=500,
total_tokens=1500
)
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
# Create ModelResponse with service_tier in the response object
response_with_service_tier = ModelResponse(
usage=usage,
@ -1367,34 +1449,36 @@ def test_completion_cost_extracts_service_tier_from_response():
)
# Set service_tier as an attribute on the response
setattr(response_with_service_tier, "service_tier", "flex")
# Test that flex pricing is used when service_tier is in response
flex_cost = completion_cost(
completion_response=response_with_service_tier,
model=model,
custom_llm_provider="openai",
)
# Create ModelResponse without service_tier (should use standard pricing)
response_without_service_tier = ModelResponse(
usage=usage,
model=model,
)
# Test that standard pricing is used when service_tier is not in response
standard_cost = completion_cost(
completion_response=response_without_service_tier,
model=model,
custom_llm_provider="openai",
)
# Flex should be approximately 50% of standard
assert flex_cost > 0, "Flex cost should be greater than 0"
assert standard_cost > 0, "Standard cost should be greater than 0"
assert flex_cost < standard_cost, "Flex cost should be less than standard cost"
flex_ratio = flex_cost / standard_cost
assert 0.45 <= flex_ratio <= 0.55, f"Flex pricing should be ~50% of standard, got {flex_ratio:.2f}"
assert (
0.45 <= flex_ratio <= 0.55
), f"Flex pricing should be ~50% of standard, got {flex_ratio:.2f}"
def test_completion_cost_extracts_service_tier_from_usage():
@ -1406,56 +1490,54 @@ def test_completion_cost_extracts_service_tier_from_usage():
# Test with gpt-5-nano which has flex pricing
model = "gpt-5-nano"
# Create usage object with service_tier
usage_with_service_tier = Usage(
prompt_tokens=1000,
completion_tokens=500,
total_tokens=1500
prompt_tokens=1000, completion_tokens=500, total_tokens=1500
)
# Set service_tier as an attribute on the usage object
setattr(usage_with_service_tier, "service_tier", "flex")
# Create ModelResponse with usage containing service_tier
response = ModelResponse(
usage=usage_with_service_tier,
model=model,
)
# Test that flex pricing is used when service_tier is in usage
flex_cost = completion_cost(
completion_response=response,
model=model,
custom_llm_provider="openai",
)
# Create usage object without service_tier
usage_without_service_tier = Usage(
prompt_tokens=1000,
completion_tokens=500,
total_tokens=1500
prompt_tokens=1000, completion_tokens=500, total_tokens=1500
)
# Create ModelResponse with usage without service_tier
response_standard = ModelResponse(
usage=usage_without_service_tier,
model=model,
)
# Test that standard pricing is used when service_tier is not in usage
standard_cost = completion_cost(
completion_response=response_standard,
model=model,
custom_llm_provider="openai",
)
# Flex should be approximately 50% of standard
assert flex_cost > 0, "Flex cost should be greater than 0"
assert standard_cost > 0, "Standard cost should be greater than 0"
assert flex_cost < standard_cost, "Flex cost should be less than standard cost"
flex_ratio = flex_cost / standard_cost
assert 0.45 <= flex_ratio <= 0.55, f"Flex pricing should be ~50% of standard, got {flex_ratio:.2f}"
assert (
0.45 <= flex_ratio <= 0.55
), f"Flex pricing should be ~50% of standard, got {flex_ratio:.2f}"
def test_completion_cost_service_tier_priority():
@ -1467,22 +1549,18 @@ def test_completion_cost_service_tier_priority():
# Test with gpt-5-nano which has flex pricing
model = "gpt-5-nano"
# Create usage object with service_tier="flex"
usage = Usage(
prompt_tokens=1000,
completion_tokens=500,
total_tokens=1500
)
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
setattr(usage, "service_tier", "flex")
# Create response with service_tier="priority"
response = ModelResponse(
usage=usage,
model=model,
)
setattr(response, "service_tier", "priority")
# Test that optional_params takes priority over response and usage
cost_from_params = completion_cost(
completion_response=response,
@ -1490,14 +1568,14 @@ def test_completion_cost_service_tier_priority():
custom_llm_provider="openai",
optional_params={"service_tier": "flex"},
)
# Test that response takes priority over usage when optional_params is not provided
cost_from_response = completion_cost(
completion_cost(
completion_response=response,
model=model,
custom_llm_provider="openai",
)
# Test that usage is used when neither optional_params nor response have service_tier
# Create a new response without service_tier attribute
response_no_tier = ModelResponse(
@ -1505,25 +1583,27 @@ def test_completion_cost_service_tier_priority():
model=model,
)
# Don't set service_tier on response, so it will fall back to usage
cost_from_usage = completion_cost(
completion_response=response_no_tier,
model=model,
custom_llm_provider="openai",
)
# All should use flex pricing (from different sources)
assert cost_from_params > 0, "Cost from params should be greater than 0"
assert cost_from_usage > 0, "Cost from usage should be greater than 0"
# Costs should be similar (all using flex)
assert abs(cost_from_params - cost_from_usage) < 1e-6, "Costs from params and usage should be similar (both flex)"
assert (
abs(cost_from_params - cost_from_usage) < 1e-6
), "Costs from params and usage should be similar (both flex)"
def test_gemini_cache_tokens_details_no_negative_values():
"""
Test for Issue #18750: Negative text_tokens with Gemini caching
When using Gemini with explicit caching, the response includes cacheTokensDetails
which breaks down cached tokens by modality. This test ensures that:
1. text_tokens is never negative
@ -1544,41 +1624,47 @@ def test_gemini_cache_tokens_details_no_negative_values():
# Total tokens by modality (includes cached + non-cached)
"promptTokensDetails": [
{"modality": "TEXT", "tokenCount": 9402},
{"modality": "IMAGE", "tokenCount": 258}
{"modality": "IMAGE", "tokenCount": 258},
],
# Breakdown of cached tokens by modality
"cacheTokensDetails": [
{"modality": "TEXT", "tokenCount": 9393},
{"modality": "IMAGE", "tokenCount": 258}
]
{"modality": "IMAGE", "tokenCount": 258},
],
}
}
usage = VertexGeminiConfig._calculate_usage(completion_response)
# Text tokens should be non-cached text only: 9402 - 9393 = 9
assert usage.prompt_tokens_details.text_tokens == 9, \
f"Expected text_tokens=9, got {usage.prompt_tokens_details.text_tokens}"
assert (
usage.prompt_tokens_details.text_tokens == 9
), f"Expected text_tokens=9, got {usage.prompt_tokens_details.text_tokens}"
# Image tokens should be non-cached image only: 258 - 258 = 0
assert usage.prompt_tokens_details.image_tokens == 0, \
f"Expected image_tokens=0, got {usage.prompt_tokens_details.image_tokens}"
assert (
usage.prompt_tokens_details.image_tokens == 0
), f"Expected image_tokens=0, got {usage.prompt_tokens_details.image_tokens}"
# Total cached should match
assert usage.prompt_tokens_details.cached_tokens == 9651, \
f"Expected cached_tokens=9651, got {usage.prompt_tokens_details.cached_tokens}"
assert (
usage.prompt_tokens_details.cached_tokens == 9651
), f"Expected cached_tokens=9651, got {usage.prompt_tokens_details.cached_tokens}"
# MOST IMPORTANT: text_tokens should NEVER be negative
assert usage.prompt_tokens_details.text_tokens >= 0, \
f"BUG: text_tokens is negative ({usage.prompt_tokens_details.text_tokens})! This was the issue in #18750"
assert (
usage.prompt_tokens_details.text_tokens >= 0
), f"BUG: text_tokens is negative ({usage.prompt_tokens_details.text_tokens})! This was the issue in #18750"
print("✅ Issue #18750 fix verified: text_tokens is correctly calculated and non-negative")
print(
"✅ Issue #18750 fix verified: text_tokens is correctly calculated and non-negative"
)
def test_gemini_without_cache_tokens_details():
"""
Test Gemini response without cacheTokensDetails (implicit caching or no cache)
When cacheTokensDetails is not present, we should use promptTokensDetails as-is
without subtracting anything.
"""
@ -1593,7 +1679,7 @@ def test_gemini_without_cache_tokens_details():
"totalTokenCount": 279,
"promptTokensDetails": [
{"modality": "TEXT", "tokenCount": 6},
{"modality": "IMAGE", "tokenCount": 258}
{"modality": "IMAGE", "tokenCount": 258},
]
# No cacheTokensDetails
}
@ -1607,3 +1693,62 @@ def test_gemini_without_cache_tokens_details():
assert usage.prompt_tokens_details.text_tokens >= 0
print("✅ Gemini without cacheTokensDetails works correctly")
def test_generic_provider_cached_token_cost():
"""
Test that the generic cost calculator correctly handles cached tokens
for providers like z.ai/deepseek that are not explicitly handled.
"""
from litellm.cost_calculator import completion_cost
from litellm.types.utils import ModelResponse, PromptTokensDetailsWrapper, Usage
# Setup model cost for a generic provider
# We use a name that will bypass complex provider mapping logic
model_name = "custom-cached-model"
litellm.model_cost[model_name] = {
"input_cost_per_token": 0.0000006,
"output_cost_per_token": 0.0000006,
"cache_read_input_token_cost": 0.0000001,
"litellm_provider": "openai",
}
# Case 1: Standard nested cached tokens (prompt_tokens_details.cached_tokens)
usage = Usage(
prompt_tokens=10000,
completion_tokens=0,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=9000),
)
response = ModelResponse(usage=usage, model=model_name)
cost = completion_cost(
completion_response=response,
model=model_name,
custom_llm_provider="openai", # Explicitly set provider to trigger generic path
)
# Expected: (1000 * 0.0000006) + (9000 * 0.0000001) = 0.0006 + 0.0009 = 0.0015
expected_cost = 0.0015
assert (
abs(cost - expected_cost) < 1e-9
), f"Nested cache cost failed. Got {cost}, expected {expected_cost}"
# Case 2: Top-level cached tokens (cache_read_input_tokens)
usage_top = Usage(
prompt_tokens=10000,
completion_tokens=0,
cache_read_input_tokens=9000,
)
response_top = ModelResponse(usage=usage_top, model=model_name)
cost_top = completion_cost(
completion_response=response_top,
model=model_name,
custom_llm_provider="openai",
)
assert (
abs(cost_top - expected_cost) < 1e-9
), f"Top-level cache cost failed. Got {cost_top}, expected {expected_cost}"
print("✅ Generic provider cached token cost verified")