refactor: expose core private helpers under public names (#44871)

* refactor: expose core private symbols with compatibility aliases

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* types: narrow core migration diagnostics

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix: preserve runtime behavior in core symbol migration

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix: preserve core private usage migration behavior

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix: preserve private value rebinding compatibility

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix: restore optional imports and cover public helpers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test: exempt router property from call coverage

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test: align recursive detector ignore names

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: mateo <mateo@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-07 00:05:27 +00:00 • committed by GitHub
parent b9ab1eee1f
commit c3a23fe499
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
461 changed files with 6719 additions and 4996 deletions

View file

@ -1,15 +1,15 @@
from typing import Dict, Literal, Type, Union
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles
from litellm_enterprise.proxy.hooks.managed_vector_stores import (
_PROXY_LiteLLMManagedVectorStores,
PROXY_LiteLLMManagedVectorStores,
)
from litellm.integrations.custom_logger import CustomLogger
ENTERPRISE_PROXY_HOOKS: Dict[str, Type[CustomLogger]] = {
"managed_files": _PROXY_LiteLLMManagedFiles,
"managed_vector_stores": _PROXY_LiteLLMManagedVectorStores,
"managed_files": PROXY_LiteLLMManagedFiles,
"managed_vector_stores": PROXY_LiteLLMManagedVectorStores,
}

View file

@ -20,7 +20,7 @@ from litellm._logging import verbose_proxy_logger
from fastapi import HTTPException
class _ENTERPRISE_BannedKeywords(CustomLogger):
class ENTERPRISE_BannedKeywords(CustomLogger):
enforces_request_content: bool = True
# Class variables or attributes
def __init__(self):
@ -114,3 +114,4 @@ class _ENTERPRISE_BannedKeywords(CustomLogger):
response: str,
):
self.test_violation(test_str=response)
_ENTERPRISE_BannedKeywords = ENTERPRISE_BannedKeywords

View file

@ -21,7 +21,7 @@ from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET
from litellm.proxy.utils import PrismaClient
class _ENTERPRISE_BlockedUserList(CustomLogger):
class ENTERPRISE_BlockedUserList(CustomLogger):
enforces_request_content: bool = True
# Class variables or attributes
def __init__(self, prisma_client: Optional[PrismaClient]):
@ -128,3 +128,4 @@ class _ENTERPRISE_BlockedUserList(CustomLogger):
str(e)
)
)
_ENTERPRISE_BlockedUserList = ENTERPRISE_BlockedUserList

View file

@ -16,7 +16,7 @@ from litellm.proxy.guardrails._content_utils import iter_message_text
from litellm.types.utils import CallTypesLiteral
class _ENTERPRISE_GoogleTextModeration(CustomLogger):
class ENTERPRISE_GoogleTextModeration(CustomLogger):
user_api_key_cache = None
confidence_categories = [
"toxic",
@ -125,9 +125,10 @@ class _ENTERPRISE_GoogleTextModeration(CustomLogger):
)
# Handle the response
return data
_ENTERPRISE_GoogleTextModeration = ENTERPRISE_GoogleTextModeration
# google_text_moderation_obj = _ENTERPRISE_GoogleTextModeration()
# google_text_moderation_obj = ENTERPRISE_GoogleTextModeration()
# asyncio.run(
# google_text_moderation_obj.async_moderation_hook(
# data={"messages": [{"role": "user", "content": "Hey, how's it going?"}]}

View file

@ -24,7 +24,7 @@ from litellm.proxy.guardrails._content_utils import iter_message_text
from litellm.types.utils import CallTypesLiteral
class _ENTERPRISE_OpenAI_Moderation(CustomLogger):
class ENTERPRISE_OpenAI_Moderation(CustomLogger):
@property
def model_name(self) -> str:
return litellm.openai_moderations_model_name or DEFAULT_OPENAI_MODERATIONS_MODEL
@ -55,3 +55,4 @@ class _ENTERPRISE_OpenAI_Moderation(CustomLogger):
status_code=403, detail={"error": "Violated content safety policy"}
)
pass
_ENTERPRISE_OpenAI_Moderation = ENTERPRISE_OpenAI_Moderation

View file

@ -3,9 +3,10 @@ Endpoints for managing email alerts on litellm
"""
import json
from typing import Dict
from typing import Dict, Final, cast
from fastapi import APIRouter, Depends, HTTPException
from pydantic import JsonValue
from litellm_enterprise.types.enterprise_callbacks.send_emails import (
DefaultEmailSettings,
EmailEvent,
@ -38,14 +39,15 @@ async def _get_email_settings(prisma_client) -> Dict[str, bool]:
and general_settings_entry.param_value is not None
):
# Get general settings value
if isinstance(general_settings_entry.param_value, str):
general_settings = json.loads(general_settings_entry.param_value)
else:
general_settings = general_settings_entry.param_value
general_settings: Final = (
cast(Dict[str, object], json.loads(general_settings_entry.param_value))
if isinstance(general_settings_entry.param_value, str)
else cast(Dict[str, object], general_settings_entry.param_value)
)
# Extract email_settings from general settings if it exists
if general_settings and "email_settings" in general_settings:
email_settings = general_settings["email_settings"]
email_settings: Final = cast(Dict[str, bool], general_settings["email_settings"])
# Update settings_dict with values from general_settings
for event_name, enabled in email_settings.items():
settings_dict[event_name] = enabled
@ -64,7 +66,7 @@ async def _save_email_settings(prisma_client, settings: Dict[str, bool]):
from litellm.proxy.proxy_server import proxy_config
proxy_config.reject_config_owned_writes(
section_name="general_settings", changed_keys={"email_settings": settings}
section_name="general_settings", changed_keys={"email_settings": cast(JsonValue, settings)}
)
try:
verbose_proxy_logger.debug(
@ -77,16 +79,15 @@ async def _save_email_settings(prisma_client, settings: Dict[str, bool]):
)
# Initialize general settings dict
if (
general_settings_entry is not None
and general_settings_entry.param_value is not None
):
if isinstance(general_settings_entry.param_value, str):
general_settings = json.loads(general_settings_entry.param_value)
else:
general_settings = dict(general_settings_entry.param_value)
else:
general_settings = {}
general_settings: Final = (
(
cast(Dict[str, object], json.loads(general_settings_entry.param_value))
if isinstance(general_settings_entry.param_value, str)
else cast(Dict[str, object], dict(general_settings_entry.param_value))
)
if general_settings_entry is not None and general_settings_entry.param_value is not None
else {}
)
# Update email_settings in general_settings
general_settings["email_settings"] = settings

View file

@ -36,7 +36,7 @@ class EnterpriseCustomGuardrailHelper:
proxy_server_request = data.get("proxy_server_request", {})
request_tags = StandardLoggingPayloadSetup._get_request_tags(
request_tags = StandardLoggingPayloadSetup.get_request_tags(
litellm_params=data,
proxy_server_request=proxy_server_request,
)

View file

@ -60,6 +60,32 @@ class _ManagedObjectRow(Protocol):
@property
def file_object(self) -> object: ...
@property
def org_id(self) -> str | None: ...
@property
def api_key(self) -> str | None: ...
@property
def team_id(self) -> str | None: ...
class _ReadableFileContent(Protocol):
async def read(self) -> bytes: ...
class _HasFileContent(Protocol):
@property
def content(self) -> bytes: ...
async def _file_content_bytes(file_content: object) -> bytes:
if hasattr(file_content, "content"):
return cast(_HasFileContent, file_content).content
if hasattr(file_content, "read"):
return await cast(_ReadableFileContent, file_content).read()
return cast(bytes, file_content)
def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]":
return ManagedObjectRepository(prisma_client).table
@ -175,17 +201,21 @@ class CheckBatchCost:
return None
async def _get_org_id(self, job: "_ManagedObjectRow", batch_id: str) -> str | None:
org_id = getattr(job, "org_id", None)
org_id: Final = getattr(job, "org_id", None)
if org_id:
return org_id
api_key = getattr(job, "api_key", None)
team_id = getattr(job, "team_id", None)
api_key: Final = getattr(job, "api_key", None)
team_id: Final = getattr(job, "team_id", None)
if api_key:
try:
key_row: prisma_models.LiteLLM_VerificationToken | None = await _token_table(
self.prisma_client
).find_unique(where={"token": api_key})
key_org_id = getattr(key_row, "organization_id", None) if key_row is not None else None
key_org_id: Final = (
cast(str | None, getattr(key_row, "organization_id", None))
if key_row is not None
else None
)
if key_org_id:
return key_org_id
except Exception as e:
@ -199,7 +229,7 @@ class CheckBatchCost:
team_row: prisma_models.LiteLLM_TeamTable | None = await _team_table(self.prisma_client).find_unique(
where={"team_id": team_id}
)
return getattr(team_row, "organization_id", None) if team_row is not None else None
return cast(str | None, getattr(team_row, "organization_id", None)) if team_row is not None else None
except Exception as e:
verbose_proxy_logger.error(f"CheckBatchCost: could not resolve the team's org for batch {batch_id}: {e}")
return None
@ -673,16 +703,17 @@ class CheckBatchCost:
from litellm.types.utils import LiteLLMBatch
file_object = job.file_object
if isinstance(file_object, str):
try:
file_object = json.loads(file_object)
except (json.JSONDecodeError, ValueError):
return None
if not isinstance(file_object, dict):
file_object: Final = job.file_object
try:
parsed_file_object: Final[object] = (
json.loads(file_object) if isinstance(file_object, str) else file_object
)
except (json.JSONDecodeError, ValueError):
return None
if not isinstance(parsed_file_object, dict):
return None
try:
return LiteLLMBatch.model_validate(file_object).input_file_id
return LiteLLMBatch.model_validate(parsed_file_object).input_file_id
except Exception:
return None
@ -704,10 +735,11 @@ class CheckBatchCost:
"""
from litellm.batches.batch_utils import (
count_error_file_failed_requests,
_get_file_content_as_dictionary,
get_file_content_as_dictionary,
calculate_batch_cost_and_usage,
)
from litellm.files.main import afile_content
from litellm.files.types import FileContentCallOptions
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.litellm_core_utils.litellm_logging import deployment_pricing_model_info, mask_api_base_credentials
@ -745,30 +777,22 @@ class CheckBatchCost:
except (IndexError, AttributeError):
pass
credentials = self.llm_router.get_deployment_credentials_with_provider(model_id) or {}
credentials: Final = self.llm_router.get_deployment_credentials_with_provider(model_id) or {}
_file_content = await afile_content(
file_id=raw_output_file_id,
file_id=cast(str, raw_output_file_id),
_litellm_internal_model_credentials=MappingProxyType(dict(credentials)),
**credentials,
**cast(FileContentCallOptions, credentials),
)
# Access content - handle both direct attribute and method call
if hasattr(_file_content, 'content'):
content_bytes = _file_content.content # type: ignore[union-attr]
elif hasattr(_file_content, 'read'):
content_bytes = await _file_content.read() # type: ignore[misc]
else:
content_bytes = _file_content # type: ignore[assignment]
content_bytes: Final = await _file_content_bytes(_file_content)
file_content_as_dict = _get_file_content_as_dictionary(
content_bytes # type: ignore[arg-type]
)
file_content_as_dict = get_file_content_as_dictionary(content_bytes)
# Record output file size
if prom_logger and content_bytes:
try:
prom_logger.record_managed_file_size(
size_bytes=len(content_bytes), # type: ignore
size_bytes=len(content_bytes),
purpose="batch",
file_type="output",
model=model_id,
@ -805,7 +829,7 @@ class CheckBatchCost:
team_id=getattr(job, "team_id", None),
)
for _file_attr in ["output_file_id", "error_file_id"]:
_raw_file_id = getattr(response, _file_attr, None)
_raw_file_id = cast(str | None, getattr(response, _file_attr, None))
if _raw_file_id and not _is_base64_encoded_unified_file_id(_raw_file_id):
try:
_unified_file_id = managed_files_hook.get_unified_output_file_id(

View file

@ -8,7 +8,7 @@ from . import ui_crud_endpoints # side-effect: registers extra UI settings
from .audit_logging_endpoints import router as audit_logging_router
from .liteadmin import router as liteadmin_router
from .management_endpoints import management_endpoints_router
from .utils import _should_block_robots
from .utils import should_block_robots
__all__ = ["router", "ui_crud_endpoints"]
@ -25,7 +25,7 @@ async def get_robots():
Block all web crawlers from indexing the proxy server endpoints
This is useful for ensuring that the API endpoints aren't indexed by search engines
"""
if _should_block_robots():
if should_block_robots():
return Response(content="User-agent: *\nDisallow: /", media_type="text/plain")
else:
return Response(status_code=404)

View file

@ -23,6 +23,7 @@ from uuid import NAMESPACE_URL, uuid5
import httpx
from fastapi import HTTPException
from pydantic import ValidationError
from typing_extensions import ReadOnly
import litellm
from litellm import Router, verbose_logger
@ -30,10 +31,12 @@ from litellm._internal_context import with_service_target
from litellm._uuid import uuid
from litellm.caching.caching import DualCache
from litellm.constants import MAX_FILE_LIST_LIMIT
from litellm.files.types import FileRetrieveCallOptions, FileRetrieveProvider
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_metadata,
)
from openai import AsyncOpenAI
from openai.types.file_deleted import FileDeleted
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
@ -190,6 +193,19 @@ class _ManagedObjectTableActions(Protocol):
async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
class _ManagedResourceDatabase(Protocol):
@property
def litellm_managedfiletable(self) -> _ManagedFileTableActions: ...
@property
def litellm_managedobjecttable(self) -> _ManagedObjectTableActions: ...
class _ManagedResourcePrismaClient(Protocol):
@property
def db(self) -> _ManagedResourceDatabase: ...
class _SchedulerWithJobLookup(Protocol):
def get_job(self, job_id: str) -> object: ...
@ -199,7 +215,12 @@ class _CursorPageArgs(TypedDict, total=False):
skip: int
def _managed_file_table(prisma_client: PrismaClient) -> _ManagedFileTableActions:
class _RouterFileCallKwargs(TypedDict, total=False):
client: ReadOnly[AsyncOpenAI | None]
custom_llm_provider: ReadOnly[str | None]
def _managed_file_table(prisma_client: _ManagedResourcePrismaClient) -> _ManagedFileTableActions:
return prisma_client.db.litellm_managedfiletable
@ -213,7 +234,7 @@ def _iter_provider_file_id_pairs(
yield provider_file_id, row.unified_file_id
def _managed_object_table(prisma_client: PrismaClient) -> _ManagedObjectTableActions:
def _managed_object_table(prisma_client: _ManagedResourcePrismaClient) -> _ManagedObjectTableActions:
return prisma_client.db.litellm_managedobjecttable
@ -233,7 +254,7 @@ def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, s
_MANAGED_FILES_TARGET: Final = "managed_files"
class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
class PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Class variables or attributes
def __init__(self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient):
self.internal_usage_cache = internal_usage_cache
@ -1282,7 +1303,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
target_model_names_list=target_model_names_list,
litellm_parent_otel_span=litellm_parent_otel_span,
)
response = await _PROXY_LiteLLMManagedFiles.return_unified_file_id(
response = await PROXY_LiteLLMManagedFiles.return_unified_file_id(
file_objects=responses,
create_file_request=create_file_request,
internal_usage_cache=self.internal_usage_cache,
@ -1459,13 +1480,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
_creds = _llm_router.get_deployment_credentials_with_provider(model_id) or {}
file_object = await litellm.afile_retrieve(
file_id=provider_file_id,
**_creds,
**cast(FileRetrieveCallOptions, _creds),
)
else:
file_object = await litellm.afile_retrieve(
custom_llm_provider=model_name.split("/")[0]
if model_name and "/" in model_name
else "openai", # type: ignore[arg-type]
custom_llm_provider=cast(
FileRetrieveProvider,
model_name.split("/")[0] if model_name and "/" in model_name else "openai",
),
file_id=provider_file_id,
)
verbose_logger.debug(
@ -1608,8 +1630,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
try:
model_id, model_file_id = next(iter(stored_file_object.model_mappings.items()))
credentials = llm_router.get_deployment_credentials_with_provider(model_id) or {}
response = await litellm.afile_retrieve(file_id=model_file_id, **credentials)
credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id) or {}
response = await litellm.afile_retrieve(
file_id=model_file_id,
**cast(FileRetrieveCallOptions, credentials),
)
response.id = file_id # Replace with unified ID
return response
except Exception as e:
@ -1904,7 +1929,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
else {}
),
}
await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data)
await llm_router.afile_delete(
model=model_id,
file_id=model_file_id,
**cast(_RouterFileCallKwargs, delete_data),
)
async def afile_content(
self,
@ -1940,7 +1969,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
data["_litellm_internal_model_credentials"] = cast(Dict, MappingProxyType(dict(credentials)))
else:
data.pop("_litellm_internal_model_credentials", None)
return await llm_router.afile_content(model=model_id, file_id=provider_file_id, **data) # type: ignore
return cast(
HttpxBinaryResponseContent,
await llm_router.afile_content(
model=model_id,
file_id=provider_file_id,
**cast(_RouterFileCallKwargs, data),
),
)
except Exception as e:
exception_dict[model_id] = str(e)
raise Exception(
@ -2068,3 +2104,4 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
verbose_logger.debug(
f"Converted file {file_id} from storage backend to base64 with format {content_type}"
)
_PROXY_LiteLLMManagedFiles = PROXY_LiteLLMManagedFiles

View file

@ -2,11 +2,11 @@
## This hook is used to manage vector stores with target_model_names support
## It allows creating vector stores across multiple models and managing them with unified IDs
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Dict, Final, List, Optional, Union, cast
from fastapi import HTTPException
import litellm
from litellm import Router, verbose_logger
from litellm._uuid import uuid
from litellm.integrations.custom_logger import CustomLogger
@ -38,7 +38,7 @@ else:
PrismaClient = Any
class _PROXY_LiteLLMManagedVectorStores(
class PROXY_LiteLLMManagedVectorStores(
CustomLogger, BaseManagedResource[VectorStoreCreateResponse]
):
"""
@ -89,10 +89,10 @@ class _PROXY_LiteLLMManagedVectorStores(
# Model ID is stored in hidden params if the response object supports it
# For TypedDict responses, we need to check if _hidden_params was added
hidden_params: Dict[str, Any] = {}
hidden_params: Mapping[str, object] = {}
if hasattr(resource_object, "_hidden_params"):
hidden_params = getattr(resource_object, "_hidden_params", {}) or {}
model_id = hidden_params.get("model_id", "")
hidden_params = cast(Mapping[str, object], getattr(resource_object, "_hidden_params", {}) or {})
model_id: Final = cast(str, hidden_params.get("model_id", ""))
return generate_unified_id_string(
resource_type=self.resource_type,
@ -106,7 +106,7 @@ class _PROXY_LiteLLMManagedVectorStores(
self,
llm_router: Router,
model: str,
request_data: Dict[str, Any],
request_data: dict[str, object] | VectorStoreCreateOptionalRequestParams,
litellm_parent_otel_span: Span,
) -> VectorStoreCreateResponse:
"""
@ -122,10 +122,8 @@ class _PROXY_LiteLLMManagedVectorStores(
VectorStoreCreateResponse from the provider
"""
# Use the router to create the vector store
response = await llm_router.avector_store_create(
model=model, **request_data
)
return response
response: Final = await llm_router.avector_store_create(model=model, **request_data)
return cast(VectorStoreCreateResponse, response)
# ============================================================================
# VECTOR STORE CRUD OPERATIONS
@ -464,3 +462,4 @@ class _PROXY_LiteLLMManagedVectorStores(
parent_otel_span=parent_otel_span,
resource_id_key="vector_store_id",
)
_PROXY_LiteLLMManagedVectorStores = PROXY_LiteLLMManagedVectorStores

View file

@ -9,7 +9,7 @@ custom_auth_settings: Optional[CustomAuthSettings] = None
class EnterpriseProxyConfig:
async def load_custom_auth_settings(
self, general_settings: dict
) -> CustomAuthSettings:
) -> Optional[CustomAuthSettings]:
custom_auth_settings = general_settings.get("custom_auth_settings", None)
if custom_auth_settings is not None:
custom_auth_settings = CustomAuthSettings(

View file

@ -3,7 +3,7 @@ from typing import Optional, Union
from litellm.secret_managers.main import str_to_bool
def _should_block_robots():
def should_block_robots():
"""
Returns True if the robots.txt file should block web crawlers
@ -33,3 +33,4 @@ def _should_block_robots():
)
return True
return False
_should_block_robots = should_block_robots

View file

@ -171,7 +171,7 @@ async def list_vector_stores(
try:
# Get vector stores from database (source of truth)
# Only return what's in the database to ensure consistency across instances
vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db(
vector_stores_from_db = await VectorStoreRegistry.get_vector_stores_from_db(
prisma_client=prisma_client
)

View file

@ -53,10 +53,10 @@ from litellm.types.integrations.pointfive import PointFiveInitParams
from litellm.types.integrations.zerobus import ZerobusInitParams
from litellm._logging import (
set_verbose,
_turn_on_debug,
turn_on_debug,
verbose_logger,
json_logs,
_turn_on_json,
turn_on_json,
log_level,
)
import re
@ -108,10 +108,10 @@ litellm_mode = os.getenv("LITELLM_MODE", "DEV") # "PRODUCTION", "DEV"
####################################################
if set_verbose:
_turn_on_debug()
turn_on_debug()
####################################################
### Callbacks /Logging / Success / Failure Handlers #####
CALLBACK_TYPES = Union[str, Callable, "CustomLogger"] # CustomLogger is lazy-loaded
CALLBACK_TYPES = Union[str, Callable[..., object], "CustomLogger"] # CustomLogger is lazy-loaded
input_callback: List[CALLBACK_TYPES] = []
success_callback: List[CALLBACK_TYPES] = []
failure_callback: List[CALLBACK_TYPES] = []
@ -177,9 +177,9 @@ _custom_logger_compatible_callbacks_literal = Literal[
]
cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None
logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None
_known_custom_logger_compatible_callbacks: List = list(get_args(_custom_logger_compatible_callbacks_literal))
_known_custom_logger_compatible_callbacks: List[str] = list(get_args(_custom_logger_compatible_callbacks_literal))
callbacks: List[
Union[Callable, _custom_logger_compatible_callbacks_literal, "CustomLogger"] # CustomLogger is lazy-loaded
Union[Callable[..., object], str, "CustomLogger"] # CustomLogger is lazy-loaded
] = []
callback_settings: Dict[str, Dict[str, Any]] = {}
initialized_langfuse_clients: int = 0
@ -194,13 +194,13 @@ datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged p
gcs_pub_sub_use_v1: Optional[bool] = False # if you want to use v1 gcs pubsub logged payload
generic_api_use_v1: Optional[bool] = False # if you want to use v1 generic api logged payload
argilla_transformation_object: Optional[Dict[str, Any]] = None
_async_input_callback: List[Union[str, Callable, "CustomLogger"]] = ( # CustomLogger is lazy-loaded
_async_input_callback: List[Union[str, Callable[..., object], "CustomLogger"]] = ( # CustomLogger is lazy-loaded
[]
) # internal variable - async custom callbacks are routed here.
_async_success_callback: List[Union[str, Callable, "CustomLogger"]] = ( # CustomLogger is lazy-loaded
_async_success_callback: List[Union[str, Callable[..., object], "CustomLogger"]] = ( # CustomLogger is lazy-loaded
[]
) # internal variable - async custom callbacks are routed here.
_async_failure_callback: List[Union[str, Callable, "CustomLogger"]] = ( # CustomLogger is lazy-loaded
_async_failure_callback: List[Union[str, Callable[..., object], "CustomLogger"]] = ( # CustomLogger is lazy-loaded
[]
) # internal variable - async custom callbacks are routed here.
pre_call_rules: List[Callable] = []
@ -1486,6 +1486,7 @@ from .realtime_api.main import (
arealtime_calls,
)
from .responses.main import _aresponses_websocket
from .fine_tuning.main import *
from .files.main import *
from .vector_store_files.main import (
@ -1527,6 +1528,9 @@ from . import rag
### CUSTOM LLMs ###
from .types.llms.custom_llm import CustomLLMItem
_turn_on_debug = turn_on_debug
_turn_on_json = turn_on_json
custom_provider_map: List[CustomLLMItem] = []
_custom_providers: List[str] = [] # internal helper util, used to track names of custom providers
disable_hf_tokenizer_download: Optional[bool] = (
@ -2240,7 +2244,9 @@ if TYPE_CHECKING:
register_model: Callable[..., None]
encode: Callable[..., list]
decode: Callable[..., str]
calculate_retry_after: Callable[..., float]
_calculate_retry_after: Callable[..., float]
should_retry: Callable[..., bool]
_should_retry: Callable[..., bool]
get_supported_openai_params: Callable[..., Optional[list]]
get_api_base: Callable[..., Optional[str]]
@ -2336,9 +2342,9 @@ def __getattr__(name: str) -> Any:
_async_client_cleanup_registered = True
# Use cached registry from _lazy_imports instead of importing tuples every time
from ._lazy_imports import _get_lazy_import_registry
from ._lazy_imports import get_lazy_import_registry
registry: Final = _get_lazy_import_registry()
registry: Final = get_lazy_import_registry()
# Check if name is in registry and call the cached handler function
if name in registry:

View file

@ -91,12 +91,15 @@ def _get_module_level_client_timeout(litellm_globals: Mapping[str, Any]) -> "flo
# They're separate from the main lazy import system because they have specific use cases
def _get_default_encoding() -> "Tokenizer":
def get_default_encoding() -> "Tokenizer":
from litellm.rust_bridge.tokenizer import get_encoding
return get_encoding("cl100k_base")
_get_default_encoding = get_default_encoding
# Lazy loader for get_modified_max_tokens to avoid importing token_counter at module import time
_get_modified_max_tokens_func: "Callable[..., int | None] | None" = None
@ -126,7 +129,7 @@ _token_counter_new_func: "Callable[..., int] | None" = None
_messages_reach_token_count_func: "Callable[..., bool] | None" = None
def _get_token_counter_new() -> "Callable[..., int]":
def get_token_counter_new() -> "Callable[..., int]":
"""
Lazily load and cache the token_counter function (aliased as token_counter_new).
@ -146,8 +149,11 @@ def _get_token_counter_new() -> "Callable[..., int]":
return _token_counter_new_func
def _get_messages_reach_token_count() -> "Callable[..., bool]":
"""Lazily load ``messages_reach_token_count`` for the same reason as ``_get_token_counter_new``."""
_get_token_counter_new = get_token_counter_new
def get_messages_reach_token_count() -> "Callable[..., bool]":
"""Lazily load ``messages_reach_token_count`` for the same reason as ``get_token_counter_new``."""
global _messages_reach_token_count_func
if _messages_reach_token_count_func is None:
from litellm.litellm_core_utils.token_counter import (
@ -158,6 +164,9 @@ def _get_messages_reach_token_count() -> "Callable[..., bool]":
return _messages_reach_token_count_func
_get_messages_reach_token_count = get_messages_reach_token_count
# ============================================================================
# MAIN LAZY IMPORT SYSTEM
# ============================================================================
@ -168,7 +177,7 @@ def _get_messages_reach_token_count() -> "Callable[..., bool]":
_LAZY_IMPORT_REGISTRY: dict[str, Callable[[str], object]] | None = None
def _get_lazy_import_registry() -> dict[str, Callable[[str], object]]:
def get_lazy_import_registry() -> dict[str, Callable[[str], object]]:
"""
Build the registry that maps attribute names to their handler functions.
@ -217,6 +226,9 @@ def _get_lazy_import_registry() -> dict[str, Callable[[str], object]]:
return _LAZY_IMPORT_REGISTRY
_get_lazy_import_registry = get_lazy_import_registry
class _AttributeView(TypedDict):
"""Holds one module attribute so the lazily fetched value is read back as ``object``."""

View file

@ -48,7 +48,9 @@ UTILS_NAMES: Final = (
"register_model",
"encode",
"decode",
"calculate_retry_after",
"_calculate_retry_after",
"should_retry",
"_should_retry",
"get_supported_openai_params",
"get_api_base",
@ -384,15 +386,18 @@ UTILS_MODULE_NAMES: Final = (
"_get_response_headers",
"get_llm_provider",
"_is_non_openai_azure_model",
"is_non_openai_azure_model",
"get_supported_openai_params",
"LiteLLMResponseObjectHandler",
"_handle_invalid_parallel_tool_calls",
"handle_invalid_parallel_tool_calls",
"convert_to_model_response_object",
"convert_to_streaming_response",
"convert_to_streaming_response_async",
"get_api_base",
"ResponseMetadata",
"_parse_content_for_reasoning",
"parse_content_for_reasoning",
"LiteLLMLoggingObject",
"redact_message_input_output_from_logging",
"CustomStreamWrapper",
@ -415,8 +420,10 @@ UTILS_MODULE_NAMES: Final = (
"delete_nested_value",
"is_nested_path",
"_get_base_model_from_litellm_call_metadata",
"get_base_model_from_litellm_call_metadata",
"get_litellm_params",
"_ensure_extra_body_is_safe",
"ensure_extra_body_is_safe",
"get_formatted_prompt",
"get_response_headers",
"update_response_metadata",
@ -475,8 +482,10 @@ _UTILS_IMPORT_MAP: Final = {
"register_model": (".utils", "register_model"),
"encode": (".utils", "encode"),
"decode": (".utils", "decode"),
"_calculate_retry_after": (".utils", "_calculate_retry_after"),
"_should_retry": (".utils", "_should_retry"),
"calculate_retry_after": (".utils", "calculate_retry_after"),
"_calculate_retry_after": (".utils", "calculate_retry_after"),
"should_retry": (".utils", "should_retry"),
"_should_retry": (".utils", "should_retry"),
"get_supported_openai_params": (".utils", "get_supported_openai_params"),
"get_api_base": (".utils", "get_api_base"),
"get_first_chars_messages": (".utils", "get_first_chars_messages"),
@ -1320,7 +1329,11 @@ _UTILS_MODULE_IMPORT_MAP: Final = {
),
"_is_non_openai_azure_model": (
"litellm.litellm_core_utils.get_llm_provider_logic",
"_is_non_openai_azure_model",
"is_non_openai_azure_model",
),
"is_non_openai_azure_model": (
"litellm.litellm_core_utils.get_llm_provider_logic",
"is_non_openai_azure_model",
),
"get_supported_openai_params": (
"litellm.litellm_core_utils.get_supported_openai_params",
@ -1332,7 +1345,11 @@ _UTILS_MODULE_IMPORT_MAP: Final = {
),
"_handle_invalid_parallel_tool_calls": (
"litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response",
"_handle_invalid_parallel_tool_calls",
"handle_invalid_parallel_tool_calls",
),
"handle_invalid_parallel_tool_calls": (
"litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response",
"handle_invalid_parallel_tool_calls",
),
"convert_to_model_response_object": (
"litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response",
@ -1356,7 +1373,11 @@ _UTILS_MODULE_IMPORT_MAP: Final = {
),
"_parse_content_for_reasoning": (
"litellm.litellm_core_utils.prompt_templates.common_utils",
"_parse_content_for_reasoning",
"parse_content_for_reasoning",
),
"parse_content_for_reasoning": (
"litellm.litellm_core_utils.prompt_templates.common_utils",
"parse_content_for_reasoning",
),
"LiteLLMLoggingObject": (
"litellm.litellm_core_utils.redact_messages",
@ -1429,7 +1450,11 @@ _UTILS_MODULE_IMPORT_MAP: Final = {
),
"_get_base_model_from_litellm_call_metadata": (
"litellm.litellm_core_utils.get_litellm_params",
"_get_base_model_from_litellm_call_metadata",
"get_base_model_from_litellm_call_metadata",
),
"get_base_model_from_litellm_call_metadata": (
"litellm.litellm_core_utils.get_litellm_params",
"get_base_model_from_litellm_call_metadata",
),
"get_litellm_params": (
"litellm.litellm_core_utils.get_litellm_params",
@ -1437,7 +1462,11 @@ _UTILS_MODULE_IMPORT_MAP: Final = {
),
"_ensure_extra_body_is_safe": (
"litellm.litellm_core_utils.llm_request_utils",
"_ensure_extra_body_is_safe",
"ensure_extra_body_is_safe",
),
"ensure_extra_body_is_safe": (
"litellm.litellm_core_utils.llm_request_utils",
"ensure_extra_body_is_safe",
),
"get_formatted_prompt": (
"litellm.litellm_core_utils.llm_response_utils.get_formatted_prompt",

View file

@ -23,10 +23,12 @@ from litellm.litellm_core_utils.env_utils import get_env_int
from litellm.litellm_core_utils.safe_json_dumps import UNSERIALIZABLE_OBJECT, safe_dumps, safe_json_structure
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.secret_redaction import (
_python_redact_string,
_python_redact_structured_value,
python_redact_string,
python_redact_structured_value,
redact_internal_details,
redact_string,
)
from litellm.litellm_core_utils.secret_redaction import (
redact_string as redact_secret_string,
)
from litellm.rust_bridge import diagnostics
@ -52,7 +54,7 @@ def _sanitize_correlation_id(value: str) -> str:
pass through credential redaction.
"""
stripped: Final = "".join(ch for ch in value if ch.isprintable())
return _redact_string(stripped)[:_MAX_CORRELATION_ID_LENGTH]
return redact_string(stripped)[:_MAX_CORRELATION_ID_LENGTH]
def set_session_id(session_id: str) -> "contextvars.Token[str]":
@ -71,10 +73,13 @@ if set_verbose is True:
_ENABLE_SECRET_REDACTION: Final = os.getenv("LITELLM_DISABLE_REDACT_SECRETS", "").lower() != "true"
def _redact_string(value: str) -> str:
def redact_string(value: str) -> str:
if not _ENABLE_SECRET_REDACTION:
return value
return redact_string(value)
return redact_secret_string(value)
_redact_string = redact_string
_REDACTED_RECORD_ATTR: Final = "litellm_redacted"
@ -114,7 +119,7 @@ def redact_secrets(value: str) -> str:
"""
if not _ENABLE_SECRET_REDACTION:
return value
return _redact_string(value)
return redact_string(value)
def redact_internal_details_from_client_message(value: str) -> str:
@ -165,7 +170,7 @@ _REDACTION_PLACEHOLDER: Final = "REDACTED"
def _hides_a_credential(value: str) -> bool:
"""Whether *value* only looks clean until it is percent-decoded."""
decoded: Final = unquote(value)
return _python_redact_string(decoded) != decoded
return python_redact_string(decoded) != decoded
def _drop_encoded_credential(scrubbed: str) -> str:
@ -193,10 +198,10 @@ def _scrub_access_arg(value: str) -> str:
pattern and would then be logged raw.
"""
if len(value) <= _MAX_SCRUBBED_ACCESS_ARG:
return _drop_encoded_credential(_python_redact_string(value))
return _drop_encoded_credential(python_redact_string(value))
head: Final = value[:_MAX_SCRUBBED_ACCESS_ARG]
kept: Final = head[: max(head.rfind("?"), head.rfind("&"))] if "?" in head else head
scrubbed: Final = _drop_encoded_credential(_python_redact_string(kept))
scrubbed: Final = _drop_encoded_credential(python_redact_string(kept))
return f"{scrubbed}... ({len(value) - len(kept)} more chars truncated) ..."
@ -383,14 +388,14 @@ def _python_process_diagnostic(
) -> tuple[str, str | None, str | None, tuple[str, ...], bool]:
def process_text(text: str) -> str:
collapsed: Final = _collapse_base64_runs(text, base64_limit) if base64_limit > 0 else text
scrubbed: Final = _python_redact_string(collapsed) if redact else collapsed
scrubbed: Final = python_redact_string(collapsed) if redact else collapsed
return _truncate_for_stdout_log(scrubbed, text_limit) if 0 < text_limit < len(scrubbed) else scrubbed
processed_message: Final = process_text(message)
processed_exception: Final = process_text(exception) if exception is not None else None
processed_stack: Final = _python_redact_string(stack) if redact and stack is not None else stack
processed_stack: Final = python_redact_string(stack) if redact and stack is not None else stack
processed_leaves: Final = tuple(
_python_redact_structured_value(key, text) if redact else text for key, text in leaves
python_redact_structured_value(key, text) if redact else text for key, text in leaves
)
changed: Final = (
processed_message != message
@ -472,7 +477,7 @@ def _process_record(record: logging.LogRecord, *, base64_limit: int, text_limit:
record.stack_info = processed_stack # rebind-ok: the Filter interface mutates the record
processed_values: Final = iter(processed_leaves[: len(extra_leaves)])
for key, original, prepared in extras:
replacement: Final = _sort_processed_sets(original, _replace_string_leaves(prepared, processed_values))
replacement = _sort_processed_sets(original, _replace_string_leaves(prepared, processed_values))
if not _scrubbing_changed_nothing(replacement, original):
setattr(record, key, replacement)
raw_color_changed: Final = (
@ -489,12 +494,12 @@ def _redact_json_record(value: object) -> object:
leaves: Final = tuple(_string_leaves(None, prepared))
candidate: Final = diagnostics.run(
lambda native: native.process_diagnostic("", None, None, leaves, (True, 0, 0))[3],
lambda: tuple(_python_redact_structured_value(key, text) for key, text in leaves),
lambda: tuple(python_redact_structured_value(key, text) for key, text in leaves),
)
replacements: Final = (
candidate
if len(candidate) == len(leaves)
else tuple(_python_redact_structured_value(key, text) for key, text in leaves)
else tuple(python_redact_structured_value(key, text) for key, text in leaves)
)
return _sort_processed_sets(value, _replace_string_leaves(prepared, iter(replacements)))
@ -791,7 +796,7 @@ class CorrelationPlainFormatter(logging.Formatter):
def format(self, record: logging.LogRecord) -> str:
rendered: Final = super().format(record)
formatted: Final = rendered if _is_redacted(record) else _redact_string(rendered)
formatted: Final = rendered if _is_redacted(record) else redact_string(rendered)
trace_id: Final = getattr(record, "trace_id", None)
session_id: Final = getattr(record, "session_id", None)
if not trace_id and not session_id:
@ -1050,7 +1055,7 @@ def _get_uvicorn_json_log_config():
return log_config
def _turn_on_json():
def turn_on_json() -> None:
"""
Turn on JSON logging
@ -1064,12 +1069,18 @@ def _turn_on_json():
_setup_json_exception_handlers(JsonFormatter())
def _turn_on_debug():
_turn_on_json = turn_on_json
def turn_on_debug() -> None:
verbose_logger.setLevel(level=logging.DEBUG) # set package log to debug
verbose_router_logger.setLevel(level=logging.DEBUG) # set router logs to debug
verbose_proxy_logger.setLevel(level=logging.DEBUG) # set proxy logs to debug
_turn_on_debug = turn_on_debug
def _disable_debugging():
"""Disable the package, router, and proxy verbose loggers."""
verbose_logger.disabled = True
@ -1093,8 +1104,11 @@ def print_verbose(print_statement):
pass
def _is_debugging_on() -> bool:
def is_debugging_on() -> bool:
"""
Returns True if debugging is on
"""
return verbose_logger.isEnabledFor(logging.DEBUG) or set_verbose is True
_is_debugging_on = is_debugging_on

View file

@ -26,7 +26,7 @@ from litellm._redis_credential_provider import (
AzureADCredentialProvider,
ElastiCacheIAMCredentialProvider,
GCPIAMCredentialProvider,
_generate_gcp_iam_access_token,
generate_gcp_iam_access_token,
)
from litellm.constants import (
REDIS_CLUSTER_HEALTH_CHECK_INTERVAL,
@ -37,6 +37,8 @@ from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from ._logging import verbose_logger
_generate_gcp_iam_access_token = generate_gcp_iam_access_token
AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default"
_AWS_IAM_KWARG_NAMES: Final = (
@ -341,7 +343,7 @@ def create_gcp_iam_redis_connect_func(
self._parser.on_connect(self)
auth_args: Final = (_generate_gcp_iam_access_token(service_account),)
auth_args: Final = (generate_gcp_iam_access_token(service_account),)
self.send_command("AUTH", *auth_args, check_health=False)
try:

View file

@ -38,7 +38,7 @@ class AzureCredential(Protocol):
def get_token(self, *scopes: str) -> AzureAccessToken: ...
def _generate_gcp_iam_access_token(service_account: str) -> str:
def generate_gcp_iam_access_token(service_account: str) -> str:
"""
Generate GCP IAM access token for Redis authentication.
@ -65,6 +65,9 @@ def _generate_gcp_iam_access_token(service_account: str) -> str:
return str(response.access_token)
_generate_gcp_iam_access_token = generate_gcp_iam_access_token
def _get_cached_gcp_iam_token(service_account: str) -> str:
"""
Return a cached GCP IAM token, refreshing only when expired.
@ -93,7 +96,7 @@ def _get_cached_gcp_iam_token(service_account: str) -> str:
if time.monotonic() < expiry:
return token
token = _generate_gcp_iam_access_token(service_account)
token = generate_gcp_iam_access_token(service_account)
_token_cache[service_account] = (
token,
time.monotonic() + _GCP_IAM_TOKEN_TTL_SECONDS,

View file

@ -3,7 +3,7 @@ from collections.abc import Iterable, Iterator, Mapping
from dataclasses import dataclass
from dataclasses import replace as dataclasses_replace
from enum import Enum
from typing import Any, Final, Literal
from typing import Any, Final, Literal, cast
import litellm
from litellm._logging import verbose_logger
@ -70,7 +70,7 @@ def batch_cost_is_final(batch: Batch) -> bool:
async def calculate_batch_cost_and_usage(
file_content_dictionary: list[dict],
file_content_dictionary: list[dict[str, object]],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
model_name: str | None = None,
model_info: ModelInfo | None = None,
@ -96,11 +96,11 @@ async def calculate_batch_cost_and_usage(
)
async def _handle_completed_batch(
async def handle_completed_batch(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
model_name: str | None = None,
litellm_params: dict | None = None,
litellm_params: dict[str, object] | None = None,
model_info: ModelInfo | None = None,
) -> BatchCostUsageResult:
"""Fetch a completed batch's output file and aggregate its cost, usage, and
@ -162,6 +162,9 @@ async def _handle_completed_batch(
)
_handle_completed_batch = handle_completed_batch
class _LineOutcome(Enum):
"""A batch output line that yielded no billable stats."""
@ -183,7 +186,7 @@ class _BatchOutputLineStats:
def _classify_output_line_stats(
entries: Iterable[dict],
entries: Iterable[dict[str, object]],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None,
model_info: ModelInfo | None,
@ -295,7 +298,7 @@ def _output_line_cost(
def _aggregate_batch_cost_usage_models(
entries: Iterable[dict],
entries: Iterable[dict[str, object]],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None = None,
model_info: ModelInfo | None = None,
@ -347,7 +350,7 @@ def _aggregate_batch_cost_usage_models(
def calculate_vertex_ai_batch_cost_and_usage(
vertex_ai_batch_responses: Iterable[dict],
vertex_ai_batch_responses: Iterable[dict[str, object]],
model_name: str | None = None,
model_info: ModelInfo | None = None,
) -> BatchCostUsageResult:
@ -436,7 +439,7 @@ def _provider_output_file_id(output_file_id: str) -> str:
async def _fetch_batch_managed_file_content(
file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai",
litellm_params: dict | None = None,
litellm_params: Mapping[str, object] | None = None,
) -> bytes:
"""
Fetch a batch's output or error file and return its raw JSONL bytes.
@ -448,17 +451,18 @@ async def _fetch_batch_managed_file_content(
Required for Azure and other providers that need authentication
"""
from litellm.files.main import afile_content
from litellm.files.types import FileContentCallOptions, FileContentRequestKwargs
# Build kwargs for afile_content with credentials from litellm_params
file_content_kwargs: Final = {
"file_id": _provider_output_file_id(file_id),
provider_output_file_id: Final = _provider_output_file_id(file_id)
credentials: Final = extract_file_access_credentials(litellm_params)
file_content_kwargs: Final[FileContentRequestKwargs] = {
"file_id": provider_output_file_id,
"custom_llm_provider": custom_llm_provider,
**cast( # cast-ok: preserve dynamic provider credentials without validation
FileContentCallOptions, credentials
),
}
# Extract and add credentials for file access
credentials: Final = _extract_file_access_credentials(litellm_params)
file_content_kwargs.update(credentials)
_file_content: Final = await afile_content(**file_content_kwargs)
return _file_content.content
@ -466,7 +470,7 @@ async def _fetch_batch_managed_file_content(
async def _fetch_batch_output_file_content(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai",
litellm_params: dict | None = None,
litellm_params: Mapping[str, object] | None = None,
) -> bytes:
"""
Fetch the batch output file and return its raw JSONL bytes
@ -488,7 +492,7 @@ async def _fetch_batch_output_file_content(
async def count_error_file_failed_requests(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
litellm_params: dict | None,
litellm_params: Mapping[str, object] | None,
) -> int:
"""Count failed requests reported only in the batch's separate error file.
@ -506,10 +510,12 @@ async def count_error_file_failed_requests(
except Exception as e: # noqa: BLE001 # a failed/missing error file must not abort cost tracking for the batch
verbose_logger.debug("Failed to fetch batch error file %s: %s", batch.error_file_id, e)
return 0
return sum(1 for _ in _iter_batch_input_lines(error_file_content))
return sum(1 for _ in iter_batch_input_lines(error_file_content))
def _extract_file_access_credentials(litellm_params: dict | None) -> dict:
def extract_file_access_credentials(
litellm_params: Mapping[str, object] | None,
) -> dict[str, object]:
"""
Extract credentials from litellm_params for file access operations.
@ -522,7 +528,7 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict:
Returns:
Dictionary containing only the credentials needed for file access
"""
credentials: Final = {}
credentials: Final[dict[str, object]] = {}
if litellm_params:
# List of credential keys that should be passed to file operations
@ -555,7 +561,10 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict:
return credentials
def _get_file_content_as_dictionary(file_content: bytes) -> list[dict]:
_extract_file_access_credentials = extract_file_access_credentials
def get_file_content_as_dictionary(file_content: bytes) -> list[dict[str, object]]:
"""
Get the file content as a list of dictionaries from JSON Lines format,
skipping malformed lines
@ -563,7 +572,10 @@ def _get_file_content_as_dictionary(file_content: bytes) -> list[dict]:
return list(_iter_batch_output_entries(file_content))
def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
_get_file_content_as_dictionary = get_file_content_as_dictionary
def iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
"""
Yield non-empty JSONL lines (unparsed) one at a time, so a caller can parse
each row in its own try/except and a single malformed line cannot abort the
@ -581,27 +593,30 @@ def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
yield line
def _iter_batch_output_entries(file_content: bytes) -> Iterator[dict]:
_iter_batch_input_lines = iter_batch_input_lines
def _iter_batch_output_entries(file_content: bytes) -> Iterator[dict[str, object]]:
"""
Yield parsed batch output JSONL entries one at a time without materializing
the whole file as a list, so peak memory stays bounded. A malformed or
non-object line is skipped with a warning so one bad line never aborts the
whole batch's cost accounting.
"""
for line in _iter_batch_input_lines(file_content):
for line in iter_batch_input_lines(file_content):
entry = _parse_batch_output_line(line)
if entry is not None:
yield entry
def _parse_batch_output_line(line: bytes) -> dict | None:
def _parse_batch_output_line(line: bytes) -> dict[str, object] | None:
try:
parsed: Final[object] = json.loads(line)
except ValueError as e:
verbose_logger.warning("skipping malformed batch output line: %s", str(e))
return None
if isinstance(parsed, dict):
return parsed
return cast("dict[str, object]", parsed) # cast-ok: JSON object keys are strings by definition
verbose_logger.warning("skipping non-object batch output line of type %s", type(parsed).__name__)
return None
@ -611,24 +626,29 @@ def _parse_batch_output_line(line: bytes) -> dict | None:
_BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN: Final = 4
def _estimate_batch_entry_tokens(raw_line: bytes) -> int:
def estimate_batch_entry_tokens(raw_line: bytes) -> int:
"""Conservative token estimate for a batch row the token counter cannot measure
(or that cannot be parsed). Keeps the batch token total non-zero so a crafted
row cannot evade the TPM limit, without hard-rejecting a legitimate batch."""
return max(1, len(raw_line) // _BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN)
def _count_entry_tokens(
entry: dict,
_estimate_batch_entry_tokens = estimate_batch_entry_tokens
def count_entry_tokens(
entry: Mapping[str, object],
model_name: str | None = None,
) -> int:
"""Token-count a single batch input entry's body (chat / text / embedding)."""
body: Final = entry.get("body", {}) or {}
model: Final = body.get("model", model_name or "")
body: Final = cast( # cast-ok: batch payload bodies come from provider JSON
Mapping[str, object], entry.get("body", {}) or {}
)
model: Final = cast(str, body.get("model", model_name or "")) # cast-ok: provider batch model names are strings
messages: Final = body.get("messages")
if messages:
return token_counter(model=model, messages=messages)
return token_counter(model=model, messages=cast(list[dict[str, object]], messages))
prompt: Final = body.get("prompt")
if prompt:
@ -641,6 +661,9 @@ def _count_entry_tokens(
return 0
_count_entry_tokens = count_entry_tokens
def _count_prompt_or_input_tokens(model: str, value: object) -> int:
"""Token-count a ``prompt`` / ``input`` field that the OpenAI batch
schema allows in four shapes:
@ -707,8 +730,8 @@ def _get_batch_job_usage_from_response_body(
from litellm.responses.utils import ResponseAPILoggingUtils
_usage_dict: Final = response_body.get("usage", None) or {}
if ResponseAPILoggingUtils._is_response_api_usage(_usage_dict):
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(_usage_dict)
if ResponseAPILoggingUtils.is_response_api_usage(_usage_dict):
return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(_usage_dict)
usage: Final[Usage] = Usage(**_usage_dict)
if custom_llm_provider == "xai":
from litellm.llms.xai.chat.transformation import XAIChatConfig
@ -737,9 +760,11 @@ def _get_response_from_batch_job_output_file(
return _get_anthropic_result_from_batch_results_line(batch_job_output_file).get("message", None) or {}
if custom_llm_provider == "bedrock":
return batch_job_output_file.get("modelOutput", None) or {}
_response: Final[dict] = batch_job_output_file.get("response", None) or {}
_response: Final = cast( # cast-ok: batch output response comes from provider JSON
Mapping[str, object], batch_job_output_file.get("response", None) or {}
)
_response_body: Final = _response.get("body", None) or {}
return _response_body
return cast(Mapping[str, object], _response_body) # cast-ok: batch response body comes from provider JSON
def _batch_response_was_successful(
@ -756,5 +781,7 @@ def _batch_response_was_successful(
return _get_anthropic_result_from_batch_results_line(batch_job_output_file).get("type") == "succeeded"
if custom_llm_provider == "bedrock":
return batch_job_output_file.get("modelOutput") is not None and batch_job_output_file.get("error") is None
_response: Final[dict] = batch_job_output_file.get("response", None) or {}
_response: Final[dict[str, object]] = cast( # cast-ok: batch response bodies come from provider JSON
dict[str, object], batch_job_output_file.get("response", None) or {}
)
return _response.get("status_code", None) == 200

View file

@ -16,7 +16,7 @@ import traceback
from collections.abc import Generator, Mapping
from contextlib import contextmanager
from enum import Enum
from typing import Any, Final, Literal
from typing import Any, Final, Literal, cast
from pydantic import BaseModel
@ -363,12 +363,12 @@ class Cache:
cache_key = ""
# verbose_logger.debug("\nGetting Cache key. Kwargs: %s", kwargs)
preset_cache_key: Final = self._get_preset_cache_key_from_kwargs(**kwargs)
preset_cache_key: Final = self.get_preset_cache_key_from_kwargs(**kwargs)
if preset_cache_key is not None:
verbose_logger.debug("\nReturning preset cache key: %s", preset_cache_key)
return preset_cache_key
combined_kwargs: Final = ModelParamHelper._get_all_llm_api_params()
combined_kwargs: Final = ModelParamHelper.get_all_llm_api_params()
is_semantic_cache: Final = self._is_semantic_cache()
scope_excluded_params: Final = self._SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS if is_semantic_cache else frozenset()
for param in kwargs:
@ -460,7 +460,7 @@ class Cache:
or litellm_params.get("file_name")
)
def _get_preset_cache_key_from_kwargs(self, **kwargs) -> str | None:
def get_preset_cache_key_from_kwargs(self, **kwargs: object) -> str | None:
"""
Get the preset cache key from kwargs["litellm_params"]
@ -469,11 +469,17 @@ class Cache:
1. optional params like max_tokens, get transformed for bedrock -> max_new_tokens
2. avoid doing duplicate / repeated work
"""
if kwargs:
if "litellm_params" in kwargs:
return kwargs["litellm_params"].get("preset_cache_key", None)
if "litellm_params" in kwargs:
litellm_params: Final = cast( # cast-ok: cache kwargs retain dynamic caller values
Mapping[str, object], kwargs["litellm_params"]
)
return cast( # cast-ok: preserve dynamically supplied cache keys
str | None, litellm_params.get("preset_cache_key", None)
)
return None
_get_preset_cache_key_from_kwargs = get_preset_cache_key_from_kwargs
def _set_preset_cache_key_in_kwargs(self, preset_cache_key: str, **kwargs) -> None:
"""
Set the calculated cache key in kwargs
@ -539,11 +545,11 @@ class Cache:
}
time.sleep(CACHED_STREAMING_CHUNK_DELAY)
def _get_cache_logic(
def get_cache_logic(
self,
cached_result: object | None,
max_age: float | None,
):
) -> object | None:
"""
Common get cache logic across sync + async implementations
"""
@ -572,6 +578,8 @@ class Cache:
return cached_response
return cached_result
_get_cache_logic = get_cache_logic
@staticmethod
def _get_safe_cache_lookup_kwargs(kwargs: Mapping[str, object]) -> dict[str, object]:
cache_lookup_kwargs: Final[dict[str, object]] = {}
@ -628,7 +636,7 @@ class Cache:
original_kwargs=kwargs,
cache_lookup_kwargs=cache_lookup_kwargs,
)
return self._get_cache_logic(cached_result=cached_result, max_age=max_age)
return self.get_cache_logic(cached_result=cached_result, max_age=max_age)
except Exception:
print_verbose(f"An exception occurred: {traceback.format_exc()}")
return None
@ -658,7 +666,7 @@ class Cache:
cached_result = await dynamic_cache_object.async_get_cache(cache_key, **kwargs)
else:
cached_result = await self.cache.async_get_cache(cache_key, **kwargs)
return self._get_cache_logic(cached_result=cached_result, max_age=max_age)
return self.get_cache_logic(cached_result=cached_result, max_age=max_age)
except Exception:
print_verbose(f"An exception occurred: {traceback.format_exc()}")
return None
@ -957,7 +965,7 @@ class Cache:
if hasattr(self.cache, "disconnect"):
await self.cache.disconnect()
def _supports_async(self) -> bool:
def supports_async(self) -> bool:
"""
Internal method to check if the cache type supports async get/set operations
@ -966,6 +974,8 @@ class Cache:
"""
return True
_supports_async = supports_async
def enable_cache(
type: LiteLLMCacheType | None = LiteLLMCacheType.LOCAL,

View file

@ -33,7 +33,7 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
update_response_metadata,
)
from litellm.litellm_core_utils.logging_utils import (
_assemble_complete_response_from_streaming_chunks,
assemble_complete_response_from_streaming_chunks,
)
from litellm.types.caching import CACHED_STREAM_EVENTS_KEY, EMBEDDING_CACHE_FORMAT_VERSION, CachedEmbedding
from litellm.types.integrations.custom_logger import converted_stream_requested
@ -66,7 +66,7 @@ _StreamResultT = TypeVar("_StreamResultT")
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
get_parent_otel_span_from_kwargs,
)
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
@ -246,14 +246,14 @@ def _current_format_embedding_entry(entry: object) -> CachedEmbedding | None:
class LLMCachingHandler:
def __init__(
self,
original_function: Callable,
original_function: Callable[..., object],
request_kwargs: dict[str, object],
start_time: datetime.datetime,
):
from litellm.caching import DualCache, RedisCache
self.async_streaming_chunks: list[ModelResponse] = []
self.sync_streaming_chunks: list[ModelResponse] = []
self.async_streaming_chunks: list[object] = []
self.sync_streaming_chunks: list[object] = []
self.request_kwargs = _drop_logging_obj_from_kwargs(request_kwargs)
self.preset_cache_key: str | None = None
self.original_function = original_function
@ -266,10 +266,10 @@ class LLMCachingHandler:
else:
self.dual_cache = None
async def _async_get_cache(
async def async_get_cache(
self,
model: str,
original_function: Callable,
original_function: Callable[..., object],
logging_obj: LiteLLMLoggingObj,
start_time: datetime.datetime,
call_type: str,
@ -313,7 +313,7 @@ class LLMCachingHandler:
cache_check_start_time: Final = time.perf_counter()
cache_check_end_time: float | None = None
#########################################################
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs)
parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs)
kwargs["parent_otel_span"] = parent_otel_span
if litellm.cache is not None and self._is_call_type_supported_by_cache(original_function=original_function):
@ -406,10 +406,12 @@ class LLMCachingHandler:
# Caching disabled - return None to indicate no caching attempted
return None
def _sync_get_cache(
_async_get_cache = async_get_cache
def sync_get_cache(
self,
model: str,
original_function: Callable,
original_function: Callable[..., object],
logging_obj: LiteLLMLoggingObj,
start_time: datetime.datetime,
call_type: str,
@ -496,6 +498,8 @@ class LLMCachingHandler:
return CachingHandlerResponse(cached_result=cached_result)
return CachingHandlerResponse(cached_result=cached_result)
_sync_get_cache = sync_get_cache
def handle_kwargs_input_list_or_str(self, kwargs: dict[str, object]) -> list[str]:
"""
Handles the input of kwargs['input'] being a list or a string
@ -715,7 +719,7 @@ class LLMCachingHandler:
except Exception:
return None
def _combine_cached_embedding_response_with_api_result(
def combine_cached_embedding_response_with_api_result(
self,
_caching_handler_response: CachingHandlerResponse,
embedding_response: EmbeddingResponse,
@ -763,6 +767,8 @@ class LLMCachingHandler:
merged._response_ms = (end_time - start_time).total_seconds() * 1000
return merged
_combine_cached_embedding_response_with_api_result = combine_cached_embedding_response_with_api_result
def _async_log_cache_hit_on_callbacks(
self,
logging_obj: LiteLLMLoggingObj,
@ -861,7 +867,7 @@ class LLMCachingHandler:
request_kwargs: Final = new_kwargs.copy()
request_cache_key: Final = _request_cache_key(request_kwargs)
request_kwargs.pop("cache_key", None)
if litellm.cache._supports_async() is True:
if litellm.cache.supports_async() is True:
## check if dual cache is supported ##
self.preset_cache_key = request_cache_key or litellm.cache.get_cache_key(**request_kwargs)
with response_cache_phase("get"):
@ -1111,7 +1117,7 @@ class LLMCachingHandler:
None
"""
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
get_parent_otel_span_from_kwargs,
)
if litellm.cache is None:
@ -1128,10 +1134,10 @@ class LLMCachingHandler:
args,
)
)
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(new_kwargs)
parent_otel_span: Final = get_parent_otel_span_from_kwargs(new_kwargs)
new_kwargs["parent_otel_span"] = parent_otel_span
# [OPTIONAL] ADD TO CACHE
if self._should_store_result_in_cache(original_function=original_function, kwargs=new_kwargs):
if self.should_store_result_in_cache(original_function=original_function, kwargs=new_kwargs):
if (
isinstance(result, litellm.ModelResponse)
or isinstance(result, litellm.EmbeddingResponse)
@ -1183,13 +1189,13 @@ class LLMCachingHandler:
)
)
if self._should_store_result_in_cache(original_function=self.original_function, kwargs=new_kwargs):
if self.should_store_result_in_cache(original_function=self.original_function, kwargs=new_kwargs):
with response_cache_phase("set"):
litellm.cache.add_cache(result, **new_kwargs)
return
def _should_store_result_in_cache(self, original_function: Callable, kwargs: dict[str, Any]) -> bool:
def should_store_result_in_cache(self, original_function: Callable[..., object], kwargs: dict[str, Any]) -> bool:
"""
Helper function to determine if the result should be stored in the cache.
@ -1200,6 +1206,8 @@ class LLMCachingHandler:
kwargs.get("cache", {}).get("no-store", False) is not True
)
_should_store_result_in_cache = should_store_result_in_cache
def wrap_streaming_result_for_cache(
self, result: _StreamResultT, call_type: str
) -> "_StreamResultT | AnthropicMessagesStreamCacheWriter":
@ -1208,7 +1216,7 @@ class LLMCachingHandler:
CallTypes.aanthropic_messages.value,
):
return result
if litellm.cache is None or not self._should_store_result_in_cache(
if litellm.cache is None or not self.should_store_result_in_cache(
original_function=self.original_function, kwargs=self.request_kwargs
):
return result
@ -1240,7 +1248,7 @@ class LLMCachingHandler:
covering_call_types: Final = ("aresponses", "responses") if call_type == "aresponses" else (call_type,)
return any(name in litellm.cache.supported_call_types for name in covering_call_types)
async def _add_streaming_response_to_cache(self, processed_chunk: ModelResponse):
async def add_streaming_response_to_cache(self, processed_chunk: ModelResponse) -> None:
"""
Internal method to add the streaming response to the cache
@ -1251,7 +1259,7 @@ class LLMCachingHandler:
"""
complete_streaming_response: Final[ModelResponse | TextCompletionResponse | None] = (
_assemble_complete_response_from_streaming_chunks(
assemble_complete_response_from_streaming_chunks(
result=processed_chunk,
start_time=self.start_time,
end_time=datetime.datetime.now(),
@ -1268,12 +1276,14 @@ class LLMCachingHandler:
kwargs=self.request_kwargs,
)
def _sync_add_streaming_response_to_cache(self, processed_chunk: ModelResponse):
_add_streaming_response_to_cache = add_streaming_response_to_cache
def sync_add_streaming_response_to_cache(self, processed_chunk: ModelResponse) -> None:
"""
Sync internal method to add the streaming response to the cache
"""
complete_streaming_response: Final[ModelResponse | TextCompletionResponse | None] = (
_assemble_complete_response_from_streaming_chunks(
assemble_complete_response_from_streaming_chunks(
result=processed_chunk,
start_time=self.start_time,
end_time=datetime.datetime.now(),
@ -1290,6 +1300,8 @@ class LLMCachingHandler:
kwargs=self.request_kwargs,
)
_sync_add_streaming_response_to_cache = sync_add_streaming_response_to_cache
def _update_litellm_logging_obj_environment(
self,
logging_obj: LiteLLMLoggingObj,
@ -1329,7 +1341,7 @@ class LLMCachingHandler:
if litellm.cache is not None:
litellm_params["preset_cache_key"] = (
self.preset_cache_key or litellm.cache._get_preset_cache_key_from_kwargs(**kwargs)
self.preset_cache_key or litellm.cache.get_preset_cache_key_from_kwargs(**kwargs)
)
else:
litellm_params["preset_cache_key"] = None

View file

@ -38,7 +38,7 @@ from litellm.constants import (
REDIS_CIRCUIT_BREAKER_TIMEOUT_MIN_DURATION,
REDIS_TIMEOUT_LOG_INTERVAL,
)
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.litellm_core_utils.core_helpers import get_parent_otel_span_from_kwargs
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
from litellm.types.caching import (
RedisPipelineIncrementOperation,
@ -1157,7 +1157,7 @@ class RedisCache(BaseCache):
error=e,
start_time=start_time,
end_time=end_time,
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
parent_otel_span=get_parent_otel_span_from_kwargs(kwargs),
call_type="async_set_cache",
caller=_get_call_stack_info(),
)
@ -1192,7 +1192,7 @@ class RedisCache(BaseCache):
caller=_get_call_stack_info(),
start_time=start_time,
end_time=end_time,
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
parent_otel_span=get_parent_otel_span_from_kwargs(kwargs),
)
)
return result
@ -1208,7 +1208,7 @@ class RedisCache(BaseCache):
caller=_get_call_stack_info(),
start_time=start_time,
end_time=end_time,
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
parent_otel_span=get_parent_otel_span_from_kwargs(kwargs),
)
)
log_redis_failure(
@ -1276,7 +1276,7 @@ class RedisCache(BaseCache):
caller=_get_call_stack_info(),
start_time=start_time,
end_time=end_time,
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
parent_otel_span=get_parent_otel_span_from_kwargs(kwargs),
)
)
return
@ -1293,7 +1293,7 @@ class RedisCache(BaseCache):
caller=_get_call_stack_info(),
start_time=start_time,
end_time=end_time,
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
parent_otel_span=get_parent_otel_span_from_kwargs(kwargs),
)
)
@ -1384,7 +1384,7 @@ class RedisCache(BaseCache):
error=e,
start_time=start_time,
end_time=end_time,
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
parent_otel_span=get_parent_otel_span_from_kwargs(kwargs),
call_type="async_set_cache_sadd",
caller=_get_call_stack_info(),
)
@ -1410,7 +1410,7 @@ class RedisCache(BaseCache):
caller=_get_call_stack_info(),
start_time=start_time,
end_time=end_time,
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
parent_otel_span=get_parent_otel_span_from_kwargs(kwargs),
)
)
except Exception as e:
@ -1425,7 +1425,7 @@ class RedisCache(BaseCache):
caller=_get_call_stack_info(),
start_time=start_time,
end_time=end_time,
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
parent_otel_span=get_parent_otel_span_from_kwargs(kwargs),
)
)
# NON blocking - notify users Redis is throwing an exception
@ -2036,7 +2036,7 @@ class RedisCache(BaseCache):
caller=_get_call_stack_info(),
start_time=start_time,
end_time=end_time,
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
parent_otel_span=get_parent_otel_span_from_kwargs(kwargs),
)
)
return results
@ -2053,7 +2053,7 @@ class RedisCache(BaseCache):
caller=_get_call_stack_info(),
start_time=start_time,
end_time=end_time,
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
parent_otel_span=get_parent_otel_span_from_kwargs(kwargs),
)
)
log_redis_failure(

View file

@ -244,7 +244,7 @@ def tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _Cha
name: Final = item.get("name") or ("custom_tool" if is_custom else "")
function_chunk: Final = ChatCompletionToolCallFunctionChunk(name=name, arguments=arguments)
tool_call_dict: Final = _ChatToolCallDict(
id=LiteLLMCompletionResponsesConfig._tool_call_id_from_responses_item(item.get("id"), item.get("call_id")),
id=LiteLLMCompletionResponsesConfig.tool_call_id_from_responses_item(item.get("id"), item.get("call_id")),
type="function",
function=function_chunk,
index=index,
@ -525,7 +525,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
elif key == "previous_response_id":
responses_api_request["previous_response_id"] = value
elif key == "reasoning_effort":
responses_api_request["reasoning"] = self._map_reasoning_effort(value)
responses_api_request["reasoning"] = self.map_reasoning_effort(value)
elif key == "web_search_options":
self._add_web_search_tool(responses_api_request, value)
@ -983,7 +983,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
setattr(
model_response,
"usage",
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_response.usage),
ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(raw_response.usage),
)
model_response.id = _upstream_response_id(raw_response.id) or raw_response.id
@ -1206,7 +1206,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return optional_params
def _map_reasoning_effort(self, reasoning_effort: object) -> Reasoning:
def map_reasoning_effort(self, reasoning_effort: object) -> Reasoning:
# If dict is passed, convert it directly to Reasoning object
if isinstance(reasoning_effort, dict):
return Reasoning(
@ -1225,6 +1225,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
else Reasoning(effort=reasoning_effort)
)
_map_reasoning_effort = map_reasoning_effort
def _add_web_search_tool(
self,
responses_api_request: ResponsesAPIOptionalRequestParams,
@ -1626,7 +1628,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
if response_data.get("usage"):
from litellm.responses.utils import ResponseAPILoggingUtils
usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(response_data.get("usage"))
usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(response_data.get("usage"))
provider_metadata: Final = _provider_metadata(response_data)
served_service_tier: Final = response_data.get("service_tier")
return ModelResponseStream(

View file

@ -1077,7 +1077,7 @@ openai_text_completion_compatible_providers: Final[list] = [ # providers that s
"hyperbolic",
"wandb",
]
_openai_like_providers: Final[list] = [
_openai_like_providers: Final[list[str]] = [
"predibase",
"databricks",
"lemonade",

View file

@ -19,7 +19,7 @@ def decode_managed_container_id_for_request(
Returns:
(original_container_id, resolved_provider, updated_litellm_params)
"""
decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id)
decoded: Final = ResponsesAPIRequestUtils.decode_container_id(container_id)
original_container_id: Final = decoded.get("response_id", container_id)
decoded_provider: Final = decoded.get("custom_llm_provider")
@ -148,7 +148,7 @@ class ContainerRequestUtils:
# Only encode if we have routing metadata
if should_encode and response_obj and hasattr(response_obj, "id"):
encoded_id: Final = ResponsesAPIRequestUtils._build_container_id(
encoded_id: Final = ResponsesAPIRequestUtils.build_container_id(
custom_llm_provider=custom_llm_provider,
model_id=model_id,
container_id=response_obj.id,

View file

@ -29,13 +29,13 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import (
_SERVICE_TIER_TO_COST_KEY_SUFFIX,
BilledTokenRates,
CostCalculatorUtils,
_generic_cost_per_character,
_get_regional_uplift_multiplier,
_get_service_tier_cost_key,
calculate_cost_component,
generic_cost_per_character,
generic_cost_per_token,
get_batch_cost_rates,
get_billable_input_tokens,
get_regional_uplift_multiplier,
get_service_tier_cost_key,
get_token_type_cost_breakdown,
parse_prompt_tokens_details,
select_cost_metric_for_model,
@ -137,7 +137,7 @@ from litellm.utils import (
ProviderConfigManager,
TextCompletionResponse,
TranscriptionResponse,
_cached_get_model_info_helper,
cached_get_model_info_helper,
token_counter,
)
@ -349,7 +349,7 @@ def _per_second_pricing_cost(
audio_seconds: float = 0.0,
) -> tuple[float, float] | None:
try:
model_info: Final = _cached_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider)
model_info: Final = cached_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider)
except Exception: # noqa: BLE001 # the lookup raises plain Exception for an unmapped model
return None
if _has_token_or_tiered_pricing(model_info) or not _bills_wall_clock_seconds(model_info):
@ -587,7 +587,7 @@ def cost_per_token(
raise ValueError(
f"prompt_characters must be provided for tts calls. prompt_characters={prompt_characters}, model={model}, custom_llm_provider={custom_llm_provider}, call_type={call_type}"
)
_prompt_cost, _completion_cost = _generic_cost_per_character(
_prompt_cost, _completion_cost = generic_cost_per_character(
model=model_without_prefix,
custom_llm_provider=custom_llm_provider,
prompt_characters=prompt_characters,
@ -749,7 +749,7 @@ def cost_per_token(
service_tier=service_tier,
)
else:
model_info: Final = _cached_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider)
model_info: Final = cached_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider)
if _has_token_or_tiered_pricing(model_info):
return generic_cost_per_token(
model=model,
@ -818,7 +818,7 @@ def _cost_map_entry_prices_anything(entry: Mapping[str, object]) -> bool:
)
def _select_model_name_for_cost_calc(
def select_model_name_for_cost_calc(
model: str | None,
completion_response: object | None,
base_model: str | None = None,
@ -889,6 +889,9 @@ def _select_model_name_for_cost_calc(
return return_model
_select_model_name_for_cost_calc = select_model_name_for_cost_calc
def _strip_unregistered_leading_segments(model: str, region_name: str | None) -> str:
"""Resolve a provider-prefixed slash alias like "vertex_ai/vertex/claude-opus-5" to the
registered cost key ("vertex_ai/claude-opus-5"), keeping the model unchanged when it already
@ -1039,9 +1042,9 @@ def get_usage_object(
elif (
usage_obj is not None
and (isinstance(usage_obj, dict) or isinstance(usage_obj, ResponseAPIUsage))
and ResponseAPILoggingUtils._is_response_api_usage(usage_obj)
and ResponseAPILoggingUtils.is_response_api_usage(usage_obj)
):
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage_obj)
return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage_obj)
elif TranscriptionUsageObjectTransformation.is_transcription_usage_object(usage_obj):
return TranscriptionUsageObjectTransformation.transform_transcription_usage_object(
cast(
@ -1070,7 +1073,7 @@ def _is_known_usage_objects(usage_obj):
)
def _infer_call_type(call_type: CallTypesLiteral | None, completion_response: object) -> CallTypesLiteral | None:
def infer_call_type(call_type: CallTypesLiteral | None, completion_response: object) -> CallTypesLiteral | None:
if call_type is not None:
return call_type
@ -1097,6 +1100,9 @@ def _infer_call_type(call_type: CallTypesLiteral | None, completion_response: ob
return call_type
_infer_call_type = infer_call_type
def _apply_cost_discount(
base_cost: float,
custom_llm_provider: str | None,
@ -1369,7 +1375,7 @@ def completion_cost(
- For un-mapped Replicate models, the cost is calculated based on the total time used for the request.
"""
try:
call_type = _infer_call_type(call_type, completion_response) or "completion"
call_type = infer_call_type(call_type, completion_response) or "completion"
if call_type == CallTypes.aresponses_websocket.value and isinstance(
completion_response, LiteLLMRealtimeStreamLoggingObject
@ -1438,7 +1444,7 @@ def completion_cost(
)
explicit_pricing: Final = custom_pricing is True or base_model is not None
selected_model: Final = _select_model_name_for_cost_calc(
selected_model: Final = select_model_name_for_cost_calc(
model=model,
completion_response=completion_response,
custom_llm_provider=custom_llm_provider,
@ -1492,10 +1498,8 @@ def completion_cost(
.calculate_usage(usage_object=_usage, reasoning_content=None)
.model_dump()
)
elif ResponseAPILoggingUtils._is_response_api_usage(_usage):
_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
_usage
).model_dump()
elif ResponseAPILoggingUtils.is_response_api_usage(_usage):
_usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(_usage).model_dump()
elif TranscriptionUsageObjectTransformation.is_transcription_usage_object(_usage):
tr_usage = TranscriptionUsageObjectTransformation.transform_transcription_usage_object(
cast(
@ -1561,7 +1565,7 @@ def completion_cost(
"litellm.cost_calculator.py::completion_cost() - Error inferring custom_llm_provider - %s",
e,
)
if CostCalculatorUtils._call_type_has_image_response(call_type) and isinstance(
if CostCalculatorUtils.call_type_has_image_response(call_type) and isinstance(
completion_response, ImageResponse
):
### IMAGE GENERATION COST CALCULATION ###
@ -1631,7 +1635,7 @@ def completion_cost(
video_resolution=video_resolution,
)
elif call_type in _SPEECH_CALL_TYPES:
prompt_characters = litellm.utils._count_characters(text=prompt)
prompt_characters = litellm.utils.count_characters(text=prompt)
elif call_type in _TRANSCRIPTION_CALL_TYPES:
# Check _hidden_params first (duration stored there to
# avoid polluting the response body), then fall back to
@ -1780,10 +1784,10 @@ def completion_cost(
data={"messages": messages}, call_type="completion"
)
prompt_characters = litellm.utils._count_characters(text=prompt_string)
prompt_characters = litellm.utils.count_characters(text=prompt_string)
if completion_response is not None and isinstance(completion_response, ModelResponse):
completion_string = litellm.utils.get_response_string(response_obj=completion_response)
completion_characters = litellm.utils._count_characters(text=completion_string)
completion_characters = litellm.utils.count_characters(text=completion_string)
# Get the original request model for router detection
request_model_for_cost = None
@ -1946,7 +1950,7 @@ def get_response_cost_from_hidden_params(
hidden_params: dict | BaseModel,
) -> float | None:
if isinstance(hidden_params, BaseModel):
_hidden_params_dict = cast(BaseModel, hidden_params).model_dump()
_hidden_params_dict = hidden_params.model_dump()
else:
_hidden_params_dict = hidden_params
@ -2132,7 +2136,7 @@ def pricing_entry_for_cost_calc(
deployment_key: Final = router_model_id or model
if deployment_entry is not None and deployment_key is not None:
return deployment_key, deployment_entry
selected_model: Final = _select_model_name_for_cost_calc(
selected_model: Final = select_model_name_for_cost_calc(
model=model,
completion_response=completion_response,
base_model=base_model,
@ -2657,7 +2661,7 @@ def batch_cost_calculator(
) # batch cost is usually half of the regular token cost
# Add cache read cost if applicable
cache_read_cost_key: Final = _get_service_tier_cost_key("cache_read_input_token_cost", None)
cache_read_cost_key: Final = 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
cache_creation_cost: Final = model_info.get("cache_creation_input_token_cost") or input_cost_per_token
@ -2672,7 +2676,7 @@ def batch_cost_calculator(
text_tokens: Final = usage.completion_tokens - image_tokens
total_completion_cost = text_tokens * text_rate + image_tokens * image_rate
uplift: Final = _get_regional_uplift_multiplier(model_info, data_residency)
uplift: Final = get_regional_uplift_multiplier(model_info, data_residency)
if uplift != 1.0:
total_prompt_cost *= uplift
total_completion_cost *= uplift
@ -2827,7 +2831,7 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor):
)
usage_objects: Final[list[Usage]] = []
for result in response_done_events:
usage_object = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage_object = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(
result["response"].get("usage", {})
)
usage_objects.append(usage_object)
@ -2885,9 +2889,7 @@ class ResponsesWebSocketTokenUsageProcessor(BaseTokenUsageProcessor):
results: Sequence[Mapping[str, object]],
) -> tuple[Usage, ...]:
return tuple(
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( # pyright: ignore[reportPrivateUsage] # same shared transform the realtime processor uses
response.usage
)
ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(response.usage)
for _, response in _billable_responses_ws_events(results)
if response.usage is not None
)

View file

@ -30,9 +30,6 @@ FileCreateProvider = Literal[
"mistral",
"xai",
]
FileRetrieveProvider = Literal[
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic", "mistral", "xai"
]
FileDeleteProvider = Literal[
"openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic", "mistral", "xai"
]
@ -40,7 +37,7 @@ FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthrop
import litellm
from litellm import get_secret_str
from litellm.files.streaming import FileContentStreamingResponse
from litellm.files.types import FileContentProvider, FileContentStreamingResult
from litellm.files.types import FileContentProvider, FileContentStreamingResult, FileRetrieveProvider
from litellm.litellm_core_utils.get_litellm_params import add_trusted_model_credentials_to_litellm_params
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj

View file

@ -1,9 +1,37 @@
from collections.abc import AsyncIterator, Iterator, Mapping
from typing import Literal, NamedTuple
from typing import Literal, NamedTuple, TypedDict
from typing_extensions import NotRequired, ReadOnly
FileContentProvider = Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "manus", "mistral"
]
FileRetrieveProvider = Literal[
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic", "mistral", "xai"
]
class FileContentCallOptions(TypedDict, total=False):
custom_llm_provider: ReadOnly[FileContentProvider]
extra_body: ReadOnly[NotRequired[dict[str, str] | None]]
extra_headers: ReadOnly[NotRequired[dict[str, str] | None]]
chunk_size: ReadOnly[NotRequired[int]]
stream: ReadOnly[NotRequired[bool]]
class FileContentRequestKwargs(TypedDict):
file_id: ReadOnly[str]
custom_llm_provider: ReadOnly[FileContentProvider]
extra_body: ReadOnly[NotRequired[dict[str, str] | None]]
extra_headers: ReadOnly[NotRequired[dict[str, str] | None]]
chunk_size: ReadOnly[NotRequired[int]]
stream: ReadOnly[NotRequired[bool]]
class FileRetrieveCallOptions(TypedDict, total=False):
custom_llm_provider: ReadOnly[FileRetrieveProvider]
extra_body: ReadOnly[NotRequired[dict[str, str] | None]]
extra_headers: ReadOnly[NotRequired[dict[str, str] | None]]
class FileContentStreamingResult(NamedTuple):

View file

@ -31,7 +31,7 @@ from litellm.integrations.SlackAlerting.hanging_request_check import (
)
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.litellm_core_utils.exception_mapping_utils import (
_add_key_name_and_team_to_alert,
add_key_name_and_team_to_alert,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
@ -293,7 +293,7 @@ class SlackAlerting(CustomBatchLogger):
# add deployment latencies to alert
if kwargs is not None and "litellm_params" in kwargs and "metadata" in kwargs["litellm_params"]:
_metadata: Final[dict] = kwargs["litellm_params"]["metadata"]
request_info = _add_key_name_and_team_to_alert(request_info=request_info, metadata=_metadata)
request_info = add_key_name_and_team_to_alert(request_info=request_info, metadata=_metadata)
_deployment_latency_map: Final = self._get_deployment_latencies_to_alert(metadata=_metadata)
if _deployment_latency_map is not None:

View file

@ -4,7 +4,7 @@ Utils used for slack alerting
import asyncio
from collections.abc import Callable
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Any, Final, cast
import litellm
from litellm.integrations.custom_logger import CustomLogger
@ -57,8 +57,8 @@ def process_slack_alerting_variables(
return alert_to_webhook_url
async def _add_langfuse_trace_id_to_alert(
request_data: dict | None = None,
async def add_langfuse_trace_id_to_alert(
request_data: dict[str, object] | None = None,
) -> str | None:
"""
Returns langfuse trace url
@ -71,7 +71,7 @@ async def _add_langfuse_trace_id_to_alert(
from litellm.integrations.langfuse.langfuse import LangFuseLogger, resolve_langfuse_host
callbacks: Final[list[CustomLogger | Callable[..., object] | str]] = (
litellm.logging_callback_manager._get_all_callbacks()
litellm.logging_callback_manager.get_all_callbacks()
)
if not any(callback == "langfuse" or isinstance(callback, LangFuseLogger) for callback in callbacks):
return None
@ -79,7 +79,9 @@ async def _add_langfuse_trace_id_to_alert(
if request_data is None or request_data.get("litellm_logging_obj", None) is None:
return None
litellm_logging_obj: Final[Logging] = request_data["litellm_logging_obj"]
litellm_logging_obj: Final = cast( # cast-ok: logging object crosses the dynamic callback payload boundary
Logging, request_data["litellm_logging_obj"]
)
instance_host: Final = next(
(callback.langfuse_host for callback in callbacks if isinstance(callback, LangFuseLogger)), None
)
@ -87,8 +89,11 @@ async def _add_langfuse_trace_id_to_alert(
litellm_logging_obj.standard_callback_dynamic_params.get("langfuse_host") or instance_host
)
for _ in range(3):
if (trace_id := litellm_logging_obj._get_trace_id(service_name="langfuse")) is not None:
if (trace_id := litellm_logging_obj.get_trace_id(service_name="langfuse")) is not None:
return f"{host}/trace/{trace_id}"
await asyncio.sleep(3) # wait 3s before retrying for trace id
return None
_add_langfuse_trace_id_to_alert = add_langfuse_trace_id_to_alert

View file

@ -118,9 +118,9 @@ class ArizePhoenixTemplateManager:
# Load prompt from Arize Phoenix if prompt_id is provided
if self.prompt_id:
self._load_prompt_from_arize(self.prompt_id)
self.load_prompt_from_arize(self.prompt_id)
def _load_prompt_from_arize(self, prompt_version_id: str) -> None:
def load_prompt_from_arize(self, prompt_version_id: str) -> None:
"""Load a specific prompt version from Arize Phoenix."""
try:
# Fetch the prompt version from Arize Phoenix
@ -134,6 +134,8 @@ class ArizePhoenixTemplateManager:
except Exception as e:
raise Exception(f"Failed to load prompt version '{prompt_version_id}' from Arize Phoenix: {e}")
_load_prompt_from_arize = load_prompt_from_arize
def _parse_prompt_data(self, data: dict[str, Any], prompt_version_id: str) -> ArizePhoenixPromptTemplate:
"""Parse Arize Phoenix prompt data and extract messages and metadata."""
template_data: Final[ArizePhoenixTemplateBody] = data.get("template", {})
@ -418,7 +420,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
try:
# Load the prompt from Arize Phoenix if not already loaded
if prompt_id not in self.prompt_manager.prompts:
self.prompt_manager._load_prompt_from_arize(prompt_id)
self.prompt_manager.load_prompt_from_arize(prompt_id)
# Get the rendered messages and metadata
rendered_messages, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables)

View file

@ -94,9 +94,9 @@ class BitBucketTemplateManager:
# Load prompts from BitBucket if prompt_id is provided
if self.prompt_id:
self._load_prompt_from_bitbucket(self.prompt_id)
self.load_prompt_from_bitbucket(self.prompt_id)
def _load_prompt_from_bitbucket(self, prompt_id: str) -> None:
def load_prompt_from_bitbucket(self, prompt_id: str) -> None:
"""Load a specific .prompt file from BitBucket."""
try:
# Fetch the .prompt file from BitBucket
@ -108,6 +108,8 @@ class BitBucketTemplateManager:
except Exception as e:
raise Exception(f"Failed to load prompt '{prompt_id}' from BitBucket: {e}")
_load_prompt_from_bitbucket = load_prompt_from_bitbucket
def _parse_prompt_file(self, content: str, prompt_id: str) -> BitBucketPromptTemplate:
"""Parse a .prompt file content and extract metadata and template."""
# Split frontmatter and content
@ -446,7 +448,7 @@ class BitBucketPromptManager(CustomPromptManagement):
try:
# Load the prompt from BitBucket if not already loaded
if prompt_id not in self.prompt_manager.prompts:
self.prompt_manager._load_prompt_from_bitbucket(prompt_id)
self.prompt_manager.load_prompt_from_bitbucket(prompt_id)
# Get the rendered prompt and metadata
rendered_prompt, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables)

View file

@ -991,7 +991,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
import litellm
from litellm._logging import verbose_logger
all_callbacks: Final = litellm.logging_callback_manager._get_all_callbacks()
all_callbacks: Final = litellm.logging_callback_manager.get_all_callbacks()
for callback_obj in all_callbacks:
if hasattr(callback_obj, "increment_callback_logging_failure"):

View file

@ -27,7 +27,7 @@ def set_global_prompt_directory(directory: str) -> None:
litellm.global_prompt_directory = directory
def _get_prompt_data_from_dotprompt_content(dotprompt_content: str) -> dict:
def get_prompt_data_from_dotprompt_content(dotprompt_content: str) -> dict[str, object]:
"""
Get the prompt data from the dotprompt content.
@ -37,12 +37,15 @@ def _get_prompt_data_from_dotprompt_content(dotprompt_content: str) -> dict:
# Parse the dotprompt content to extract frontmatter and content
temp_manager: Final = PromptManager()
metadata, content = temp_manager._parse_frontmatter(dotprompt_content)
metadata, content = temp_manager.parse_frontmatter(dotprompt_content)
# Convert to prompt_data format
return {"content": content.strip(), "metadata": metadata}
_get_prompt_data_from_dotprompt_content = get_prompt_data_from_dotprompt_content
def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec") -> "CustomPromptManagement":
"""
Initialize a prompt from a .prompt file.
@ -60,7 +63,7 @@ def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "Prom
# Handle dotprompt_content from database
dotprompt_content: Final = getattr(litellm_params, "dotprompt_content", None)
if dotprompt_content and not prompt_data and not prompt_file:
prompt_data = _get_prompt_data_from_dotprompt_content(dotprompt_content)
prompt_data = get_prompt_data_from_dotprompt_content(dotprompt_content)
from .prompt_manager import strip_version_suffix

View file

@ -168,7 +168,7 @@ class PromptManager:
content: Final = file_path.read_text(encoding="utf-8")
# Split frontmatter and content
frontmatter, template_content = self._parse_frontmatter(content)
frontmatter, template_content = self.parse_frontmatter(content)
return PromptTemplate(
content=template_content.strip(),
@ -176,7 +176,7 @@ class PromptManager:
template_id=prompt_id,
)
def _parse_frontmatter(self, content: str) -> tuple[dict[str, object], str]:
def parse_frontmatter(self, content: str) -> tuple[dict[str, object], str]:
"""Parse YAML frontmatter from prompt content."""
# Match YAML frontmatter between --- delimiters
frontmatter_pattern: Final = r"^---\s*\n(.*?)\n---\s*\n(.*)$"
@ -197,6 +197,8 @@ class PromptManager:
return frontmatter, template_content
_parse_frontmatter = parse_frontmatter
def render(
self,
prompt_id: str,
@ -329,7 +331,7 @@ class PromptManager:
content: Final = file_path.read_text(encoding="utf-8")
# Parse frontmatter and content
frontmatter, template_content = self._parse_frontmatter(content)
frontmatter, template_content = self.parse_frontmatter(content)
return {"content": template_content.strip(), "metadata": frontmatter}

View file

@ -3,7 +3,8 @@
import os
import traceback
from collections.abc import Mapping
from collections.abc import Callable, Mapping
from datetime import datetime
from typing import Final, Protocol
import litellm
@ -32,9 +33,18 @@ class DyanmoDBLogger:
)
self.table_name = litellm.dynamodb_table_name
async def _async_log_event(self, kwargs, response_obj, start_time, end_time, print_verbose):
async def async_log_event(
self,
kwargs: Mapping[str, object],
response_obj: object,
start_time: datetime,
end_time: datetime,
print_verbose: Callable[[str], object],
) -> None:
self.log_event(kwargs, response_obj, start_time, end_time, print_verbose)
_async_log_event = async_log_event
def log_event(self, kwargs, response_obj, start_time, end_time, print_verbose):
try:
print_verbose(f"DynamoDB Logging - Enters logging function for model {kwargs}")

View file

@ -2,14 +2,15 @@
from __future__ import annotations
from typing import Any, Final
from collections.abc import Callable
from typing import Any, Final, SupportsFloat, SupportsIndex, SupportsInt, cast
import polars as pl
from litellm._logging import verbose_logger
from .database import FocusLiteLLMDatabase
from .destinations import FocusDestinationFactory, FocusTimeWindow
from .destinations import FocusDestination, FocusDestinationFactory, FocusTimeWindow
from .serializers import FocusCsvSerializer, FocusParquetSerializer, FocusSerializer
from .transformer import FocusTransformer
@ -28,14 +29,46 @@ class FocusExportEngine:
self.provider = provider
self.export_format = export_format
self.prefix = prefix
self._destination = FocusDestinationFactory.create(
self.destination = FocusDestinationFactory.create(
provider=self.provider,
prefix=self.prefix,
config=destination_config,
)
self._serializer = self._init_serializer()
self._transformer = FocusTransformer()
self._database = FocusLiteLLMDatabase()
self.serializer = self._init_serializer()
self.transformer = FocusTransformer()
self.database = FocusLiteLLMDatabase()
@property
def _destination(self) -> FocusDestination:
return self.destination
@_destination.setter
def _destination(self, value: FocusDestination) -> None:
self.destination = value
@property
def _serializer(self) -> FocusSerializer:
return self.serializer
@_serializer.setter
def _serializer(self, value: FocusSerializer) -> None:
self.serializer = value
@property
def _transformer(self) -> FocusTransformer:
return self.transformer
@_transformer.setter
def _transformer(self, value: FocusTransformer) -> None:
self.transformer = value
@property
def _database(self) -> FocusLiteLLMDatabase:
return self.database
@_database.setter
def _database(self, value: FocusLiteLLMDatabase) -> None:
self.database = value
def _init_serializer(self) -> FocusSerializer:
if self.export_format == "csv":
@ -45,18 +78,18 @@ class FocusExportEngine:
raise NotImplementedError(f"Export format '{self.export_format}' not supported. Use 'parquet' or 'csv'.")
async def dry_run_export_usage_data(self, limit: int | None) -> dict[str, Any]:
data: Final = await self._database.get_usage_data(limit=limit)
normalized: Final = self._transformer.transform(data)
data: Final = await self.database.get_usage_data(limit=limit)
normalized: Final = self.transformer.transform(data)
usage_sample: Final = data.head(min(50, len(data))).to_dicts()
normalized_sample: Final = normalized.head(min(50, len(normalized))).to_dicts()
summary: Final = {
"total_records": len(normalized),
"total_spend": self._sum_column(data, "spend"),
"total_tokens": self._sum_column(data, "total_tokens"),
"unique_teams": self._count_unique(data, "team_id"),
"unique_models": self._count_unique(data, "model"),
"total_spend": self.sum_column(data, "spend"),
"total_tokens": self.sum_column(data, "total_tokens"),
"unique_teams": self.count_unique(data, "team_id"),
"unique_models": self.count_unique(data, "model"),
}
return {
@ -71,12 +104,12 @@ class FocusExportEngine:
limit: int | None,
) -> None:
"""Export all available data without time-window filtering."""
data: Final = await self._database.get_usage_data(limit=limit)
data: Final = await self.database.get_usage_data(limit=limit)
if data.is_empty():
verbose_logger.debug("Focus export: no usage data available")
return
normalized: Final = self._transformer.transform(data)
normalized: Final = self.transformer.transform(data)
if normalized.is_empty():
verbose_logger.debug("Focus export: normalized data empty")
return
@ -98,7 +131,7 @@ class FocusExportEngine:
window: FocusTimeWindow,
limit: int | None,
) -> None:
data: Final = await self._database.get_usage_data(
data: Final = await self.database.get_usage_data(
limit=limit,
start_time_utc=window.start_time,
end_time_utc=window.end_time,
@ -107,7 +140,7 @@ class FocusExportEngine:
verbose_logger.debug("Focus export: no usage data for window %s", window)
return
normalized: Final = self._transformer.transform(data)
normalized: Final = self.transformer.transform(data)
if normalized.is_empty():
verbose_logger.debug("Focus export: normalized data empty for window %s", window)
return
@ -115,39 +148,45 @@ class FocusExportEngine:
await self._serialize_and_upload(normalized, window)
async def _serialize_and_upload(self, frame: pl.DataFrame, window: FocusTimeWindow) -> None:
payload: Final = self._serializer.serialize(frame)
payload: Final = self.serializer.serialize(frame)
if not payload:
verbose_logger.debug("Focus export: serializer returned empty payload")
return
await self._destination.deliver(
await self.destination.deliver(
content=payload,
time_window=window,
filename=self._build_filename(window),
filename=self.build_filename(window),
)
def _build_filename(self, window: FocusTimeWindow) -> str:
if not self._serializer.extension:
def build_filename(self, window: FocusTimeWindow) -> str:
if not self.serializer.extension:
raise ValueError("Serializer must declare a file extension")
# Include time window in filename so Vantage (which deduplicates
# by filename) doesn't overwrite previous uploads.
start_str: Final = window.start_time.strftime("%Y%m%dT%H%M%SZ")
end_str: Final = window.end_time.strftime("%Y%m%dT%H%M%SZ")
return f"usage_{start_str}_{end_str}.{self._serializer.extension}"
return f"usage_{start_str}_{end_str}.{self.serializer.extension}"
_build_filename = build_filename
@staticmethod
def _sum_column(frame: pl.DataFrame, column: str) -> float:
def sum_column(frame: pl.DataFrame, column: str) -> float:
if frame.is_empty() or column not in frame.columns:
return 0.0
value: Final = frame.select(pl.col(column).sum().alias("sum")).row(0)[0]
value: Final[object] = frame.select(pl.col(column).sum().alias("sum")).row(0)[0]
if value is None:
return 0.0
return float(value)
return float(cast(SupportsFloat | SupportsIndex | str | bytes | bytearray, value))
_sum_column: Final[Callable[[pl.DataFrame, str], float]] = sum_column
@staticmethod
def _count_unique(frame: pl.DataFrame, column: str) -> int:
def count_unique(frame: pl.DataFrame, column: str) -> int:
if frame.is_empty() or column not in frame.columns:
return 0
value: Final = frame.select(pl.col(column).n_unique().alias("unique")).row(0)[0]
value: Final[object] = frame.select(pl.col(column).n_unique().alias("unique")).row(0)[0]
if value is None:
return 0
return int(value)
return int(cast(SupportsInt | SupportsIndex | str | bytes | bytearray, value))
_count_unique: Final[Callable[[pl.DataFrame, str], int]] = count_unique

View file

@ -120,17 +120,19 @@ class GitLabTemplateManager:
)
if self.prompt_id:
self._load_prompt_from_gitlab(self.prompt_id)
self.load_prompt_from_gitlab(self.prompt_id)
# ---------- path helpers ----------
def _id_to_repo_path(self, prompt_id: str) -> str:
def id_to_repo_path(self, prompt_id: str) -> str:
"""Map a prompt_id to a repo path (respects prompts_path and adds .prompt)."""
prompt_id = decode_prompt_id(prompt_id)
if self.prompts_path:
return f"{self.prompts_path}/{prompt_id}.prompt"
return f"{prompt_id}.prompt"
_id_to_repo_path = id_to_repo_path
def _repo_path_to_id(self, repo_path: str) -> str:
"""
Map a repo path like 'prompts/chat/greeting.prompt' to an ID relative
@ -144,11 +146,11 @@ class GitLabTemplateManager:
# ---------- loading ----------
def _load_prompt_from_gitlab(self, prompt_id: str, *, ref: str | None = None) -> None:
def load_prompt_from_gitlab(self, prompt_id: str, *, ref: str | None = None) -> None:
"""Load a specific .prompt file from GitLab (scoped under prompts_path if set)."""
try:
# prompt_id = decode_prompt_id(prompt_id)
file_path: Final = self._id_to_repo_path(prompt_id)
file_path: Final = self.id_to_repo_path(prompt_id)
prompt_content: Final = self.gitlab_client.get_file_content(file_path, ref=ref)
if prompt_content:
template: Final = self._parse_prompt_file(prompt_content, prompt_id)
@ -156,6 +158,8 @@ class GitLabTemplateManager:
except Exception as e:
raise Exception(f"Failed to load prompt '{encode_prompt_id(prompt_id)}' from GitLab: {e}")
_load_prompt_from_gitlab = load_prompt_from_gitlab
def load_all_prompts(self, *, recursive: bool = True) -> list[str]:
"""
Eagerly load all .prompt files from prompts_path. Returns loaded IDs.
@ -164,7 +168,7 @@ class GitLabTemplateManager:
loaded: Final[list[str]] = []
for pid in files:
if pid not in self.prompts:
self._load_prompt_from_gitlab(pid)
self.load_prompt_from_gitlab(pid)
loaded.append(pid)
return loaded
@ -333,7 +337,7 @@ class GitLabPromptManager(CustomPromptManagement):
ref: str | None = None,
) -> tuple[str, dict[str, Any]]:
if prompt_id not in self.prompt_manager.prompts:
self.prompt_manager._load_prompt_from_gitlab(prompt_id, ref=ref)
self.prompt_manager.load_prompt_from_gitlab(prompt_id, ref=ref)
template: Final = self.prompt_manager.get_template(prompt_id)
if not template:
@ -506,7 +510,7 @@ class GitLabPromptManager(CustomPromptManagement):
if hasattr(dynamic_callback_params, "extra")
else None
)
self.prompt_manager._load_prompt_from_gitlab(decoded_id, ref=git_ref)
self.prompt_manager.load_prompt_from_gitlab(decoded_id, ref=git_ref)
rendered_prompt, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables)
@ -689,17 +693,17 @@ class GitLabPromptCache:
for pid in ids:
# Ensure template is loaded into TemplateManager
if pid not in self.template_manager.prompts:
self.template_manager._load_prompt_from_gitlab(pid)
self.template_manager.load_prompt_from_gitlab(pid)
tmpl = self.template_manager.get_template(pid)
if tmpl is None:
# If something raced/failed, try once more
self.template_manager._load_prompt_from_gitlab(pid)
self.template_manager.load_prompt_from_gitlab(pid)
tmpl = self.template_manager.get_template(pid)
if tmpl is None:
continue
file_path = self.template_manager._id_to_repo_path(pid) # "prompts/chat/..../file.prompt"
file_path = self.template_manager.id_to_repo_path(pid) # "prompts/chat/..../file.prompt"
entry = self._template_to_json(pid, tmpl)
self._by_file[file_path] = entry
@ -758,7 +762,7 @@ class GitLabPromptCache:
return {
"id": prompt_id, # e.g. "greet/hi"
"path": self.template_manager._id_to_repo_path(prompt_id), # e.g. "prompts/chat/greet/hi.prompt"
"path": self.template_manager.id_to_repo_path(prompt_id), # e.g. "prompts/chat/greet/hi.prompt"
"content": tmpl.content, # rendered content (without frontmatter)
"metadata": md, # parsed frontmatter
"model": model,

View file

@ -1023,7 +1023,7 @@ class LangFuseLogger:
_cache_key = _hidden_params.get("cache_key", None)
if _cache_key is None and litellm.cache is not None:
# fallback to using "preset_cache_key"
_preset_cache_key: Final = litellm.cache._get_preset_cache_key_from_kwargs(**kwargs) # pyright: ignore[reportPrivateUsage] # kwargs-ok: no public preset-cache-key accessor
_preset_cache_key: Final = litellm.cache.get_preset_cache_key_from_kwargs(**kwargs)
_cache_key = _preset_cache_key
tags.append(f"cache_key:{_cache_key}")
return tags

View file

@ -116,7 +116,7 @@ class MavvrikFocusLogger(FocusLogger):
"""Export with Mavvrik row cap applied when no explicit limit is passed."""
effective_limit: Final = limit if limit is not None else self._max_rows
engine: Final = self._ensure_engine()
data: Final = await engine._database.get_usage_data(
data: Final = await engine.database.get_usage_data(
limit=effective_limit,
start_time_utc=window.start_time,
end_time_utc=window.end_time,
@ -134,13 +134,13 @@ class MavvrikFocusLogger(FocusLogger):
if data.is_empty():
verbose_proxy_logger.debug("Mavvrik FOCUS export: no usage data for window %s", window)
else:
normalized: Final = engine._transformer.transform(data)
normalized: Final = engine.transformer.transform(data)
if not normalized.is_empty():
payload = engine._serializer.serialize(normalized)
await engine._destination.deliver(
payload = engine.serializer.serialize(normalized)
await engine.destination.deliver(
content=payload or b"",
time_window=window,
filename=engine._build_filename(window),
filename=engine.build_filename(window),
)
# Maximum number of days to catch up in a single run. Prevents runaway
@ -165,7 +165,7 @@ class MavvrikFocusLogger(FocusLogger):
FocusMavvrikDestination,
)
destination: Final = engine._destination
destination: Final = engine.destination
if not isinstance(destination, FocusMavvrikDestination):
await super()._run_scheduled_export()
return

View file

@ -181,7 +181,7 @@ class OTELMetricAttributeFilter:
exclude_list: list[str] | None = None
def _build_metric_attribute_filter(value: object) -> OTELMetricAttributeFilter:
def build_metric_attribute_filter(value: object) -> OTELMetricAttributeFilter:
if isinstance(value, OTELMetricAttributeFilter):
return value
if not isinstance(value, dict):
@ -195,7 +195,10 @@ def _build_metric_attribute_filter(value: object) -> OTELMetricAttributeFilter:
)
def _resolve_metric_attribute_filter(
_build_metric_attribute_filter = build_metric_attribute_filter
def resolve_metric_attribute_filter(
attributes: OTELMetricAttributeFilter | None,
) -> tuple[frozenset[str] | None, frozenset[str] | None]:
if attributes is None:
@ -220,6 +223,9 @@ def _resolve_metric_attribute_filter(
)
_resolve_metric_attribute_filter = resolve_metric_attribute_filter
def _provider_label(custom_llm_provider: object) -> str | None:
"""The provider label for one call's metrics and events, or None when the
call carries no provider.
@ -408,7 +414,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
if metadata_keys_override is not None:
config.baggage_metadata_keys = _normalize_team_metadata_keys(metadata_keys_override)
if metric_attributes_override is not None:
config.attributes = _build_metric_attribute_filter(metric_attributes_override)
config.attributes = build_metric_attribute_filter(metric_attributes_override)
self.config = config
self.callback_name = callback_name
@ -1643,11 +1649,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
otel_settings: Final = (litellm.callback_settings or {}).get("otel") or {}
raw: Final[object] = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None
if raw is not None:
attributes = _build_metric_attribute_filter(raw)
attributes = build_metric_attribute_filter(raw)
(
self._metric_attr_include,
self._metric_attr_exclude,
) = _resolve_metric_attribute_filter(attributes)
) = resolve_metric_attribute_filter(attributes)
self._metric_attr_filter_resolved = True
def _filter_metric_attributes(self, attrs: Mapping[str, str | None]) -> dict[str, str]:

View file

@ -21,8 +21,8 @@ from litellm._logging import verbose_logger
from litellm.integrations.opentelemetry import (
METRIC_METADATA_KEYS,
TOKEN_TYPE_ATTRIBUTE,
_build_metric_attribute_filter,
_resolve_metric_attribute_filter,
build_metric_attribute_filter,
resolve_metric_attribute_filter,
)
from litellm.integrations.otel.model.metadata import time_to_first_chunk_seconds
from litellm.integrations.otel.model.semconv import (
@ -324,13 +324,13 @@ class GenAIMetricRecorder:
otel_settings: Final = (litellm.callback_settings or {}).get("otel") or {}
raw: Final[object] = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None
if raw is not None:
attributes = _build_metric_attribute_filter(raw)
attributes = build_metric_attribute_filter(raw)
# A bad filter (include_list + exclude_list both set, an unfilterable name)
# raises here; the caller (logger._record_metrics) surfaces it once at ERROR
# so the operator-fixable config error is visible. Not cached on the raise
# path -- _filter_resolved stays False -- so a corrected config takes effect
# without reconstructing the recorder.
self._include, self._exclude = _resolve_metric_attribute_filter(attributes)
self._include, self._exclude = resolve_metric_attribute_filter(attributes)
self._filter_resolved = True
self._warn_about_metric_ineligible_names()

View file

@ -12,7 +12,7 @@ from litellm.integrations.otel.presets.utils import (
ensure_mappers,
)
from litellm.integrations.weave.weave_otel import (
_get_weave_authorization_header,
get_weave_authorization_header,
get_weave_otel_config,
)
from litellm.types.utils import StandardCallbackDynamicParams
@ -58,7 +58,7 @@ def weave_dynamic_headers(params: StandardCallbackDynamicParams) -> dict[str, st
headers: Final[dict[str, str]] = {}
api_key: Final = params.get("wandb_api_key")
if api_key:
headers["Authorization"] = _get_weave_authorization_header(api_key=api_key)
headers["Authorization"] = get_weave_authorization_header(api_key=api_key)
project_id: Final = params.get("weave_project_id")
if project_id:
headers["project_id"] = project_id

View file

@ -32,7 +32,7 @@ from litellm.exceptions import (
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.prometheus_helpers import (
PrometheusLabelFactoryContext,
_get_cached_end_user_id_for_cost_tracking,
get_cached_end_user_id_for_cost_tracking,
)
from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker import (
BoundedPrometheusSeriesTracker,
@ -65,8 +65,8 @@ from litellm.repositories.user_repository import UserRepository
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.integrations.prometheus import *
from litellm.types.integrations.prometheus import (
_sanitize_prometheus_label_name,
_sanitize_prometheus_label_value,
sanitize_prometheus_label_name,
sanitize_prometheus_label_value,
validate_prometheus_deployment_and_latency_caller_identity,
)
from litellm.types.proxy.carried_budget_state import (
@ -1069,10 +1069,10 @@ class PrometheusLogger(CustomLogger):
builtin_labels: Final = frozenset(label.value for label in UserAPIKeyLabelNames)
custom_metadata_labels: Final = frozenset(
_sanitize_prometheus_label_name(label) for label in litellm.custom_prometheus_metadata_labels
sanitize_prometheus_label_name(label) for label in litellm.custom_prometheus_metadata_labels
)
custom_tag_labels: Final = frozenset(
_sanitize_prometheus_label_name(f"tag_{tag}") for tag in litellm.custom_prometheus_tags
sanitize_prometheus_label_name(f"tag_{tag}") for tag in litellm.custom_prometheus_tags
)
return builtin_labels | _NON_ENUM_METRIC_LABELS | custom_metadata_labels | custom_tag_labels
@ -1508,7 +1508,7 @@ class PrometheusLogger(CustomLogger):
model: Final = kwargs.get("model", "")
litellm_params: Final = kwargs.get("litellm_params", {}) or {}
_metadata: Final = litellm_params.get("metadata") or {}
get_end_user_id_for_cost_tracking: Final = _get_cached_end_user_id_for_cost_tracking()
get_end_user_id_for_cost_tracking: Final = get_cached_end_user_id_for_cost_tracking()
end_user_id: Final = get_end_user_id_for_cost_tracking(litellm_params, service_type="prometheus")
user_id: Final = standard_logging_payload["metadata"]["user_api_key_user_id"]
@ -2523,7 +2523,7 @@ class PrometheusLogger(CustomLogger):
model: Final = kwargs.get("model", "")
litellm_params: Final = kwargs.get("litellm_params", {}) or {}
get_end_user_id_for_cost_tracking: Final = _get_cached_end_user_id_for_cost_tracking()
get_end_user_id_for_cost_tracking: Final = get_cached_end_user_id_for_cost_tracking()
end_user_id: Final = get_end_user_id_for_cost_tracking(litellm_params, service_type="prometheus")
user_id: Final = standard_logging_payload["metadata"]["user_api_key_user_id"]
@ -2775,7 +2775,7 @@ class PrometheusLogger(CustomLogger):
status_code: Final = self._extract_status_code(exception=original_exception)
try:
_tags: Final = StandardLoggingPayloadSetup._get_request_tags(
_tags: Final = StandardLoggingPayloadSetup.get_request_tags(
litellm_params=request_data,
proxy_server_request=request_data.get("proxy_server_request", {}),
)
@ -3725,11 +3725,11 @@ class PrometheusLogger(CustomLogger):
increment metric when litellm.Router / load balancing logic places a deployment in cool down
"""
self.litellm_deployment_cooled_down.labels(
_sanitize_prometheus_label_value(litellm_model_name),
_sanitize_prometheus_label_value(model_id),
_sanitize_prometheus_label_value(api_base),
_sanitize_prometheus_label_value(api_provider),
_sanitize_prometheus_label_value(exception_status),
sanitize_prometheus_label_value(litellm_model_name),
sanitize_prometheus_label_value(model_id),
sanitize_prometheus_label_value(api_base),
sanitize_prometheus_label_value(api_provider),
sanitize_prometheus_label_value(exception_status),
).inc()
def increment_callback_logging_failure(
@ -4750,17 +4750,17 @@ def _prometheus_labels_from_context(
ctx: PrometheusLabelFactoryContext,
) -> dict[str, str | None]:
filtered_labels: Final[dict[str, str | None]] = {
label: ctx._sanitized_enum[label] for label in supported_enum_labels if label in ctx._sanitized_enum
label: ctx.sanitized_enum[label] for label in supported_enum_labels if label in ctx.sanitized_enum
}
if UserAPIKeyLabelNames.END_USER.value in filtered_labels:
filtered_labels[UserAPIKeyLabelNames.END_USER.value] = ctx.get_resolved_end_user()
for sk, val in ctx._custom_by_sanitized_key.items():
for sk, val in ctx.custom_by_sanitized_key.items():
if sk in supported_enum_labels:
filtered_labels[sk] = val
for k, v in ctx._tag_labels.items():
for k, v in ctx.tag_labels.items():
if k in supported_enum_labels:
filtered_labels[k] = v
@ -4797,13 +4797,13 @@ def prometheus_label_factory(
# Filter supported labels and sanitize values to prevent breaking
# the Prometheus text format (e.g. U+2028 Line Separator in label values)
filtered_labels: Final = {
label: _sanitize_prometheus_label_value(value)
label: sanitize_prometheus_label_value(value)
for label, value in enum_dict.items()
if label in supported_enum_labels
}
if UserAPIKeyLabelNames.END_USER.value in filtered_labels:
get_end_user_id_for_cost_tracking: Final = _get_cached_end_user_id_for_cost_tracking()
get_end_user_id_for_cost_tracking: Final = get_cached_end_user_id_for_cost_tracking()
filtered_labels["end_user"] = get_end_user_id_for_cost_tracking(
litellm_params={"user_api_key_end_user_id": enum_values.end_user},
@ -4813,16 +4813,16 @@ def prometheus_label_factory(
if enum_values.custom_metadata_labels is not None:
for key, value in enum_values.custom_metadata_labels.items():
# check sanitized key
sanitized_key = _sanitize_prometheus_label_name(key)
sanitized_key = sanitize_prometheus_label_name(key)
if sanitized_key in supported_enum_labels:
filtered_labels[sanitized_key] = _sanitize_prometheus_label_value(value)
filtered_labels[sanitized_key] = sanitize_prometheus_label_value(value)
# Add custom tags if configured
if enum_values.tags is not None:
custom_tag_labels: Final = get_custom_labels_from_tags(enum_values.tags)
for key, value in custom_tag_labels.items():
if key in supported_enum_labels:
filtered_labels[key] = _sanitize_prometheus_label_value(value)
filtered_labels[key] = sanitize_prometheus_label_value(value)
for label in supported_enum_labels:
if label not in filtered_labels:
@ -4919,7 +4919,7 @@ def _tag_matches_wildcard_configured_pattern(tags: Sequence[str], configured_tag
from litellm.router_utils.pattern_match_deployments import PatternMatchRouter
pattern_router: Final = PatternMatchRouter()
regex_pattern: Final = pattern_router._pattern_to_regex(configured_tag)
regex_pattern: Final = pattern_router.pattern_to_regex(configured_tag)
return any(re.match(pattern=regex_pattern, string=tag) for tag in tags)
@ -4945,7 +4945,7 @@ def get_custom_labels_from_tags(tags: Sequence[str]) -> dict[str, str]:
}
"""
from litellm.types.integrations.prometheus import _sanitize_prometheus_label_name
from litellm.types.integrations.prometheus import sanitize_prometheus_label_name
configured_tags: Final = litellm.custom_prometheus_tags
if configured_tags is None or len(configured_tags) == 0:
@ -4954,7 +4954,7 @@ def get_custom_labels_from_tags(tags: Sequence[str]) -> dict[str, str]:
result: Final[dict[str, str]] = {}
for configured_tag in configured_tags:
label_name = _sanitize_prometheus_label_name(f"tag_{configured_tag}")
label_name = sanitize_prometheus_label_name(f"tag_{configured_tag}")
# Check for exact match first (backwards compatibility)
if configured_tag in tags:

View file

@ -6,18 +6,26 @@ Helpers for the Prometheus integration (extracted to keep ``prometheus.py`` smal
from __future__ import annotations
from typing import Final, cast
from typing import Final, Literal, Protocol, cast
from litellm.types.integrations.prometheus import (
UserAPIKeyLabelValues,
_sanitize_prometheus_label_name,
_sanitize_prometheus_label_value,
sanitize_prometheus_label_name,
sanitize_prometheus_label_value,
)
_get_end_user_id_for_cost_tracking = None
def _get_cached_end_user_id_for_cost_tracking():
class _EndUserIdGetter(Protocol):
def __call__(
self,
litellm_params: dict[str, object],
service_type: Literal["litellm_logging", "prometheus"] = "litellm_logging",
) -> str | None: ...
def get_cached_end_user_id_for_cost_tracking() -> _EndUserIdGetter:
"""
Get cached get_end_user_id_for_cost_tracking function.
Lazy imports on first call to avoid loading utils.py at import time (60MB saved).
@ -31,6 +39,9 @@ def _get_cached_end_user_id_for_cost_tracking():
return _get_end_user_id_for_cost_tracking
_get_cached_end_user_id_for_cost_tracking = get_cached_end_user_id_for_cost_tracking
class PrometheusLabelFactoryContext:
"""
Precomputes per-request label inputs so prometheus_label_factory can subset
@ -38,11 +49,11 @@ class PrometheusLabelFactoryContext:
"""
__slots__ = (
"_custom_by_sanitized_key",
"_resolved_end_user",
"_sanitized_enum",
"_tag_labels",
"custom_by_sanitized_key",
"enum_values",
"sanitized_enum",
"tag_labels",
)
_END_USER_NOT_COMPUTED = object()
@ -50,27 +61,51 @@ class PrometheusLabelFactoryContext:
def __init__(self, enum_values: UserAPIKeyLabelValues) -> None:
self.enum_values = enum_values
enum_dict: Final = enum_values.model_dump()
self._sanitized_enum: dict[str, str | None] = {
k: _sanitize_prometheus_label_value(v) for k, v in enum_dict.items()
self.sanitized_enum: dict[str, str | None] = {
k: sanitize_prometheus_label_value(v) for k, v in enum_dict.items()
}
self._custom_by_sanitized_key: dict[str, str | None] = {}
self.custom_by_sanitized_key: dict[str, str | None] = {}
if enum_values.custom_metadata_labels is not None:
for key, value in enum_values.custom_metadata_labels.items():
sk = _sanitize_prometheus_label_name(key)
self._custom_by_sanitized_key[sk] = _sanitize_prometheus_label_value(value)
self._tag_labels: dict[str, str | None] = {}
sk = sanitize_prometheus_label_name(key)
self.custom_by_sanitized_key[sk] = sanitize_prometheus_label_value(value)
self.tag_labels: dict[str, str | None] = {}
if enum_values.tags is not None:
# Late import avoids circular import: ``prometheus`` imports this module.
from litellm.integrations.prometheus import get_custom_labels_from_tags
for k, v in get_custom_labels_from_tags(enum_values.tags).items():
self._tag_labels[k] = _sanitize_prometheus_label_value(v)
self.tag_labels[k] = sanitize_prometheus_label_value(v)
# Use a dedicated sentinel so `None` can be cached as a computed result.
self._resolved_end_user: object = self._END_USER_NOT_COMPUTED
@property
def _custom_by_sanitized_key(self) -> dict[str, str | None]:
return self.custom_by_sanitized_key
@_custom_by_sanitized_key.setter
def _custom_by_sanitized_key(self, value: dict[str, str | None]) -> None:
self.custom_by_sanitized_key = value
@property
def _sanitized_enum(self) -> dict[str, str | None]:
return self.sanitized_enum
@_sanitized_enum.setter
def _sanitized_enum(self, value: dict[str, str | None]) -> None:
self.sanitized_enum = value
@property
def _tag_labels(self) -> dict[str, str | None]:
return self.tag_labels
@_tag_labels.setter
def _tag_labels(self, value: dict[str, str | None]) -> None:
self.tag_labels = value
def get_resolved_end_user(self) -> str | None:
if self._resolved_end_user is self._END_USER_NOT_COMPUTED:
fn: Final = _get_cached_end_user_id_for_cost_tracking()
fn: Final = get_cached_end_user_id_for_cost_tracking()
self._resolved_end_user = fn(
litellm_params={"user_api_key_end_user_id": self.enum_values.end_user},
service_type="prometheus",

View file

@ -113,7 +113,9 @@ class PrometheusServicesLogger:
"""
Helper function to get a metric from the registry by name.
"""
return self.REGISTRY._names_to_collectors.get(metric_name)
return self.REGISTRY._names_to_collectors.get( # pyright: ignore[reportPrivateUsage] # Registry lookup has no public API
metric_name
)
def create_histogram(self, service: str, type_of_request: str):
metric_name: Final = f"litellm_{service}_{type_of_request}"
@ -196,7 +198,7 @@ class PrometheusServicesLogger:
labels=payload.service.value,
amount=payload.duration,
)
elif isinstance(obj, self.Counter) and "total_requests" in obj._name:
elif isinstance(obj, self.Counter) and "total_requests" in obj._name: # pyright: ignore[reportPrivateUsage] # Metric names have no public accessor
self.increment_counter(
counter=obj,
labels=payload.service.value,
@ -233,7 +235,7 @@ class PrometheusServicesLogger:
labels=payload.service.value,
amount=payload.duration,
)
elif isinstance(obj, self.Counter) and "total_requests" in obj._name:
elif isinstance(obj, self.Counter) and "total_requests" in obj._name: # pyright: ignore[reportPrivateUsage] # Metric names have no public accessor
self.increment_counter(
counter=obj,
labels=payload.service.value,
@ -262,7 +264,7 @@ class PrometheusServicesLogger:
for obj in prom_objects:
# increment both failed and total requests
if isinstance(obj, self.Counter):
if "failed_requests" in obj._name:
if "failed_requests" in obj._name: # pyright: ignore[reportPrivateUsage] # Metric names have no public accessor
self.increment_counter(
counter=obj,
labels=payload.service.value,

View file

@ -106,7 +106,7 @@ def _set_weave_specific_attributes(span: Span, kwargs: Mapping[str, Any], respon
safe_set_attribute(span, OpenInferenceSpanAttributes.OUTPUT_VALUE, safe_dumps(output_dict))
def _get_weave_authorization_header(api_key: str) -> str:
def get_weave_authorization_header(api_key: str) -> str:
"""
Get the authorization header for Weave OpenTelemetry.
@ -117,6 +117,9 @@ def _get_weave_authorization_header(api_key: str) -> str:
return f"Basic {auth_header}"
_get_weave_authorization_header = get_weave_authorization_header
def weave_otel_endpoint(host: str | None) -> str:
"""The OTLP traces endpoint for a self-managed ``host``, else Weave cloud."""
if not host:
@ -155,7 +158,7 @@ def get_weave_otel_config() -> WeaveOtelConfig:
verbose_logger.debug("Using Weave OTEL endpoint: %s", endpoint)
# Weave uses Basic auth with format: api:<WANDB_API_KEY>
auth_header: Final = _get_weave_authorization_header(api_key=api_key)
auth_header: Final = get_weave_authorization_header(api_key=api_key)
otlp_auth_headers: Final = f"Authorization={auth_header},project_id={project_id}"
# Set standard OTEL environment variables
@ -320,7 +323,7 @@ class WeaveOtelLogger(OpenTelemetry):
dynamic_weave_project_id: Final = standard_callback_dynamic_params.get("weave_project_id")
if dynamic_wandb_api_key:
auth_header: Final = _get_weave_authorization_header(
auth_header: Final = get_weave_authorization_header(
api_key=dynamic_wandb_api_key,
)
dynamic_headers["Authorization"] = auth_header

View file

@ -72,7 +72,7 @@ class BaseInteractionsAPIStreamingIterator:
return None
# Handle SSE format (data: {...})
stripped_chunk: Final = CustomStreamWrapper._strip_sse_data_from_chunk(chunk)
stripped_chunk: Final = CustomStreamWrapper.strip_sse_data_from_chunk(chunk)
if stripped_chunk is None:
return None

View file

@ -5,7 +5,7 @@ import logging
import re
from collections.abc import Collection, Iterable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast
import httpx
from pydantic import TypeAdapter, ValidationError
@ -427,30 +427,43 @@ def reconstruct_model_name(
# Helper functions used for OTEL logging
def _get_parent_otel_span_from_kwargs(
kwargs: dict | None = None,
def get_parent_otel_span_from_kwargs(
kwargs: dict[str, object] | None = None,
) -> Span | None:
try:
if kwargs is None:
return None
litellm_params: Final = kwargs.get("litellm_params")
_metadata: Final = kwargs.get("metadata") or {}
if "litellm_parent_otel_span" in _metadata:
return _metadata["litellm_parent_otel_span"]
metadata: Final = cast( # cast-ok: metadata is caller-provided request data
Mapping[str, object], kwargs.get("metadata") or {}
)
if "litellm_parent_otel_span" in metadata:
return cast( # cast-ok: tracing metadata crosses an external boundary
Span | None, metadata["litellm_parent_otel_span"]
)
elif (
litellm_params is not None
and litellm_params.get("metadata") is not None
and "litellm_parent_otel_span" in litellm_params.get("metadata", {})
and cast(Mapping[str, object], litellm_params).get("metadata") is not None
and "litellm_parent_otel_span"
in cast(
Mapping[str, object],
cast(Mapping[str, object], litellm_params).get("metadata", {}),
)
):
return litellm_params["metadata"]["litellm_parent_otel_span"]
typed_litellm_params: Final = cast(Mapping[str, object], litellm_params)
litellm_metadata: Final = cast(Mapping[str, object], typed_litellm_params["metadata"])
return cast(Span | None, litellm_metadata["litellm_parent_otel_span"])
elif "litellm_parent_otel_span" in kwargs:
return kwargs["litellm_parent_otel_span"]
return cast(Span | None, kwargs["litellm_parent_otel_span"])
return None
except Exception as e:
verbose_logger.exception("Error in _get_parent_otel_span_from_kwargs: " + str(e))
return None
_get_parent_otel_span_from_kwargs = get_parent_otel_span_from_kwargs
def process_response_headers(
response_headers: httpx.Headers | dict,
preserve_litellm_internal_headers: bool = False,

View file

@ -57,11 +57,14 @@ def _should_use_dd_tracer():
return get_secret_bool("USE_DDTRACE", False) is True
def _should_use_dd_profiler():
def should_use_dd_profiler() -> bool:
"""Returns True if `USE_DDPROFILER` is set to True in .env"""
return get_secret_bool("USE_DDPROFILER", False) is True
_should_use_dd_profiler = should_use_dd_profiler
# Initialize tracer
should_use_dd_tracer: Final = _should_use_dd_tracer()
tracer: NullTracer | DD_TRACER = NullTracer()

View file

@ -9,13 +9,13 @@ from typing import Final, Protocol, cast
import httpx
import litellm
from litellm._logging import _ENABLE_SECRET_REDACTION, _redact_string, verbose_logger
from litellm._logging import _ENABLE_SECRET_REDACTION, redact_string, verbose_logger
from litellm.litellm_core_utils.bug_report import (
bug_report_notice,
build_bug_report,
should_report_bug,
)
from litellm.litellm_core_utils.secret_redaction import redact_string
from litellm.litellm_core_utils.secret_redaction import redact_string as redact_secret_string
from litellm.types.utils import LlmProviders
from ..exceptions import (
@ -2128,7 +2128,7 @@ def _map_azure_exception(
else:
# if no status code then it is an APIConnectionError: https://github.com/openai/openai-python#handling-errors
raise APIConnectionError(
message=f"{exception_provider} APIConnectionError - {message}\n{_redact_string(traceback.format_exc())}",
message=f"{exception_provider} APIConnectionError - {message}\n{redact_string(traceback.format_exc())}",
llm_provider="azure",
model=model,
litellm_debug_info=extra_information,
@ -2373,18 +2373,20 @@ def exception_type(
"\033[1;31mGive Feedback / Get Help: https://github.com/BerriAI/litellm/issues/new\033[0m"
)
print( # noqa: T201
"LiteLLM.Info: If you need to debug this error, use `litellm._turn_on_debug()'."
"LiteLLM.Info: If you need to debug this error, use `litellm.turn_on_debug()'."
)
print() # noqa: T201
litellm_response_headers: Final = _get_response_headers(original_exception=original_exception)
try:
error_str = redact_string(str(original_exception)) if _ENABLE_SECRET_REDACTION else str(original_exception)
error_str = (
redact_secret_string(str(original_exception)) if _ENABLE_SECRET_REDACTION else str(original_exception)
)
extra_information = ""
if model or custom_llm_provider:
if hasattr(original_exception, "message"):
error_str = (
redact_string(str(original_exception.message))
redact_secret_string(str(original_exception.message))
if _ENABLE_SECRET_REDACTION
else str(original_exception.message)
)
@ -2425,7 +2427,7 @@ def exception_type(
extra_information += f"\nvertex_location: `{_vertex_location}`\n"
# on litellm proxy add key name + team to exceptions
extra_information = _add_key_name_and_team_to_alert(request_info=extra_information, metadata=_metadata)
extra_information = add_key_name_and_team_to_alert(request_info=extra_information, metadata=_metadata)
except Exception:
# DO NOT LET this Block raising the original exception
pass
@ -2687,7 +2689,7 @@ def exception_type(
else:
raise APIConnectionError(
message=(
f"{original_exception}\n{_redact_string(traceback.format_exc())}"
f"{original_exception}\n{redact_string(traceback.format_exc())}"
+ (
"\n"
+ bug_report_notice(
@ -2726,7 +2728,7 @@ def exception_type(
setattr(e, "litellm_response_headers", litellm_response_headers)
raise e # it's already mapped
raised_exc: Final = APIConnectionError(
message=f"{original_exception}\n{_redact_string(traceback.format_exc())}",
message=f"{original_exception}\n{redact_string(traceback.format_exc())}",
llm_provider="",
model="",
)
@ -2766,7 +2768,7 @@ def exception_logging(
)
def _add_key_name_and_team_to_alert(request_info: str, metadata: dict) -> str:
def add_key_name_and_team_to_alert(request_info: str, metadata: Mapping[str, object]) -> str:
"""
Internal helper function for litellm proxy
Add the Key Name + Team Name to the error
@ -2783,3 +2785,6 @@ def _add_key_name_and_team_to_alert(request_info: str, metadata: dict) -> str:
return request_info
except Exception:
return request_info
_add_key_name_and_team_to_alert = add_key_name_and_team_to_alert

View file

@ -2,7 +2,7 @@ import reprlib
from collections.abc import Mapping, MutableMapping
from dataclasses import dataclass, fields
from types import MappingProxyType
from typing import Final
from typing import Final, cast
from pydantic import TypeAdapter, ValidationError
@ -123,17 +123,25 @@ def with_control_options(litellm_params: Mapping[str, object], control: ControlO
return {**litellm_params, CONTROL_OPTIONS_KEY: control}
def _get_base_model_from_litellm_call_metadata(
metadata: dict | None,
def get_base_model_from_litellm_call_metadata(
metadata: Mapping[str, object] | None,
) -> str | None:
if metadata is None:
return None
model_info: Final = metadata.get("model_info")
if model_info:
return model_info.get("base_model")
model_info_mapping: Final = cast( # cast-ok: model metadata is caller-provided and preserves its mapping shape
Mapping[str, object], model_info
)
return cast( # cast-ok: model metadata values are caller-provided
str | None, model_info_mapping.get("base_model")
)
return None
_get_base_model_from_litellm_call_metadata = get_base_model_from_litellm_call_metadata
def get_litellm_params(
api_key: str | None = None,
force_timeout=600,
@ -229,7 +237,7 @@ def get_litellm_params(
"azure_ad_token_provider": azure_ad_token_provider,
"user_continue_message": user_continue_message,
"base_model": base_model
or (_get_base_model_from_litellm_call_metadata(metadata=metadata) if metadata else None),
or (get_base_model_from_litellm_call_metadata(metadata=metadata) if metadata else None),
"litellm_trace_id": litellm_trace_id,
"litellm_session_id": litellm_session_id,
"hf_model_name": hf_model_name,

View file

@ -53,7 +53,7 @@ def _endpoint_matches_api_base(endpoint: str, api_base: str) -> bool:
return url_path == endpoint_path or url_path.startswith(endpoint_path + "/")
def _is_non_openai_azure_model(model: str) -> bool:
def is_non_openai_azure_model(model: str) -> bool:
try:
model_name: Final = model.split("/", 1)[1]
if model_name in litellm.cohere_chat_models or f"mistral/{model_name}" in litellm.mistral_chat_models:
@ -63,6 +63,9 @@ def _is_non_openai_azure_model(model: str) -> bool:
return False
_is_non_openai_azure_model = is_non_openai_azure_model
def _is_azure_claude_model(model: str) -> bool:
"""
Check if a model name contains 'claude' (case-insensitive).
@ -195,7 +198,7 @@ def get_llm_provider(
# AZURE AI-Studio Logic - Azure AI Studio supports AZURE/Cohere
# If User passes azure/command-r-plus -> we should send it to cohere_chat/command-r-plus
if model.split("/", 1)[0] == "azure":
if _is_non_openai_azure_model(model):
if is_non_openai_azure_model(model):
custom_llm_provider = "openai"
return model, custom_llm_provider, dynamic_api_key, api_base

View file

@ -105,13 +105,15 @@ class GetModelCostMap:
cls._loaded_catalog = MappingProxyType({key: MappingProxyType(entry) for key, entry in raw.items()})
@classmethod
def _get_backup_model_count(cls) -> int:
def get_backup_model_count(cls) -> int:
"""Return the number of models in the local backup (cached int)."""
if cls._backup_model_count < 0:
backup: Final = cls.load_local_model_cost_map()
cls._backup_model_count = _count_model_entries(backup)
return cls._backup_model_count
_get_backup_model_count = get_backup_model_count
@staticmethod
def _check_is_valid_dict(fetched_map: dict) -> bool:
"""Check 1: fetched map is a non-empty dict."""
@ -402,7 +404,7 @@ async def refetch_model_cost_map(
return result
if not GetModelCostMap.validate_model_cost_map(
fetched_map=result.model_cost_map,
backup_model_count=GetModelCostMap._get_backup_model_count(),
backup_model_count=GetModelCostMap.get_backup_model_count(),
):
return ModelCostMapReloadUnavailable(reason=f"model cost map from {url} failed integrity validation")
_cost_map_source_info.loaded_at = datetime.now(timezone.utc)
@ -600,7 +602,7 @@ def _retry_remote_fetch_in_background(
_litellm_import_complete.wait()
if not GetModelCostMap.validate_model_cost_map(
fetched_map=result.model_cost_map,
backup_model_count=GetModelCostMap._get_backup_model_count(), # pyright: ignore[reportPrivateUsage] # integrity cache
backup_model_count=GetModelCostMap.get_backup_model_count(),
):
verbose_logger.warning(
"LiteLLM: Fetched model cost map failed integrity check. Using local backup instead. url=%s",
@ -683,7 +685,7 @@ def get_model_cost_map(
# Validate using cached count (cheap int comparison, no file I/O)
if not GetModelCostMap.validate_model_cost_map(
fetched_map=content,
backup_model_count=GetModelCostMap._get_backup_model_count(),
backup_model_count=GetModelCostMap.get_backup_model_count(),
):
verbose_logger.warning(
"LiteLLM: Fetched model cost map failed integrity check. Using local backup instead. url=%s",

View file

@ -96,7 +96,7 @@ class HealthCheckHelpers:
return {}
@staticmethod
def _update_model_params_with_health_check_tracking_information(
def update_model_params_with_health_check_tracking_information(
model_params: dict,
) -> dict:
"""
@ -120,6 +120,10 @@ class HealthCheckHelpers:
)
return model_params
_update_model_params_with_health_check_tracking_information = (
update_model_params_with_health_check_tracking_information
)
@staticmethod
def _get_metadata_for_health_check_call():
"""
@ -219,7 +223,7 @@ class HealthCheckHelpers:
get_audio_file_for_health_check,
)
from litellm.litellm_core_utils.health_check_utils import DECISIONS_CALL_PARAMS, _filter_model_params
from litellm.realtime_api.main import _realtime_health_check
from litellm.realtime_api.main import realtime_health_check
return {
"chat": lambda: litellm.acompletion(
@ -264,7 +268,7 @@ class HealthCheckHelpers:
query=prompt or "",
documents=["my sample text"],
),
"realtime": lambda: _realtime_health_check(
"realtime": lambda: realtime_health_check(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=model_params.get("api_base", None),

View file

@ -2,6 +2,7 @@
Utils used for litellm.ahealth_check()
"""
from collections.abc import Mapping
from typing import Final
from pydantic import TypeAdapter
@ -17,8 +18,8 @@ def _filter_model_params(model_params: dict) -> dict:
return {k: v for k, v in model_params.items() if k != "messages"}
def _create_health_check_response(response_headers: dict) -> dict:
response: Final = {}
def create_health_check_response(response_headers: Mapping[str, object]) -> dict[str, object]:
response: Final[dict[str, object]] = {}
if response_headers.get("x-ratelimit-remaining-requests", None) is not None: # not provided for dall-e requests
response["x-ratelimit-remaining-requests"] = response_headers["x-ratelimit-remaining-requests"]
@ -29,3 +30,6 @@ def _create_health_check_response(response_headers: dict) -> dict:
if response_headers.get("x-ms-region", None) is not None:
response["x-ms-region"] = response_headers["x-ms-region"]
return response
_create_health_check_response = create_health_check_response

View file

@ -11,11 +11,11 @@ import subprocess
import sys
import time
import traceback
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from collections.abc import Awaitable, Callable, Iterable, Iterator, Mapping, Sequence
from datetime import datetime as dt_object
from functools import lru_cache
from types import MappingProxyType, TracebackType
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Union, cast
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast
from httpx import Response
from pydantic import BaseModel, JsonValue
@ -24,8 +24,8 @@ import litellm
from litellm import _custom_logger_compatible_callbacks_literal
from litellm._internal_context import post_response_phase
from litellm._logging import (
_is_debugging_on,
_redact_string,
is_debugging_on,
redact_string,
session_id_var,
set_session_id,
set_trace_id,
@ -33,7 +33,7 @@ from litellm._logging import (
verbose_logger,
)
from litellm._uuid import uuid
from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final
from litellm.batches.batch_utils import batch_cost_is_final, handle_completed_batch
from litellm.caching.caching import DualCache
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.caching.redis_batch import flush_post_call_redis_batches
@ -47,9 +47,9 @@ from litellm.constants import (
from litellm.cost_calculator import (
RealtimeAPITokenUsageProcessor,
ResponsesWebSocketTokenUsageProcessor,
_select_model_name_for_cost_calc,
get_usage_object,
pricing_entry_for_cost_calc,
select_model_name_for_cost_calc,
)
from litellm.exceptions import (
BudgetExceededError,
@ -178,7 +178,7 @@ from litellm.types.utils import (
Usage,
)
from litellm.types.videos.main import VideoObject
from litellm.utils import _get_base_model_from_metadata, print_verbose
from litellm.utils import get_base_model_from_metadata, print_verbose
from ..integrations.argilla import ArgillaLogger
from ..integrations.arize.arize_phoenix import ArizePhoenixLogger
@ -575,6 +575,38 @@ class Logging(LiteLLMLoggingBaseClass):
baseline_cache_context: "BaselineCacheContext | None" = None
baseline_observation: "CapturedBaselineObservation | None" = None
@property
def _defer_async_logging(self) -> bool:
return self.defer_async_logging
@_defer_async_logging.setter
def _defer_async_logging(self, value: bool) -> None:
self.defer_async_logging = value
@property
def _enqueue_deferred_logging(self) -> Callable[[], None] | None:
return self.enqueue_deferred_logging
@_enqueue_deferred_logging.setter
def _enqueue_deferred_logging(self, value: Callable[[], None] | None) -> None:
self.enqueue_deferred_logging = value
@property
def _llm_caching_handler(self) -> LLMCachingHandler | None:
return self.llm_caching_handler
@_llm_caching_handler.setter
def _llm_caching_handler(self, value: LLMCachingHandler | None) -> None:
self.llm_caching_handler = value
@property
def _on_detached_stream_failure(self) -> Callable[[Exception], Awaitable[None]] | None:
return self.on_detached_stream_failure
@_on_detached_stream_failure.setter
def _on_detached_stream_failure(self, value: Callable[[Exception], Awaitable[None]] | None) -> None:
self.on_detached_stream_failure = value
def __init__(
self,
model: str,
@ -686,7 +718,7 @@ class Logging(LiteLLMLoggingBaseClass):
# once that response is priced, the same way a non-streamed response is logged
self.client_facing_stream_model: str | None = None
self.zero_cost_warned: bool = False
self._llm_caching_handler: LLMCachingHandler | None = None
self.llm_caching_handler: LLMCachingHandler | None = None
# INITIAL LITELLM_PARAMS
litellm_params = {}
@ -721,10 +753,10 @@ class Logging(LiteLLMLoggingBaseClass):
# Set by proxy request handlers to defer spend-log fire until after
# post_call guardrails have run; the @client decorator then stores the
# enqueue closure here instead of firing it immediately.
self._defer_async_logging: bool = False
self._enqueue_deferred_logging: Callable[[], None] | None = None
self.defer_async_logging: bool = False
self.enqueue_deferred_logging: Callable[[], None] | None = None
self._async_success_scheduled: bool = False
self._on_detached_stream_failure: Callable[[Exception], Awaitable[None]] | None = None
self.on_detached_stream_failure: Callable[[Exception], Awaitable[None]] | None = None
self.shadow_eval_request_snapshot: GuardrailRequestSnapshot | None = None
def set_response_timing_metrics(self, timing_metrics: Mapping[str, float]) -> None:
@ -826,7 +858,7 @@ class Logging(LiteLLMLoggingBaseClass):
trusted-vars channel is the only way credentials reach a per-team logger.
"""
_trusted_var_prefix: Final = "dd_" if callback == "datadog" else "newrelic_" if callback == "newrelic" else None
_custom_logger_init_args: Final[dict | None] = (
_custom_logger_init_args: Final[dict[str, object] | None] = (
{k: v for k, v in self._trusted_callback_vars if k.startswith(_trusted_var_prefix)}
if _trusted_var_prefix is not None
else None
@ -869,8 +901,8 @@ class Logging(LiteLLMLoggingBaseClass):
checks if web_search_options in kwargs or tools and sets the corresponding attribute in StandardBuiltInToolsParams
"""
return StandardBuiltInToolsParams(
web_search_options=StandardBuiltInToolCostTracking._get_web_search_options(kwargs or {}),
file_search=StandardBuiltInToolCostTracking._get_file_search_tool_call(kwargs or {}),
web_search_options=StandardBuiltInToolCostTracking.get_web_search_options(kwargs or {}),
file_search=StandardBuiltInToolCostTracking.get_file_search_tool_call(kwargs or {}),
)
def get_router_model_id(self) -> str | None:
@ -930,7 +962,7 @@ class Logging(LiteLLMLoggingBaseClass):
}
self.litellm_request_debug = litellm_params.get("litellm_request_debug", False)
self.logger_fn = litellm_params.get("logger_fn", None)
if _is_debugging_on() or self.litellm_request_debug:
if is_debugging_on() or self.litellm_request_debug:
verbose_logger.debug("self.optional_params: %s", self.optional_params)
self.model_call_details.update(
@ -1393,12 +1425,12 @@ class Logging(LiteLLMLoggingBaseClass):
additional_args=additional_args,
data=additional_args.get("complete_input_dict", {}),
)
_metadata["raw_request"] = _redact_string(str(curl_command))
_metadata["raw_request"] = redact_string(str(curl_command))
except Exception as e:
self.model_call_details["raw_request_typed_dict"] = RawRequestTypedDict(
error=str(e),
)
_metadata["raw_request"] = _redact_string(
_metadata["raw_request"] = redact_string(
f"Unable to Log \
raw request: {e}"
)
@ -1495,7 +1527,7 @@ class Logging(LiteLLMLoggingBaseClass):
Prints the RAW curl command sent from LiteLLM
"""
if _is_debugging_on() or self.litellm_request_debug:
if is_debugging_on() or self.litellm_request_debug:
if litellm.json_logs:
masked_headers: Final = self._get_masked_headers(headers or {})
masked_api_base: Final = self._get_masked_api_base(str(api_base or ""))
@ -1555,7 +1587,7 @@ class Logging(LiteLLMLoggingBaseClass):
Masks the headers of the request sent from LiteLLM
"""
return _get_masked_values(headers, ignore_sensitive_values=ignore_sensitive_headers)
return get_masked_values(headers, ignore_sensitive_values=ignore_sensitive_headers)
def post_call(self, original_response, input=None, api_key=None, additional_args={}):
# Log the exact result from the LLM API, for streaming - log the type of response received
@ -1800,29 +1832,9 @@ class Logging(LiteLLMLoggingBaseClass):
if margin_total_amount is not None:
self.cost_breakdown["margin_total_amount"] = margin_total_amount
def _response_cost_calculator(
def response_cost_calculator(
self,
result: Union[
ModelResponse,
ModelResponseStream,
EmbeddingResponse,
ImageResponse,
TranscriptionResponse,
TextCompletionResponse,
HttpxBinaryResponseContent,
RerankResponse,
Batch,
FineTuningJob,
ResponsesAPIResponse,
ResponseCompletedEvent,
OpenAIFileObject,
LiteLLMRealtimeStreamLoggingObject,
OpenAIModerationResponse,
"SearchResponse",
DecisionsResponse,
dict,
list,
],
result: object,
cache_hit: bool | None = None,
litellm_model_name: str | None = None,
router_model_id: str | None = None,
@ -1845,13 +1857,12 @@ class Logging(LiteLLMLoggingBaseClass):
return 0.0
transformed_result: Final = self._generate_content_result_as_model_response(result)
if transformed_result is not None:
result = transformed_result
response_result: Final[object] = transformed_result if transformed_result is not None else result
priced_result: Final = (
result.response
if isinstance(result, (ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent))
else result
response_result.response
if isinstance(response_result, (ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent))
else response_result
)
result_hidden_params: Final = getattr(priced_result, "_hidden_params", None) or MappingProxyType({})
@ -1881,18 +1892,28 @@ class Logging(LiteLLMLoggingBaseClass):
## RESPONSE COST ##
custom_pricing: Final = self._custom_pricing_for(priced_result)
prompt = self._prompt_for_cost_calculation()
prompt: Final = self._prompt_for_cost_calculation()
if cache_hit is None:
cache_hit = self.model_call_details.get("cache_hit", False)
model_value: Final = litellm_model_name or self.model
cost_cache_hit_value: Final = (
self.model_call_details.get("cache_hit", False) if cache_hit is None else cache_hit
)
cost_cache_hit: Final[bool | None] = cast( # cast-ok: callback metadata is caller-provided
bool | None, cost_cache_hit_value
)
provider_value: Final = self.model_call_details.get("custom_llm_provider", None)
cost_custom_llm_provider: Final[str | None] = cast( # cast-ok: callback metadata is caller-provided
str | None, provider_value
)
base_model: Final = get_base_model_from_metadata(model_call_details=self.model_call_details)
try:
response_cost_calculator_kwargs: Final = {
"response_object": priced_result,
"model": litellm_model_name or self.model,
"cache_hit": cache_hit,
"custom_llm_provider": self.model_call_details.get("custom_llm_provider", None),
"base_model": _get_base_model_from_metadata(model_call_details=self.model_call_details),
"model": model_value,
"cache_hit": cost_cache_hit,
"custom_llm_provider": cost_custom_llm_provider,
"base_model": base_model,
"call_type": self.call_type,
"optional_params": self.optional_params,
"custom_pricing": custom_pricing,
@ -1906,13 +1927,13 @@ class Logging(LiteLLMLoggingBaseClass):
else None
),
"vertex_location": _resolve_vertex_location_for_cost(
custom_llm_provider=self.model_call_details.get("custom_llm_provider", None),
custom_llm_provider=cost_custom_llm_provider,
litellm_params=(self.litellm_params if hasattr(self, "litellm_params") else None),
optional_params=self.optional_params,
model=litellm_model_name or self.model,
model=model_value,
),
"region_name": _resolve_mantle_region_for_cost(
custom_llm_provider=self.model_call_details.get("custom_llm_provider", None),
custom_llm_provider=cost_custom_llm_provider,
litellm_params=self.model_call_details.get("litellm_params"),
),
}
@ -1946,12 +1967,12 @@ class Logging(LiteLLMLoggingBaseClass):
debug_info = StandardLoggingModelCostFailureDebugInformation(
error_str=str(e),
traceback_str=_get_traceback_str_for_error(str(e)),
model=response_cost_calculator_kwargs["model"],
cache_hit=response_cost_calculator_kwargs["cache_hit"],
custom_llm_provider=response_cost_calculator_kwargs["custom_llm_provider"],
base_model=response_cost_calculator_kwargs["base_model"],
call_type=response_cost_calculator_kwargs["call_type"],
custom_pricing=response_cost_calculator_kwargs["custom_pricing"],
model=model_value,
cache_hit=cost_cache_hit,
custom_llm_provider=cost_custom_llm_provider,
base_model=base_model,
call_type=self.call_type,
custom_pricing=custom_pricing,
)
verbose_logger.debug("response_cost_failure_debug_information: %s", debug_info)
self.model_call_details["response_cost_failure_debug_information"] = debug_info
@ -1965,6 +1986,8 @@ class Logging(LiteLLMLoggingBaseClass):
return None
_response_cost_calculator = response_cost_calculator
def _record_zero_cost_diagnostic(
self,
result: object,
@ -2018,7 +2041,7 @@ class Logging(LiteLLMLoggingBaseClass):
completion_response=result,
custom_llm_provider=custom_llm_provider,
custom_pricing=self._custom_pricing_for(result),
base_model=_get_base_model_from_metadata(model_call_details=self.model_call_details),
base_model=get_base_model_from_metadata(model_call_details=self.model_call_details),
router_model_id=router_model_id or self.get_router_model_id(),
region_name=_resolve_mantle_region_for_cost(
custom_llm_provider=custom_llm_provider,
@ -2117,10 +2140,10 @@ class Logging(LiteLLMLoggingBaseClass):
| FineTuningJob,
cache_hit: bool | None = None,
) -> float | None:
return self._response_cost_calculator(result=result, cache_hit=cache_hit)
return self.response_cost_calculator(result=result, cache_hit=cache_hit)
@staticmethod
def _is_sync_litellm_request(litellm_params: dict) -> bool:
def is_sync_litellm_request(litellm_params: Mapping[str, object]) -> bool:
"""True for sync SDK entrypoints (``completion``), false for async (``acompletion``, etc.)."""
return (
litellm_params.get(CallTypes.acompletion.value, False) is not True
@ -2135,6 +2158,8 @@ class Logging(LiteLLMLoggingBaseClass):
and litellm_params.get(CallTypes.arealtime.value, False) is not True
)
_is_sync_litellm_request = is_sync_litellm_request
def _is_assembled_stream_success(self, result=None) -> bool:
"""Final assembled stream export (not a per-chunk success call).
@ -2176,7 +2201,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["has_dispatched_final_stream_success"] = True
litellm_params: Final = self.model_call_details.get("litellm_params", {}) or {}
sync_sdk: Final = self._is_sync_litellm_request(litellm_params)
sync_sdk: Final = self.is_sync_litellm_request(litellm_params)
passthrough: Final = self.call_type == CallTypes.pass_through.value
if sync_sdk and not prefer_async_handlers and not passthrough:
self.success_handler(
@ -2217,7 +2242,7 @@ class Logging(LiteLLMLoggingBaseClass):
"""Bill a fully streamed response on the failure log when a post-call hook rejects it."""
usage: Final = getattr(assembled, "usage", None)
if isinstance(usage, Usage):
self.record_partial_usage_for_failure(usage, self._response_cost_calculator(result=assembled) or 0.0)
self.record_partial_usage_for_failure(usage, self.response_cost_calculator(result=assembled) or 0.0)
async def dispatch_failure_handlers(
self,
@ -2237,7 +2262,7 @@ class Logging(LiteLLMLoggingBaseClass):
the request failed).
"""
litellm_params: Final = self.model_call_details.get("litellm_params", {}) or {}
sync_sdk: Final = self._is_sync_litellm_request(litellm_params)
sync_sdk: Final = self.is_sync_litellm_request(litellm_params)
passthrough: Final = self.call_type == CallTypes.pass_through.value
if sync_sdk and not prefer_async_handlers and not passthrough:
self.failure_handler(exception, traceback_exception)
@ -2314,10 +2339,12 @@ class Logging(LiteLLMLoggingBaseClass):
return True
def _update_completion_start_time(self, completion_start_time: datetime.datetime):
def update_completion_start_time(self, completion_start_time: datetime.datetime) -> None:
self.completion_start_time = completion_start_time
self.model_call_details["completion_start_time"] = self.completion_start_time
_update_completion_start_time = update_completion_start_time
def normalize_logging_result(self, result: object) -> object:
"""
Some endpoints return a different type of result than what is expected by the logging system.
@ -2434,7 +2461,7 @@ class Logging(LiteLLMLoggingBaseClass):
# Do not preserve 0 from failure_handler on intermediate router retries.
pass
else:
self.model_call_details["response_cost"] = self._response_cost_calculator(result=logging_result)
self.model_call_details["response_cost"] = self.response_cost_calculator(result=logging_result)
if not build_logging_payload:
return
@ -2482,7 +2509,7 @@ class Logging(LiteLLMLoggingBaseClass):
def _transform_usage_objects(self, result):
if isinstance(result, ResponsesAPIResponse):
result = result.model_copy()
transformed_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(result.usage)
transformed_usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(result.usage)
setattr(result, "usage", transformed_usage)
if (standard_logging_payload := self.model_call_details.get("standard_logging_object")) is not None:
response_dict: Final = result.model_dump() if hasattr(result, "model_dump") else dict(result)
@ -2808,7 +2835,7 @@ class Logging(LiteLLMLoggingBaseClass):
standard_logging_object=kwargs.get("standard_logging_object", None),
)
litellm_params = self.model_call_details.get("litellm_params", {})
is_sync_request: Final = self._is_sync_litellm_request(litellm_params)
is_sync_request: Final = self.is_sync_litellm_request(litellm_params)
try:
## BUILD COMPLETE STREAMED RESPONSE
complete_streaming_response: (
@ -2827,7 +2854,7 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug("Logging Details LiteLLM-Success Call streaming complete")
self.model_call_details["complete_streaming_response"] = complete_streaming_response
self._surface_response_headers_from_result(complete_streaming_response)
self.model_call_details["response_cost"] = self._response_cost_calculator(
self.model_call_details["response_cost"] = self.response_cost_calculator(
result=complete_streaming_response
)
self._merge_hidden_params_from_response_into_metadata(complete_streaming_response)
@ -3285,7 +3312,7 @@ class Logging(LiteLLMLoggingBaseClass):
)
elif should_compute_batch_data:
batch_result: Final = await _handle_completed_batch(
batch_result: Final = await handle_completed_batch(
batch=result,
custom_llm_provider=self.custom_llm_provider,
model_name=self.get_deployment_model_for_cost(),
@ -3351,9 +3378,9 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["response_cost"] = 0.0
else:
# check if base_model set on azure
_get_base_model_from_metadata(model_call_details=self.model_call_details)
get_base_model_from_metadata(model_call_details=self.model_call_details)
# base_model defaults to None if not set on model_info
self.model_call_details["response_cost"] = self._response_cost_calculator(
self.model_call_details["response_cost"] = self.response_cost_calculator(
result=complete_streaming_response
)
@ -3576,7 +3603,7 @@ class Logging(LiteLLMLoggingBaseClass):
if self.stream:
if "async_complete_streaming_response" in self.model_call_details:
print_verbose("DynamoDB Logger: Got Stream Event - Completed Stream Response")
await dynamoLogger._async_log_event(
await dynamoLogger.async_log_event(
kwargs=self.model_call_details,
response_obj=self.model_call_details["async_complete_streaming_response"],
start_time=start_time,
@ -3586,7 +3613,7 @@ class Logging(LiteLLMLoggingBaseClass):
else:
print_verbose("DynamoDB Logger: Got Stream Event - No complete stream response as yet")
else:
await dynamoLogger._async_log_event(
await dynamoLogger.async_log_event(
kwargs=self.model_call_details,
response_obj=result,
start_time=start_time,
@ -3613,7 +3640,7 @@ class Logging(LiteLLMLoggingBaseClass):
try:
callback_name: Final = self._get_callback_name(callback)
all_callbacks: Final = litellm.logging_callback_manager._get_all_callbacks()
all_callbacks: Final = litellm.logging_callback_manager.get_all_callbacks()
for callback_obj in all_callbacks:
if hasattr(callback_obj, "increment_callback_logging_failure"):
@ -3642,7 +3669,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["log_event_type"] = "failed_api_call"
self.model_call_details["exception"] = exception
self.model_call_details["traceback_exception"] = (
_redact_string(traceback_exception) if isinstance(traceback_exception, str) else traceback_exception
redact_string(traceback_exception) if isinstance(traceback_exception, str) else traceback_exception
)
self.model_call_details["end_time"] = end_time
self.model_call_details.setdefault("original_response", None)
@ -3667,7 +3694,7 @@ class Logging(LiteLLMLoggingBaseClass):
end_time=end_time,
logging_obj=self,
status="failure",
error_str=_redact_string(str(exception)),
error_str=redact_string(str(exception)),
original_exception=exception,
standard_built_in_tools_params=self.standard_built_in_tools_params,
)
@ -3735,7 +3762,7 @@ class Logging(LiteLLMLoggingBaseClass):
if not self.should_run_logging(event_type="sync_failure"): # prevent double logging
return
litellm_params: Final = self.model_call_details.get("litellm_params", {})
is_sync_request: Final = self._is_sync_litellm_request(litellm_params)
is_sync_request: Final = self.is_sync_litellm_request(litellm_params)
try:
start_time, end_time = self._failure_handler_helper_fn(
@ -3987,7 +4014,7 @@ class Logging(LiteLLMLoggingBaseClass):
self._handle_callback_failure(callback=callback)
await flush_post_call_redis_batches()
def _get_trace_id(self, service_name: Literal["langfuse"]) -> str | None:
def get_trace_id(self, service_name: Literal["langfuse"]) -> str | None:
"""
For the given service (e.g. langfuse), return the trace_id actually logged.
@ -4005,6 +4032,8 @@ class Logging(LiteLLMLoggingBaseClass):
return trace_id
_get_trace_id = get_trace_id
def handle_sync_success_callbacks_for_async_calls(
self,
result: object,
@ -4149,7 +4178,7 @@ class Logging(LiteLLMLoggingBaseClass):
## return unified Usage object
if isinstance(result.response.usage, ResponseAPIUsage):
set_response_cost_in_hidden_params(result.response, result.response.usage.cost)
transformed_usage: Final = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
transformed_usage: Final = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(
result.response.usage
)
# Set as dict instead of Usage object so model_dump() serializes it correctly
@ -4314,11 +4343,11 @@ class Logging(LiteLLMLoggingBaseClass):
model_response: Final = litellm.ModelResponse(id=served_id)
model_response.model = self.model
usage: Final = getattr(result, "usage", None)
if usage is not None and ResponseAPILoggingUtils._is_response_api_usage(usage):
if usage is not None and ResponseAPILoggingUtils.is_response_api_usage(usage):
setattr(
model_response,
"usage",
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage),
ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage),
)
return model_response
@ -4368,15 +4397,15 @@ class Logging(LiteLLMLoggingBaseClass):
return result_copy
def _get_masked_values(
sensitive_object: dict,
def get_masked_values(
sensitive_object: Mapping[str, object],
ignore_sensitive_values: bool = False,
mask_all_values: bool = False,
unmasked_length: int = 4,
number_of_asterisks: int | None = 4,
_depth: int = 0,
_max_depth: int = 20,
) -> dict:
) -> dict[str, object]:
"""
Internal debugging helper function
@ -4386,7 +4415,7 @@ def _get_masked_values(
masked_length: Optional length for the masked portion (number of *). If set, will use exactly this many *
regardless of original string length. The total length will be unmasked_length + masked_length.
"""
sensitive_keywords: Final = [
sensitive_keywords: Final = (
"authorization",
"token",
"key",
@ -4395,14 +4424,17 @@ def _get_masked_values(
"credentials",
"password",
"passwd",
]
)
def _mask_value(v: object) -> object:
if isinstance(v, dict):
typed_value: Final = cast( # cast-ok: request values can contain arbitrary header keys
dict[object, object], v
)
if _depth >= _max_depth:
return v
return _get_masked_values(
v,
return get_masked_values(
cast(dict[str, object], typed_value), # cast-ok: preserve dynamic key behavior
ignore_sensitive_values=ignore_sensitive_values,
mask_all_values=mask_all_values,
unmasked_length=unmasked_length,
@ -4429,7 +4461,13 @@ def _get_masked_values(
}
def set_callbacks(callback_list, function_id=None):
_get_masked_values = get_masked_values
def set_callbacks(
callback_list: Iterable[str | Callable[..., object] | CustomLogger],
function_id: str | None = None,
) -> None:
"""
Globally sets the callback client
"""
@ -4527,14 +4565,14 @@ def set_callbacks(callback_list, function_id=None):
def _init_custom_logger_compatible_class(
logging_integration: _custom_logger_compatible_callbacks_literal,
internal_usage_cache: DualCache | None,
llm_router: object, # expect litellm.Router, but typing errors due to circular import
custom_logger_init_args: dict | None = {},
llm_router: object,
custom_logger_init_args: dict[str, object] | None = {},
) -> CustomLogger | None:
"""
Initialize a custom logger compatible class
"""
try:
custom_logger_init_args = custom_logger_init_args or {}
custom_logger_init_args_value: Final[dict[str, object]] = custom_logger_init_args or {}
if logging_integration == "agentops": # Add AgentOps initialization
_v2 = _maybe_construct_otel_v2("agentops", _in_memory_loggers)
if _v2 is not None:
@ -4624,10 +4662,10 @@ def _init_custom_logger_compatible_class(
return _prometheus_logger
elif logging_integration == "datadog":
# Check if team-scoped credentials are provided
_dd_api_key: Final = custom_logger_init_args.get("dd_api_key")
_dd_site: Final = custom_logger_init_args.get("dd_site")
_dd_agent_host: Final = custom_logger_init_args.get("dd_agent_host")
_dd_agent_port: Final = custom_logger_init_args.get("dd_agent_port")
_dd_api_key: Final = custom_logger_init_args_value.get("dd_api_key")
_dd_site: Final = custom_logger_init_args_value.get("dd_site")
_dd_agent_host: Final = custom_logger_init_args_value.get("dd_agent_host")
_dd_agent_port: Final = custom_logger_init_args_value.get("dd_agent_port")
if _dd_api_key or _dd_site or _dd_agent_host:
# Team-scoped credentials: use DynamicLoggingCache for per-credential isolation
@ -4636,7 +4674,9 @@ def _init_custom_logger_compatible_class(
)
return DataDogHandler.get_datadog_logger_for_request(
standard_callback_dynamic_params=custom_logger_init_args,
standard_callback_dynamic_params=cast( # cast-ok: callback options come from dynamic config
StandardCallbackDynamicParams, custom_logger_init_args_value
),
in_memory_dynamic_logger_cache=in_memory_dynamic_logger_cache,
)
@ -5078,7 +5118,7 @@ def _init_custom_logger_compatible_class(
for callback in _in_memory_loggers:
if isinstance(callback, pagerduty_loggers.pagerduty):
return callback
pagerduty_logger: Final = pagerduty_loggers.pagerduty(**custom_logger_init_args)
pagerduty_logger: Final = pagerduty_loggers.pagerduty(**custom_logger_init_args_value)
_in_memory_loggers.append(pagerduty_logger)
return pagerduty_logger
elif logging_integration == "anthropic_cache_control_hook":
@ -5188,7 +5228,7 @@ def _init_custom_logger_compatible_class(
_in_memory_loggers.append(gitlab_logger)
return gitlab_logger
elif logging_integration == "newrelic":
if custom_logger_init_args.get("newrelic_api_key"):
if custom_logger_init_args_value.get("newrelic_api_key"):
# Team-scoped credentials: per-team METRICS logger, isolated per
# credential set via DynamicLoggingCache. The trace logger for
# this name stays on the global path below.
@ -5197,7 +5237,9 @@ def _init_custom_logger_compatible_class(
)
return NewRelicHandler.get_newrelic_logger_for_request(
standard_callback_dynamic_params=custom_logger_init_args,
standard_callback_dynamic_params=cast( # cast-ok: callback options come from dynamic config
StandardCallbackDynamicParams, custom_logger_init_args_value
),
in_memory_dynamic_logger_cache=in_memory_dynamic_logger_cache,
)
@ -5360,7 +5402,7 @@ def _maybe_auto_initialize_arize_phoenix(_in_memory_loggers: list[CustomLogger])
def get_custom_logger_compatible_class(
logging_integration: _custom_logger_compatible_callbacks_literal,
logging_integration: str,
) -> CustomLogger | None:
try:
if logging_integration == "lago":
@ -5932,10 +5974,10 @@ class StandardLoggingPayloadSetup:
elif isinstance(usage, Usage):
return usage
elif isinstance(usage, ResponseAPIUsage):
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage)
elif isinstance(usage, dict):
if ResponseAPILoggingUtils._is_response_api_usage(usage):
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
if ResponseAPILoggingUtils.is_response_api_usage(usage):
return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage)
if InteractionsUsageObjectTransformation.is_interactions_usage_object(usage):
return InteractionsUsageObjectTransformation.transform_interactions_usage_object(usage)
return Usage(**usage)
@ -5960,10 +6002,10 @@ class StandardLoggingPayloadSetup:
if _raw is None:
return _empty
if isinstance(_raw, ResponseAPIUsage):
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(_raw).model_dump()
return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(_raw).model_dump()
if isinstance(_raw, dict):
if ResponseAPILoggingUtils._is_response_api_usage(_raw):
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(_raw).model_dump()
if ResponseAPILoggingUtils.is_response_api_usage(_raw):
return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(_raw).model_dump()
if InteractionsUsageObjectTransformation.is_interactions_usage_object(_raw):
return InteractionsUsageObjectTransformation.transform_interactions_usage_object(_raw).model_dump()
return _raw
@ -5979,7 +6021,7 @@ class StandardLoggingPayloadSetup:
init_response_obj: object,
api_base: str | None = None,
) -> StandardLoggingModelInformation:
model_cost_name: Final = _select_model_name_for_cost_calc(
model_cost_name: Final = select_model_name_for_cost_calc(
model=base_model if custom_pricing else None,
completion_response=init_response_obj,
base_model=base_model,
@ -6206,8 +6248,8 @@ class StandardLoggingPayloadSetup:
error_code=error_status,
error_class=error_class,
llm_provider=_llm_provider_in_exception,
traceback=_redact_string(traceback_info),
error_message=_redact_string(error_message),
traceback=redact_string(traceback_info),
error_message=redact_string(error_message),
error_provider_request_id=provider_request_id,
error_rate_limit_category=rate_limit_category,
error_rate_limit_type=rate_limit_type,
@ -6331,7 +6373,7 @@ class StandardLoggingPayloadSetup:
return logging_obj.litellm_session_id
@staticmethod
def _get_user_agent_tags(proxy_server_request: dict) -> list[str] | None:
def _get_user_agent_tags(proxy_server_request: Mapping[str, object]) -> list[str] | None:
"""
Return the user agent tags from the proxy server request for spend tracking
"""
@ -6340,8 +6382,11 @@ class StandardLoggingPayloadSetup:
user_agent_tags: list[str] | None = None
headers: Final = proxy_server_request.get("headers", {})
if headers is not None and isinstance(headers, dict):
if "user-agent" in headers:
user_agent: Final = headers["user-agent"]
request_headers: Final = cast( # cast-ok: request headers are untyped framework data
dict[str, str | None], headers
)
if "user-agent" in request_headers:
user_agent: Final = request_headers["user-agent"]
if user_agent is not None:
if user_agent_tags is None:
user_agent_tags = []
@ -6350,12 +6395,11 @@ class StandardLoggingPayloadSetup:
user_agent_part = user_agent.split("/")[0]
if user_agent_part is not None:
user_agent_tags.append("User-Agent: " + user_agent_part)
if user_agent is not None:
user_agent_tags.append("User-Agent: " + user_agent)
user_agent_tags.append("User-Agent: " + user_agent)
return user_agent_tags
@staticmethod
def _get_extra_header_tags(proxy_server_request: dict) -> list[str] | None:
def _get_extra_header_tags(proxy_server_request: Mapping[str, object]) -> list[str] | None:
"""
Extract additional header tags for spend tracking based on config.
"""
@ -6366,24 +6410,33 @@ class StandardLoggingPayloadSetup:
headers: Final = proxy_server_request.get("headers", {})
if not isinstance(headers, dict):
return None
request_headers: Final = cast( # cast-ok: request headers are untyped framework data
dict[str, str], headers
)
header_tags: Final = []
for header_name in extra_headers:
header_value = headers.get(header_name)
header_value = request_headers.get(header_name)
if header_value:
header_tags.append(f"{header_name}: {header_value}")
return header_tags if header_tags else None
@staticmethod
def _get_request_tags(litellm_params: dict, proxy_server_request: dict) -> list[str]:
# check for 'tags' in both 'metadata' and 'litellm_metadata'
metadata: Final = litellm_params.get("metadata") or {}
litellm_metadata: Final = litellm_params.get("litellm_metadata") or {}
def get_request_tags(
litellm_params: dict[str, object],
proxy_server_request: dict[str, object],
) -> list[str]:
metadata: Final = cast( # cast-ok: request metadata is caller-provided
Mapping[str, object], litellm_params.get("metadata") or {}
)
litellm_metadata: Final = cast( # cast-ok: request metadata is caller-provided
Mapping[str, object], litellm_params.get("litellm_metadata") or {}
)
if metadata.get("tags", []):
request_tags = metadata.get("tags", []).copy()
request_tags = cast(list[str], metadata.get("tags", [])).copy() # cast-ok: tags are caller-provided
elif litellm_metadata.get("tags", []):
request_tags = litellm_metadata.get("tags", []).copy()
request_tags = cast(list[str], litellm_metadata.get("tags", [])).copy() # cast-ok: tags are caller-provided
else:
request_tags = []
user_agent_tags: Final = StandardLoggingPayloadSetup._get_user_agent_tags(proxy_server_request)
@ -6394,6 +6447,8 @@ class StandardLoggingPayloadSetup:
request_tags.extend(additional_header_tags)
return request_tags
_get_request_tags = get_request_tags
def _get_status_fields(
status: StandardLoggingPayloadStatus,
@ -6575,7 +6630,7 @@ def get_standard_logging_object_payload(
_model_id: Final = metadata.get("model_info", {}).get("id", "")
_model_group: Final = metadata.get("model_group", "")
request_tags: Final = StandardLoggingPayloadSetup._get_request_tags(
request_tags: Final = StandardLoggingPayloadSetup.get_request_tags(
litellm_params=litellm_params, proxy_server_request=proxy_server_request
)
request_model_access_groups: Final = request_model_access_groups_from_litellm_params(litellm_params)
@ -6620,7 +6675,7 @@ def get_standard_logging_object_payload(
if cache_hit is True:
id = f"{id}_cache_hit{time.time()}" # do not duplicate the request id
saved_cache_cost = (
logging_obj._response_cost_calculator(
logging_obj.response_cost_calculator(
result=init_response_obj,
cache_hit=False,
)
@ -6628,7 +6683,7 @@ def get_standard_logging_object_payload(
)
## Get model cost information ##
base_model = _get_base_model_from_metadata(model_call_details=kwargs)
base_model = get_base_model_from_metadata(model_call_details=kwargs)
# The router overrides completion_response.model to the model-group alias before
# this payload is built, so cost-map lookup via that alias always misses.
# Fall back to the actual deployment model set by the router in metadata.
@ -6681,7 +6736,9 @@ def get_standard_logging_object_payload(
# Reconstruct full model name with provider prefix for logging
# This ensures Bedrock models like "us.anthropic.claude-3-5-sonnet-20240620-v1:0"
# are logged as "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0"
custom_llm_provider: Final = cast(str | None, kwargs.get("custom_llm_provider"))
custom_llm_provider: Final = cast( # cast-ok: provider name is caller-provided
str | None, kwargs.get("custom_llm_provider")
)
model_name = reconstruct_model_name(kwargs.get("model", "") or "", custom_llm_provider, metadata)
response_model_name: str | None = None
if isinstance(final_response_obj, dict):
@ -6933,7 +6990,7 @@ def _get_traceback_str_for_error(error_str: str) -> str:
from decimal import Decimal
# used for unit testing
from typing import Any, Union
from typing import Any
def create_dummy_standard_logging_payload() -> StandardLoggingPayload:

View file

@ -3,7 +3,7 @@ Helper utilities for tracking the cost of built-in tools.
"""
from collections.abc import Mapping
from typing import Final, Literal
from typing import Final, Literal, cast
from pydantic import ValidationError
@ -739,9 +739,10 @@ class StandardBuiltInToolCostTracking:
return False
@staticmethod
def _get_web_search_options(kwargs: dict) -> WebSearchOptions | None:
def get_web_search_options(kwargs: Mapping[str, object]) -> WebSearchOptions | None:
if "web_search_options" in kwargs:
return WebSearchOptions(**kwargs.get("web_search_options", {}))
web_search_options: Final = cast(WebSearchOptions, kwargs.get("web_search_options", {}))
return WebSearchOptions(**web_search_options)
tools: Final = StandardBuiltInToolCostTracking._get_tools_from_kwargs(
kwargs=kwargs, tool_type="web_search_preview"
@ -751,27 +752,31 @@ class StandardBuiltInToolCostTracking:
for tool in tools:
if isinstance(tool, dict):
if StandardBuiltInToolCostTracking._is_web_search_tool_call(tool):
return WebSearchOptions(**tool)
return WebSearchOptions(**cast(WebSearchOptions, tool))
return None
_get_web_search_options = get_web_search_options
@staticmethod
def _get_tools_from_kwargs(kwargs: dict, tool_type: str) -> list[dict] | None:
def _get_tools_from_kwargs(kwargs: Mapping[str, object], tool_type: str) -> list[object] | None:
if "tools" in kwargs:
return kwargs.get("tools", [])
return cast(list[object], kwargs.get("tools", []))
return None
@staticmethod
def _get_file_search_tool_call(kwargs: dict) -> FileSearchTool | None:
def get_file_search_tool_call(kwargs: Mapping[str, object]) -> FileSearchTool | None:
tools: Final = StandardBuiltInToolCostTracking._get_tools_from_kwargs(kwargs, "file_search")
if tools:
for tool in tools:
if isinstance(tool, dict):
if StandardBuiltInToolCostTracking._is_file_search_tool_call(tool):
return FileSearchTool(**tool)
return FileSearchTool(**cast(FileSearchTool, tool))
return None
_get_file_search_tool_call = get_file_search_tool_call
@staticmethod
def _is_web_search_tool_call(tool: dict) -> bool:
def _is_web_search_tool_call(tool: Mapping[str, object]) -> bool:
if tool.get("type", None) == "web_search_preview":
return True
if tool.get("type", None) == "web_search":
@ -781,7 +786,7 @@ class StandardBuiltInToolCostTracking:
return False
@staticmethod
def _is_file_search_tool_call(tool: dict) -> bool:
def _is_file_search_tool_call(tool: Mapping[str, object]) -> bool:
if tool.get("type", None) == "file_search":
return True
return False

View file

@ -153,12 +153,15 @@ def get_web_search_requests_from_usage(usage: Usage) -> int | None:
return get_web_search_requests(getattr(usage, "server_tool_use", None))
def _is_above_128k(tokens: float) -> bool:
def is_above_128k(tokens: float) -> bool:
if tokens > 128000:
return True
return False
_is_above_128k = is_above_128k
def get_billable_input_tokens(usage: Usage) -> int:
"""
Returns the number of billable input tokens.
@ -185,7 +188,7 @@ def select_cost_metric_for_model(
)
def _generic_cost_per_character(
def generic_cost_per_character(
model: str,
custom_llm_provider: str,
prompt_characters: float,
@ -250,7 +253,10 @@ def _generic_cost_per_character(
return prompt_cost, completion_cost
def _get_service_tier_cost_key(base_key: str, service_tier: str | None) -> str:
_generic_cost_per_character = generic_cost_per_character
def get_service_tier_cost_key(base_key: str, service_tier: str | None) -> str:
"""
Get the appropriate cost key based on service tier.
@ -271,6 +277,9 @@ def _get_service_tier_cost_key(base_key: str, service_tier: str | None) -> str:
return f"{base_key}_{suffix}"
_get_service_tier_cost_key = get_service_tier_cost_key
def _parse_token_threshold(threshold: str) -> float:
return float(threshold.replace("k", "")) * (1000 if "k" in threshold else 1)
@ -386,7 +395,7 @@ def _get_tiered_base_costs(model_info: ModelInfo, usage: Usage) -> tuple[float,
completion_cost: Final = (
tier_rate(tier, "output_cost_per_token")
if "output_cost_per_token" in tier
else _get_cost_per_unit(model_info, "output_cost_per_token") or 0.0
else get_cost_per_unit(model_info, "output_cost_per_token") or 0.0
)
return (
tier_rate(tier, "input_cost_per_token"),
@ -644,30 +653,34 @@ def _get_token_base_cost(
return _apply_off_peak_to_base_costs(model_info, current_time, tiered_base_costs)
# Get service tier aware cost keys
input_cost_key: Final = _get_service_tier_cost_key("input_cost_per_token", service_tier)
output_cost_key: Final = _get_service_tier_cost_key("output_cost_per_token", service_tier)
cache_creation_cost_key: Final = _get_service_tier_cost_key("cache_creation_input_token_cost", service_tier)
cache_read_cost_key: Final = _get_service_tier_cost_key("cache_read_input_token_cost", service_tier)
input_cost_key: Final = get_service_tier_cost_key("input_cost_per_token", service_tier)
output_cost_key: Final = get_service_tier_cost_key("output_cost_per_token", service_tier)
cache_creation_cost_key: Final = get_service_tier_cost_key("cache_creation_input_token_cost", service_tier)
cache_read_cost_key: Final = get_service_tier_cost_key("cache_read_input_token_cost", service_tier)
prompt_base_cost = cast(float, _get_cost_per_unit(model_info, input_cost_key))
completion_base_cost = cast(float, _get_cost_per_unit(model_info, output_cost_key))
prompt_base_cost = cast( # cast-ok: model pricing data is external
float, get_cost_per_unit(model_info, input_cost_key)
)
completion_base_cost = cast( # cast-ok: model pricing data is external
float, get_cost_per_unit(model_info, output_cost_key)
)
# For image generation models that don't have output_cost_per_token,
# use output_cost_per_image_token as the base cost (all output tokens are image tokens)
if completion_base_cost == 0.0 or completion_base_cost is None:
output_image_cost: Final = _get_cost_per_unit(model_info, "output_cost_per_image_token", None)
output_image_cost: Final = get_cost_per_unit(model_info, "output_cost_per_image_token", None)
if output_image_cost is not None:
completion_base_cost = cast(float, output_image_cost)
cache_creation_cost = _get_cost_per_unit(model_info, cache_creation_cost_key, default_value=None)
cache_creation_cost_above_1hr = _get_cost_per_unit(
completion_base_cost = output_image_cost
cache_creation_cost = get_cost_per_unit(model_info, cache_creation_cost_key, default_value=None)
cache_creation_cost_above_1hr = get_cost_per_unit(
model_info, "cache_creation_input_token_cost_above_1hr", default_value=None
)
cache_read_cost = _get_cost_per_unit(model_info, cache_read_cost_key, default_value=None)
cache_read_cost = get_cost_per_unit(model_info, cache_read_cost_key, default_value=None)
## CHECK IF ABOVE THRESHOLD
# Optimization: collect threshold keys first to avoid sorting all model_info keys.
# Standard thresholds and thresholds suffixed for this request's service tier both count.
tier_key_suffix: Final = _get_service_tier_cost_key("", service_tier)
tier_key_suffix: Final = get_service_tier_cost_key("", service_tier)
threshold_keys: Final = [
k
for k in model_info
@ -692,7 +705,7 @@ def _get_token_base_cost(
# ON_DEMAND_PRIORITY. Falls back to the standard key automatically
# via _get_cost_per_unit's service_tier fallback logic.
tiered_input_key = (
_get_service_tier_cost_key(
get_service_tier_cost_key(
f"input_cost_per_token_above_{threshold_str}_tokens",
service_tier,
)
@ -701,10 +714,10 @@ def _get_token_base_cost(
)
prompt_base_cost = cast(
float,
_get_cost_per_unit(model_info, tiered_input_key, prompt_base_cost),
get_cost_per_unit(model_info, tiered_input_key, prompt_base_cost),
)
tiered_output_key = (
_get_service_tier_cost_key(
get_service_tier_cost_key(
f"output_cost_per_token_above_{threshold_str}_tokens",
service_tier,
)
@ -713,7 +726,7 @@ def _get_token_base_cost(
)
completion_base_cost = cast(
float,
_get_cost_per_unit(
get_cost_per_unit(
model_info,
tiered_output_key,
completion_base_cost,
@ -722,7 +735,7 @@ def _get_token_base_cost(
# Apply tiered pricing to cache costs
cache_creation_tiered_key = (
_get_service_tier_cost_key(
get_service_tier_cost_key(
f"cache_creation_input_token_cost_above_{threshold_str}_tokens",
service_tier,
)
@ -730,7 +743,7 @@ def _get_token_base_cost(
else f"cache_creation_input_token_cost_above_{threshold_str}_tokens"
)
cache_creation_1hr_tiered_key = (
_get_service_tier_cost_key(
get_service_tier_cost_key(
f"cache_creation_input_token_cost_above_1hr_above_{threshold_str}_tokens",
service_tier,
)
@ -738,7 +751,7 @@ def _get_token_base_cost(
else f"cache_creation_input_token_cost_above_1hr_above_{threshold_str}_tokens"
)
cache_read_tiered_key = (
_get_service_tier_cost_key(
get_service_tier_cost_key(
f"cache_read_input_token_cost_above_{threshold_str}_tokens",
service_tier,
)
@ -746,13 +759,13 @@ def _get_token_base_cost(
else f"cache_read_input_token_cost_above_{threshold_str}_tokens"
)
cache_creation_cost = _get_cost_per_unit(model_info, cache_creation_tiered_key, cache_creation_cost)
cache_creation_cost = get_cost_per_unit(model_info, cache_creation_tiered_key, cache_creation_cost)
cache_creation_cost_above_1hr = _get_cost_per_unit(
cache_creation_cost_above_1hr = get_cost_per_unit(
model_info, cache_creation_1hr_tiered_key, cache_creation_cost_above_1hr
)
cache_read_cost = _get_cost_per_unit(model_info, cache_read_tiered_key, cache_read_cost)
cache_read_cost = get_cost_per_unit(model_info, cache_read_tiered_key, cache_read_cost)
break
except (IndexError, ValueError):
@ -795,13 +808,13 @@ def calculate_cost_component(model_info: ModelInfo, cost_key: str, usage_value:
Returns:
float: The calculated cost
"""
cost_per_unit: Final = _get_cost_per_unit(model_info, cost_key)
cost_per_unit: Final = get_cost_per_unit(model_info, cost_key)
if cost_per_unit is not None and isinstance(cost_per_unit, float) and usage_value is not None and usage_value > 0:
return float(usage_value) * cost_per_unit
return 0.0
def _get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: float | None = 0.0) -> float | None:
def get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: float | None = 0.0) -> float | None:
# Sometimes the cost per unit is a string (e.g.: If a value like "3e-7" was read from the config.yaml)
cost_per_unit: Final = model_info.get(cost_key)
if isinstance(cost_per_unit, float):
@ -834,7 +847,7 @@ def _get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: floa
return float(fallback_cost)
except ValueError:
verbose_logger.exception(
"litellm.litellm_core_utils.llm_cost_calc.utils.py::_get_cost_per_unit(): Exception occured - %s\nDefaulting to 0.0",
"litellm.litellm_core_utils.llm_cost_calc.utils.py::get_cost_per_unit(): Exception occured - %s\nDefaulting to 0.0",
fallback_cost,
)
break # Only try the first matching suffix
@ -842,6 +855,9 @@ def _get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: floa
return default_value
_get_cost_per_unit = get_cost_per_unit
def deployment_pricing(model_info: ModelInfo | None) -> ModelInfo | None:
"""The prices a deployment sets itself, as floats; None when it sets none that parse."""
if model_info is None:
@ -851,7 +867,7 @@ def deployment_pricing(model_info: ModelInfo | None) -> ModelInfo | None:
{
key: price
for key in priced_keys
if (price := _get_cost_per_unit(model_info, key, default_value=None)) is not None
if (price := get_cost_per_unit(model_info, key, default_value=None)) is not None
}
)
if not pricing:
@ -868,7 +884,7 @@ def flat_image_cost(model_info: ModelInfo | None, image_response: ImageResponse)
"""The per-image price times the images returned; 0.0 when the table sets no per-image price."""
if model_info is None:
return 0.0
output_cost_per_image: Final = _get_cost_per_unit(model_info, "output_cost_per_image", default_value=None) or 0.0
output_cost_per_image: Final = get_cost_per_unit(model_info, "output_cost_per_image", default_value=None) or 0.0
num_images: Final = len(image_response.data) if image_response.data else 0
return output_cost_per_image * num_images
@ -1086,9 +1102,9 @@ def _calculate_input_cost(
### CACHE READ COST - Now uses tiered pricing
cache_hit_audio_tokens: Final = prompt_tokens_details["cache_hit_audio_tokens"]
audio_cache_read_rate: Final = _get_cost_per_unit(
audio_cache_read_rate: Final = get_cost_per_unit(
model_info,
_get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier),
get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier),
None,
)
prompt_cost += float(prompt_tokens_details["cache_hit_tokens"] - cache_hit_audio_tokens) * cache_read_cost
@ -1100,7 +1116,7 @@ def _calculate_input_cost(
if prompt_tokens_details["audio_tokens"] and not (
prompt_tokens_details["audio_length_seconds"] and model_info.get("input_cost_per_audio_per_second") is not None
):
audio_cost_key: Final = _get_service_tier_cost_key("input_cost_per_audio_token", service_tier)
audio_cost_key: Final = get_service_tier_cost_key("input_cost_per_audio_token", service_tier)
prompt_cost += calculate_cost_component(model_info, audio_cost_key, prompt_tokens_details["audio_tokens"])
### IMAGE TOKEN COST
@ -1173,7 +1189,7 @@ def _calculate_input_cost(
return prompt_cost
def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str | None) -> float:
def get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str | None) -> float:
"""
Resolve the per-model regional-processing uplift multiplier for a given
data-residency region.
@ -1195,7 +1211,7 @@ def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str |
if multiplier is None:
return 1.0
try:
return float(cast(float, multiplier))
return float(multiplier)
except (TypeError, ValueError):
verbose_logger.exception(
"Invalid regional_processing_uplift_multiplier_%s for model; defaulting to 1.0",
@ -1204,6 +1220,9 @@ def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str |
return 1.0
_get_regional_uplift_multiplier = get_regional_uplift_multiplier
def get_vertex_regional_endpoint_uplift(model_info: ModelInfo, vertex_location: str | None) -> float:
"""
Resolve the per-model uplift multiplier for Vertex AI non-global (regional and
@ -1223,7 +1242,7 @@ def get_vertex_regional_endpoint_uplift(model_info: ModelInfo, vertex_location:
if multiplier is None:
return 1.0
try:
return float(cast(float, multiplier))
return float(cast(float, multiplier)) # cast-ok: pricing multiplier is external model data
except (TypeError, ValueError):
verbose_logger.exception(
"Invalid regional_endpoint_uplift_multiplier for model; defaulting to 1.0",
@ -1253,15 +1272,15 @@ def _resolve_reasoning_token_cost(
service_tier: str | None,
completion_base_cost: float,
) -> float:
tier_reasoning_key: Final = _get_service_tier_cost_key("output_cost_per_reasoning_token", service_tier)
tier_reasoning_key: Final = get_service_tier_cost_key("output_cost_per_reasoning_token", service_tier)
if model_info.get(tier_reasoning_key) is not None:
tier_reasoning_cost: Final = _get_cost_per_unit(model_info, tier_reasoning_key, None)
tier_reasoning_cost: Final = get_cost_per_unit(model_info, tier_reasoning_key, None)
if tier_reasoning_cost is not None:
return tier_reasoning_cost
tier_output_key: Final = _get_service_tier_cost_key("output_cost_per_token", service_tier)
tier_output_key: Final = get_service_tier_cost_key("output_cost_per_token", service_tier)
if tier_output_key != "output_cost_per_token" and model_info.get(tier_output_key) is not None:
return completion_base_cost
standard_reasoning_cost: Final = _get_cost_per_unit(model_info, "output_cost_per_reasoning_token", None)
standard_reasoning_cost: Final = get_cost_per_unit(model_info, "output_cost_per_reasoning_token", None)
return standard_reasoning_cost if standard_reasoning_cost is not None else completion_base_cost
@ -1444,7 +1463,7 @@ def generic_cost_per_token(
## AUDIO COST
if not is_text_tokens_total and audio_tokens is not None and audio_tokens > 0:
_output_cost_per_audio_token = _get_cost_per_unit(resolved_model_info, "output_cost_per_audio_token", None)
_output_cost_per_audio_token = get_cost_per_unit(resolved_model_info, "output_cost_per_audio_token", None)
_output_cost_per_audio_token = (
_output_cost_per_audio_token if _output_cost_per_audio_token is not None else completion_base_cost
)
@ -1462,7 +1481,7 @@ def generic_cost_per_token(
## IMAGE COST
if not is_text_tokens_total and image_tokens and image_tokens > 0:
_output_cost_per_image_token = _get_cost_per_unit(resolved_model_info, "output_cost_per_image_token", None)
_output_cost_per_image_token = get_cost_per_unit(resolved_model_info, "output_cost_per_image_token", None)
_output_cost_per_image_token = (
_output_cost_per_image_token if _output_cost_per_image_token is not None else completion_base_cost
)
@ -1470,7 +1489,7 @@ def generic_cost_per_token(
## VIDEO COST
if not is_text_tokens_total and video_tokens and video_tokens > 0:
_output_cost_per_video_token = _get_cost_per_unit(resolved_model_info, "output_cost_per_video_token", None)
_output_cost_per_video_token = get_cost_per_unit(resolved_model_info, "output_cost_per_video_token", None)
_output_cost_per_video_token = (
_output_cost_per_video_token if _output_cost_per_video_token is not None else completion_base_cost
)
@ -1479,7 +1498,7 @@ def generic_cost_per_token(
## REGIONAL DATA-RESIDENCY UPLIFT
# Applied as a flat multiplier across all token costs for the request
# when the upstream is a regionalized OpenAI host (eu./us.api.openai.com).
uplift: Final = _get_regional_uplift_multiplier(resolved_model_info, data_residency)
uplift: Final = get_regional_uplift_multiplier(resolved_model_info, data_residency)
if uplift != 1.0:
prompt_cost *= uplift
completion_cost *= uplift
@ -1603,13 +1622,13 @@ def _cost_map_billed_rates(
completion_base_cost=completion_base_cost,
current_time=billing_time,
)
audio_cache_read_rate: Final = _get_cost_per_unit(
audio_cache_read_rate: Final = get_cost_per_unit(
model_info,
_get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier),
get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier),
None,
)
multiplier: Final = (
_get_regional_uplift_multiplier(model_info, data_residency)
get_regional_uplift_multiplier(model_info, data_residency)
* get_vertex_regional_endpoint_uplift(model_info, vertex_location)
* get_provider_specific_geo_multiplier(model_info=model_info, usage=usage)
)
@ -1764,7 +1783,7 @@ def calculate_prompt_caching_savings(
cache_creation_cost_above_1hr=write_rate_1h - prompt_base_cost,
cache_creation_cost=write_rate - prompt_base_cost,
)
uplift: Final = _get_regional_uplift_multiplier(model_info, data_residency) * get_vertex_regional_endpoint_uplift(
uplift: Final = get_regional_uplift_multiplier(model_info, data_residency) * get_vertex_regional_endpoint_uplift(
model_info, vertex_location
)
return (read_discount - write_premium) * uplift
@ -1899,7 +1918,7 @@ def calculate_image_response_web_search_cost(
class CostCalculatorUtils:
@staticmethod
def _call_type_has_image_response(call_type: str) -> bool:
def call_type_has_image_response(call_type: str) -> bool:
"""
Returns True if the call type has an image response
@ -1910,6 +1929,8 @@ class CostCalculatorUtils:
"""
return call_type in _IMAGE_RESPONSE_CALL_TYPES
_call_type_has_image_response = call_type_has_image_response
@staticmethod
def route_image_generation_cost_calculator(
model: str,

View file

@ -1,5 +1,5 @@
from collections.abc import Mapping
from typing import Final
from typing import Final, cast
import litellm
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
@ -106,7 +106,7 @@ def serialize_multipart_form_fields(data: Mapping[str, object]) -> tuple[tuple[s
)
def _ensure_extra_body_is_safe(extra_body: dict | None) -> dict | None:
def ensure_extra_body_is_safe(extra_body: dict[str, object] | None) -> dict[str, object] | None:
"""
Ensure that the extra_body sent in the request is safe, otherwise users will see this error
@ -117,22 +117,28 @@ def _ensure_extra_body_is_safe(extra_body: dict | None) -> dict | None:
"""
if extra_body is None:
return None
if not isinstance(extra_body, dict):
return extra_body
if "metadata" in extra_body and isinstance(extra_body["metadata"], dict):
if "prompt" in extra_body["metadata"]:
_prompt: Final = extra_body["metadata"].get("prompt")
if "prompt" in cast(dict[str, object], extra_body["metadata"]):
prompt: Final = cast( # cast-ok: request metadata is caller-provided
dict[str, object], extra_body["metadata"]
).get("prompt")
# users can send Langfuse TextPromptClient objects, so we need to convert them to dicts
# Langfuse TextPromptClients have .__dict__ attribute
if _prompt is not None and hasattr(_prompt, "__dict__"):
extra_body["metadata"]["prompt"] = _prompt.__dict__
if prompt is not None and hasattr(prompt, "__dict__"):
cast(dict[str, object], extra_body["metadata"])["prompt"] = (
cast( # cast-ok: prompt is an external SDK object
object, getattr(prompt, "__dict__")
)
)
return extra_body
_ensure_extra_body_is_safe = ensure_extra_body_is_safe
def pick_cheapest_chat_models_from_llm_provider(custom_llm_provider: str, n=1):
"""
Pick the n cheapest chat models from the LLM provider.

View file

@ -10,7 +10,7 @@ import litellm
from litellm._logging import verbose_logger
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.litellm_core_utils.prompt_templates.common_utils import (
_extract_reasoning_content,
extract_reasoning_content,
)
from litellm.types.llms.databricks import DatabricksTool
from litellm.types.llms.openai import (
@ -71,7 +71,7 @@ def _normalize_images_for_message(
return normalized
def _safe_convert_created_field(created_value) -> int:
def safe_convert_created_field(created_value: object) -> int:
"""
Safely convert a 'created' field value to an integer.
@ -91,19 +91,20 @@ def _safe_convert_created_field(created_value) -> int:
elif isinstance(created_value, float):
return int(created_value)
else:
# for strings, etc
try:
return int(float(created_value))
return int(float(cast(float | str, created_value)))
except (ValueError, TypeError):
# Fallback to current time if conversion fails
return int(time.time())
_safe_convert_created_field = safe_convert_created_field
def convert_tool_call_to_json_mode(
tool_calls: list[ChatCompletionMessageToolCall],
convert_tool_call_to_json_mode: bool,
) -> tuple[Message | None, str | None]:
if _should_convert_tool_call_to_json_mode(
if should_convert_tool_call_to_json_mode(
tool_calls=tool_calls,
convert_tool_call_to_json_mode=convert_tool_call_to_json_mode,
):
@ -245,7 +246,7 @@ async def convert_to_streaming_response_async(
model_response_object.id = response_object["id"]
if "created" in response_object:
model_response_object.created = _safe_convert_created_field(response_object["created"])
model_response_object.created = safe_convert_created_field(response_object["created"])
if "system_fingerprint" in response_object:
model_response_object.system_fingerprint = response_object["system_fingerprint"]
@ -334,7 +335,7 @@ def convert_to_streaming_response(
model_response_object.id = response_object["id"]
if "created" in response_object:
model_response_object.created = _safe_convert_created_field(response_object["created"])
model_response_object.created = safe_convert_created_field(response_object["created"])
if "system_fingerprint" in response_object:
model_response_object.system_fingerprint = response_object["system_fingerprint"]
@ -371,9 +372,9 @@ def convert_to_streaming_response(
from collections import defaultdict
def _handle_invalid_parallel_tool_calls(
def handle_invalid_parallel_tool_calls(
tool_calls: list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall],
):
) -> list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall] | None:
"""
Handle hallucinated parallel tool call from openai - https://community.openai.com/t/model-tries-to-call-unknown-function-multi-tool-use-parallel/490653
@ -414,6 +415,9 @@ def _handle_invalid_parallel_tool_calls(
return tool_calls
_handle_invalid_parallel_tool_calls = handle_invalid_parallel_tool_calls
class LiteLLMResponseObjectHandler:
@staticmethod
def convert_to_image_response(
@ -530,7 +534,7 @@ class LiteLLMResponseObjectHandler:
return transformed_logprobs
def _should_convert_tool_call_to_json_mode(
def should_convert_tool_call_to_json_mode(
tool_calls: (
Sequence[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall] | Sequence[DatabricksTool] | None
) = None,
@ -545,6 +549,9 @@ def _should_convert_tool_call_to_json_mode(
return False
_should_convert_tool_call_to_json_mode = should_convert_tool_call_to_json_mode
def convert_to_model_response_object(
response_object: dict | None = None,
model_response_object: ModelResponse
@ -645,14 +652,14 @@ def convert_to_model_response_object(
for _tc in tool_calls:
_openai_tc = chat_completion_tool_call_from_dict(_tc)
_openai_tool_calls.append(_openai_tc)
fixed_tool_calls = _handle_invalid_parallel_tool_calls(_openai_tool_calls)
fixed_tool_calls = handle_invalid_parallel_tool_calls(_openai_tool_calls)
if fixed_tool_calls is not None:
tool_calls = fixed_tool_calls
message: Message | None = None
finish_reason: str | None = None
if tool_calls is not None and _should_convert_tool_call_to_json_mode(
if tool_calls is not None and should_convert_tool_call_to_json_mode(
tool_calls=tool_calls,
convert_tool_call_to_json_mode=convert_tool_call_to_json_mode,
):
@ -669,7 +676,7 @@ def convert_to_model_response_object(
provider_specific_fields[f] = choice["message"][f]
# Handle reasoning models that display `reasoning_content` within `content`
reasoning_content, content = _extract_reasoning_content(choice["message"])
reasoning_content, content = extract_reasoning_content(choice["message"])
# Handle thinking models that display `thinking_blocks` within `content`
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = (
@ -718,7 +725,7 @@ def convert_to_model_response_object(
usage_object: Final = litellm.Usage(**response_object["usage"])
setattr(model_response_object, "usage", usage_object)
if "created" in response_object:
model_response_object.created = _safe_convert_created_field(response_object["created"])
model_response_object.created = safe_convert_created_field(response_object["created"])
if "id" in response_object:
# Preserve the auto-generated id from ModelResponse.__init__

View file

@ -125,7 +125,7 @@ class ResponseMetadata:
"litellm_call_id": getattr(logging_obj, "litellm_call_id", None),
"api_base": get_api_base(model=model or "", optional_params=kwargs),
"model_id": model_id,
"response_cost": logging_obj._response_cost_calculator(
"response_cost": logging_obj.response_cost_calculator(
result=self.result, litellm_model_name=model, router_model_id=model_id
),
"additional_headers": process_response_headers(

View file

@ -38,7 +38,7 @@ class LoggingCallbackManager:
except Exception:
return False
def add_litellm_input_callback(self, callback: CustomLogger | str | Callable):
def add_litellm_input_callback(self, callback: CustomLogger | str | Callable[..., object]):
"""
Add a input callback to litellm.input_callback.
Auto-routes async callbacks to litellm._async_input_callback.
@ -48,13 +48,13 @@ class LoggingCallbackManager:
else:
self._safe_add_callback_to_list(callback=callback, parent_list=litellm.input_callback)
def add_litellm_service_callback(self, callback: CustomLogger | str | Callable):
def add_litellm_service_callback(self, callback: CustomLogger | str | Callable[..., object]):
"""
Add a service callback to litellm.service_callback
"""
self._safe_add_callback_to_list(callback=callback, parent_list=litellm.service_callback)
def add_litellm_callback(self, callback: CustomLogger | str | Callable):
def add_litellm_callback(self, callback: CustomLogger | str | Callable[..., object]):
"""
Add a callback to litellm.callbacks
@ -65,7 +65,7 @@ class LoggingCallbackManager:
parent_list=litellm.callbacks,
)
def add_litellm_success_callback(self, callback: CustomLogger | str | Callable):
def add_litellm_success_callback(self, callback: CustomLogger | str | Callable[..., object]):
"""
Add a success callback to `litellm.success_callback`.
Auto-routes async callbacks to litellm._async_success_callback.
@ -81,7 +81,7 @@ class LoggingCallbackManager:
else:
self._safe_add_callback_to_list(callback=callback, parent_list=litellm.success_callback)
def add_litellm_failure_callback(self, callback: CustomLogger | str | Callable):
def add_litellm_failure_callback(self, callback: CustomLogger | str | Callable[..., object]):
"""
Add a failure callback to `litellm.failure_callback`.
Auto-routes async callbacks to litellm._async_failure_callback.
@ -91,13 +91,13 @@ class LoggingCallbackManager:
else:
self._safe_add_callback_to_list(callback=callback, parent_list=litellm.failure_callback)
def add_litellm_async_success_callback(self, callback: CustomLogger | Callable | str):
def add_litellm_async_success_callback(self, callback: CustomLogger | Callable[..., object] | str):
"""
Add a success callback to litellm._async_success_callback
"""
self._safe_add_callback_to_list(callback=callback, parent_list=litellm._async_success_callback)
def add_litellm_async_failure_callback(self, callback: CustomLogger | Callable | str):
def add_litellm_async_failure_callback(self, callback: CustomLogger | Callable[..., object] | str):
"""
Add a failure callback to litellm._async_failure_callback
"""
@ -139,7 +139,9 @@ class LoggingCallbackManager:
for c in remove_list:
callback_list.remove(c)
def _add_string_callback_to_list(self, callback: str, parent_list: list[CustomLogger | Callable | str]):
def _add_string_callback_to_list(
self, callback: str, parent_list: list[CustomLogger | Callable[..., object] | str]
):
"""
Add a string callback to a list, if the callback is already in the list, do not add it again.
"""
@ -148,7 +150,7 @@ class LoggingCallbackManager:
else:
verbose_logger.debug("Callback %s already exists in %s, not adding again..", callback, parent_list)
def _check_callback_list_size(self, parent_list: list[CustomLogger | Callable | str]) -> bool:
def _check_callback_list_size(self, parent_list: list[CustomLogger | Callable[..., object] | str]) -> bool:
"""
Check if adding another callback would exceed MAX_CALLBACKS
Returns True if safe to add, False if would exceed limit
@ -163,7 +165,7 @@ class LoggingCallbackManager:
return True
@staticmethod
def _add_custom_callback_generic_api_str(
def add_custom_callback_generic_api_str(
callback: str,
) -> GenericAPILogger | str:
"""
@ -244,10 +246,12 @@ class LoggingCallbackManager:
return callback
_add_custom_callback_generic_api_str = add_custom_callback_generic_api_str
def _safe_add_callback_to_list(
self,
callback: CustomLogger | Callable | str,
parent_list: list[CustomLogger | Callable | str],
callback: CustomLogger | Callable[..., object] | str,
parent_list: list[CustomLogger | Callable[..., object] | str],
):
"""
Safe add a callback to a list, if the callback is already in the list, do not add it again.
@ -261,7 +265,7 @@ class LoggingCallbackManager:
# Check if the callback is a custom callback
if isinstance(callback, str):
callback = LoggingCallbackManager._add_custom_callback_generic_api_str(callback)
callback = LoggingCallbackManager.add_custom_callback_generic_api_str(callback)
if isinstance(callback, str):
self._add_string_callback_to_list(callback=callback, parent_list=parent_list)
@ -274,7 +278,11 @@ class LoggingCallbackManager:
elif callable(callback):
self._add_callback_function_to_list(callback=callback, parent_list=parent_list)
def _add_callback_function_to_list(self, callback: Callable, parent_list: list[CustomLogger | Callable | str]):
def _add_callback_function_to_list(
self,
callback: Callable[..., object],
parent_list: list[CustomLogger | Callable[..., object] | str],
):
"""
Add a callback function to a list, if the callback is already in the list, do not add it again.
"""
@ -289,7 +297,7 @@ class LoggingCallbackManager:
def _add_custom_logger_to_list(
self,
custom_logger: CustomLogger,
parent_list: list[CustomLogger | Callable | str],
parent_list: list[CustomLogger | Callable[..., object] | str],
):
"""
Add a custom logger to a list, if another instance of the same custom logger exists in the list, do not add it again.
@ -341,7 +349,7 @@ class LoggingCallbackManager:
litellm._async_failure_callback = []
litellm.callbacks = []
def _get_all_callbacks(self) -> list[CustomLogger | Callable | str]:
def get_all_callbacks(self) -> list[CustomLogger | Callable[..., object] | str]:
"""
Get all callbacks from litellm.callbacks, litellm.success_callback, litellm.failure_callback, litellm._async_success_callback, litellm._async_failure_callback
"""
@ -353,6 +361,8 @@ class LoggingCallbackManager:
+ litellm._async_failure_callback
)
_get_all_callbacks = get_all_callbacks
def remove_callback_from_all_lists(self, obj, require_self=False) -> None:
"""
Remove a callback object from every callback list it may have been
@ -379,7 +389,7 @@ class LoggingCallbackManager:
Returns:
Set[CustomLogger]: Set of custom loggers that are instances of the given class type
"""
all_callbacks: Final = self._get_all_callbacks()
all_callbacks: Final = self.get_all_callbacks()
matched_callbacks: Final[set[AdditionalLoggingUtils]] = set()
for callback in all_callbacks:
if isinstance(callback, CustomLogger) and isinstance(callback, AdditionalLoggingUtils):
@ -392,7 +402,7 @@ class LoggingCallbackManager:
"""
# ensure we don't have duplicate instances
all_callbacks: Final = []
for callback in self._get_all_callbacks():
for callback in self.get_all_callbacks():
if isinstance(callback, callback_type) and callback not in all_callbacks:
all_callbacks.append(callback)
return all_callbacks
@ -401,7 +411,7 @@ class LoggingCallbackManager:
"""
Returns True if any of the active callbacks are of the given type
"""
return any(isinstance(callback, callback_type) for callback in self._get_all_callbacks())
return any(isinstance(callback, callback_type) for callback in self.get_all_callbacks())
def get_callbacks_by_type(self) -> CallbacksByType:
"""
@ -444,11 +454,11 @@ class LoggingCallbackManager:
def get_callback_objects(self) -> tuple[tuple[str, CustomLogger | Callable], ...]:
return tuple(
(self._get_callback_string(callback), callback)
for callback in self._get_all_callbacks()
for callback in self.get_all_callbacks()
if not isinstance(callback, str)
)
def _get_callback_string(self, callback: CustomLogger | Callable | str) -> str:
def _get_callback_string(self, callback: CustomLogger | Callable[..., object] | str) -> str:
from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.litellm_core_utils.custom_logger_registry import (
CustomLoggerRegistry,

View file

@ -5,7 +5,7 @@ import re
import time
from collections.abc import Iterator, Mapping, Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Any, Final, cast
from litellm._logging import format_base64_size, verbose_logger
from litellm.constants import (
@ -199,10 +199,10 @@ def _get_parent_otel_span_from_logging_obj(
# Reuse existing function by passing model_call_details as kwargs
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
get_parent_otel_span_from_kwargs,
)
return _get_parent_otel_span_from_kwargs(logging_obj.model_call_details)
return get_parent_otel_span_from_kwargs(logging_obj.model_call_details)
except Exception as e:
verbose_logger.exception("Error in _get_parent_otel_span_from_logging_obj: %s", e)
@ -227,14 +227,14 @@ def convert_litellm_response_object_to_str(
return None
def _assemble_complete_response_from_streaming_chunks(
def assemble_complete_response_from_streaming_chunks(
result: ModelResponse | TextCompletionResponse | ModelResponseStream,
start_time: datetime,
end_time: datetime,
request_kwargs: dict,
streaming_chunks: list[Any],
request_kwargs: Mapping[str, object],
streaming_chunks: list[object],
is_async: bool,
):
) -> ModelResponse | TextCompletionResponse | None:
"""
Assemble a complete response from a streaming chunks
@ -262,9 +262,10 @@ def _assemble_complete_response_from_streaming_chunks(
if result.choices[0].finish_reason is not None: # if it's the last chunk
streaming_chunks.append(result)
try:
messages: Final = cast(list[dict[str, object]] | None, request_kwargs.get("messages", None))
complete_streaming_response = litellm.stream_chunk_builder(
chunks=streaming_chunks,
messages=request_kwargs.get("messages", None),
messages=messages,
start_time=start_time,
end_time=end_time,
)
@ -279,6 +280,9 @@ def _assemble_complete_response_from_streaming_chunks(
return complete_streaming_response
_assemble_complete_response_from_streaming_chunks = assemble_complete_response_from_streaming_chunks
def _set_duration_in_model_call_details(
logging_obj: Any, # we're not guaranteed this will be `LiteLLMLoggingObject`
start_time: datetime,

View file

@ -1,5 +1,5 @@
from functools import lru_cache
from typing import Final
from typing import ClassVar, Final
from openai.types.chat.completion_create_params import (
CompletionCreateParamsNonStreaming,
@ -24,7 +24,7 @@ from litellm.types.rerank import RerankRequest
class ModelParamHelper:
# Cached at class level — deterministic set built from static OpenAI type annotations
_relevant_logging_args: frozenset = frozenset()
relevant_logging_args: ClassVar[frozenset[str]] = frozenset()
@staticmethod
def get_standard_logging_model_parameters(
@ -32,7 +32,7 @@ class ModelParamHelper:
) -> dict:
""" """
standard_logging_model_parameters: Final[dict] = {}
supported_model_parameters: Final = ModelParamHelper._relevant_logging_args
supported_model_parameters: Final = ModelParamHelper.relevant_logging_args
for key, value in model_parameters.items():
if key in supported_model_parameters:
@ -44,20 +44,22 @@ class ModelParamHelper:
return set(["messages", "prompt", "input", "system"])
@staticmethod
def _get_relevant_args_to_use_for_logging() -> set[str]:
def get_relevant_args_to_use_for_logging() -> set[str]:
"""
Gets all relevant llm api params besides the ones with prompt content
"""
all_openai_llm_api_params: Final = ModelParamHelper._get_all_llm_api_params()
all_openai_llm_api_params: Final = ModelParamHelper.get_all_llm_api_params()
# Exclude parameters that contain prompt content
combined_kwargs: Final = all_openai_llm_api_params.difference(
set(ModelParamHelper.get_exclude_params_for_model_parameters())
)
return combined_kwargs
_get_relevant_args_to_use_for_logging = get_relevant_args_to_use_for_logging
@staticmethod
@lru_cache(maxsize=1)
def _get_all_llm_api_params() -> set[str]:
def get_all_llm_api_params() -> set[str]:
"""
Gets the supported kwargs for each call type and combines them.
@ -88,6 +90,8 @@ class ModelParamHelper:
combined_kwargs = combined_kwargs.difference(exclude_kwargs)
return combined_kwargs
_get_all_llm_api_params = get_all_llm_api_params
@staticmethod
def get_litellm_provider_specific_params_for_chat_params() -> set[str]:
return set(["thinking"])
@ -185,4 +189,5 @@ class ModelParamHelper:
return set(["metadata", "litellm_metadata"])
ModelParamHelper._relevant_logging_args = frozenset(ModelParamHelper._get_relevant_args_to_use_for_logging())
ModelParamHelper.relevant_logging_args = frozenset(ModelParamHelper.get_relevant_args_to_use_for_logging())
ModelParamHelper._relevant_logging_args = ModelParamHelper.relevant_logging_args

View file

@ -271,7 +271,7 @@ def request_contains_image_content(messages: Sequence[Mapping[str, object]]) ->
)
def _audio_or_image_in_message_content(message: AllMessageValues) -> bool:
def audio_or_image_in_message_content(message: AllMessageValues) -> bool:
"""
Checks if message content contains an image or audio
"""
@ -284,6 +284,9 @@ def _audio_or_image_in_message_content(message: AllMessageValues) -> bool:
return False
_audio_or_image_in_message_content = audio_or_image_in_message_content
def convert_openai_message_to_only_content_messages(
messages: list[AllMessageValues],
) -> list[dict[str, str]]:
@ -1503,7 +1506,7 @@ def tool_with_sanitized_parameters(
return tool if sanitized_schema is input_schema else {**tool, "input_schema": sanitized_schema}
def _get_image_mime_type_from_url(url: str) -> str | None:
def get_image_mime_type_from_url(url: str) -> str | None:
"""
Get mime type for common image URLs
See gemini mime types: https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/image-understanding#image-requirements
@ -1568,6 +1571,9 @@ def _get_image_mime_type_from_url(url: str) -> str | None:
return None
_get_image_mime_type_from_url = get_image_mime_type_from_url
def infer_content_type_from_url_and_content(
url: str,
content: bytes,
@ -1968,7 +1974,11 @@ def convert_prefix_message_to_non_prefix_messages(
return new_messages
def _extract_reasoning_content(message: dict) -> tuple[str | None, str | None]:
def _provider_text_or_none(value: object) -> str | None:
return cast(str | None, value) # cast-ok: reasoning fields arrive in untyped provider messages
def extract_reasoning_content(message: Mapping[str, object]) -> tuple[str | None, str | None]:
"""
Extract reasoning content and main content from a message.
@ -1980,12 +1990,15 @@ def _extract_reasoning_content(message: dict) -> tuple[str | None, str | None]:
"""
message_content: Final = message.get("content")
if "reasoning_content" in message:
return message["reasoning_content"], message_content
return _provider_text_or_none(message["reasoning_content"]), _provider_text_or_none(message_content)
elif "reasoning" in message:
return message["reasoning"], message_content
return _provider_text_or_none(message["reasoning"]), _provider_text_or_none(message_content)
elif isinstance(message_content, str):
return _parse_content_for_reasoning(message_content)
return None, message_content
return parse_content_for_reasoning(message_content)
return None, _provider_text_or_none(message_content)
_extract_reasoning_content = extract_reasoning_content
def _readable_thinking_text(block: Mapping[str, object]) -> str:
@ -2163,7 +2176,7 @@ def responses_reasoning_items_from_thinking_blocks(
)
def _parse_content_for_reasoning(
def parse_content_for_reasoning(
message_text: str | None,
) -> tuple[str | None, str | None]:
"""
@ -2188,6 +2201,9 @@ def _parse_content_for_reasoning(
return None, message_text
_parse_content_for_reasoning = parse_content_for_reasoning
def _extract_base64_data(image_url: str) -> str:
"""
Extract pure base64 data from an image URL.

View file

@ -6,7 +6,7 @@ import json
import mimetypes
import re
import xml.etree.ElementTree as ET
from collections.abc import Iterator, Mapping, Sequence
from collections.abc import Container, Iterator, Mapping, Sequence
from enum import Enum
from types import MappingProxyType
from typing import Any, Final, TypeAlias, TypedDict, cast, overload
@ -459,7 +459,7 @@ async def _afetch_and_extract_template(
Returns: (chat_template, bos_token, eos_token)
"""
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
_extract_token_value,
extract_token_value,
)
bos_token = ""
@ -481,8 +481,8 @@ async def _afetch_and_extract_template(
and "chat_template" in tokenizer_config["tokenizer"]
):
tokenizer_data: dict = tokenizer_config["tokenizer"]
bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token"))
eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token"))
bos_token = extract_token_value(token_value=tokenizer_data.get("bos_token"))
eos_token = extract_token_value(token_value=tokenizer_data.get("eos_token"))
chat_template = tokenizer_data["chat_template"]
else:
# Fallback: Try to fetch chat template from separate .jinja file
@ -496,8 +496,8 @@ async def _afetch_and_extract_template(
and isinstance(tokenizer_config["tokenizer"], dict)
):
tokenizer_data: dict = tokenizer_config["tokenizer"]
bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token"))
eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token"))
bos_token = extract_token_value(token_value=tokenizer_data.get("bos_token"))
eos_token = extract_token_value(token_value=tokenizer_data.get("eos_token"))
else:
raise Exception("No chat template found")
@ -513,7 +513,7 @@ def _fetch_and_extract_template(
Returns: (chat_template, bos_token, eos_token)
"""
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
_extract_token_value,
extract_token_value,
)
bos_token = ""
@ -535,8 +535,8 @@ def _fetch_and_extract_template(
and "chat_template" in tokenizer_config["tokenizer"]
):
tokenizer_data: dict = tokenizer_config["tokenizer"]
bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token"))
eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token"))
bos_token = extract_token_value(token_value=tokenizer_data.get("bos_token"))
eos_token = extract_token_value(token_value=tokenizer_data.get("eos_token"))
chat_template = tokenizer_data["chat_template"]
else:
# Fallback: Try to fetch chat template from separate .jinja file
@ -550,8 +550,8 @@ def _fetch_and_extract_template(
and isinstance(tokenizer_config["tokenizer"], dict)
):
tokenizer_data: dict = tokenizer_config["tokenizer"]
bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token"))
eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token"))
bos_token = extract_token_value(token_value=tokenizer_data.get("bos_token"))
eos_token = extract_token_value(token_value=tokenizer_data.get("eos_token"))
else:
raise Exception("No chat template found")
@ -561,8 +561,8 @@ def _fetch_and_extract_template(
async def ahf_chat_template(model: str, messages: list, chat_template: str | None = None):
"""HuggingFace chat template (async version)"""
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
_aget_chat_template_file,
_aget_tokenizer_config,
aget_chat_template_file,
aget_tokenizer_config,
strftime_now,
)
@ -573,8 +573,8 @@ async def ahf_chat_template(model: str, messages: list, chat_template: str | Non
template, bos_token, eos_token = await _afetch_and_extract_template(
model=model,
chat_template=chat_template,
get_config_fn=_aget_tokenizer_config,
get_template_fn=_aget_chat_template_file,
get_config_fn=aget_tokenizer_config,
get_template_fn=aget_chat_template_file,
)
return _render_chat_template(
env=env,
@ -588,8 +588,8 @@ async def ahf_chat_template(model: str, messages: list, chat_template: str | Non
def hf_chat_template(model: str, messages: list, chat_template: str | None = None):
"""HuggingFace chat template (sync version)"""
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
_get_chat_template_file,
_get_tokenizer_config,
get_chat_template_file,
get_tokenizer_config,
strftime_now,
)
@ -600,8 +600,8 @@ def hf_chat_template(model: str, messages: list, chat_template: str | None = Non
template, bos_token, eos_token = _fetch_and_extract_template(
model=model,
chat_template=chat_template,
get_config_fn=_get_tokenizer_config,
get_template_fn=_get_chat_template_file,
get_config_fn=get_tokenizer_config,
get_template_fn=get_chat_template_file,
)
return _render_chat_template(
env=env,
@ -1161,7 +1161,7 @@ def _gemini_tool_call_invoke_helper(
return function_call
def _encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: str | None) -> str:
def encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: str | None) -> str:
"""
Embed thought signature into tool call ID for OpenAI client compatibility.
@ -1180,7 +1180,10 @@ def _encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: st
return tool_call_id
def _get_thought_signature_from_tool(tool: dict) -> str | None:
_encode_tool_call_id_with_signature = encode_tool_call_id_with_signature
def get_thought_signature_from_tool(tool: Mapping[str, object]) -> str | None:
"""Extract thought signature from tool call's provider_specific_fields.
If not provided try to extract thought signature from tool call id
@ -1192,26 +1195,41 @@ def _get_thought_signature_from_tool(tool: dict) -> str | None:
# First check tool's provider_specific_fields
provider_fields: Final = tool.get("provider_specific_fields") or {}
if isinstance(provider_fields, dict):
signature = provider_fields.get("thought_signature")
if signature:
return signature
typed_provider_fields: Final = cast( # cast-ok: preserve dynamic provider response fields
dict[str, object], provider_fields
)
signature_from_tool_fields: Final = typed_provider_fields.get("thought_signature")
if signature_from_tool_fields:
return cast(str, signature_from_tool_fields) # cast-ok: untyped provider response field
# Then check function's provider_specific_fields
function: Final = tool.get("function")
if function:
if isinstance(function, dict):
func_provider_fields: Final = function.get("provider_specific_fields") or {}
function_dict: Final = cast( # cast-ok: preserve dynamic provider response fields
dict[str, object], function
)
func_provider_fields: Final = function_dict.get("provider_specific_fields") or {}
if isinstance(func_provider_fields, dict):
signature = func_provider_fields.get("thought_signature")
if signature:
return signature
elif hasattr(function, "provider_specific_fields") and function.provider_specific_fields:
if isinstance(function.provider_specific_fields, dict):
signature = function.provider_specific_fields.get("thought_signature")
if signature:
return signature
typed_func_provider_fields: Final = cast( # cast-ok: preserve dynamic provider response fields
dict[str, object], func_provider_fields
)
signature_from_function_fields: Final = typed_func_provider_fields.get("thought_signature")
if signature_from_function_fields:
return cast(str, signature_from_function_fields) # cast-ok: untyped provider response field
elif hasattr(function, "provider_specific_fields") and getattr(function, "provider_specific_fields"):
function_provider_fields: Final[object] = getattr(function, "provider_specific_fields")
if isinstance(function_provider_fields, dict):
typed_function_provider_fields: Final = cast( # cast-ok: provider fields are dynamic
dict[str, object], function_provider_fields
)
signature_from_model_fields: Final = typed_function_provider_fields.get("thought_signature")
if signature_from_model_fields:
return cast(str, signature_from_model_fields) # cast-ok: untyped provider response field
# Check if thought signature is embedded in tool call ID
tool_call_id: Final = tool.get("id")
tool_call_id: Final = cast( # cast-ok: tool IDs come from model responses
str, tool.get("id")
)
if tool_call_id and THOUGHT_SIGNATURE_SEPARATOR in tool_call_id:
parts: Final = tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)
if len(parts) == 2:
@ -1220,6 +1238,9 @@ def _get_thought_signature_from_tool(tool: dict) -> str | None:
return None
_get_thought_signature_from_tool = get_thought_signature_from_tool
def _get_dummy_thought_signature() -> str:
"""Generate a dummy thought signature for models that require it.
@ -1301,7 +1322,7 @@ def convert_to_gemini_tool_call_invoke(
)
if gemini_function_call is not None:
part_dict: VertexPartType = {"function_call": gemini_function_call}
thought_signature = _get_thought_signature_from_tool(dict(tool))
thought_signature = get_thought_signature_from_tool(dict(tool))
# Gemini signs only the first functionCall part of a parallel batch, so scope the
# placeholder fallback to that part instead of fabricating one per sibling call:
# https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/thinking/thought-signatures#parallel_function_calling_example
@ -1921,7 +1942,7 @@ def anthropic_infer_file_id_content_type(
def anthropic_process_openai_file_message(
message: ChatCompletionFileObject,
) -> AnthropicMessagesDocumentParam | AnthropicMessagesImageParam | AnthropicMessagesContainerUploadParam:
file_message: Final = cast(ChatCompletionFileObject, message)
file_message: Final = message
file_sub: Final = file_message.get("file")
if file_sub is None:
raise litellm.BadRequestError(
@ -3383,7 +3404,7 @@ def _parse_content_type(content_type: str) -> str:
return m.get_content_type()
def _parse_mime_type(base64_data: str) -> str | None:
def parse_mime_type(base64_data: str) -> str | None:
mime_type_match: Final = re.match(r"data:(.*?);base64", base64_data)
if mime_type_match:
return mime_type_match.group(1)
@ -3391,6 +3412,9 @@ def _parse_mime_type(base64_data: str) -> str | None:
return None
_parse_mime_type = parse_mime_type
class BedrockImageProcessor:
"""Handles both sync and async image processing for Bedrock conversations."""
@ -3461,7 +3485,7 @@ class BedrockImageProcessor:
return img_without_base_64, mime_type, image_format
@staticmethod
def _validate_format(mime_type: str, image_format: str) -> str:
def validate_format(mime_type: str, image_format: str) -> str:
"""Validate image format and mime type for both images and documents."""
supported_image_formats: Final = litellm.AmazonConverseConfig().get_supported_image_types()
@ -3488,6 +3512,8 @@ class BedrockImageProcessor:
)
return image_format
_validate_format = validate_format
@staticmethod
def _get_document_format(mime_type: str, supported_doc_formats: list[str]) -> str:
"""
@ -3598,7 +3624,7 @@ class BedrockImageProcessor:
mime_type = format
image_format = mime_type.split("/")[1]
image_format = cls._validate_format(mime_type, image_format)
image_format = cls.validate_format(mime_type, image_format)
return cls._create_bedrock_block(img_bytes, mime_type, image_format)
@classmethod
@ -3617,7 +3643,7 @@ class BedrockImageProcessor:
mime_type = format
image_format = mime_type.split("/")[1]
image_format = cls._validate_format(mime_type, image_format)
image_format = cls.validate_format(mime_type, image_format)
return cls._create_bedrock_block(img_bytes, mime_type, image_format)
@ -3889,7 +3915,7 @@ def _convert_to_bedrock_tool_call_result(
tool_result: Final = BedrockToolResultBlock(content=tool_result_content_blocks, toolUseId=id)
if used_search_results:
tool_result["status"] = cast(Literal["success"], "success")
tool_result["status"] = "success"
content_block: Final = BedrockContentBlock(toolResult=tool_result)
@ -4082,7 +4108,7 @@ def _insert_assistant_continue_message(
)
)
elif litellm.modify_params:
text = convert_content_list_to_str(cast(ChatCompletionAssistantMessage, DEFAULT_ASSISTANT_CONTINUE_MESSAGE))
text = convert_content_list_to_str(DEFAULT_ASSISTANT_CONTINUE_MESSAGE)
messages.append(
BedrockMessageBlock(
role="assistant",
@ -4407,14 +4433,14 @@ class BedrockConverseMessagesProcessor:
_parts.append(_part)
elif element["type"] == "file":
_part = await BedrockConverseMessagesProcessor._async_process_file_message(
message=cast(ChatCompletionFileObject, element)
message=element
)
_parts.append(_part)
elif element["type"] == "document":
_part = BedrockConverseMessagesProcessor._process_document_message(element)
_parts.append(_part)
_cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block(
message_block=cast(OpenAIMessageContentListBlock, element),
message_block=element,
block_type="content_block",
model=model,
)
@ -4528,7 +4554,7 @@ class BedrockConverseMessagesProcessor:
if isinstance(element, dict):
if element["type"] == "thinking":
thinking_block = BedrockConverseMessagesProcessor.translate_thinking_blocks_to_reasoning_content_blocks(
thinking_blocks=[cast(ChatCompletionThinkingBlock, element)]
thinking_blocks=[element]
)
assistants_parts = (
BedrockConverseMessagesProcessor.add_thinking_blocks_to_assistant_content(
@ -4611,7 +4637,7 @@ class BedrockConverseMessagesProcessor:
return reasoning_content_blocks
@staticmethod
def _process_file_message(message: ChatCompletionFileObject) -> BedrockContentBlock:
def process_file_message(message: ChatCompletionFileObject) -> BedrockContentBlock:
file_message: Final = message.get("file")
if file_message is None:
raise litellm.BadRequestError(
@ -4631,6 +4657,8 @@ class BedrockConverseMessagesProcessor:
format: Final = file_message.get("format")
return BedrockImageProcessor.process_image_sync(image_url=cast(str, file_id or file_data), format=format)
_process_file_message = process_file_message
@staticmethod
async def _async_process_file_message(
message: ChatCompletionFileObject,
@ -4669,7 +4697,7 @@ class BedrockConverseMessagesProcessor:
)
media_type: Final[str] = source["media_type"]
data: Final[str] = source["data"]
doc_format = BedrockImageProcessor._validate_format(mime_type=media_type, image_format=media_type.split("/")[1])
doc_format = BedrockImageProcessor.validate_format(mime_type=media_type, image_format=media_type.split("/")[1])
# Deterministic name using the same hashing pattern as _create_bedrock_block
HASH_SAMPLE_BYTES: Final = 64 * 1024
@ -4780,15 +4808,13 @@ def _bedrock_converse_messages_pt(
)
_parts.append(_part)
elif element["type"] == "file":
_part = BedrockConverseMessagesProcessor._process_file_message(
message=cast(ChatCompletionFileObject, element)
)
_part = BedrockConverseMessagesProcessor.process_file_message(message=element)
_parts.append(_part)
elif element["type"] == "document":
_part = BedrockConverseMessagesProcessor._process_document_message(element)
_parts.append(_part)
_cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block(
message_block=cast(OpenAIMessageContentListBlock, element),
message_block=element,
block_type="content_block",
model=model,
)
@ -4905,7 +4931,7 @@ def _bedrock_converse_messages_pt(
if element["type"] == "thinking":
thinking_block = (
BedrockConverseMessagesProcessor.translate_thinking_blocks_to_reasoning_content_blocks(
thinking_blocks=[cast(ChatCompletionThinkingBlock, element)]
thinking_blocks=[element]
)
)
assistants_parts = (
@ -5166,18 +5192,33 @@ def _bedrock_tools_pt(tools: list, model: str | None = None) -> list[BedrockTool
# Function call template
def function_call_prompt(messages: list, functions: list):
def function_call_prompt(
messages: list[dict[str, object]],
functions: list[object],
) -> list[dict[str, object]]:
function_prompt = """Produce JSON OUTPUT ONLY! Adhere to this format {"name": "function_name", "arguments":{"argument_name": "argument_value"}} The following functions are available to you:"""
for function in functions:
function_prompt += f"""\n{function}\n"""
def _append_function_prompt(message: dict[str, object]) -> bool:
role: Final = cast( # cast-ok: preserve dynamic role membership behavior
Container[object], message["role"]
)
if "system" not in role:
return False
content: Final = message["content"]
if isinstance(content, str):
message["content"] = f"{content} {function_prompt}"
else:
cast( # cast-ok: preserve dynamic content append behavior
list[object], content
).append({"type": "text", "text": f""" {function_prompt}"""})
return True
function_added_to_prompt = False
for message in messages:
if "system" in message["role"]:
if isinstance(message["content"], str):
message["content"] += f""" {function_prompt}"""
else:
message["content"].append({"type": "text", "text": f""" {function_prompt}"""})
if _append_function_prompt(message):
function_added_to_prompt = True
if function_added_to_prompt is False:

View file

@ -38,7 +38,7 @@ def strftime_now(fmt: str) -> str:
return datetime.now().strftime(fmt)
def _get_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult:
def get_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult:
"""
Fetch tokenizer_config.json from HuggingFace (sync)
@ -61,7 +61,10 @@ def _get_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult:
return {"status": "failure"}
async def _aget_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult:
_get_tokenizer_config = get_tokenizer_config
async def aget_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult:
"""
Fetch tokenizer_config.json from HuggingFace (async)
@ -86,7 +89,10 @@ async def _aget_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult:
return {"status": "failure"}
def _get_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult:
_aget_tokenizer_config = aget_tokenizer_config
def get_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult:
"""
Fetch chat template from separate .jinja file (sync)
@ -114,7 +120,10 @@ def _get_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult:
return {"status": "failure"}
async def _aget_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult:
_get_chat_template_file = get_chat_template_file
async def aget_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult:
"""
Fetch chat template from separate .jinja file (async)
@ -144,7 +153,10 @@ async def _aget_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResul
return {"status": "failure"}
def _extract_token_value(token_value: None | str | dict[str, Any]) -> str:
_aget_chat_template_file = aget_chat_template_file
def extract_token_value(token_value: None | str | dict[str, Any]) -> str:
"""
Extract token string from various formats (string, dict, etc.)
@ -159,3 +171,6 @@ def _extract_token_value(token_value: None | str | dict[str, Any]) -> str:
if isinstance(token_value, dict):
return token_value.get("content", "")
return ""
_extract_token_value = extract_token_value

View file

@ -150,7 +150,7 @@ class RealTimeStreaming:
self.tool_calls: list[dict] = []
# Detect whether the client is explicitly opting into the beta protocol.
self._client_wants_beta = self._detect_beta_header(websocket)
self._client_wants_beta = self.detect_beta_header(websocket)
self._backend_uses_beta_protocol = (
self._client_wants_beta if backend_uses_beta_protocol is None else backend_uses_beta_protocol
)
@ -178,7 +178,7 @@ class RealTimeStreaming:
self._pending_guardrail_message: str | None = None
# Track whether session.created has already been sent to the client
# (e.g. synthetic event in deferred setup mode).
self._session_created_sent_to_client: bool = False
self.session_created_sent_to_client: bool = False
# Track whether we have already sent the guardrail turn-detection update
# that disables provider auto-response for transcription guardrails.
self._guardrail_turn_detection_update_sent: bool = False
@ -199,6 +199,14 @@ class RealTimeStreaming:
# Optional per-provider GA event normalizer (e.g. XAIRealtimeNormalizer).
self._event_normalizer = event_normalizer
@property
def _session_created_sent_to_client(self) -> bool:
return self.session_created_sent_to_client
@_session_created_sent_to_client.setter
def _session_created_sent_to_client(self, value: bool) -> None:
self.session_created_sent_to_client = value
# Per-connection caps for pre-setup audio frames (message count + total bytes).
_MAX_BUFFERED_MESSAGES: int = 200
_MAX_BUFFERED_BYTES: int = 10 * 1024 * 1024 # 10 MB
@ -252,7 +260,7 @@ class RealTimeStreaming:
# TypedDict union members do not narrow to plain dict for mypy.
message_obj: dict[str, Any] = cast(dict[str, Any], message)
else:
message_obj = cast(dict[str, Any], json.loads(cast(str, message)))
message_obj = cast(dict[str, Any], json.loads(message))
self._collect_tool_calls_from_response_done(cast(dict, message_obj))
if not self._should_store_message(message_obj):
return
@ -457,7 +465,7 @@ class RealTimeStreaming:
if self._content_sent_after_setup:
verbose_logger.debug("Dropping follow-up setup after content was already sent to backend")
continue
msg = self._maybe_inject_guardrail_auto_response_disable(msg)
msg = self.maybe_inject_guardrail_auto_response_disable(msg)
await self.backend_ws.send(msg)
self._cache_session_configuration_request(msg)
sent = True
@ -734,7 +742,7 @@ class RealTimeStreaming:
if sent:
self._guardrail_turn_detection_update_sent = True
def _maybe_inject_guardrail_auto_response_disable(self, setup_message: str) -> str:
def maybe_inject_guardrail_auto_response_disable(self, setup_message: str) -> str:
"""Fold the transcription-guardrail auto-response disable into the setup.
Gemini/Vertex Live reject a second ``setup`` (1007), so the guardrail's
@ -764,6 +772,8 @@ class RealTimeStreaming:
)
return json.dumps(obj)
_maybe_inject_guardrail_auto_response_disable = maybe_inject_guardrail_auto_response_disable
def _has_realtime_guardrails_for_event_hooks(
self,
event_hooks: Sequence["GuardrailEventHooks"],
@ -981,7 +991,7 @@ class RealTimeStreaming:
break
finally:
self._flushing_pending_messages_until_setup = False
if self._session_created_sent_to_client:
if self.session_created_sent_to_client:
# A synthetic session.created (with placeholder defaults) was
# already forwarded to the client when we connected. The
# provider's real session.created (e.g. emitted from Gemini
@ -991,7 +1001,7 @@ class RealTimeStreaming:
# configuration without seeing two `session.created` events.
event = {**event, "type": "session.updated"}
else:
self._session_created_sent_to_client = True
self.session_created_sent_to_client = True
event_str = json.dumps(event)
## For audio/VAD guardrail path: forward the (possibly retyped)
## session.created first, then invoke the one-time guardrail
@ -1154,7 +1164,7 @@ class RealTimeStreaming:
self.logging_obj.model_call_details[REALTIME_SESSION_FAILURE_LOGGED_KEY] = True
@staticmethod
def _detect_beta_header(websocket: ScopedWebSocket) -> bool:
def detect_beta_header(websocket: ScopedWebSocket) -> bool:
"""Return True if the client sent 'OpenAI-Beta: realtime=v1'.
Checks the raw ASGI scope headers so it works for both FastAPI WebSocket
@ -1173,6 +1183,8 @@ class RealTimeStreaming:
pass
return False
_detect_beta_header = detect_beta_header
@staticmethod
def _remap_beta_session_to_ga(session: dict) -> dict:
"""
@ -1591,4 +1603,4 @@ class RealTimeStreaming:
def client_sent_openai_beta_realtime_header(websocket: ScopedWebSocket) -> bool:
"""True when the client WebSocket includes ``OpenAI-Beta: realtime=v1``."""
return RealTimeStreaming._detect_beta_header(websocket)
return RealTimeStreaming.detect_beta_header(websocket)

View file

@ -88,13 +88,16 @@ def _build_secret_patterns() -> "re.Pattern[str]":
_SECRET_RE: Final = _build_secret_patterns()
def _python_redact_string(value: str) -> str:
def python_redact_string(value: str) -> str:
return _SECRET_RE.sub(REDACTED, value)
_python_redact_string = python_redact_string
def redact_string(value: str) -> str:
"""Scrub known secret/credential patterns from *value* and return the result."""
return diagnostics.run(lambda native: native.redact_text(value), lambda: _python_redact_string(value))
return diagnostics.run(lambda native: native.redact_text(value), lambda: python_redact_string(value))
_UNIX_SYSTEM_PATH: Final = r"/(?:etc|var|opt|usr|home|root|private|Users|tmp|mnt|srv)/[^\s'\"\)\]}>,]+"
@ -115,7 +118,7 @@ def _python_redact_internal_details(value: str) -> str:
on top of redact_string(). For client-facing messages only: server logs keep this detail."""
marker_index: Final = value.find(_TRACEBACK_MARKER)
without_traceback: Final = value[:marker_index].rstrip() if marker_index != -1 else value
return _INTERNAL_DETAIL_RE.sub(REDACTED, _python_redact_string(without_traceback))
return _INTERNAL_DETAIL_RE.sub(REDACTED, python_redact_string(without_traceback))
def redact_internal_details(value: str) -> str:
@ -124,7 +127,7 @@ def redact_internal_details(value: str) -> str:
)
def _python_redact_structured_value(key: str | None, value: str) -> str:
def python_redact_structured_value(key: str | None, value: str) -> str:
"""Scrub *value* as it appeared under *key* inside a structured record.
redact_string() replaces a whole ``key: value`` span with REDACTED, which is
@ -133,15 +136,18 @@ def _python_redact_structured_value(key: str | None, value: str) -> str:
repr would, so the key-name patterns still fire, but collapses only the value
so the caller's structure survives.
"""
scrubbed: Final = _python_redact_string(value)
scrubbed: Final = python_redact_string(value)
if scrubbed != value or key is None:
return scrubbed
rendered: Final = f"'{key}': '{value}'"
return REDACTED if _python_redact_string(rendered) != rendered else value
return REDACTED if python_redact_string(rendered) != rendered else value
_python_redact_structured_value = python_redact_structured_value
def redact_structured_value(key: str | None, value: str) -> str:
return diagnostics.run(
lambda native: native.redact_structured_text(key, value),
lambda: _python_redact_structured_value(key, value),
lambda: python_redact_structured_value(key, value),
)

View file

@ -54,7 +54,7 @@ class SensitiveDataMasker:
self.mask_char = mask_char
self.mask_short_values = mask_short_values
def _mask_value(self, value: str) -> str:
def mask_value(self, value: str) -> str:
value_str: Final = str(value)
if not value_str:
return value
@ -71,6 +71,8 @@ class SensitiveDataMasker:
f"{value_str[: self.visible_prefix]}{self.mask_char * masked_length}{value_str[-self.visible_suffix :]}"
)
_mask_value = mask_value
def is_sensitive_key(self, key: str, excluded_keys: set[str] | None = None) -> bool:
# Check if key is in excluded_keys first (exact match)
if excluded_keys and key in excluded_keys:
@ -109,7 +111,7 @@ class SensitiveDataMasker:
elif isinstance(item, list):
masked_items.append(self._mask_sequence(item, depth + 1, max_depth, excluded_keys, key_is_sensitive))
elif key_is_sensitive and isinstance(item, str):
masked_items.append(self._mask_value(item))
masked_items.append(self.mask_value(item))
else:
masked_items.append(item if isinstance(item, (int, float, bool, str, list)) else str(item))
return masked_items
@ -136,7 +138,7 @@ class SensitiveDataMasker:
masked_data[k] = self.mask_dict(vars(v), depth + 1, max_depth, excluded_keys)
elif key_is_sensitive:
str_value = str(v) if v is not None else ""
masked_data[k] = self._mask_value(str_value)
masked_data[k] = self.mask_value(str_value)
else:
masked_data[k] = v if isinstance(v, (int, float, bool, str, list)) else str(v)
except Exception:
@ -198,7 +200,7 @@ class _PayloadWalker:
def walk(self, node: object, key_is_sensitive: bool, depth: int) -> object:
if not isinstance(node, (Mapping, list, tuple, BaseModel)):
return _default_masker._mask_value(node) if key_is_sensitive and isinstance(node, str) and node else node
return _default_masker.mask_value(node) if key_is_sensitive and isinstance(node, str) and node else node
if depth >= DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER:
return REDACTED
memo_key: Final = (id(node), key_is_sensitive and not isinstance(node, Mapping))
@ -242,7 +244,7 @@ def mask_sensitive_keys(data: Mapping[str, object], sensitive_fields: set[str])
if len(value) < min_visible:
masked[key] = mask_char * len(value) if value else value
else:
masked[key] = _default_masker._mask_value(value)
masked[key] = _default_masker.mask_value(value)
else:
masked[key] = value
return masked

View file

@ -698,7 +698,7 @@ class CustomStreamWrapper:
if isinstance(chunk, bytes):
chunk = chunk.decode("utf-8")
if "text_output" in chunk:
response = CustomStreamWrapper._strip_sse_data_from_chunk(chunk) or ""
response = CustomStreamWrapper.strip_sse_data_from_chunk(chunk) or ""
response = response.strip()
parsed_response = json.loads(response)
else:
@ -1676,15 +1676,15 @@ class CustomStreamWrapper:
"""
Caches the streaming response
"""
if not cache_hit and self.logging_obj._llm_caching_handler is not None:
self.logging_obj._llm_caching_handler._sync_add_streaming_response_to_cache(processed_chunk)
if not cache_hit and self.logging_obj.llm_caching_handler is not None:
self.logging_obj.llm_caching_handler.sync_add_streaming_response_to_cache(processed_chunk)
async def async_cache_streaming_response(self, processed_chunk, cache_hit: bool):
"""
Caches the streaming response
"""
if not cache_hit and self.logging_obj._llm_caching_handler is not None:
await self.logging_obj._llm_caching_handler._add_streaming_response_to_cache(processed_chunk)
if not cache_hit and self.logging_obj.llm_caching_handler is not None:
await self.logging_obj.llm_caching_handler.add_streaming_response_to_cache(processed_chunk)
def run_success_logging_and_cache_storage(self, processed_chunk, cache_hit: bool):
"""
@ -1710,7 +1710,7 @@ class CustomStreamWrapper:
asyncio.run(self.logging_obj.async_success_handler(processed_chunk, None, None, cache_hit))
## SYNC LOGGING — only for sync SDK entrypoints; async proxy paths export via async_success_handler
litellm_params: Final = self.logging_obj.model_call_details.get("litellm_params", {})
if self.logging_obj._is_sync_litellm_request(litellm_params):
if self.logging_obj.is_sync_litellm_request(litellm_params):
self.logging_obj.success_handler(processed_chunk, None, None, cache_hit)
def finish_reason_handler(self):
@ -1794,7 +1794,7 @@ class CustomStreamWrapper:
if response is None:
continue
if self.logging_obj.completion_start_time is None:
self.logging_obj._update_completion_start_time(completion_start_time=datetime.datetime.now())
self.logging_obj.update_completion_start_time(completion_start_time=datetime.datetime.now())
## LOGGING
if not litellm.disable_streaming_logging:
executor.submit(
@ -2005,7 +2005,7 @@ class CustomStreamWrapper:
continue
if self.logging_obj.completion_start_time is None:
self.logging_obj._update_completion_start_time(completion_start_time=datetime.datetime.now())
self.logging_obj.update_completion_start_time(completion_start_time=datetime.datetime.now())
if processed_chunk.choices:
choice = processed_chunk.choices[0]
@ -2252,7 +2252,7 @@ class CustomStreamWrapper:
backfill_missing_cache_usage_fields(usage)
self.logging_obj.model_call_details["combined_usage_object"] = usage
self.logging_obj.model_call_details["response_cost"] = (
self.logging_obj._response_cost_calculator(result=partial_response) or 0.0
self.logging_obj.response_cost_calculator(result=partial_response) or 0.0
)
except Exception as recover_error:
verbose_logger.debug(
@ -2334,7 +2334,7 @@ class CustomStreamWrapper:
)
@staticmethod
def _strip_sse_data_from_chunk(chunk: str | None) -> str | None:
def strip_sse_data_from_chunk(chunk: str | None) -> str | None:
"""
Strips the 'data: ' prefix from Server-Sent Events (SSE) chunks.
@ -2369,6 +2369,8 @@ class CustomStreamWrapper:
return chunk
_strip_sse_data_from_chunk = strip_sse_data_from_chunk
def _cache_token_count(details: PromptTokensDetailsWrapper | None, keys: tuple[str, ...]) -> int:
for key in keys:

View file

@ -15,7 +15,7 @@ from typing_extensions import ParamSpec, TypeVar
import litellm
from litellm import verbose_logger
from litellm._lazy_imports import _get_default_encoding
from litellm._lazy_imports import get_default_encoding
from litellm.constants import (
DEFAULT_IMAGE_HEIGHT,
DEFAULT_IMAGE_TOKEN_COUNT,
@ -650,35 +650,34 @@ def _get_exact_count_function(
) -> TokenCounterFunction:
"""
Get the function to count tokens based on the model and custom tokenizer."""
from litellm.utils import _select_tokenizer
from litellm.utils import select_tokenizer
if model is not None or custom_tokenizer is not None:
tokenizer_json: Final = custom_tokenizer or _select_tokenizer(model)
if tokenizer_json["type"] == "huggingface_tokenizer":
tokenizer: Final[HuggingFace] = tokenizer_json["tokenizer"]
def count_tokens(text: str) -> int:
if isinstance(tokenizer, HuggingFaceTokenizer):
return tokenizer.count(text)
return len(tokenizer.encode_batch_fast([text])[0])
return count_tokens
elif tokenizer_json["type"] == "openai_tokenizer":
encoding: Final = openai_tokenizer_encoding(model)
def encode_length(text: str) -> int:
return _encoding_count(encoding, text)
return _get_tiktoken_count_function(encode_length)
else:
raise ValueError("Unsupported tokenizer type")
tokenizer_json: Final = custom_tokenizer or select_tokenizer(model)
else:
default_encoding: Final = _get_default_encoding()
default_encoding: Final = get_default_encoding()
def encode_length(text: str) -> int:
return _encoding_count(default_encoding, text)
return _get_tiktoken_count_function(encode_length)
if tokenizer_json["type"] == "huggingface_tokenizer":
tokenizer: Final[HuggingFace] = tokenizer_json["tokenizer"]
def count_tokens(text: str) -> int:
if isinstance(tokenizer, HuggingFaceTokenizer):
return tokenizer.count(text)
return len(tokenizer.encode_batch_fast([text])[0])
return count_tokens
if tokenizer_json["type"] == "openai_tokenizer":
encoding: Final = openai_tokenizer_encoding(model)
def encode_length(text: str) -> int:
return _encoding_count(encoding, text)
return _get_tiktoken_count_function(encode_length)
raise ValueError("Unsupported tokenizer type")
def _encoding_count(encoding: Encoding, text: str) -> int:

View file

@ -52,7 +52,7 @@ from litellm.types.utils import (
ModelResponseStream,
StreamingChoices,
Usage,
_generate_id,
generate_id,
)
from ...base import BaseLLM
@ -566,7 +566,7 @@ class ModelResponseIterator:
# common case (no '/' or other invalid chars in any tool name).
self.tool_name_reverse_map: dict[str, str] = tool_name_reverse_map or {}
# Generate response ID once per stream to match OpenAI-compatible behavior
self.response_id = _generate_id()
self.response_id = generate_id()
self.served_model: str | None = None
# Track if we're currently streaming a response_format tool

View file

@ -754,11 +754,11 @@ class AnthropicModelInfo(BaseLLMModelInfo):
def _get_model_capability(model: str, key: str) -> bool | None:
"""Read boolean capability ``key`` from the model map, or None when
no entry declares it."""
from litellm.utils import _get_bundled_model_cost_map
from litellm.utils import get_bundled_model_cost_map
try:
candidates: Final = AnthropicModelInfo._model_map_lookup_candidates(model)
for model_cost in (litellm.model_cost, _get_bundled_model_cost_map()):
for model_cost in (litellm.model_cost, get_bundled_model_cost_map()):
for cand in candidates:
value = model_cost.get(cand, {}).get(key)
if isinstance(value, bool):
@ -793,13 +793,13 @@ class AnthropicModelInfo(BaseLLMModelInfo):
model does not resolve under that provider or the resolved entry has no
opinion on ``key``.
"""
from litellm.utils import _get_model_info_helper
from litellm.utils import get_model_info_helper
try:
resolved_model, resolved_provider, _, _ = litellm.get_llm_provider(
model=model, custom_llm_provider=custom_llm_provider
)
value: Final = _get_model_info_helper(model=resolved_model, custom_llm_provider=resolved_provider).get(key)
value: Final = get_model_info_helper(model=resolved_model, custom_llm_provider=resolved_provider).get(key)
except Exception: # noqa: BLE001 # _get_model_info_helper raises bare Exception for unmapped models
return None
return value if isinstance(value, bool) else None
@ -813,13 +813,13 @@ class AnthropicModelInfo(BaseLLMModelInfo):
Otherwise ``_supports_factory``'s provider-level fallbacks and the raw
model-map walk remain as backstops for alias forms the lookup misses.
"""
from litellm.utils import _supports_factory
from litellm.utils import supports_factory
resolved: Final = AnthropicModelInfo._get_provider_resolved_capability(model, key, custom_llm_provider)
if resolved is not None:
return resolved
try:
if _supports_factory(
if supports_factory(
model=model,
custom_llm_provider=custom_llm_provider,
key=key,

View file

@ -538,7 +538,7 @@ def anthropic_messages_handler(
LiteLLM_Proxy_MCP_Handler,
)
if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools):
if LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway(tools=tools):
return anthropic_messages_with_mcp(
max_tokens=max_tokens,
messages=messages,

View file

@ -83,7 +83,7 @@ async def anthropic_messages_with_mcp(
LiteLLM_Proxy_MCP_Handler,
)
mcp_references, other_tools = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(tools)
mcp_references, other_tools = await LiteLLM_Proxy_MCP_Handler.split_mcp_tools(tools)
if not mcp_references:
return await _AnthropicMessagesCall(fn=litellm.anthropic_messages).fn(
@ -100,7 +100,7 @@ async def anthropic_messages_with_mcp(
(
deduplicated_mcp_tools,
tool_server_map,
) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
) = await LiteLLM_Proxy_MCP_Handler.process_mcp_tools_without_openai_transform(
context.user_api_key_auth,
mcp_references,
litellm_trace_id=context.litellm_trace_id,
@ -114,7 +114,7 @@ async def anthropic_messages_with_mcp(
)
all_tools: Final = [*anthropic_tools, *(other_tools or ())]
should_auto_execute: Final = LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(
should_auto_execute: Final = LiteLLM_Proxy_MCP_Handler.should_auto_execute_tools(
mcp_tools_with_litellm_proxy=mcp_references
)
stream: Final = bool(kwargs.pop("stream", False))
@ -145,7 +145,7 @@ async def anthropic_messages_with_mcp(
if not tool_use_blocks:
break
tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_results = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls(
tool_server_map=tool_server_map,
served_tools=deduplicated_mcp_tools,
tool_calls=list(tool_use_blocks),

View file

@ -94,7 +94,7 @@ class AnthropicMessagesStreamCacheWriter:
return
self.persisted = True
if not self.caching_handler._should_store_result_in_cache(
if not self.caching_handler.should_store_result_in_cache(
original_function=self.caching_handler.original_function,
kwargs=self.caching_handler.request_kwargs,
):

View file

@ -699,7 +699,7 @@ class BaseAnthropicMessagesStreamingIterator:
async def _fire_detached_failure_hook(self, exc: Exception) -> None:
from litellm._logging import verbose_proxy_logger
on_detached_failure: Final = getattr(self.litellm_logging_obj, "_on_detached_stream_failure", None)
on_detached_failure: Final = getattr(self.litellm_logging_obj, "on_detached_stream_failure", None)
if on_detached_failure is None:
return
try:

View file

@ -208,7 +208,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
Subclasses whose upstream rejects the role opt in by calling this from
their ``transform_anthropic_messages_request``; the first-party Anthropic
path forwards ``messages`` untouched and never calls it."""
from litellm.utils import _supports_factory
from litellm.utils import supports_factory
messages: Final = anthropic_messages_request.get("messages")
if not isinstance(messages, list):
@ -220,7 +220,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
hoisted: Final = messages[:leading_count]
remaining: Final = (
messages[leading_count:]
if _supports_factory(
if supports_factory(
model=model,
custom_llm_provider=self.custom_llm_provider,
key="supports_mid_conversation_system",

View file

@ -74,7 +74,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
)
from litellm.responses.utils import ResponseAPILoggingUtils
chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_usage)
chat_usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(raw_usage)
return LiteLLMAnthropicMessagesAdapter._translate_openai_usage_to_anthropic_usage(chat_usage)
# ------------------------------------------------------------------ #

View file

@ -22,7 +22,7 @@ from litellm.secret_managers.get_azure_ad_token_provider import (
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import _add_path_to_api_base
from litellm.utils import add_path_to_api_base
azure_ad_cache: Final = DualCache()
@ -814,7 +814,7 @@ class BaseAzureLLM(BaseOpenAILLM):
# Add the path to the base URL
if route not in api_base:
new_url = _add_path_to_api_base(api_base=api_base, ending_path=route)
new_url = add_path_to_api_base(api_base=api_base, ending_path=route)
else:
new_url = api_base

View file

@ -7,7 +7,7 @@ from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import _add_path_to_api_base
from litellm.utils import add_path_to_api_base
class AzureImageEditConfig(OpenAIImageEditConfig):
@ -122,7 +122,7 @@ class AzureImageEditConfig(OpenAIImageEditConfig):
# Add the path to the base URL using the model as deployment name
if "/openai/deployments/" not in api_base:
new_url = _add_path_to_api_base(
new_url = add_path_to_api_base(
api_base=api_base,
ending_path=f"/openai/deployments/{model}/images/edits",
)

View file

@ -9,7 +9,7 @@ from httpx import Response
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
_audio_or_image_in_message_content,
audio_or_image_in_message_content,
convert_content_list_to_str,
filter_value_from_dict,
)
@ -27,7 +27,7 @@ from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import ModelResponse, ProviderField
from litellm.utils import _add_path_to_api_base, supports_tool_choice
from litellm.utils import add_path_to_api_base, supports_tool_choice
if TYPE_CHECKING:
from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer
@ -193,9 +193,9 @@ class AzureAIStudioConfig(OpenAIConfig):
# Add the path to the base URL
if "services.ai.azure.com" in api_base:
new_url = _add_path_to_api_base(api_base=api_base, ending_path="/models/chat/completions")
new_url = add_path_to_api_base(api_base=api_base, ending_path="/models/chat/completions")
else:
new_url = _add_path_to_api_base(api_base=api_base, ending_path="/chat/completions")
new_url = add_path_to_api_base(api_base=api_base, ending_path="/chat/completions")
# Use the new query_params dictionary
final_url: Final = httpx.URL(new_url).copy_with(params=query_params)
@ -245,7 +245,7 @@ class AzureAIStudioConfig(OpenAIConfig):
filter_value_from_dict(message_dict, field)
# Do nothing if the message contains an image or audio
if _audio_or_image_in_message_content(message):
if audio_or_image_in_message_content(message):
continue
texts = convert_content_list_to_str(message=message)

View file

@ -9,7 +9,7 @@ from litellm.llms.azure_ai.common_utils import (
)
from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig
from litellm.secret_managers.main import get_secret_str
from litellm.utils import _add_path_to_api_base
from litellm.utils import add_path_to_api_base
class AzureFoundryFluxImageEditConfig(OpenAIImageEditConfig):
@ -81,12 +81,12 @@ class AzureFoundryFluxImageEditConfig(OpenAIImageEditConfig):
# Add the path to the base URL using the model as deployment name
# Azure AI Foundry FLUX models use /images/edits for editing
if "/openai/deployments/" in api_base:
new_url = _add_path_to_api_base(
new_url = add_path_to_api_base(
api_base=api_base,
ending_path="/images/edits",
)
else:
new_url = _add_path_to_api_base(
new_url = add_path_to_api_base(
api_base=api_base,
ending_path=f"/openai/deployments/{model}/images/edits",
)

View file

@ -3,15 +3,15 @@ from typing import Final
import litellm
from litellm.litellm_core_utils.llm_cost_calc.utils import (
_get_cost_per_unit,
calculate_image_response_cost_from_usage,
get_cost_per_unit,
resolve_image_model_info,
)
from litellm.types.utils import ImageResponse, ModelInfo
def _input_cost_per_pixel(resolved: ModelInfo) -> float:
deployment_price: Final = _get_cost_per_unit(resolved, "input_cost_per_pixel", default_value=None)
deployment_price: Final = get_cost_per_unit(resolved, "input_cost_per_pixel", default_value=None)
if deployment_price is not None:
return deployment_price
model_cost_key: Final = resolved.get("key")

View file

@ -13,7 +13,7 @@ from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers
from litellm.llms.cohere.rerank.transformation import CohereRerankConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.utils import RerankResponse
from litellm.utils import _add_path_to_api_base
from litellm.utils import add_path_to_api_base
class AzureAIRerankConfig(CohereRerankConfig):
@ -52,13 +52,13 @@ class AzureAIRerankConfig(CohereRerankConfig):
or normalized_path.endswith("/v2")
or normalized_path.endswith("/providers/cohere/v2")
):
return _add_path_to_api_base(
return add_path_to_api_base(
api_base=str(original_url.copy_with(path=normalized_path or "/")),
ending_path="/rerank",
)
# Backwards compatible default: Azure AI rerank was originally exposed under /v1/rerank
return _add_path_to_api_base(api_base=api_base, ending_path="/v1/rerank")
return add_path_to_api_base(api_base=api_base, ending_path="/v1/rerank")
def validate_environment(
self,

View file

@ -98,7 +98,7 @@ class BaseModelResponseIterator:
@staticmethod
def _string_to_dict_parser(str_line: str) -> dict | None:
stripped_json_chunk: dict | None = None
stripped_chunk: Final = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(str_line)
stripped_chunk: Final = litellm.CustomStreamWrapper.strip_sse_data_from_chunk(str_line)
try:
if stripped_chunk is not None:
stripped_json_chunk = json.loads(stripped_chunk)

View file

@ -28,8 +28,8 @@ from litellm.litellm_core_utils.core_helpers import (
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.prompt_templates.common_utils import (
_parse_content_for_reasoning,
drop_lookaround_regex_patterns,
parse_content_for_reasoning,
tool_with_sanitized_parameters,
)
from litellm.litellm_core_utils.prompt_templates.factory import (
@ -2387,7 +2387,7 @@ class AmazonConverseConfig(BaseConfig):
(
extracted_reasoning_content_str,
_content_str,
) = _parse_content_for_reasoning(content["text"])
) = parse_content_for_reasoning(content["text"])
if _content_str is not None:
content_str += _content_str
if "toolUse" in content:

View file

@ -4,7 +4,7 @@ from httpx import Response
from litellm import verbose_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
_parse_content_for_reasoning,
parse_content_for_reasoning,
)
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
@ -63,7 +63,7 @@ class AmazonDeepSeekR1Config(AmazonLlamaConfig):
message_content: Final = cast(str | None, cast(Choices, response.choices[0]).message.get("content"))
if prompt and prompt.strip().endswith("<think>") and message_content:
message_content_with_reasoning_token: Final = "<think>" + message_content
reasoning, content = _parse_content_for_reasoning(message_content_with_reasoning_token)
reasoning, content = parse_content_for_reasoning(message_content_with_reasoning_token)
provider_specific_fields: Final = cast(Choices, response.choices[0]).message.provider_specific_fields or {}
if reasoning:
provider_specific_fields["reasoning_content"] = reasoning

View file

@ -246,9 +246,9 @@ def convert_bedrock_invoke_output_format_to_inline_schema(
def _bedrock_model_supports(model: str, key: str) -> bool:
from litellm.utils import _supports_factory
from litellm.utils import supports_factory
return _supports_factory(model=model, custom_llm_provider="bedrock", key=key)
return supports_factory(model=model, custom_llm_provider="bedrock", key=key)
def apply_bedrock_invoke_structured_output(

View file

@ -25,9 +25,9 @@ from litellm.types.images.main import ImageEditOptionalRequestParams
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import FileTypes, ImageObject, ImageResponse
from litellm.utils import (
_get_model_cost_key,
_get_potential_model_names,
get_model_cost_key,
get_model_info,
get_potential_model_names,
)
if TYPE_CHECKING:
@ -192,7 +192,7 @@ def _supports_nova_canvas_image_edit_from_model_cost(model: str) -> bool:
pass
try:
potential: Final = _get_potential_model_names(model=model, custom_llm_provider=None)
potential: Final = get_potential_model_names(model=model, custom_llm_provider=None)
for field in (
"combined_model_name",
"combined_stripped_model_name",
@ -206,7 +206,7 @@ def _supports_nova_canvas_image_edit_from_model_cost(model: str) -> bool:
pass
for name in candidates:
key = _get_model_cost_key(name)
key = get_model_cost_key(name)
if key is None:
continue
entry = _litellm.model_cost.get(key) or {}

View file

@ -16,7 +16,7 @@ from typing import Final, NoReturn, Protocol, runtime_checkable
from pydantic import JsonValue, TypeAdapter
import litellm
from litellm._logging import _redact_string, verbose_proxy_logger
from litellm._logging import redact_string, verbose_proxy_logger
from litellm.constants import (
BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY,
BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY,
@ -192,9 +192,9 @@ def _pending_session_update(scope: Mapping[str, object]) -> str | None:
def _raise_provider_failure(scope: MutableMapping[str, object], failure: BaseException) -> NoReturn:
error: Final = _as_bedrock_error(failure)
verbose_proxy_logger.error("Bedrock Realtime: provider stream failed: %s", _redact_string(str(error)))
verbose_proxy_logger.error("Bedrock Realtime: provider stream failed: %s", redact_string(str(error)))
if scope.get(BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY) is True:
scope[BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY] = _redact_string(str(error))
scope[BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY] = redact_string(str(error))
raise error from failure

View file

@ -6,7 +6,7 @@ import httpx
from litellm.exceptions import AuthenticationError
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
_safe_convert_created_field,
safe_convert_created_field,
)
from litellm.llms.openai.common_utils import OpenAIError
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
@ -216,7 +216,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
if not response_payload.get("output") and streamed_output_items:
response_payload["output"] = [item for _, item in sorted(streamed_output_items.items())]
if "created_at" in response_payload:
response_payload["created_at"] = _safe_convert_created_field(response_payload["created_at"])
response_payload["created_at"] = safe_convert_created_field(response_payload["created_at"])
try:
return ResponsesAPIResponse(**response_payload)
except Exception:

View file

@ -83,7 +83,7 @@ class CodestralTextCompletionConfig(OpenAITextCompletionConfig):
finish_reason = None
logprobs = None
chunk_data = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(chunk_data) or ""
chunk_data = litellm.CustomStreamWrapper.strip_sse_data_from_chunk(chunk_data) or ""
chunk_data = chunk_data.strip()
if len(chunk_data) == 0 or chunk_data == "[DONE]":
return {

View file

@ -33,7 +33,7 @@ import litellm
import litellm.litellm_core_utils
import litellm.types
import litellm.types.utils
from litellm._logging import _redact_string, verbose_logger
from litellm._logging import redact_string, verbose_logger
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
from litellm.constants import MAX_FILE_LIST_LIMIT, REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.files.types import FileContentStreamingResult
@ -411,12 +411,12 @@ def _mask_presigned_request_headers(transformed_request: bytes | str | dict) ->
return transformed_request
from litellm.litellm_core_utils.litellm_logging import (
_get_masked_values, # pyright: ignore[reportPrivateUsage] # the shared header-masking helper has no public name
get_masked_values,
)
return {
**transformed_request,
"headers": _get_masked_values(request_headers),
"headers": get_masked_values(request_headers),
}
@ -6342,7 +6342,7 @@ class BaseLLMHTTPHandler:
if provider_config.requires_session_configuration():
_session_config = provider_config.session_configuration_request(model)
if _session_config:
_session_config = realtime_streaming._maybe_inject_guardrail_auto_response_disable(
_session_config = realtime_streaming.maybe_inject_guardrail_auto_response_disable(
_session_config
)
await backend_ws.send(_session_config)
@ -6364,7 +6364,7 @@ class BaseLLMHTTPHandler:
# success_handler / async_success_handler payloads.
realtime_streaming.store_message(synthetic_session_str)
await websocket.send_text(synthetic_session_str)
realtime_streaming._session_created_sent_to_client = True
realtime_streaming.session_created_sent_to_client = True
verbose_logger.debug("Sent synthetic session.created to client to unblock connection")
await realtime_streaming.bidirectional_forward()
@ -6374,7 +6374,7 @@ class BaseLLMHTTPHandler:
await close_after_upstream_handshake_refusal(websocket, e.response.status_code)
except Exception as e:
verbose_logger.exception("Error connecting to backend: %s", e)
redacted_error: Final = _redact_string(str(e))
redacted_error: Final = redact_string(str(e))
try:
await websocket.send_text(realtime_error_event(redacted_error, error_type="server_error"))
except Exception: # noqa: BLE001 # best-effort notice: a dead client socket must not skip the close below
@ -6383,7 +6383,7 @@ class BaseLLMHTTPHandler:
await websocket.close(
code=1011,
reason=websocket_close_reason(
_redact_string(f"Internal server error: {e}"),
redact_string(f"Internal server error: {e}"),
fallback="Internal server error",
),
)
@ -6790,7 +6790,7 @@ class BaseLLMHTTPHandler:
except Exception as e:
verbose_logger.exception("Error in responses WS: %s", e)
try:
await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {e}"))
await websocket.close(code=1011, reason=redact_string(f"Internal server error: {e}"))
except RuntimeError as close_error:
if "already completed" in str(close_error) or "websocket.close" in str(close_error):
pass

View file

@ -11,11 +11,11 @@ from pydantic import BaseModel
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
_handle_invalid_parallel_tool_calls,
_should_convert_tool_call_to_json_mode,
handle_invalid_parallel_tool_calls,
should_convert_tool_call_to_json_mode,
)
from litellm.litellm_core_utils.prompt_templates.common_utils import (
_extract_reasoning_content, # pyright: ignore[reportPrivateUsage] # same import as the OpenAI transformation
extract_reasoning_content,
merge_consecutive_system_messages,
strip_litellm_internal_message_fields,
strip_name_from_message,
@ -561,7 +561,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
content_str: Final = DatabricksConfig.extract_content_str(message["content"])
if block_reasoning_content is not None:
return block_reasoning_content, content_str
return _extract_reasoning_content({**message, "content": content_str})
return extract_reasoning_content({**message, "content": content_str})
@staticmethod
def extract_citations(
@ -588,14 +588,14 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
for _tc in tool_calls:
_openai_tc = ChatCompletionMessageToolCall(**_tc)
_openai_tool_calls.append(_openai_tc)
fixed_tool_calls = _handle_invalid_parallel_tool_calls(_openai_tool_calls)
fixed_tool_calls = handle_invalid_parallel_tool_calls(_openai_tool_calls)
if fixed_tool_calls is not None:
tool_calls = fixed_tool_calls
translated_message: Message | None = None
finish_reason: str | None = None
if tool_calls and _should_convert_tool_call_to_json_mode(
if tool_calls and should_convert_tool_call_to_json_mode(
tool_calls=tool_calls,
convert_tool_call_to_json_mode=json_mode,
):

View file

@ -105,7 +105,7 @@ class ModelResponseIterator:
raise RuntimeError(f"Error receiving chunk from stream: {e}")
try:
chunk = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(chunk) or ""
chunk = litellm.CustomStreamWrapper.strip_sse_data_from_chunk(chunk) or ""
chunk = chunk.strip()
if len(chunk) > 0:
json_chunk: Final = json.loads(chunk)
@ -150,7 +150,7 @@ class ModelResponseIterator:
raise RuntimeError(f"Error receiving chunk from stream: {e}")
try:
chunk = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(chunk) or ""
chunk = litellm.CustomStreamWrapper.strip_sse_data_from_chunk(chunk) or ""
chunk = chunk.strip()
if chunk == "[DONE]":
raise StopAsyncIteration

Some files were not shown because too many files have changed in this diff Show more