mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
b9ab1eee1f
commit
c3a23fe499
461 changed files with 6719 additions and 4996 deletions
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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?"}]}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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``."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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__
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue