mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
refactor(types): replace Any with proven types in 8 files (#43551)
* refactor(types): replace Any with proven types in 8 files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(types): keep email logger untyped where its alert types differ Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
90e4962c81
commit
2c9b0e00ac
8 changed files with 112 additions and 79 deletions
|
|
@ -2287,7 +2287,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.completion_start_time = completion_start_time
|
||||
self.model_call_details["completion_start_time"] = self.completion_start_time
|
||||
|
||||
def normalize_logging_result(self, result: Any) -> object:
|
||||
def normalize_logging_result(self, result: object) -> object:
|
||||
"""
|
||||
Some endpoints return a different type of result than what is expected by the logging system.
|
||||
This function is used to normalize the result to the expected type.
|
||||
|
|
@ -2736,7 +2736,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
start_time: datetime.datetime | None = None,
|
||||
end_time: datetime.datetime | None = None,
|
||||
cache_hit: bool | None = None,
|
||||
**kwargs: Any, # kwargs-ok: forwarded to _success_handler_body
|
||||
**kwargs: object, # kwargs-ok: forwarded to _success_handler_body
|
||||
) -> None:
|
||||
"""Restores trace_id/session_id contextvars once this attempt's own success
|
||||
logging (including any nested calls its callbacks trigger) is fully done."""
|
||||
|
|
@ -3175,7 +3175,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
start_time: datetime.datetime | None = None,
|
||||
end_time: datetime.datetime | None = None,
|
||||
cache_hit: bool | None = None,
|
||||
**kwargs: Any, # kwargs-ok: forwarded to _async_success_handler_body
|
||||
**kwargs: object, # kwargs-ok: forwarded to _async_success_handler_body
|
||||
) -> None:
|
||||
"""Restores trace_id/session_id contextvars once this attempt's own success
|
||||
logging (including any nested calls its callbacks trigger) is fully done."""
|
||||
|
|
@ -3554,7 +3554,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
self._handle_callback_failure(callback=callback)
|
||||
|
||||
def _handle_callback_failure(self, callback: Any):
|
||||
def _handle_callback_failure(self, callback: object):
|
||||
"""
|
||||
Handle callback logging failures by incrementing Prometheus metrics.
|
||||
|
||||
|
|
@ -3959,10 +3959,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
|
||||
def handle_sync_success_callbacks_for_async_calls(
|
||||
self,
|
||||
result: Any,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
cache_hit: Any | None = None,
|
||||
cache_hit: object | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Handles calling success callbacks for Async calls.
|
||||
|
|
@ -4132,10 +4132,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
response so cost calculation and spend tracking see one shape.
|
||||
"""
|
||||
if result.event_type == "interaction.completed" and result.interaction is not None:
|
||||
return InteractionsAPIResponse(**result.interaction)
|
||||
return InteractionsAPIResponse.model_validate(result.interaction)
|
||||
if result.status == "completed":
|
||||
return InteractionsAPIResponse(
|
||||
**result.model_dump(
|
||||
return InteractionsAPIResponse.model_validate(
|
||||
result.model_dump(
|
||||
exclude={ # mutable-ok: pydantic types exclude as set[str], which a frozenset does not satisfy
|
||||
"event_type",
|
||||
"delta",
|
||||
|
|
@ -4165,7 +4165,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
return logged
|
||||
return logged.model_copy(update={"id": streamed_message_id})
|
||||
|
||||
def _handle_anthropic_messages_response_logging(self, result: Any) -> ModelResponse:
|
||||
def _handle_anthropic_messages_response_logging(self, result: object) -> ModelResponse:
|
||||
"""
|
||||
Handles logging for Anthropic messages responses.
|
||||
|
||||
|
|
@ -4274,7 +4274,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
return model_response
|
||||
|
||||
def _handle_non_streaming_google_genai_generate_content_response_logging(self, result: Any) -> ModelResponse:
|
||||
def _handle_non_streaming_google_genai_generate_content_response_logging(self, result: object) -> ModelResponse:
|
||||
"""
|
||||
Handles logging for Google GenAI generate content responses.
|
||||
"""
|
||||
|
|
@ -4349,7 +4349,7 @@ def _get_masked_values(
|
|||
"passwd",
|
||||
]
|
||||
|
||||
def _mask_value(v: Any) -> Any:
|
||||
def _mask_value(v: object) -> object:
|
||||
if isinstance(v, dict):
|
||||
if _depth >= _max_depth:
|
||||
return v
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
import asyncio
|
||||
import json
|
||||
import ssl
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Coroutine, Iterator, Mapping, Sequence
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Coroutine, Iterator, Mapping, Sequence
|
||||
from contextlib import asynccontextmanager
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType, ModuleType
|
||||
from types import MappingProxyType
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
|
|
@ -239,6 +239,26 @@ class _ResponsesClientWebSocket(Protocol):
|
|||
async def close(self, code: int = ..., reason: str | None = ...) -> None: ...
|
||||
|
||||
|
||||
class _WebsocketsExceptions(Protocol):
|
||||
@property
|
||||
def WebSocketException(self) -> type[Exception]: ...
|
||||
|
||||
|
||||
class _WebsocketsModule(Protocol):
|
||||
@property
|
||||
def exceptions(self) -> _WebsocketsExceptions: ...
|
||||
|
||||
def connect(
|
||||
self,
|
||||
uri: str,
|
||||
*,
|
||||
additional_headers: Mapping[str, str],
|
||||
max_size: int | None,
|
||||
ssl: bool | str | ssl.SSLContext,
|
||||
open_timeout: float,
|
||||
) -> Awaitable["ClientConnection"]: ...
|
||||
|
||||
|
||||
_ResponseT = TypeVar("_ResponseT")
|
||||
|
||||
|
||||
|
|
@ -5468,7 +5488,7 @@ class BaseLLMHTTPHandler:
|
|||
max_loops: int,
|
||||
fingerprints: list[str],
|
||||
fingerprint: str,
|
||||
) -> Any:
|
||||
) -> ModelResponse | CustomStreamWrapper:
|
||||
patch: Final = plan.request_patch or AgenticLoopRequestPatch()
|
||||
if patch.messages is None:
|
||||
raise ValueError("Agentic loop plan missing patched messages")
|
||||
|
|
@ -5978,7 +5998,7 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
@staticmethod
|
||||
async def _open_realtime_backend_ws(
|
||||
websockets_module: ModuleType,
|
||||
websockets_module: _WebsocketsModule,
|
||||
url: str,
|
||||
headers: dict,
|
||||
ssl_context: bool | str | ssl.SSLContext,
|
||||
|
|
@ -6037,7 +6057,7 @@ class BaseLLMHTTPHandler:
|
|||
headers: dict,
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
client: Any | None = None,
|
||||
client: object | None = None,
|
||||
timeout: float | None = None,
|
||||
user_api_key_dict: object | None = None,
|
||||
litellm_metadata: dict[str, object] | None = None,
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import contextlib
|
|||
import json
|
||||
import logging
|
||||
import math
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
|
|
@ -974,7 +974,7 @@ _NO_GENERAL_SETTINGS: Final[Mapping[str, object]] = MappingProxyType({})
|
|||
|
||||
|
||||
async def create_response(
|
||||
generator: AsyncGenerator[str, None],
|
||||
generator: AsyncGenerator[str, None] | Coroutine[object, object, AsyncGenerator[str, None]],
|
||||
media_type: str,
|
||||
headers: Mapping[str, str],
|
||||
default_status_code: int = status.HTTP_200_OK,
|
||||
|
|
@ -4166,7 +4166,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _stream_usage_for_event(obj: Mapping[str, object], usage: Mapping[str, Any]) -> Usage | None:
|
||||
def _stream_usage_for_event(obj: Mapping[str, object], usage: Mapping[str, object]) -> Usage | None:
|
||||
# Anthropic reports input_tokens excluding cache tokens, so reuse the non-streaming
|
||||
# transformation to total the prompt and keep the 5m/1h cache creation split
|
||||
if obj.get("type") == "message_delta":
|
||||
|
|
|
|||
|
|
@ -2550,7 +2550,7 @@ def general_settings_view() -> Mapping[str, object]:
|
|||
return _GENERAL_SETTINGS_VIEW.validate_python(general_settings)
|
||||
|
||||
|
||||
config_passthrough_endpoints: list[dict[str, Any]] | None = None
|
||||
config_passthrough_endpoints: list[dict[str, object]] | None = None
|
||||
log_file: Final = "api_log.json"
|
||||
worker_config: Final = None
|
||||
master_key: str | None = None
|
||||
|
|
@ -2584,7 +2584,7 @@ use_queue = False
|
|||
health_check_interval = None
|
||||
health_check_concurrency = None
|
||||
health_check_details = None
|
||||
health_check_results: dict[str, int | list[dict[str, Any]]] = {}
|
||||
health_check_results: dict[str, int | list[dict[str, object]]] = {}
|
||||
background_health_check_loop_active = False
|
||||
background_health_check_cycle_seq = 0
|
||||
queue: Final[list] = []
|
||||
|
|
@ -13640,7 +13640,7 @@ async def _try_provider_token_count(
|
|||
model_to_use: str,
|
||||
messages: list | None,
|
||||
contents: list | None,
|
||||
deployment: dict[str, Any] | None,
|
||||
deployment: dict[str, object] | None,
|
||||
request_model: str,
|
||||
tools: list | None = None,
|
||||
system: str | None = None,
|
||||
|
|
@ -14439,7 +14439,7 @@ async def _fetch_db_models_for_search(
|
|||
sort_by: str | None,
|
||||
is_byok_outside_caller_teams: Callable[[dict[str, JsonValue]], bool],
|
||||
model_name: str | None = None,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
) -> tuple[list[dict[str, object]], int]:
|
||||
"""
|
||||
Run the bounded DB query that backs `/v2/model/info?search=`. Returns
|
||||
`(decrypted_models, total_count)` where `total_count` is the cheap
|
||||
|
|
@ -14586,7 +14586,7 @@ async def _apply_search_filter_to_models(
|
|||
router_models_count: Final = config_models_count + db_models_in_router_count
|
||||
|
||||
# Query database for additional models with search term
|
||||
db_models: list[dict[str, Any]] = []
|
||||
db_models: list[dict[str, object]] = []
|
||||
exact_name_can_match: Final = model_name is None or search_lower in model_name.lower()
|
||||
if prisma_client is not None and exact_name_can_match:
|
||||
try:
|
||||
|
|
@ -14781,7 +14781,7 @@ def _matches_model_info_filters(
|
|||
|
||||
|
||||
def _paginate_models_response(
|
||||
all_models: list[dict[str, Any]],
|
||||
all_models: Sequence[Mapping[str, object]],
|
||||
page: int,
|
||||
size: int,
|
||||
total_count: int | None,
|
||||
|
|
|
|||
|
|
@ -303,6 +303,29 @@ class _EndUserSpendBatch(Protocol):
|
|||
def litellm_endusertable(self) -> _EndUserBatchTable: ...
|
||||
|
||||
|
||||
class _BatchUpdateTable(Protocol):
|
||||
def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...
|
||||
|
||||
|
||||
class _UpdateManyBatch(Protocol):
|
||||
@property
|
||||
def litellm_verificationtoken(self) -> _BatchUpdateTable: ...
|
||||
|
||||
@property
|
||||
def litellm_usertable(self) -> _EndUserBatchTable: ...
|
||||
|
||||
@property
|
||||
def litellm_endusertable(self) -> _EndUserBatchTable: ...
|
||||
|
||||
@property
|
||||
def litellm_budgettable(self) -> _EndUserBatchTable: ...
|
||||
|
||||
@property
|
||||
def litellm_teamtable(self) -> _EndUserBatchTable: ...
|
||||
|
||||
async def commit(self) -> None: ...
|
||||
|
||||
|
||||
unified_guardrail: Final = UnifiedLLMGuardrails()
|
||||
|
||||
NON_OPENAI_STREAM_GUARDRAIL_TRANSLATION_CALL_TYPES: "frozenset[CallTypes]" = frozenset({CallTypes.anthropic_messages})
|
||||
|
|
@ -1111,7 +1134,7 @@ class _CallbackCapabilities:
|
|||
# Tuple[(resolved_callback, "override" | "apply_guardrail"), ...]
|
||||
# Ordered the same as ``litellm.callbacks``; used to build the streaming
|
||||
# iterator chain without re-scanning per request.
|
||||
iterator_overrides: tuple[tuple[Any, str], ...] = field(default_factory=tuple)
|
||||
iterator_overrides: tuple[tuple[CustomLogger, str], ...] = field(default_factory=tuple)
|
||||
# Resolved CustomLogger callbacks in original order. Pre-resolving once
|
||||
# avoids the per-request ``get_custom_logger_compatible_class`` walk for
|
||||
# every string entry in ``litellm.callbacks``.
|
||||
|
|
@ -2761,7 +2784,7 @@ class ProxyLogging:
|
|||
has_pre_call_override = False
|
||||
has_content_enforcer = False
|
||||
has_moderation_override = False
|
||||
iterator_overrides: Final[list[tuple[Any, str]]] = [] # (callback, kind)
|
||||
iterator_overrides: Final[list[tuple[CustomLogger, str]]] = [] # (callback, kind)
|
||||
resolved_callbacks: Final[list[CustomLogger]] = []
|
||||
|
||||
for callback in callbacks:
|
||||
|
|
@ -5481,7 +5504,7 @@ class PrismaClient:
|
|||
"""
|
||||
Batch write update queries
|
||||
"""
|
||||
batcher = self.db.batch_()
|
||||
batcher: _UpdateManyBatch = self.db.batch_()
|
||||
for idx, t in enumerate(data_list):
|
||||
# check if plain text or hash
|
||||
if t.token.startswith("sk-"):
|
||||
|
|
@ -6638,7 +6661,7 @@ class PrismaClient:
|
|||
read-only as a whole (replica, failover in progress) does not get its
|
||||
engine killed on every watchdog cycle or failed write."""
|
||||
backoff_seconds: Final = min(
|
||||
self._db_reconnect_cooldown_seconds * 2 ** min(self._db_read_only_recreate_streak, 10),
|
||||
self._db_reconnect_cooldown_seconds * (1 << min(self._db_read_only_recreate_streak, 10)),
|
||||
_READ_ONLY_RECREATE_BACKOFF_CAP_SECONDS,
|
||||
)
|
||||
if time.time() - self._db_read_only_recreate_ts < backoff_seconds:
|
||||
|
|
@ -7344,7 +7367,7 @@ class ProxyUpdateSpend:
|
|||
if i >= n_retry_times:
|
||||
await requeue_spend_logs(prisma_client, proxy_logging_obj, logs_to_process)
|
||||
raise
|
||||
await asyncio.sleep(2**i)
|
||||
await asyncio.sleep(1 << i)
|
||||
except Exception as e:
|
||||
_raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj)
|
||||
finally:
|
||||
|
|
|
|||
|
|
@ -981,7 +981,7 @@ def _responses_try_dispatch_mcp_gateway(
|
|||
background: bool | None,
|
||||
stream: bool | None,
|
||||
temperature: float | None,
|
||||
text: Any,
|
||||
text: Optional["ResponseText"],
|
||||
tool_choice: ToolChoice | None,
|
||||
top_p: float | None,
|
||||
truncation: Literal["auto", "disabled"] | None,
|
||||
|
|
@ -1054,7 +1054,7 @@ def _responses_try_dispatch_emulated_file_search(
|
|||
background: bool | None,
|
||||
stream: bool | None,
|
||||
temperature: float | None,
|
||||
text: Any,
|
||||
text: Optional["ResponseText"],
|
||||
tool_choice: ToolChoice | None,
|
||||
top_p: float | None,
|
||||
truncation: Literal["auto", "disabled"] | None,
|
||||
|
|
|
|||
|
|
@ -433,6 +433,7 @@ def model_info_is_active_for_environment(model_info: Mapping[str, object] | None
|
|||
|
||||
|
||||
_PreRoutingStrategyT = TypeVar("_PreRoutingStrategyT")
|
||||
_CallbackT = TypeVar("_CallbackT")
|
||||
|
||||
_ALIAS_PARAMS_NEVER_FORWARDED: Final = frozenset({"model", "api_base", "api_key", "api_version"})
|
||||
_ALIAS_MARKER_FORWARDED_PARAMS_KWARG: Final = "_alias_marker_forwarded_params"
|
||||
|
|
@ -2217,8 +2218,8 @@ class Router:
|
|||
|
||||
def _add_encrypted_content_affinity_check(self, enable_global_affinity: bool) -> None:
|
||||
def _move_before_deployment_affinity(
|
||||
callback_list: list[Any],
|
||||
callback_to_move: EncryptedContentAffinityCheck,
|
||||
callback_list: list[_CallbackT],
|
||||
callback_to_move: _CallbackT,
|
||||
) -> None:
|
||||
if callback_to_move not in callback_list:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -681,7 +681,7 @@ def load_credentials_from_list(kwargs: dict):
|
|||
Updates kwargs with the credentials if credential_name in kwarg
|
||||
"""
|
||||
# Access CredentialAccessor via module to trigger lazy loading if needed
|
||||
credential_accessor: Final[type[CredentialAccessor]] = getattr(sys.modules[__name__], "CredentialAccessor")
|
||||
credential_accessor: Final[type[CredentialAccessor]] = litellm_utils.CredentialAccessor
|
||||
|
||||
credential_name: Final = kwargs.get("litellm_credential_name")
|
||||
if not credential_name:
|
||||
|
|
@ -1015,9 +1015,7 @@ def function_setup(
|
|||
len(litellm.input_callback) > 0 or len(litellm.success_callback) > 0 or len(litellm.failure_callback) > 0
|
||||
) and len(callback_list) == 0:
|
||||
callback_list = list(set(litellm.input_callback + litellm.success_callback + litellm.failure_callback))
|
||||
get_set_callbacks: Final[Callable[[], Callable[..., None]]] = getattr(
|
||||
sys.modules[__name__], "get_set_callbacks"
|
||||
)
|
||||
get_set_callbacks: Final[Callable[[], Callable[..., None]]] = litellm_utils.get_set_callbacks
|
||||
get_set_callbacks()(callback_list=callback_list, function_id=function_id)
|
||||
## ASYNC CALLBACKS - safety net for callbacks added via direct append
|
||||
if len(litellm.input_callback) > 0:
|
||||
|
|
@ -1262,9 +1260,7 @@ def function_setup(
|
|||
call_type=call_type,
|
||||
):
|
||||
stream = True
|
||||
get_litellm_logging_class: Final[_LoggingClassGetter] = getattr(
|
||||
sys.modules[__name__], "get_litellm_logging_class"
|
||||
)
|
||||
get_litellm_logging_class: Final[_LoggingClassGetter] = litellm_utils.get_litellm_logging_class
|
||||
# Victim for object pool
|
||||
logging_obj = get_litellm_logging_class()( # rebind-ok: 2nd assignment to logging_obj (see initial None above)
|
||||
model=model,
|
||||
|
|
@ -1413,8 +1409,8 @@ def _get_wrapper_num_retries(kwargs: dict[str, Any], exception: Exception) -> tu
|
|||
if num_retries is None:
|
||||
num_retries = litellm.num_retries
|
||||
if kwargs.get("retry_policy", None):
|
||||
get_num_retries_from_retry_policy: Final[Callable[..., int | None]] = getattr(
|
||||
sys.modules[__name__], "get_num_retries_from_retry_policy"
|
||||
get_num_retries_from_retry_policy: Final[Callable[..., int | None]] = (
|
||||
litellm_utils.get_num_retries_from_retry_policy
|
||||
)
|
||||
reset_retry_policy: Final = litellm_utils.reset_retry_policy
|
||||
retry_policy_num_retries: Final[int | None] = get_num_retries_from_retry_policy(
|
||||
|
|
@ -1802,9 +1798,7 @@ def client(original_function):
|
|||
return litellm.stream_chunk_builder(chunks, messages=kwargs.get("messages", None))
|
||||
else:
|
||||
# RETURN RESULT
|
||||
update_response_metadata: _ResponseMetadataUpdater = getattr(
|
||||
sys.modules[__name__], "update_response_metadata"
|
||||
)
|
||||
update_response_metadata: _ResponseMetadataUpdater = litellm_utils.update_response_metadata
|
||||
update_response_metadata(
|
||||
result=result,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -1845,9 +1839,7 @@ def client(original_function):
|
|||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
_update_response_metadata: Final[_ResponseMetadataUpdater] = getattr(
|
||||
sys.modules[__name__], "update_response_metadata"
|
||||
)
|
||||
_update_response_metadata: Final[_ResponseMetadataUpdater] = litellm_utils.update_response_metadata
|
||||
_update_response_metadata(
|
||||
result=result,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -1877,8 +1869,8 @@ def client(original_function):
|
|||
if call_type == CallTypes.completion.value:
|
||||
num_retries = kwargs.get("num_retries", None) or litellm.num_retries or None
|
||||
if kwargs.get("retry_policy", None):
|
||||
get_num_retries_from_retry_policy: Callable[..., int | None] = getattr(
|
||||
sys.modules[__name__], "get_num_retries_from_retry_policy"
|
||||
get_num_retries_from_retry_policy: Callable[..., int | None] = (
|
||||
litellm_utils.get_num_retries_from_retry_policy
|
||||
)
|
||||
reset_retry_policy = litellm_utils.reset_retry_policy
|
||||
num_retries = get_num_retries_from_retry_policy(
|
||||
|
|
@ -1916,9 +1908,7 @@ def client(original_function):
|
|||
elif call_type == CallTypes.responses.value:
|
||||
num_retries = kwargs.get("num_retries", None) or litellm.num_retries or None
|
||||
if kwargs.get("retry_policy", None):
|
||||
get_num_retries_from_retry_policy = getattr(
|
||||
sys.modules[__name__], "get_num_retries_from_retry_policy"
|
||||
)
|
||||
get_num_retries_from_retry_policy = litellm_utils.get_num_retries_from_retry_policy
|
||||
reset_retry_policy = litellm_utils.reset_retry_policy
|
||||
num_retries = get_num_retries_from_retry_policy(
|
||||
exception=e,
|
||||
|
|
@ -1955,9 +1945,7 @@ def client(original_function):
|
|||
print_args_passed_to_litellm(original_function, args, kwargs)
|
||||
start_time: Final = datetime.datetime.now()
|
||||
result = None
|
||||
_update_response_metadata: Final[_ResponseMetadataUpdater] = getattr(
|
||||
sys.modules[__name__], "update_response_metadata"
|
||||
)
|
||||
_update_response_metadata: Final[_ResponseMetadataUpdater] = litellm_utils.update_response_metadata
|
||||
logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None)
|
||||
LLMCachingHandler: Final = _get_cached_llm_caching_handler()
|
||||
_llm_caching_handler: Final[LLMCachingHandler] = LLMCachingHandler(
|
||||
|
|
@ -3496,7 +3484,8 @@ def get_optional_params_transcription(
|
|||
|
||||
passed_params.pop("OPENAI_TRANSCRIPTION_PARAMS")
|
||||
custom_llm_provider = passed_params.pop("custom_llm_provider")
|
||||
drop_params = normalize_drop_params(passed_params.pop("drop_params"))
|
||||
passed_params.pop("drop_params")
|
||||
drop_params = normalize_drop_params(drop_params)
|
||||
special_params: Final[Mapping[str, object]] = passed_params.pop("kwargs")
|
||||
for k, v in special_params.items():
|
||||
passed_params[k] = v
|
||||
|
|
@ -3597,16 +3586,18 @@ def get_optional_params_image_gen(
|
|||
additional_drop_params: list | None = None,
|
||||
provider_config: BaseImageGenerationConfig | None = None,
|
||||
drop_params: bool | None = None,
|
||||
**kwargs,
|
||||
**kwargs: object,
|
||||
):
|
||||
# retrieve all parameters passed to the function
|
||||
passed_params: Final = locals()
|
||||
model = passed_params.pop("model", None)
|
||||
custom_llm_provider = passed_params.pop("custom_llm_provider")
|
||||
provider_config = passed_params.pop("provider_config", None)
|
||||
drop_params = normalize_drop_params(passed_params.pop("drop_params", None))
|
||||
passed_params.pop("model", None)
|
||||
passed_params.pop("custom_llm_provider")
|
||||
passed_params.pop("provider_config", None)
|
||||
passed_params.pop("drop_params", None)
|
||||
drop_params = normalize_drop_params(drop_params)
|
||||
additional_drop_params = passed_params.pop("additional_drop_params", None)
|
||||
special_params: Final[Mapping[str, object]] = passed_params.pop("kwargs")
|
||||
passed_params.pop("kwargs")
|
||||
special_params: Final[Mapping[str, object]] = kwargs
|
||||
for k, v in special_params.items():
|
||||
if (
|
||||
k.startswith("aws_")
|
||||
|
|
@ -3725,18 +3716,18 @@ def get_optional_params_embeddings(
|
|||
**kwargs,
|
||||
):
|
||||
# Lazy load get_supported_openai_params
|
||||
get_supported_openai_params: Final[_SupportedOpenAIParamsGetter] = getattr(
|
||||
sys.modules[__name__], "get_supported_openai_params"
|
||||
)
|
||||
get_supported_openai_params: Final[_SupportedOpenAIParamsGetter] = litellm_utils.get_supported_openai_params
|
||||
|
||||
# retrieve all parameters passed to the function
|
||||
passed_params: Final = locals()
|
||||
custom_llm_provider = passed_params.pop("custom_llm_provider", None)
|
||||
special_params: Final = passed_params.pop("kwargs")
|
||||
|
||||
drop_params = normalize_drop_params(passed_params.pop("drop_params", None))
|
||||
additional_drop_params = passed_params.pop("additional_drop_params", None)
|
||||
allowed_openai_params = passed_params.pop("allowed_openai_params", None) or []
|
||||
passed_params.pop("drop_params", None)
|
||||
drop_params = normalize_drop_params(drop_params)
|
||||
passed_params.pop("additional_drop_params", None)
|
||||
passed_params.pop("allowed_openai_params", None)
|
||||
allowed_openai_params = allowed_openai_params or []
|
||||
# Remove function objects from passed_params to avoid JSON serialization errors
|
||||
passed_params.pop("get_supported_openai_params", None)
|
||||
|
||||
|
|
@ -4502,9 +4493,7 @@ def get_optional_params(
|
|||
message=f"{custom_llm_provider} does not support parameters: {list(unsupported_params.keys())}, for model={model}. To drop these, set `litellm.drop_params=True` or for proxy:\n\n`litellm_settings:\n drop_params: true`\n. \n If you want to use these params dynamically send allowed_openai_params={list(unsupported_params.keys())} in your request.",
|
||||
)
|
||||
|
||||
get_supported_openai_params: Final[_SupportedOpenAIParamsGetter] = getattr(
|
||||
sys.modules[__name__], "get_supported_openai_params"
|
||||
)
|
||||
get_supported_openai_params: Final[_SupportedOpenAIParamsGetter] = litellm_utils.get_supported_openai_params
|
||||
supported_params = get_supported_openai_params(
|
||||
model=model, custom_llm_provider=custom_llm_provider, base_model=base_model
|
||||
)
|
||||
|
|
@ -4675,7 +4664,7 @@ def get_optional_params(
|
|||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "bedrock":
|
||||
bedrock_model_info: Final[type[BedrockModelInfo]] = getattr(sys.modules[__name__], "BedrockModelInfo")
|
||||
bedrock_model_info: Final[type[BedrockModelInfo]] = litellm_utils.BedrockModelInfo
|
||||
bedrock_route: Final = bedrock_model_info.get_bedrock_route(model)
|
||||
bedrock_base_model: Final = bedrock_model_info.get_base_model(model)
|
||||
if bedrock_route == "converse" or bedrock_route == "converse_like":
|
||||
|
|
@ -4996,8 +4985,8 @@ def get_optional_params(
|
|||
|
||||
# Apply nested drops from additional_drop_params
|
||||
if additional_drop_params:
|
||||
is_nested_path: Final[_NestedPathChecker] = getattr(sys.modules[__name__], "is_nested_path")
|
||||
delete_nested_value: Final[_NestedValueDeleter] = getattr(sys.modules[__name__], "delete_nested_value")
|
||||
is_nested_path: Final[_NestedPathChecker] = litellm_utils.is_nested_path
|
||||
delete_nested_value: Final[_NestedValueDeleter] = litellm_utils.delete_nested_value
|
||||
nested_paths: Final = [p for p in additional_drop_params if is_nested_path(p)]
|
||||
for path in nested_paths:
|
||||
optional_params = delete_nested_value(optional_params, path)
|
||||
|
|
@ -7328,7 +7317,7 @@ class TextCompletionStreamWrapper:
|
|||
except StopIteration:
|
||||
raise StopIteration
|
||||
except Exception as e:
|
||||
exception_type: Final = getattr(sys.modules[__name__], "exception_type")
|
||||
exception_type: Final = litellm_utils.exception_type
|
||||
raise exception_type(
|
||||
model=self.model,
|
||||
custom_llm_provider=self.custom_llm_provider or "",
|
||||
|
|
@ -9974,7 +9963,7 @@ def get_end_user_id_for_cost_tracking(
|
|||
|
||||
service_type: "litellm_logging" or "prometheus" - used to allow prometheus only disable cost tracking.
|
||||
"""
|
||||
get_litellm_metadata_from_kwargs: Final = getattr(sys.modules[__name__], "get_litellm_metadata_from_kwargs")
|
||||
get_litellm_metadata_from_kwargs: Final = litellm_utils.get_litellm_metadata_from_kwargs
|
||||
_metadata: Final = cast(dict, get_litellm_metadata_from_kwargs(dict(litellm_params=litellm_params)))
|
||||
|
||||
end_user_id: Final = cast(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue