mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/elated-margulis-7f300f
This commit is contained in:
commit
560a4ac891
1606 changed files with 37539 additions and 36732 deletions
|
|
@ -47,7 +47,7 @@ When you fix violations gated by `ruff-strict-budget.json`, `type-discipline-bud
|
|||
|
||||
If you're trying to create a new function that relies on untyped stuff, instead of adding more Any's and pushing `reportAny` / `reportExplicitAny` closer to their basedpyright ceilings, just validate it in the caller with Pydantic (a model or `TypeAdapter` that returns the typed thing or raises will do) and then pass the now typed variable in
|
||||
|
||||
If you get an LIT001 or LIT002 fail, refactor the code to follow functional programming best practices rather than introducing mutable data structures. For example, build values in one shot with comprehensions or generators wrapped in `tuple()` / `frozenset()` instead of seeding an empty `list`/`dict`/`set` and mutating it over time. Ideally `# mutable-ok` is never used; reach for it only as a genuine last resort when an immutable rewrite is truly impossible, and always pair it with a real reason
|
||||
If you get an LIT001 or LIT002 fail, refactor the code to follow functional programming best practices rather than introducing mutable data structures. For example, build values in one shot with comprehensions or generators wrapped in `tuple()` / `frozenset()` instead of seeding an empty `list`/`dict`/`set` and mutating it over time. Ideally, `# mutable-ok` is never used; reach for it only as a genuine last resort when an immutable rewrite is truly impossible, and always pair it with a real reason
|
||||
|
||||
Every lint or type suppression must name the exact rule inside brackets and carry a reason comment, e.g. `# pyright: ignore[reportArgumentType] # stubs lack async overload` or `# noqa: TID251 # <reason>`. `# type: ignore` is banned (LIT009): pyrightconfig.json sets `enableTypeIgnoreComments` to false, so it silently does nothing
|
||||
|
||||
|
|
@ -75,6 +75,7 @@ Follow these coding conventions for new/updated code (a three-line fix in a lega
|
|||
- Never-nester: early returns over deep nesting
|
||||
- Don't throw; model failures as values (One function (e.g., raise_public) maps error union to existing public exception contracts via exhaustive match + assert_never)
|
||||
- No mutation; don't reassign variables, global or local. Instead of mutable lists and dicts, prefer tuples, frozen dataclasses (with slots=True), etc.
|
||||
- Annotate every variable with `: Final` (LIT010). Unpacking and walrus targets cannot carry the annotation, so they are implicitly final. Don't rebind them. Never rebind or mutate function parameters (LIT011); `self`/`cls` attribute stores are the exception. If rebinding or in-place mutation is truly unavoidable, suppress with `# rebind-ok: <reason>` explaining why
|
||||
- Use dependency injection
|
||||
- Fully typed; no `Any` or coarse types like `dict[str, Any]` or just `dict`. Every function parameter must be strongly typed
|
||||
- Use tagged unions + match
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@
|
|||
"limit": 11
|
||||
},
|
||||
"reportGeneralTypeIssues": {
|
||||
"limit": 227
|
||||
"limit": 157
|
||||
},
|
||||
"reportIncompatibleMethodOverride": {
|
||||
"limit": 77
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.81"
|
||||
version = "0.4.82"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.81"
|
||||
version = "0.4.82"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -27,18 +27,19 @@ if os.getenv("LITELLM_MODE", "DEV") == "DEV":
|
|||
_dotenv.load_dotenv(override=_dev_env_hot_reload_enabled())
|
||||
|
||||
from typing import (
|
||||
Callable,
|
||||
List,
|
||||
Optional,
|
||||
Dict,
|
||||
Union,
|
||||
Any,
|
||||
Literal,
|
||||
Callable,
|
||||
Dict,
|
||||
Final,
|
||||
get_args,
|
||||
TYPE_CHECKING,
|
||||
Tuple,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
overload,
|
||||
Tuple,
|
||||
Type,
|
||||
TYPE_CHECKING,
|
||||
Union,
|
||||
)
|
||||
from litellm.types.integrations.datadog import DatadogInitParams
|
||||
from litellm.types.integrations.newrelic import NewRelicInitParams
|
||||
|
|
@ -681,12 +682,12 @@ def is_bedrock_pricing_only_model(key: str) -> bool:
|
|||
bool: True if the key matches the Bedrock pattern, False otherwise.
|
||||
"""
|
||||
# Regex to match 'bedrock/<region>/<model>'
|
||||
bedrock_pattern = re.compile(r"^bedrock/[a-zA-Z0-9_-]+/.+$")
|
||||
bedrock_pattern: Final = re.compile(r"^bedrock/[a-zA-Z0-9_-]+/.+$")
|
||||
|
||||
if "month-commitment" in key:
|
||||
return True
|
||||
|
||||
is_match = bedrock_pattern.match(key)
|
||||
is_match: Final = bedrock_pattern.match(key)
|
||||
return is_match is not None
|
||||
|
||||
|
||||
|
|
@ -704,7 +705,7 @@ def is_openai_finetune_model(key: str) -> bool:
|
|||
|
||||
|
||||
def add_known_models(model_cost_map: Optional[Dict] = None):
|
||||
_map = model_cost_map if model_cost_map is not None else model_cost
|
||||
_map: Final = model_cost_map if model_cost_map is not None else model_cost
|
||||
for key, value in _map.items():
|
||||
if value.get("litellm_provider") == "openai" and not is_openai_finetune_model(key):
|
||||
open_ai_chat_completion_models.add(key)
|
||||
|
|
@ -2140,11 +2141,11 @@ def __getattr__(name: str) -> Any:
|
|||
# Use cached registry from _lazy_imports instead of importing tuples every time
|
||||
from ._lazy_imports import _get_lazy_import_registry
|
||||
|
||||
registry = _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:
|
||||
handler_func = registry[name]
|
||||
handler_func: Final = registry[name]
|
||||
return handler_func(name)
|
||||
|
||||
# Lazy load encoding from main.py to avoid heavy tiktoken import
|
||||
|
|
@ -2197,7 +2198,7 @@ def __getattr__(name: str) -> Any:
|
|||
return _globals["openaiOSeriesConfig"]
|
||||
|
||||
# Lazy load other config instances
|
||||
_config_instances = {
|
||||
_config_instances: Final = {
|
||||
"openAIGPTConfig": "OpenAIGPTConfig",
|
||||
"openAIGPTAudioConfig": "OpenAIGPTAudioConfig",
|
||||
"openAIGPT5Config": "OpenAIGPT5Config",
|
||||
|
|
@ -2239,7 +2240,7 @@ def __getattr__(name: str) -> Any:
|
|||
# Check if already cached
|
||||
if "priority_reservation_settings" not in _globals:
|
||||
# Import the class and instantiate it
|
||||
PriorityReservationSettings = __getattr__("PriorityReservationSettings")
|
||||
PriorityReservationSettings: Final = __getattr__("PriorityReservationSettings")
|
||||
_globals["priority_reservation_settings"] = PriorityReservationSettings()
|
||||
return _globals["priority_reservation_settings"]
|
||||
|
||||
|
|
@ -2251,7 +2252,7 @@ def __getattr__(name: str) -> Any:
|
|||
# Check if already cached
|
||||
if "logging_callback_manager" not in _globals:
|
||||
# Import the class and instantiate it
|
||||
LoggingCallbackManager = __getattr__("LoggingCallbackManager")
|
||||
LoggingCallbackManager: Final = __getattr__("LoggingCallbackManager")
|
||||
_globals["logging_callback_manager"] = LoggingCallbackManager()
|
||||
return _globals["logging_callback_manager"]
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,8 @@ asyncio task and cannot be injected via HTTP request bodies.
|
|||
"""
|
||||
|
||||
from contextvars import ContextVar
|
||||
from typing import Final
|
||||
|
||||
# When True, suppresses async logging and billing for internal sub-calls
|
||||
# (e.g., emulated file-search steps that make nested LLM calls).
|
||||
is_internal_call: ContextVar[bool] = ContextVar("is_internal_call", default=False)
|
||||
is_internal_call: Final[ContextVar[bool]] = ContextVar("is_internal_call", default=False)
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ until they're actually needed.
|
|||
import importlib
|
||||
import sys
|
||||
from collections.abc import Callable
|
||||
from typing import Any, cast
|
||||
from typing import Any, Final, cast
|
||||
|
||||
# Import all the data structures that define what can be lazy-loaded
|
||||
# These are just lists of names and maps of where to find them
|
||||
|
|
@ -233,7 +233,7 @@ def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], cate
|
|||
raise AttributeError(f"{category} lazy import: unknown attribute {name!r}")
|
||||
|
||||
# Step 2: Get the cache (where we store imported things)
|
||||
_globals = _get_litellm_globals()
|
||||
_globals: Final = _get_litellm_globals()
|
||||
|
||||
# Step 3: If we've already imported it, just return the cached version
|
||||
if name in _globals:
|
||||
|
|
@ -255,7 +255,7 @@ def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], cate
|
|||
|
||||
# Step 6: Get the actual attribute from the module
|
||||
# Example: getattr(utils_module, "ModelResponse") returns the ModelResponse class
|
||||
value = getattr(module, attr_name)
|
||||
value: Final = getattr(module, attr_name)
|
||||
|
||||
# Step 7: Cache it so we don't have to import again next time
|
||||
_globals[name] = value
|
||||
|
|
@ -339,7 +339,7 @@ def _lazy_import_utils_module(name: str) -> Any:
|
|||
raise AttributeError(f"Utils module lazy import: unknown attribute {name!r}")
|
||||
|
||||
# Get the cache (where we store imported things) - use utils globals
|
||||
_globals = _get_utils_globals()
|
||||
_globals: Final = _get_utils_globals()
|
||||
|
||||
# If we've already imported it, just return the cached version
|
||||
if name in _globals:
|
||||
|
|
@ -355,7 +355,7 @@ def _lazy_import_utils_module(name: str) -> Any:
|
|||
module = importlib.import_module(module_path)
|
||||
|
||||
# Get the actual attribute from the module
|
||||
value = getattr(module, attr_name)
|
||||
value: Final = getattr(module, attr_name)
|
||||
|
||||
# Cache it so we don't have to import again next time
|
||||
_globals[name] = value
|
||||
|
|
@ -379,15 +379,15 @@ def _lazy_import_llm_client_cache(name: str) -> Any:
|
|||
- "in_memory_llm_clients_cache" is a singleton instance of that class
|
||||
So we need custom logic to handle both cases.
|
||||
"""
|
||||
_globals = _get_litellm_globals()
|
||||
_globals: Final = _get_litellm_globals()
|
||||
|
||||
# If already cached, return it
|
||||
if name in _globals:
|
||||
return _globals[name]
|
||||
|
||||
# Import the class
|
||||
module = importlib.import_module("litellm.caching.llm_caching_handler")
|
||||
LLMClientCache = getattr(module, "LLMClientCache")
|
||||
module: Final = importlib.import_module("litellm.caching.llm_caching_handler")
|
||||
LLMClientCache: Final = getattr(module, "LLMClientCache")
|
||||
|
||||
# If they want the class itself, return it
|
||||
if name == "LLMClientCache":
|
||||
|
|
@ -396,7 +396,7 @@ def _lazy_import_llm_client_cache(name: str) -> Any:
|
|||
|
||||
# If they want the singleton instance, create it (only once)
|
||||
if name == "in_memory_llm_clients_cache":
|
||||
instance = LLMClientCache()
|
||||
instance: Final = LLMClientCache()
|
||||
_globals["in_memory_llm_clients_cache"] = instance
|
||||
return instance
|
||||
|
||||
|
|
@ -412,7 +412,7 @@ def _lazy_import_http_handlers(name: str) -> Any:
|
|||
- They need configuration (timeout, etc.) from the module globals
|
||||
- They use factory functions instead of direct instantiation
|
||||
"""
|
||||
_globals = _get_litellm_globals()
|
||||
_globals: Final = _get_litellm_globals()
|
||||
|
||||
if name == "module_level_aclient":
|
||||
# Create an async HTTP client using the factory function
|
||||
|
|
@ -420,11 +420,11 @@ def _lazy_import_http_handlers(name: str) -> Any:
|
|||
|
||||
# Get timeout from module config (if set)
|
||||
timeout = _globals.get("request_timeout")
|
||||
params = {"timeout": timeout, "client_alias": "module level aclient"}
|
||||
params: Final = {"timeout": timeout, "client_alias": "module level aclient"}
|
||||
|
||||
# Create the client instance
|
||||
provider_id = cast(Any, "litellm_module_level_client")
|
||||
async_client = get_async_httpx_client(
|
||||
provider_id: Final = cast(Any, "litellm_module_level_client")
|
||||
async_client: Final = get_async_httpx_client(
|
||||
llm_provider=provider_id,
|
||||
params=params,
|
||||
)
|
||||
|
|
@ -438,7 +438,7 @@ def _lazy_import_http_handlers(name: str) -> Any:
|
|||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
timeout = _globals.get("request_timeout")
|
||||
sync_client = HTTPHandler(timeout=timeout)
|
||||
sync_client: Final = HTTPHandler(timeout=timeout)
|
||||
|
||||
# Cache it
|
||||
_globals["module_level_client"] = sync_client
|
||||
|
|
|
|||
|
|
@ -5,21 +5,23 @@ This module contains all the name tuples and import maps used by the lazy import
|
|||
Separated from the handler functions for better organization.
|
||||
"""
|
||||
|
||||
from typing import Final
|
||||
|
||||
# Cost calculator names that support lazy loading via _lazy_import_cost_calculator
|
||||
COST_CALCULATOR_NAMES = (
|
||||
COST_CALCULATOR_NAMES: Final = (
|
||||
"completion_cost",
|
||||
"cost_per_token",
|
||||
"response_cost_calculator",
|
||||
)
|
||||
|
||||
# Litellm logging names that support lazy loading via _lazy_import_litellm_logging
|
||||
LITELLM_LOGGING_NAMES = (
|
||||
LITELLM_LOGGING_NAMES: Final = (
|
||||
"Logging",
|
||||
"modify_integration",
|
||||
)
|
||||
|
||||
# Utils names that support lazy loading via _lazy_import_utils
|
||||
UTILS_NAMES = (
|
||||
UTILS_NAMES: Final = (
|
||||
"exception_type",
|
||||
"get_optional_params",
|
||||
"get_response_string",
|
||||
|
|
@ -66,20 +68,20 @@ UTILS_NAMES = (
|
|||
)
|
||||
|
||||
# Token counter names that support lazy loading via _lazy_import_token_counter
|
||||
TOKEN_COUNTER_NAMES = ("get_modified_max_tokens",)
|
||||
TOKEN_COUNTER_NAMES: Final = ("get_modified_max_tokens",)
|
||||
|
||||
# LLM client cache names that support lazy loading via _lazy_import_llm_client_cache
|
||||
LLM_CLIENT_CACHE_NAMES = (
|
||||
LLM_CLIENT_CACHE_NAMES: Final = (
|
||||
"LLMClientCache",
|
||||
"in_memory_llm_clients_cache",
|
||||
)
|
||||
|
||||
# Bedrock type names that support lazy loading via _lazy_import_bedrock_types
|
||||
BEDROCK_TYPES_NAMES = ("COHERE_EMBEDDING_INPUT_TYPES",)
|
||||
BEDROCK_TYPES_NAMES: Final = ("COHERE_EMBEDDING_INPUT_TYPES",)
|
||||
|
||||
# Common types from litellm.types.utils that support lazy loading via
|
||||
# _lazy_import_types_utils
|
||||
TYPES_UTILS_NAMES = (
|
||||
TYPES_UTILS_NAMES: Final = (
|
||||
"ImageObject",
|
||||
"BudgetConfig",
|
||||
"all_litellm_params",
|
||||
|
|
@ -92,7 +94,7 @@ TYPES_UTILS_NAMES = (
|
|||
)
|
||||
|
||||
# Caching / cache classes that support lazy loading via _lazy_import_caching
|
||||
CACHING_NAMES = (
|
||||
CACHING_NAMES: Final = (
|
||||
"Cache",
|
||||
"DualCache",
|
||||
"RedisCache",
|
||||
|
|
@ -100,20 +102,20 @@ CACHING_NAMES = (
|
|||
)
|
||||
|
||||
# HTTP handler names that support lazy loading via _lazy_import_http_handlers
|
||||
HTTP_HANDLER_NAMES = (
|
||||
HTTP_HANDLER_NAMES: Final = (
|
||||
"module_level_aclient",
|
||||
"module_level_client",
|
||||
)
|
||||
|
||||
# Dotprompt integration names that support lazy loading via _lazy_import_dotprompt
|
||||
DOTPROMPT_NAMES = (
|
||||
DOTPROMPT_NAMES: Final = (
|
||||
"global_prompt_manager",
|
||||
"global_prompt_directory",
|
||||
"set_global_prompt_directory",
|
||||
)
|
||||
|
||||
# LLM config classes that support lazy loading via _lazy_import_llm_configs
|
||||
LLM_CONFIG_NAMES = (
|
||||
LLM_CONFIG_NAMES: Final = (
|
||||
"AmazonConverseConfig",
|
||||
"OpenAILikeChatConfig",
|
||||
"GaladrielChatConfig",
|
||||
|
|
@ -328,7 +330,7 @@ LLM_CONFIG_NAMES = (
|
|||
)
|
||||
|
||||
# Types that support lazy loading via _lazy_import_types
|
||||
TYPES_NAMES = (
|
||||
TYPES_NAMES: Final = (
|
||||
"GuardrailItem",
|
||||
"DefaultTeamSSOParams",
|
||||
"LiteLLM_UpperboundKeyGenerateParams",
|
||||
|
|
@ -344,14 +346,14 @@ TYPES_NAMES = (
|
|||
)
|
||||
|
||||
# LLM provider logic names that support lazy loading via _lazy_import_llm_provider_logic
|
||||
LLM_PROVIDER_LOGIC_NAMES = (
|
||||
LLM_PROVIDER_LOGIC_NAMES: Final = (
|
||||
"get_llm_provider",
|
||||
"remove_index_from_tool_calls",
|
||||
)
|
||||
|
||||
# Utils module names that support lazy loading via _lazy_import_utils_module
|
||||
# These are attributes accessed from litellm.utils module
|
||||
UTILS_MODULE_NAMES = (
|
||||
UTILS_MODULE_NAMES: Final = (
|
||||
"encoding",
|
||||
"BaseVectorStore",
|
||||
"CredentialAccessor",
|
||||
|
|
@ -423,7 +425,7 @@ UTILS_MODULE_NAMES = (
|
|||
)
|
||||
|
||||
# Import maps for registry pattern - reduces repetition
|
||||
_UTILS_IMPORT_MAP = {
|
||||
_UTILS_IMPORT_MAP: Final = {
|
||||
"exception_type": (".utils", "exception_type"),
|
||||
"get_optional_params": (".utils", "get_optional_params"),
|
||||
"get_response_string": (".utils", "get_response_string"),
|
||||
|
|
@ -478,13 +480,13 @@ _UTILS_IMPORT_MAP = {
|
|||
),
|
||||
}
|
||||
|
||||
_COST_CALCULATOR_IMPORT_MAP = {
|
||||
_COST_CALCULATOR_IMPORT_MAP: Final = {
|
||||
"completion_cost": (".cost_calculator", "completion_cost"),
|
||||
"cost_per_token": (".cost_calculator", "cost_per_token"),
|
||||
"response_cost_calculator": (".cost_calculator", "response_cost_calculator"),
|
||||
}
|
||||
|
||||
_TYPES_UTILS_IMPORT_MAP = {
|
||||
_TYPES_UTILS_IMPORT_MAP: Final = {
|
||||
"ImageObject": (".types.utils", "ImageObject"),
|
||||
"BudgetConfig": (".types.utils", "BudgetConfig"),
|
||||
"all_litellm_params": (".types.utils", "all_litellm_params"),
|
||||
|
|
@ -496,28 +498,28 @@ _TYPES_UTILS_IMPORT_MAP = {
|
|||
"GenericStreamingChunk": (".types.utils", "GenericStreamingChunk"),
|
||||
}
|
||||
|
||||
_TOKEN_COUNTER_IMPORT_MAP = {
|
||||
_TOKEN_COUNTER_IMPORT_MAP: Final = {
|
||||
"get_modified_max_tokens": (
|
||||
"litellm.litellm_core_utils.token_counter",
|
||||
"get_modified_max_tokens",
|
||||
),
|
||||
}
|
||||
|
||||
_BEDROCK_TYPES_IMPORT_MAP = {
|
||||
_BEDROCK_TYPES_IMPORT_MAP: Final = {
|
||||
"COHERE_EMBEDDING_INPUT_TYPES": (
|
||||
"litellm.types.llms.bedrock",
|
||||
"COHERE_EMBEDDING_INPUT_TYPES",
|
||||
),
|
||||
}
|
||||
|
||||
_CACHING_IMPORT_MAP = {
|
||||
_CACHING_IMPORT_MAP: Final = {
|
||||
"Cache": ("litellm.caching.caching", "Cache"),
|
||||
"DualCache": ("litellm.caching.caching", "DualCache"),
|
||||
"RedisCache": ("litellm.caching.caching", "RedisCache"),
|
||||
"InMemoryCache": ("litellm.caching.caching", "InMemoryCache"),
|
||||
}
|
||||
|
||||
_LITELLM_LOGGING_IMPORT_MAP = {
|
||||
_LITELLM_LOGGING_IMPORT_MAP: Final = {
|
||||
"Logging": ("litellm.litellm_core_utils.litellm_logging", "Logging"),
|
||||
"modify_integration": (
|
||||
"litellm.litellm_core_utils.litellm_logging",
|
||||
|
|
@ -525,7 +527,7 @@ _LITELLM_LOGGING_IMPORT_MAP = {
|
|||
),
|
||||
}
|
||||
|
||||
_DOTPROMPT_IMPORT_MAP = {
|
||||
_DOTPROMPT_IMPORT_MAP: Final = {
|
||||
"global_prompt_manager": (
|
||||
"litellm.integrations.dotprompt",
|
||||
"global_prompt_manager",
|
||||
|
|
@ -540,7 +542,7 @@ _DOTPROMPT_IMPORT_MAP = {
|
|||
),
|
||||
}
|
||||
|
||||
_TYPES_IMPORT_MAP = {
|
||||
_TYPES_IMPORT_MAP: Final = {
|
||||
"GuardrailItem": ("litellm.types.guardrails", "GuardrailItem"),
|
||||
"DefaultTeamSSOParams": (
|
||||
"litellm.types.proxy.management_endpoints.ui_sso",
|
||||
|
|
@ -569,7 +571,7 @@ _TYPES_IMPORT_MAP = {
|
|||
),
|
||||
}
|
||||
|
||||
_LLM_PROVIDER_LOGIC_IMPORT_MAP = {
|
||||
_LLM_PROVIDER_LOGIC_IMPORT_MAP: Final = {
|
||||
"get_llm_provider": (
|
||||
"litellm.litellm_core_utils.get_llm_provider_logic",
|
||||
"get_llm_provider",
|
||||
|
|
@ -580,7 +582,7 @@ _LLM_PROVIDER_LOGIC_IMPORT_MAP = {
|
|||
),
|
||||
}
|
||||
|
||||
_LLM_CONFIGS_IMPORT_MAP = {
|
||||
_LLM_CONFIGS_IMPORT_MAP: Final = {
|
||||
"AmazonConverseConfig": (
|
||||
".llms.bedrock.chat.converse_transformation",
|
||||
"AmazonConverseConfig",
|
||||
|
|
@ -1215,7 +1217,7 @@ _LLM_CONFIGS_IMPORT_MAP = {
|
|||
}
|
||||
|
||||
# Import map for utils module lazy imports
|
||||
_UTILS_MODULE_IMPORT_MAP = {
|
||||
_UTILS_MODULE_IMPORT_MAP: Final = {
|
||||
"encoding": ("litellm.main", "encoding"),
|
||||
"BaseVectorStore": (
|
||||
"litellm.integrations.vector_store_integrations.base_vector_store",
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import os
|
|||
import sys
|
||||
from datetime import datetime
|
||||
from logging import Formatter
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
|
|
@ -17,7 +17,7 @@ if set_verbose is True:
|
|||
"`litellm.set_verbose` is deprecated. Please set `os.environ['LITELLM_LOG'] = 'DEBUG'` for debug logs."
|
||||
)
|
||||
|
||||
_ENABLE_SECRET_REDACTION = os.getenv("LITELLM_DISABLE_REDACT_SECRETS", "").lower() != "true"
|
||||
_ENABLE_SECRET_REDACTION: Final = os.getenv("LITELLM_DISABLE_REDACT_SECRETS", "").lower() != "true"
|
||||
|
||||
|
||||
def _redact_string(value: str) -> str:
|
||||
|
|
@ -74,14 +74,14 @@ class SecretRedactionFilter(logging.Filter):
|
|||
return True
|
||||
|
||||
|
||||
_secret_filter = SecretRedactionFilter()
|
||||
_secret_filter: Final = SecretRedactionFilter()
|
||||
|
||||
|
||||
json_logs = bool(os.getenv("JSON_LOGS", False))
|
||||
# Create a handler for the logger (you may need to adapt this based on your needs)
|
||||
log_level = os.getenv("LITELLM_LOG", "DEBUG")
|
||||
numeric_level: str = getattr(logging, log_level.upper())
|
||||
handler = logging.StreamHandler()
|
||||
log_level: Final = os.getenv("LITELLM_LOG", "DEBUG")
|
||||
numeric_level: Final[str] = getattr(logging, log_level.upper())
|
||||
handler: Final = logging.StreamHandler()
|
||||
handler.setLevel(numeric_level)
|
||||
handler.addFilter(_secret_filter)
|
||||
|
||||
|
|
@ -94,10 +94,10 @@ def _try_parse_json_message(message: str) -> dict[str, Any] | None:
|
|||
"""
|
||||
if not message or not isinstance(message, str):
|
||||
return None
|
||||
msg_stripped = message.strip()
|
||||
msg_stripped: Final = message.strip()
|
||||
if not (msg_stripped.startswith("{") or msg_stripped.startswith("[")):
|
||||
return None
|
||||
parsed = safe_json_loads(message, default=None)
|
||||
parsed: Final = safe_json_loads(message, default=None)
|
||||
if parsed is None or not isinstance(parsed, dict):
|
||||
return None
|
||||
return parsed
|
||||
|
|
@ -144,7 +144,7 @@ def _get_standard_record_attrs() -> frozenset:
|
|||
return frozenset(logging.LogRecord("", 0, "", 0, "", (), None).__dict__.keys())
|
||||
|
||||
|
||||
_STANDARD_RECORD_ATTRS = _get_standard_record_attrs()
|
||||
_STANDARD_RECORD_ATTRS: Final = _get_standard_record_attrs()
|
||||
|
||||
|
||||
class JsonFormatter(Formatter):
|
||||
|
|
@ -153,12 +153,12 @@ class JsonFormatter(Formatter):
|
|||
|
||||
def formatTime(self, record, datefmt=None):
|
||||
# Use datetime to format the timestamp in ISO 8601 format
|
||||
dt = datetime.fromtimestamp(record.created)
|
||||
dt: Final = datetime.fromtimestamp(record.created)
|
||||
return dt.isoformat()
|
||||
|
||||
def format(self, record):
|
||||
message_str = record.getMessage()
|
||||
json_record: dict[str, Any] = {
|
||||
message_str: Final = record.getMessage()
|
||||
json_record: Final[dict[str, Any]] = {
|
||||
"message": message_str,
|
||||
"level": record.levelname,
|
||||
"timestamp": self.formatTime(record),
|
||||
|
|
@ -193,13 +193,13 @@ class JsonFormatter(Formatter):
|
|||
# Function to set up exception handlers for JSON logging
|
||||
def _setup_json_exception_handlers(formatter):
|
||||
# Create a handler with JSON formatting for exceptions
|
||||
error_handler = logging.StreamHandler()
|
||||
error_handler: Final = logging.StreamHandler()
|
||||
error_handler.setFormatter(formatter)
|
||||
error_handler.addFilter(_secret_filter)
|
||||
|
||||
# Setup excepthook for uncaught exceptions
|
||||
def json_excepthook(exc_type, exc_value, exc_traceback):
|
||||
record = logging.LogRecord(
|
||||
record: Final = logging.LogRecord(
|
||||
name="LiteLLM",
|
||||
level=logging.ERROR,
|
||||
pathname="",
|
||||
|
|
@ -217,10 +217,10 @@ def _setup_json_exception_handlers(formatter):
|
|||
import asyncio
|
||||
|
||||
def async_json_exception_handler(loop, context):
|
||||
exception = context.get("exception")
|
||||
exception: Final = context.get("exception")
|
||||
if exception:
|
||||
exc_type = type(exception)
|
||||
record = logging.LogRecord(
|
||||
exc_type: Final = type(exception)
|
||||
record: Final = logging.LogRecord(
|
||||
name="LiteLLM",
|
||||
level=logging.ERROR,
|
||||
pathname="",
|
||||
|
|
@ -243,7 +243,7 @@ if json_logs:
|
|||
handler.setFormatter(JsonFormatter())
|
||||
_setup_json_exception_handlers(JsonFormatter())
|
||||
else:
|
||||
formatter = logging.Formatter(
|
||||
formatter: Final = logging.Formatter(
|
||||
"\033[92m%(asctime)s - %(name)s:%(levelname)s\033[0m: %(filename)s:%(lineno)s - %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
|
|
@ -263,20 +263,20 @@ verbose_logger.addHandler(handler)
|
|||
def _suppress_loggers():
|
||||
"""Suppress noisy loggers at INFO level"""
|
||||
# Suppress httpx request logging at INFO level
|
||||
httpx_logger = logging.getLogger("httpx")
|
||||
httpx_logger: Final = logging.getLogger("httpx")
|
||||
httpx_logger.setLevel(logging.WARNING)
|
||||
|
||||
# Suppress APScheduler logging at INFO level
|
||||
apscheduler_executors_logger = logging.getLogger("apscheduler.executors.default")
|
||||
apscheduler_executors_logger: Final = logging.getLogger("apscheduler.executors.default")
|
||||
apscheduler_executors_logger.setLevel(logging.WARNING)
|
||||
apscheduler_scheduler_logger = logging.getLogger("apscheduler.scheduler")
|
||||
apscheduler_scheduler_logger: Final = logging.getLogger("apscheduler.scheduler")
|
||||
apscheduler_scheduler_logger.setLevel(logging.WARNING)
|
||||
|
||||
|
||||
# Call the suppression function
|
||||
_suppress_loggers()
|
||||
|
||||
ALL_LOGGERS = [
|
||||
ALL_LOGGERS: Final = [
|
||||
logging.getLogger(),
|
||||
verbose_logger,
|
||||
verbose_router_logger,
|
||||
|
|
@ -293,11 +293,11 @@ def _get_loggers_to_initialize():
|
|||
"""
|
||||
import litellm
|
||||
|
||||
loggers = list(ALL_LOGGERS)
|
||||
loggers: Final = list(ALL_LOGGERS)
|
||||
|
||||
# Add langfuse logger if langfuse is being used as a callback
|
||||
langfuse_callbacks = {"langfuse", "langfuse_otel"}
|
||||
all_callbacks = set(litellm.success_callback + litellm.failure_callback)
|
||||
langfuse_callbacks: Final = {"langfuse", "langfuse_otel"}
|
||||
all_callbacks: Final = set(litellm.success_callback + litellm.failure_callback)
|
||||
if langfuse_callbacks & all_callbacks:
|
||||
loggers.append(logging.getLogger("langfuse"))
|
||||
|
||||
|
|
@ -325,12 +325,12 @@ def _get_uvicorn_json_log_config():
|
|||
This ensures that uvicorn's access logs, error logs, and all application logs
|
||||
are formatted as JSON when json_logs is enabled.
|
||||
"""
|
||||
json_formatter_class = "litellm._logging.JsonFormatter"
|
||||
json_formatter_class: Final = "litellm._logging.JsonFormatter"
|
||||
|
||||
# Use the module-level log_level variable for consistency
|
||||
uvicorn_log_level = log_level.upper()
|
||||
uvicorn_log_level: Final = log_level.upper()
|
||||
|
||||
log_config = {
|
||||
log_config: Final = {
|
||||
"version": 1,
|
||||
"disable_existing_loggers": False,
|
||||
"formatters": {
|
||||
|
|
@ -384,7 +384,7 @@ def _turn_on_json():
|
|||
|
||||
- Adds a JSON formatter to all loggers
|
||||
"""
|
||||
handler = logging.StreamHandler()
|
||||
handler: Final = logging.StreamHandler()
|
||||
handler.setFormatter(JsonFormatter())
|
||||
_initialize_loggers_with_handler(handler)
|
||||
# Set up exception handlers
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import json
|
|||
# s/o [@Frank Colson](https://www.linkedin.com/in/frank-colson-422b9b183/) for this redis implementation
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from typing import Final
|
||||
|
||||
import redis # type: ignore
|
||||
import redis.asyncio as async_redis # type: ignore
|
||||
|
|
@ -32,20 +33,20 @@ from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
|||
|
||||
from ._logging import verbose_logger
|
||||
|
||||
AZURE_REDIS_SCOPE = "https://redis.azure.com/.default"
|
||||
AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default"
|
||||
|
||||
|
||||
def _get_redis_kwargs():
|
||||
arg_spec = inspect.getfullargspec(redis.Redis)
|
||||
arg_spec: Final = inspect.getfullargspec(redis.Redis)
|
||||
|
||||
# Only allow primitive arguments
|
||||
exclude_args = {
|
||||
exclude_args: Final = {
|
||||
"self",
|
||||
"connection_pool",
|
||||
"retry",
|
||||
}
|
||||
|
||||
include_args = {
|
||||
include_args: Final = {
|
||||
"url",
|
||||
"redis_connect_func",
|
||||
"gcp_service_account",
|
||||
|
|
@ -56,7 +57,7 @@ def _get_redis_kwargs():
|
|||
"azure_client_secret",
|
||||
}
|
||||
|
||||
available_args = {x for x in arg_spec.args if x not in exclude_args} | include_args
|
||||
available_args: Final = {x for x in arg_spec.args if x not in exclude_args} | include_args
|
||||
|
||||
return available_args
|
||||
|
||||
|
|
@ -92,9 +93,9 @@ def _get_redis_url_kwargs(client: type | None = None) -> tuple[str, ...]:
|
|||
"""
|
||||
if client is None:
|
||||
client = redis.Redis
|
||||
connection_cls = async_redis.Connection if client is async_redis.Redis else redis.Connection
|
||||
connection_cls: Final = async_redis.Connection if client is async_redis.Redis else redis.Connection
|
||||
|
||||
exclude_args = frozenset(
|
||||
exclude_args: Final = frozenset(
|
||||
{
|
||||
"self",
|
||||
"connection_pool",
|
||||
|
|
@ -103,7 +104,7 @@ def _get_redis_url_kwargs(client: type | None = None) -> tuple[str, ...]:
|
|||
)
|
||||
|
||||
# Only allow primitive arguments
|
||||
include_args = ("url", "max_connections")
|
||||
include_args: Final = ("url", "max_connections")
|
||||
|
||||
return tuple(x for x in _init_arg_names(connection_cls) if x not in exclude_args) + include_args
|
||||
|
||||
|
|
@ -111,10 +112,10 @@ def _get_redis_url_kwargs(client: type | None = None) -> tuple[str, ...]:
|
|||
def _get_redis_cluster_kwargs(client=None):
|
||||
if client is None:
|
||||
client = redis.Redis.from_url
|
||||
arg_spec = inspect.getfullargspec(redis.RedisCluster)
|
||||
arg_spec: Final = inspect.getfullargspec(redis.RedisCluster)
|
||||
|
||||
# Only allow primitive arguments
|
||||
exclude_args = {"self", "connection_pool", "retry", "host", "port", "startup_nodes"}
|
||||
exclude_args: Final = {"self", "connection_pool", "retry", "host", "port", "startup_nodes"}
|
||||
|
||||
available_args = {x for x in arg_spec.args if x not in exclude_args}
|
||||
available_args |= {
|
||||
|
|
@ -142,15 +143,15 @@ def _get_redis_cluster_kwargs(client=None):
|
|||
|
||||
|
||||
def _get_redis_env_kwarg_mapping():
|
||||
PREFIX = "REDIS_"
|
||||
PREFIX: Final = "REDIS_"
|
||||
|
||||
return {f"{PREFIX}{x.upper()}": x for x in _get_redis_kwargs()}
|
||||
|
||||
|
||||
def _redis_kwargs_from_environment():
|
||||
mapping = _get_redis_env_kwarg_mapping()
|
||||
mapping: Final = _get_redis_env_kwarg_mapping()
|
||||
|
||||
return_dict = {}
|
||||
return_dict: Final = {}
|
||||
for k, v in mapping.items():
|
||||
value = get_secret(k, default_value=None) # type: ignore
|
||||
if value is not None:
|
||||
|
|
@ -183,7 +184,7 @@ def create_gcp_iam_redis_connect_func(
|
|||
|
||||
self._parser.on_connect(self)
|
||||
|
||||
auth_args = (_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:
|
||||
|
|
@ -224,9 +225,9 @@ def _build_azure_credential(
|
|||
"azure-identity is required for Azure AD Redis authentication. Install it with: pip install azure-identity"
|
||||
)
|
||||
|
||||
_client_id = azure_client_id or os.environ.get("AZURE_CLIENT_ID")
|
||||
_tenant_id = azure_tenant_id or os.environ.get("AZURE_TENANT_ID")
|
||||
_client_secret = azure_client_secret or os.environ.get("AZURE_CLIENT_SECRET")
|
||||
_client_id: Final = azure_client_id or os.environ.get("AZURE_CLIENT_ID")
|
||||
_tenant_id: Final = azure_tenant_id or os.environ.get("AZURE_TENANT_ID")
|
||||
_client_secret: Final = azure_client_secret or os.environ.get("AZURE_CLIENT_SECRET")
|
||||
|
||||
if _client_id and _tenant_id and _client_secret:
|
||||
return ClientSecretCredential(
|
||||
|
|
@ -253,12 +254,12 @@ def _generate_azure_ad_redis_token(
|
|||
(``AzureADCredentialProvider``) keep the credential alive across
|
||||
connections so the Azure SDK's internal cache + silent refresh apply.
|
||||
"""
|
||||
credential = _build_azure_credential(
|
||||
credential: Final = _build_azure_credential(
|
||||
azure_client_id=azure_client_id,
|
||||
azure_tenant_id=azure_tenant_id,
|
||||
azure_client_secret=azure_client_secret,
|
||||
)
|
||||
token = credential.get_token(AZURE_REDIS_SCOPE)
|
||||
token: Final = credential.get_token(AZURE_REDIS_SCOPE)
|
||||
return token.token
|
||||
|
||||
|
||||
|
|
@ -274,7 +275,7 @@ def create_azure_ad_redis_connect_func(
|
|||
closure) and reused across connections — the Azure SDK handles token caching
|
||||
and silent renewal internally. Only ``get_token`` is called per connection.
|
||||
"""
|
||||
credential = _build_azure_credential(
|
||||
credential: Final = _build_azure_credential(
|
||||
azure_client_id=azure_client_id,
|
||||
azure_tenant_id=azure_tenant_id,
|
||||
azure_client_secret=azure_client_secret,
|
||||
|
|
@ -290,11 +291,11 @@ def create_azure_ad_redis_connect_func(
|
|||
|
||||
self._parser.on_connect(self)
|
||||
|
||||
access_token = credential.get_token(AZURE_REDIS_SCOPE).token
|
||||
access_token: Final = credential.get_token(AZURE_REDIS_SCOPE).token
|
||||
|
||||
# Only include username when explicitly set — sending AUTH "" <token>
|
||||
# is invalid for most ACL-configured Azure Redis instances.
|
||||
username = os.environ.get("REDIS_USERNAME", "")
|
||||
username: Final = os.environ.get("REDIS_USERNAME", "")
|
||||
if username:
|
||||
auth_args = (username, access_token)
|
||||
else:
|
||||
|
|
@ -353,23 +354,23 @@ def _get_redis_client_logic(**env_overrides):
|
|||
value = get_secret(v) # type: ignore
|
||||
env_overrides[k] = value
|
||||
|
||||
environment_kwargs = _redis_kwargs_from_environment()
|
||||
environment_kwargs: Final = _redis_kwargs_from_environment()
|
||||
|
||||
# An explicitly configured connection target outranks REDIS_URL from the
|
||||
# environment. Without this, the url branch below strips the caller's
|
||||
# host/port/password and silently connects to whatever REDIS_URL names.
|
||||
caller_named_a_target = any(
|
||||
caller_named_a_target: Final = any(
|
||||
env_overrides.get(key) is not None for key in ("host", "startup_nodes", "sentinel_nodes")
|
||||
)
|
||||
if caller_named_a_target and env_overrides.get("url") is None:
|
||||
environment_kwargs.pop("url", None)
|
||||
|
||||
redis_kwargs = {
|
||||
redis_kwargs: Final = {
|
||||
**environment_kwargs,
|
||||
**env_overrides,
|
||||
}
|
||||
|
||||
_startup_nodes: str | list | None = redis_kwargs.get("startup_nodes", None) or get_secret( # type: ignore
|
||||
_startup_nodes: Final[str | list | None] = redis_kwargs.get("startup_nodes", None) or get_secret( # type: ignore
|
||||
"REDIS_CLUSTER_NODES"
|
||||
)
|
||||
|
||||
|
|
@ -380,21 +381,21 @@ def _get_redis_client_logic(**env_overrides):
|
|||
elif _startup_nodes is None:
|
||||
redis_kwargs.pop("startup_nodes", None)
|
||||
|
||||
_sentinel_nodes: str | list | None = redis_kwargs.get("sentinel_nodes", None) or get_secret( # type: ignore
|
||||
_sentinel_nodes: Final[str | list | None] = redis_kwargs.get("sentinel_nodes", None) or get_secret( # type: ignore
|
||||
"REDIS_SENTINEL_NODES"
|
||||
)
|
||||
|
||||
if _sentinel_nodes is not None and isinstance(_sentinel_nodes, str):
|
||||
redis_kwargs["sentinel_nodes"] = json.loads(_sentinel_nodes)
|
||||
|
||||
_sentinel_password: str | None = redis_kwargs.get("sentinel_password", None) or get_secret_str(
|
||||
_sentinel_password: Final[str | None] = redis_kwargs.get("sentinel_password", None) or get_secret_str(
|
||||
"REDIS_SENTINEL_PASSWORD"
|
||||
)
|
||||
|
||||
if _sentinel_password is not None:
|
||||
redis_kwargs["sentinel_password"] = _sentinel_password
|
||||
|
||||
_service_name: str | None = redis_kwargs.get("service_name", None) or get_secret( # type: ignore
|
||||
_service_name: Final[str | None] = redis_kwargs.get("service_name", None) or get_secret( # type: ignore
|
||||
"REDIS_SERVICE_NAME"
|
||||
)
|
||||
|
||||
|
|
@ -402,8 +403,8 @@ def _get_redis_client_logic(**env_overrides):
|
|||
redis_kwargs["service_name"] = _service_name
|
||||
|
||||
# Handle GCP IAM authentication
|
||||
_gcp_service_account = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT")
|
||||
_gcp_ssl_ca_certs = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str("REDIS_GCP_SSL_CA_CERTS")
|
||||
_gcp_service_account: Final = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT")
|
||||
_gcp_ssl_ca_certs: Final = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str("REDIS_GCP_SSL_CA_CERTS")
|
||||
|
||||
if _gcp_service_account is not None:
|
||||
verbose_logger.debug("Setting up GCP IAM authentication for Redis with service account.")
|
||||
|
|
@ -422,9 +423,9 @@ def _get_redis_client_logic(**env_overrides):
|
|||
redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs
|
||||
|
||||
# Handle Azure AD authentication (after GCP IAM block)
|
||||
_azure_redis_ad_token = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN")
|
||||
_azure_redis_ad_token: Final = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN")
|
||||
|
||||
_azure_ad_enabled = _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true"
|
||||
_azure_ad_enabled: Final = _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true"
|
||||
|
||||
if _azure_ad_enabled and _gcp_service_account is not None:
|
||||
verbose_logger.warning(
|
||||
|
|
@ -433,9 +434,9 @@ def _get_redis_client_logic(**env_overrides):
|
|||
)
|
||||
|
||||
if _azure_ad_enabled and _gcp_service_account is None:
|
||||
_azure_client_id = redis_kwargs.get("azure_client_id") or get_secret_str("AZURE_CLIENT_ID")
|
||||
_azure_tenant_id = redis_kwargs.get("azure_tenant_id") or get_secret_str("AZURE_TENANT_ID")
|
||||
_azure_client_secret = redis_kwargs.get("azure_client_secret") or get_secret_str("AZURE_CLIENT_SECRET")
|
||||
_azure_client_id: Final = redis_kwargs.get("azure_client_id") or get_secret_str("AZURE_CLIENT_ID")
|
||||
_azure_tenant_id: Final = redis_kwargs.get("azure_tenant_id") or get_secret_str("AZURE_TENANT_ID")
|
||||
_azure_client_secret: Final = redis_kwargs.get("azure_client_secret") or get_secret_str("AZURE_CLIENT_SECRET")
|
||||
|
||||
verbose_logger.debug("Setting up Azure AD authentication for Redis.")
|
||||
redis_kwargs["redis_connect_func"] = create_azure_ad_redis_connect_func(
|
||||
|
|
@ -480,7 +481,7 @@ def _get_redis_client_logic(**env_overrides):
|
|||
|
||||
|
||||
def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
|
||||
_redis_cluster_nodes_in_env: str | None = get_secret("REDIS_CLUSTER_NODES") # type: ignore
|
||||
_redis_cluster_nodes_in_env: Final[str | None] = get_secret("REDIS_CLUSTER_NODES") # type: ignore
|
||||
if _redis_cluster_nodes_in_env is not None:
|
||||
try:
|
||||
redis_kwargs["startup_nodes"] = json.loads(_redis_cluster_nodes_in_env)
|
||||
|
|
@ -492,13 +493,13 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
|
|||
verbose_logger.debug("init_redis_cluster: startup nodes are being initialized.")
|
||||
from redis.cluster import ClusterNode
|
||||
|
||||
args = _get_redis_cluster_kwargs()
|
||||
cluster_kwargs = {}
|
||||
args: Final = _get_redis_cluster_kwargs()
|
||||
cluster_kwargs: Final = {}
|
||||
for arg in redis_kwargs:
|
||||
if arg in args:
|
||||
cluster_kwargs[arg] = redis_kwargs[arg]
|
||||
|
||||
new_startup_nodes: list[ClusterNode] = []
|
||||
new_startup_nodes: Final[list[ClusterNode]] = []
|
||||
|
||||
for item in redis_kwargs["startup_nodes"]:
|
||||
new_startup_nodes.append(ClusterNode(**item))
|
||||
|
|
@ -508,8 +509,8 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
|
|||
|
||||
|
||||
def _get_redis_sentinel_connection_kwargs(redis_kwargs: dict) -> dict:
|
||||
connection_kwargs = {}
|
||||
args = _get_redis_kwargs()
|
||||
connection_kwargs: Final = {}
|
||||
args: Final = _get_redis_kwargs()
|
||||
for arg in redis_kwargs:
|
||||
if arg in args:
|
||||
connection_kwargs[arg] = redis_kwargs[arg]
|
||||
|
|
@ -518,12 +519,12 @@ def _get_redis_sentinel_connection_kwargs(redis_kwargs: dict) -> dict:
|
|||
|
||||
|
||||
def _init_redis_sentinel(redis_kwargs) -> redis.Redis:
|
||||
sentinel_nodes = redis_kwargs.get("sentinel_nodes")
|
||||
sentinel_password = redis_kwargs.get("sentinel_password")
|
||||
service_name = redis_kwargs.get("service_name")
|
||||
connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs)
|
||||
sentinel_nodes: Final = redis_kwargs.get("sentinel_nodes")
|
||||
sentinel_password: Final = redis_kwargs.get("sentinel_password")
|
||||
service_name: Final = redis_kwargs.get("service_name")
|
||||
connection_kwargs: Final = _get_redis_sentinel_connection_kwargs(redis_kwargs)
|
||||
connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT)
|
||||
sentinel_kwargs = dict(connection_kwargs)
|
||||
sentinel_kwargs: Final = dict(connection_kwargs)
|
||||
sentinel_kwargs["password"] = sentinel_password
|
||||
|
||||
if not sentinel_nodes or not service_name:
|
||||
|
|
@ -532,7 +533,7 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis:
|
|||
verbose_logger.debug("init_redis_sentinel: sentinel nodes are being initialized.")
|
||||
|
||||
# Set up the Sentinel client
|
||||
sentinel = redis.Sentinel(
|
||||
sentinel: Final = redis.Sentinel(
|
||||
sentinel_nodes,
|
||||
sentinel_kwargs=sentinel_kwargs,
|
||||
)
|
||||
|
|
@ -543,12 +544,12 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis:
|
|||
|
||||
|
||||
def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis:
|
||||
sentinel_nodes = redis_kwargs.get("sentinel_nodes")
|
||||
sentinel_password = redis_kwargs.get("sentinel_password")
|
||||
service_name = redis_kwargs.get("service_name")
|
||||
connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs)
|
||||
sentinel_nodes: Final = redis_kwargs.get("sentinel_nodes")
|
||||
sentinel_password: Final = redis_kwargs.get("sentinel_password")
|
||||
service_name: Final = redis_kwargs.get("service_name")
|
||||
connection_kwargs: Final = _get_redis_sentinel_connection_kwargs(redis_kwargs)
|
||||
connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT)
|
||||
sentinel_kwargs = dict(connection_kwargs)
|
||||
sentinel_kwargs: Final = dict(connection_kwargs)
|
||||
sentinel_kwargs["password"] = sentinel_password
|
||||
|
||||
if not sentinel_nodes or not service_name:
|
||||
|
|
@ -557,7 +558,7 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis:
|
|||
verbose_logger.debug("init_redis_sentinel: sentinel nodes are being initialized.")
|
||||
|
||||
# Set up the Sentinel client
|
||||
sentinel = async_redis.Sentinel(
|
||||
sentinel: Final = async_redis.Sentinel(
|
||||
sentinel_nodes,
|
||||
sentinel_kwargs=sentinel_kwargs,
|
||||
)
|
||||
|
|
@ -568,14 +569,14 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis:
|
|||
|
||||
|
||||
def get_redis_client(**env_overrides):
|
||||
redis_kwargs = _get_redis_client_logic(**env_overrides)
|
||||
redis_kwargs: Final = _get_redis_client_logic(**env_overrides)
|
||||
|
||||
if "startup_nodes" in redis_kwargs:
|
||||
return init_redis_cluster(redis_kwargs)
|
||||
|
||||
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
|
||||
args = _get_redis_url_kwargs()
|
||||
url_kwargs = {}
|
||||
args: Final = _get_redis_url_kwargs()
|
||||
url_kwargs: Final = {}
|
||||
for arg in redis_kwargs:
|
||||
if arg in args:
|
||||
url_kwargs[arg] = redis_kwargs[arg]
|
||||
|
|
@ -593,13 +594,13 @@ def get_redis_async_client(
|
|||
connection_pool: async_redis.BlockingConnectionPool | None = None,
|
||||
**env_overrides,
|
||||
) -> async_redis.Redis | async_redis.RedisCluster:
|
||||
redis_kwargs = _get_redis_client_logic(**env_overrides)
|
||||
redis_kwargs: Final = _get_redis_client_logic(**env_overrides)
|
||||
|
||||
if "startup_nodes" in redis_kwargs:
|
||||
from redis.cluster import ClusterNode
|
||||
|
||||
args = _get_redis_cluster_kwargs()
|
||||
cluster_kwargs = {}
|
||||
cluster_kwargs: Final = {}
|
||||
for arg in redis_kwargs:
|
||||
if arg in args:
|
||||
cluster_kwargs[arg] = redis_kwargs[arg]
|
||||
|
|
@ -621,7 +622,7 @@ def get_redis_async_client(
|
|||
username=os.environ.get("REDIS_USERNAME") or None,
|
||||
)
|
||||
|
||||
new_startup_nodes: list[ClusterNode] = []
|
||||
new_startup_nodes: Final[list[ClusterNode]] = []
|
||||
|
||||
for item in redis_kwargs["startup_nodes"]:
|
||||
new_startup_nodes.append(ClusterNode(**item))
|
||||
|
|
@ -635,7 +636,7 @@ def get_redis_async_client(
|
|||
cluster_kwargs.setdefault("socket_keepalive", True)
|
||||
|
||||
# Create async RedisCluster with IAM token as password if available
|
||||
cluster_client = async_redis.RedisCluster(
|
||||
cluster_client: Final = async_redis.RedisCluster(
|
||||
startup_nodes=new_startup_nodes,
|
||||
**cluster_kwargs, # type: ignore
|
||||
)
|
||||
|
|
@ -646,7 +647,7 @@ def get_redis_async_client(
|
|||
if connection_pool is not None:
|
||||
return async_redis.Redis(connection_pool=connection_pool)
|
||||
args = _get_redis_url_kwargs(client=async_redis.Redis)
|
||||
url_kwargs = {}
|
||||
url_kwargs: Final = {}
|
||||
for arg in redis_kwargs:
|
||||
if arg in args:
|
||||
url_kwargs[arg] = redis_kwargs[arg]
|
||||
|
|
@ -686,15 +687,15 @@ def get_redis_async_client(
|
|||
def get_redis_connection_pool(
|
||||
**env_overrides,
|
||||
) -> async_redis.BlockingConnectionPool | None:
|
||||
redis_kwargs = _get_redis_client_logic(**env_overrides)
|
||||
redis_kwargs: Final = _get_redis_client_logic(**env_overrides)
|
||||
verbose_logger.debug("get_redis_connection_pool: redis_kwargs", redis_kwargs)
|
||||
|
||||
if "startup_nodes" in redis_kwargs:
|
||||
return None
|
||||
|
||||
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
|
||||
allowed_args = _get_redis_url_kwargs(client=async_redis.Redis)
|
||||
pool_kwargs = {k: v for k, v in redis_kwargs.items() if k in allowed_args and k != "max_connections"}
|
||||
allowed_args: Final = _get_redis_url_kwargs(client=async_redis.Redis)
|
||||
pool_kwargs: Final = {k: v for k, v in redis_kwargs.items() if k in allowed_args and k != "max_connections"}
|
||||
pool_kwargs["timeout"] = REDIS_CONNECTION_POOL_TIMEOUT
|
||||
pool_kwargs["url"] = redis_kwargs["url"]
|
||||
if "max_connections" in redis_kwargs:
|
||||
|
|
@ -710,7 +711,7 @@ def get_redis_connection_pool(
|
|||
# Wrap GCP / Azure AD auth in a CredentialProvider so pool-managed
|
||||
# connections re-fetch tokens via the SDK's internal cache + silent refresh
|
||||
# rather than reusing a single token captured at pool creation.
|
||||
redis_connect_func = redis_kwargs.pop("redis_connect_func", None)
|
||||
redis_connect_func: Final = redis_kwargs.pop("redis_connect_func", None)
|
||||
if redis_connect_func and hasattr(redis_connect_func, "_azure_credential"):
|
||||
redis_kwargs["credential_provider"] = AzureADCredentialProvider(
|
||||
redis_connect_func._azure_credential,
|
||||
|
|
@ -737,7 +738,7 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None:
|
|||
if not verbose_logger.isEnabledFor(logging.DEBUG):
|
||||
return
|
||||
|
||||
console = Console()
|
||||
console: Final = Console()
|
||||
|
||||
# Initialize the sensitive data masker
|
||||
masker = SensitiveDataMasker()
|
||||
|
|
@ -746,10 +747,10 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None:
|
|||
masked_redis_kwargs = masker.mask_dict(redis_kwargs)
|
||||
|
||||
# Create main panel title
|
||||
title = Text("Redis Configuration", style="bold blue")
|
||||
title: Final = Text("Redis Configuration", style="bold blue")
|
||||
|
||||
# Create configuration table
|
||||
config_table = Table(
|
||||
config_table: Final = Table(
|
||||
title="🔧 Redis Connection Parameters",
|
||||
show_header=True,
|
||||
header_style="bold magenta",
|
||||
|
|
@ -786,7 +787,7 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None:
|
|||
connection_type = "Redis (URL-based)"
|
||||
|
||||
# Create connection type info
|
||||
info_table = Table(
|
||||
info_table: Final = Table(
|
||||
title="📊 Connection Info",
|
||||
show_header=True,
|
||||
header_style="bold green",
|
||||
|
|
|
|||
|
|
@ -1,21 +1,21 @@
|
|||
import asyncio
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from redis.credentials import CredentialProvider # type: ignore[attr-defined]
|
||||
|
||||
# Azure AD scope for Redis Cache for Azure.
|
||||
AZURE_REDIS_SCOPE = "https://redis.azure.com/.default"
|
||||
AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default"
|
||||
|
||||
# GCP IAM tokens are valid for 1 hour. Cache for 55 minutes to refresh before expiry.
|
||||
_GCP_IAM_TOKEN_TTL_SECONDS = 3300
|
||||
_GCP_IAM_TOKEN_TTL_SECONDS: Final = 3300
|
||||
|
||||
# Module-level cache shared across all GCPIAMCredentialProvider instances for the
|
||||
# same service account, so multiple Redis connections on the same pod share one token.
|
||||
# Keyed by service_account → (token, expiry_monotonic_timestamp).
|
||||
_token_cache: dict[str, tuple[str, float]] = {}
|
||||
_token_cache_lock = threading.Lock()
|
||||
_token_cache: Final[dict[str, tuple[str, float]]] = {}
|
||||
_token_cache_lock: Final = threading.Lock()
|
||||
|
||||
|
||||
def _generate_gcp_iam_access_token(service_account: str) -> str:
|
||||
|
|
@ -36,12 +36,12 @@ def _generate_gcp_iam_access_token(service_account: str) -> str:
|
|||
"Install it with: pip install google-cloud-iam"
|
||||
)
|
||||
|
||||
client = iam_credentials_v1.IAMCredentialsClient()
|
||||
request = iam_credentials_v1.GenerateAccessTokenRequest(
|
||||
client: Final = iam_credentials_v1.IAMCredentialsClient()
|
||||
request: Final = iam_credentials_v1.GenerateAccessTokenRequest(
|
||||
name=service_account,
|
||||
scope=["https://www.googleapis.com/auth/cloud-platform"],
|
||||
)
|
||||
response = client.generate_access_token(request=request)
|
||||
response: Final = client.generate_access_token(request=request)
|
||||
return str(response.access_token)
|
||||
|
||||
|
||||
|
|
@ -96,11 +96,11 @@ class GCPIAMCredentialProvider(CredentialProvider):
|
|||
self._gcp_service_account = gcp_service_account
|
||||
|
||||
def get_credentials(self) -> tuple[str]:
|
||||
token = _get_cached_gcp_iam_token(self._gcp_service_account)
|
||||
token: Final = _get_cached_gcp_iam_token(self._gcp_service_account)
|
||||
return (token,)
|
||||
|
||||
async def get_credentials_async(self) -> tuple[str]:
|
||||
token = await asyncio.to_thread(_get_cached_gcp_iam_token, self._gcp_service_account)
|
||||
token: Final = await asyncio.to_thread(_get_cached_gcp_iam_token, self._gcp_service_account)
|
||||
return (token,)
|
||||
|
||||
|
||||
|
|
@ -120,13 +120,13 @@ class AzureADCredentialProvider(CredentialProvider):
|
|||
self._username = username
|
||||
|
||||
def get_credentials(self) -> tuple[str] | tuple[str, str]:
|
||||
token = self._credential.get_token(AZURE_REDIS_SCOPE).token
|
||||
token: Final = self._credential.get_token(AZURE_REDIS_SCOPE).token
|
||||
if self._username:
|
||||
return (self._username, token)
|
||||
return (token,)
|
||||
|
||||
async def get_credentials_async(self) -> tuple[str] | tuple[str, str]:
|
||||
token_obj = await asyncio.to_thread(self._credential.get_token, AZURE_REDIS_SCOPE)
|
||||
token_obj: Final = await asyncio.to_thread(self._credential.get_token, AZURE_REDIS_SCOPE)
|
||||
if self._username:
|
||||
return (self._username, token_obj.token)
|
||||
return (token_obj.token,)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import asyncio
|
||||
from datetime import datetime, timedelta
|
||||
from typing import TYPE_CHECKING, Any, Union
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -67,7 +67,7 @@ class ServiceLogging(CustomLogger):
|
|||
whether the callback is the logger instance itself or the ``"otel"`` string
|
||||
(which routes to the proxy's registered ``open_telemetry_logger``).
|
||||
"""
|
||||
otel_v2_cls = _get_otel_v2_class()
|
||||
otel_v2_cls: Final = _get_otel_v2_class()
|
||||
|
||||
def _is_otel_logger(obj: Any) -> bool:
|
||||
if isinstance(obj, OpenTelemetry):
|
||||
|
|
@ -101,7 +101,7 @@ class ServiceLogging(CustomLogger):
|
|||
|
||||
try:
|
||||
# Try to get the current event loop
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
# Check if the loop is running
|
||||
if loop.is_running():
|
||||
# If we're in a running loop, create a task
|
||||
|
|
@ -163,7 +163,7 @@ class ServiceLogging(CustomLogger):
|
|||
if self.mock_testing:
|
||||
self.mock_testing_async_success_hook += 1
|
||||
|
||||
payload = ServiceLoggerPayload(
|
||||
payload: Final = ServiceLoggerPayload(
|
||||
is_error=False,
|
||||
error=None,
|
||||
service=service,
|
||||
|
|
@ -178,7 +178,7 @@ class ServiceLogging(CustomLogger):
|
|||
# (the V2 logger self-registers its instance even when the string is
|
||||
# present, unlike V1). Without this guard each such reference emits its own
|
||||
# span, so a single DB call shows up as duplicate ``postgres ...`` spans.
|
||||
emitted_otel_logger_ids: set = set()
|
||||
emitted_otel_logger_ids: Final[set] = set()
|
||||
for callback in litellm.service_callback:
|
||||
if callback == "prometheus_system":
|
||||
await self.init_prometheus_services_logger_if_none()
|
||||
|
|
@ -267,7 +267,7 @@ class ServiceLogging(CustomLogger):
|
|||
elif isinstance(error, str):
|
||||
error_message = error
|
||||
|
||||
payload = ServiceLoggerPayload(
|
||||
payload: Final = ServiceLoggerPayload(
|
||||
is_error=True,
|
||||
error=error_message,
|
||||
service=service,
|
||||
|
|
@ -278,7 +278,7 @@ class ServiceLogging(CustomLogger):
|
|||
|
||||
# Dedupe OTel loggers per event — see ``async_service_success_hook`` for why
|
||||
# the same logger can be referenced twice in ``service_callback``.
|
||||
emitted_otel_logger_ids: set = set()
|
||||
emitted_otel_logger_ids: Final[set] = set()
|
||||
for callback in litellm.service_callback:
|
||||
if callback == "prometheus_system":
|
||||
await self.init_prometheus_services_logger_if_none()
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Custom A2A Card Resolver for LiteLLM.
|
|||
Extends the A2A SDK's card resolver to support multiple well-known paths.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import LOCALHOST_URL_PATTERNS
|
||||
|
|
@ -43,18 +43,18 @@ def is_localhost_or_internal_url(url: str | None) -> bool:
|
|||
if not url:
|
||||
return False
|
||||
|
||||
url_lower = url.lower()
|
||||
url_lower: Final = url.lower()
|
||||
|
||||
return any(pattern in url_lower for pattern in LOCALHOST_URL_PATTERNS)
|
||||
|
||||
|
||||
def get_agent_card_url(agent_card: "AgentCard") -> str | None:
|
||||
"""Return the agent endpoint URL from the resolved SDK card."""
|
||||
url = getattr(agent_card, "url", None)
|
||||
url: Final = getattr(agent_card, "url", None)
|
||||
if url:
|
||||
return url
|
||||
|
||||
interfaces = getattr(agent_card, "supported_interfaces", None)
|
||||
interfaces: Final = getattr(agent_card, "supported_interfaces", None)
|
||||
if interfaces:
|
||||
return getattr(interfaces[0], "url", None)
|
||||
return None
|
||||
|
|
@ -62,11 +62,11 @@ def get_agent_card_url(agent_card: "AgentCard") -> str | None:
|
|||
|
||||
def set_agent_card_url(agent_card: "AgentCard", url: str) -> None:
|
||||
"""Set the agent endpoint URL on the resolved SDK card."""
|
||||
normalized = url.rstrip("/") + "/"
|
||||
normalized: Final = url.rstrip("/") + "/"
|
||||
if hasattr(agent_card, "url"):
|
||||
agent_card.url = normalized
|
||||
|
||||
interfaces = getattr(agent_card, "supported_interfaces", None)
|
||||
interfaces: Final = getattr(agent_card, "supported_interfaces", None)
|
||||
if interfaces:
|
||||
interfaces[0].url = normalized
|
||||
|
||||
|
|
@ -86,16 +86,16 @@ def fix_agent_card_url(agent_card: "AgentCard", base_url: str) -> "AgentCard":
|
|||
Returns:
|
||||
The agent card with the URL fixed if necessary
|
||||
"""
|
||||
card_url = getattr(agent_card, "url", None)
|
||||
card_url: Final = getattr(agent_card, "url", None)
|
||||
|
||||
if card_url and is_localhost_or_internal_url(card_url):
|
||||
# Normalize base_url to ensure it ends with /
|
||||
fixed_url = base_url.rstrip("/") + "/"
|
||||
fixed_url: Final = base_url.rstrip("/") + "/"
|
||||
agent_card.url = fixed_url
|
||||
|
||||
interfaces = getattr(agent_card, "supported_interfaces", None)
|
||||
interfaces: Final = getattr(agent_card, "supported_interfaces", None)
|
||||
if interfaces:
|
||||
interface_url = getattr(interfaces[0], "url", None)
|
||||
interface_url: Final = getattr(interfaces[0], "url", None)
|
||||
if interface_url and is_localhost_or_internal_url(interface_url):
|
||||
interfaces[0].url = base_url.rstrip("/") + "/"
|
||||
|
||||
|
|
@ -140,7 +140,7 @@ class LiteLLMA2ACardResolver(_A2ACardResolver): # type: ignore[misc]
|
|||
)
|
||||
|
||||
# Try both well-known paths
|
||||
paths = [
|
||||
paths: Final = [
|
||||
AGENT_CARD_WELL_KNOWN_PATH,
|
||||
PREV_AGENT_CARD_WELL_KNOWN_PATH,
|
||||
]
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Provides a class-based interface for A2A agent invocation.
|
|||
"""
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm.types.agents import LiteLLMSendMessageResponse
|
||||
|
||||
|
|
@ -92,7 +92,7 @@ class A2AClient:
|
|||
"""Send a message to the A2A agent."""
|
||||
from litellm.a2a_protocol.main import asend_message
|
||||
|
||||
a2a_client = await self._get_client()
|
||||
a2a_client: Final = await self._get_client()
|
||||
return await asend_message(a2a_client=a2a_client, request=request)
|
||||
|
||||
async def send_message_streaming(
|
||||
|
|
@ -101,6 +101,6 @@ class A2AClient:
|
|||
"""Send a streaming message to the A2A agent."""
|
||||
from litellm.a2a_protocol.main import asend_message_streaming
|
||||
|
||||
a2a_client = await self._get_client()
|
||||
a2a_client: Final = await self._get_client()
|
||||
async for chunk in asend_message_streaming(a2a_client=a2a_client, request=request):
|
||||
yield chunk
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Supports dynamic cost parameters that allow platform owners
|
|||
to define custom costs per agent query or per token.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
|
|
@ -42,23 +42,23 @@ class A2ACostCalculator:
|
|||
if litellm_logging_obj is None:
|
||||
return 0.0
|
||||
|
||||
model_call_details = litellm_logging_obj.model_call_details
|
||||
model_call_details: Final = litellm_logging_obj.model_call_details
|
||||
|
||||
# Check if user set a custom response cost (backward compatibility)
|
||||
response_cost = model_call_details.get("response_cost", None)
|
||||
response_cost: Final = model_call_details.get("response_cost", None)
|
||||
if response_cost is not None:
|
||||
return float(response_cost)
|
||||
|
||||
# Get litellm_params for cost parameters
|
||||
litellm_params = model_call_details.get("litellm_params", {}) or {}
|
||||
litellm_params: Final = model_call_details.get("litellm_params", {}) or {}
|
||||
|
||||
# Check for cost_per_query (fixed cost per query)
|
||||
if litellm_params.get("cost_per_query") is not None:
|
||||
return float(litellm_params["cost_per_query"])
|
||||
|
||||
# Check for token-based pricing
|
||||
input_cost_per_token = litellm_params.get("input_cost_per_token")
|
||||
output_cost_per_token = litellm_params.get("output_cost_per_token")
|
||||
input_cost_per_token: Final = litellm_params.get("input_cost_per_token")
|
||||
output_cost_per_token: Final = litellm_params.get("output_cost_per_token")
|
||||
|
||||
if input_cost_per_token is not None or output_cost_per_token is not None:
|
||||
return A2ACostCalculator._calculate_token_based_cost(
|
||||
|
|
@ -88,16 +88,16 @@ class A2ACostCalculator:
|
|||
float: The calculated cost
|
||||
"""
|
||||
# Get usage from model_call_details
|
||||
usage = model_call_details.get("usage")
|
||||
usage: Final = model_call_details.get("usage")
|
||||
if usage is None:
|
||||
return 0.0
|
||||
|
||||
# Get token counts
|
||||
prompt_tokens = getattr(usage, "prompt_tokens", 0) or 0
|
||||
completion_tokens = getattr(usage, "completion_tokens", 0) or 0
|
||||
prompt_tokens: Final = getattr(usage, "prompt_tokens", 0) or 0
|
||||
completion_tokens: Final = getattr(usage, "completion_tokens", 0) or 0
|
||||
|
||||
# Calculate costs
|
||||
input_cost = prompt_tokens * (float(input_cost_per_token) if input_cost_per_token else 0.0)
|
||||
output_cost = completion_tokens * (float(output_cost_per_token) if output_cost_per_token else 0.0)
|
||||
input_cost: Final = prompt_tokens * (float(input_cost_per_token) if input_cost_per_token else 0.0)
|
||||
output_cost: Final = completion_tokens * (float(output_cost_per_token) if output_cost_per_token else 0.0)
|
||||
|
||||
return input_cost + output_cost
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ A2A Protocol Exception Mapping Utils.
|
|||
Maps A2A SDK exceptions to LiteLLM A2A exception types.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.a2a_protocol.card_resolver import (
|
||||
|
|
@ -53,7 +53,7 @@ class A2AExceptionCheckers:
|
|||
if not isinstance(error_str, str):
|
||||
return False
|
||||
|
||||
error_str_lower = error_str.lower()
|
||||
error_str_lower: Final = error_str.lower()
|
||||
return any(pattern in error_str_lower for pattern in CONNECTION_ERROR_PATTERNS)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -83,8 +83,8 @@ class A2AExceptionCheckers:
|
|||
if not isinstance(error_str, str):
|
||||
return False
|
||||
|
||||
error_str_lower = error_str.lower()
|
||||
agent_card_patterns = [
|
||||
error_str_lower: Final = error_str.lower()
|
||||
agent_card_patterns: Final = [
|
||||
"agent card",
|
||||
"agent-card",
|
||||
".well-known",
|
||||
|
|
@ -118,7 +118,7 @@ def map_a2a_exception(
|
|||
A2AAgentCardError: If the error is related to agent card issues
|
||||
A2AError: For other A2A-related errors
|
||||
"""
|
||||
error_str = str(original_exception)
|
||||
error_str: Final = str(original_exception)
|
||||
|
||||
# Check for localhost URL connection error (special case - retryable)
|
||||
if (
|
||||
|
|
@ -190,7 +190,7 @@ async def handle_a2a_localhost_retry(
|
|||
"rewrite, so the upstream URL cannot be corrected."
|
||||
)
|
||||
|
||||
request_type = "streaming " if is_streaming else ""
|
||||
request_type: Final = "streaming " if is_streaming else ""
|
||||
verbose_logger.warning(
|
||||
"A2A %srequest to '%s' failed: %s. Agent card contains localhost/internal URL. Retrying with base_url '%s'.",
|
||||
request_type,
|
||||
|
|
@ -205,14 +205,14 @@ async def handle_a2a_localhost_retry(
|
|||
# Reuse the httpx client LiteLLM attached at creation. It carries this agent's
|
||||
# trace-id and auth headers, so a fresh client would drop them. Only clients built
|
||||
# by ``create_a2a_client`` have it; an externally-supplied client cannot be retried.
|
||||
httpx_client = getattr(a2a_client, "_litellm_httpx_client", None)
|
||||
httpx_client: Final = getattr(a2a_client, "_litellm_httpx_client", None)
|
||||
if httpx_client is None:
|
||||
raise RuntimeError(
|
||||
"Cannot retry A2A localhost URL fix: the client was not created by "
|
||||
"create_a2a_client, so no LiteLLM httpx client is attached."
|
||||
)
|
||||
|
||||
new_client = await create_client( # pyright: ignore[reportOptionalCall]
|
||||
new_client: Final = await create_client( # pyright: ignore[reportOptionalCall]
|
||||
agent_card,
|
||||
client_config=ClientConfig( # pyright: ignore[reportOptionalCall]
|
||||
httpx_client=httpx_client,
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ A2A Streaming Events (in order):
|
|||
"""
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -25,10 +25,10 @@ from litellm.interactions.agents.utils import merge_agent_headers
|
|||
# litellm_params key carrying the authenticated principal (hashed virtual key) so
|
||||
# A2A provider configs can scope provider-side state (e.g. LangFlow session memory)
|
||||
# per key instead of trusting the client-supplied A2A contextId.
|
||||
A2A_USER_API_KEY_HASH_PARAM = "litellm_a2a_user_api_key_hash"
|
||||
A2A_USER_API_KEY_HASH_PARAM: Final = "litellm_a2a_user_api_key_hash"
|
||||
|
||||
# Agent metadata fields stored in litellm_params that are not valid litellm.acompletion() kwargs
|
||||
_AGENT_ONLY_PARAMS = frozenset(
|
||||
_AGENT_ONLY_PARAMS: Final = frozenset(
|
||||
{
|
||||
"is_public",
|
||||
"agent_name",
|
||||
|
|
@ -70,7 +70,7 @@ class A2ACompletionBridgeHandler:
|
|||
"""
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
if not _skip_a2a_provider_routing:
|
||||
a2a_provider_config = A2AProviderConfigManager.get_provider_config(
|
||||
a2a_provider_config: Final = A2AProviderConfigManager.get_provider_config(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model=litellm_params.get("model"),
|
||||
)
|
||||
|
|
@ -87,14 +87,14 @@ class A2ACompletionBridgeHandler:
|
|||
)
|
||||
|
||||
# Extract message from params
|
||||
message = params.get("message", {})
|
||||
message: Final = params.get("message", {})
|
||||
|
||||
# Transform A2A message to OpenAI format
|
||||
openai_messages = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
openai_messages: Final = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
|
||||
# Get completion params
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
model = litellm_params.get("model", "agent")
|
||||
model: Final = litellm_params.get("model", "agent")
|
||||
|
||||
# Build full model string if provider specified
|
||||
# Skip prepending if model already starts with the provider prefix
|
||||
|
|
@ -106,14 +106,14 @@ class A2ACompletionBridgeHandler:
|
|||
verbose_logger.info("A2A completion bridge: model=%s, api_base=%s", full_model, api_base)
|
||||
|
||||
# Build completion params dict
|
||||
completion_params: dict[str, Any] = {
|
||||
completion_params: Final[dict[str, Any]] = {
|
||||
"model": full_model,
|
||||
"messages": openai_messages,
|
||||
"api_base": api_base,
|
||||
"stream": False,
|
||||
}
|
||||
# Add litellm_params (contains api_key, client_id, client_secret, tenant_id, etc.)
|
||||
litellm_params_to_add = {
|
||||
litellm_params_to_add: Final = {
|
||||
k: v
|
||||
for k, v in litellm_params.items()
|
||||
if k not in ("model", "custom_llm_provider") and k not in _AGENT_ONLY_PARAMS
|
||||
|
|
@ -135,10 +135,10 @@ class A2ACompletionBridgeHandler:
|
|||
)
|
||||
|
||||
# Call litellm.acompletion
|
||||
response = await litellm.acompletion(**completion_params)
|
||||
response: Final = await litellm.acompletion(**completion_params)
|
||||
|
||||
# Transform response to A2A format
|
||||
a2a_response = A2ACompletionBridgeTransformation.openai_response_to_a2a_response(
|
||||
a2a_response: Final = A2ACompletionBridgeTransformation.openai_response_to_a2a_response(
|
||||
response=response,
|
||||
request_id=request_id,
|
||||
)
|
||||
|
|
@ -179,7 +179,7 @@ class A2ACompletionBridgeHandler:
|
|||
"""
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
if not _skip_a2a_provider_routing:
|
||||
a2a_provider_config = A2AProviderConfigManager.get_provider_config(
|
||||
a2a_provider_config: Final = A2AProviderConfigManager.get_provider_config(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model=litellm_params.get("model"),
|
||||
)
|
||||
|
|
@ -199,20 +199,20 @@ class A2ACompletionBridgeHandler:
|
|||
return
|
||||
|
||||
# Extract message from params
|
||||
message = params.get("message", {})
|
||||
message: Final = params.get("message", {})
|
||||
|
||||
# Create streaming context
|
||||
ctx = A2AStreamingContext(
|
||||
ctx: Final = A2AStreamingContext(
|
||||
request_id=request_id,
|
||||
input_message=message,
|
||||
)
|
||||
|
||||
# Transform A2A message to OpenAI format
|
||||
openai_messages = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
openai_messages: Final = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
|
||||
# Get completion params
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
model = litellm_params.get("model", "agent")
|
||||
model: Final = litellm_params.get("model", "agent")
|
||||
|
||||
# Build full model string if provider specified
|
||||
# Skip prepending if model already starts with the provider prefix
|
||||
|
|
@ -224,14 +224,14 @@ class A2ACompletionBridgeHandler:
|
|||
verbose_logger.info("A2A completion bridge streaming: model=%s, api_base=%s", full_model, api_base)
|
||||
|
||||
# Build completion params dict
|
||||
completion_params: dict[str, Any] = {
|
||||
completion_params: Final[dict[str, Any]] = {
|
||||
"model": full_model,
|
||||
"messages": openai_messages,
|
||||
"api_base": api_base,
|
||||
"stream": True,
|
||||
}
|
||||
# Add litellm_params (contains api_key, client_id, client_secret, tenant_id, etc.)
|
||||
litellm_params_to_add = {
|
||||
litellm_params_to_add: Final = {
|
||||
k: v
|
||||
for k, v in litellm_params.items()
|
||||
if k not in ("model", "custom_llm_provider") and k not in _AGENT_ONLY_PARAMS
|
||||
|
|
@ -253,11 +253,11 @@ class A2ACompletionBridgeHandler:
|
|||
)
|
||||
|
||||
# 1. Emit initial task event (kind: "task", status: "submitted")
|
||||
task_event = A2ACompletionBridgeTransformation.create_task_event(ctx)
|
||||
task_event: Final = A2ACompletionBridgeTransformation.create_task_event(ctx)
|
||||
yield task_event
|
||||
|
||||
# 2. Emit status update (kind: "status-update", status: "working")
|
||||
working_event = A2ACompletionBridgeTransformation.create_status_update_event(
|
||||
working_event: Final = A2ACompletionBridgeTransformation.create_status_update_event(
|
||||
ctx=ctx,
|
||||
state="working",
|
||||
final=False,
|
||||
|
|
@ -266,7 +266,7 @@ class A2ACompletionBridgeHandler:
|
|||
yield working_event
|
||||
|
||||
# Call litellm.acompletion with streaming
|
||||
response = await litellm.acompletion(**completion_params)
|
||||
response: Final = await litellm.acompletion(**completion_params)
|
||||
|
||||
# 3. Accumulate content and emit artifact update
|
||||
accumulated_text = ""
|
||||
|
|
@ -286,14 +286,14 @@ class A2ACompletionBridgeHandler:
|
|||
|
||||
# Emit artifact update with accumulated content
|
||||
if accumulated_text:
|
||||
artifact_event = A2ACompletionBridgeTransformation.create_artifact_update_event(
|
||||
artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event(
|
||||
ctx=ctx,
|
||||
text=accumulated_text,
|
||||
)
|
||||
yield artifact_event
|
||||
|
||||
# 4. Emit final status update (kind: "status-update", status: "completed", final: true)
|
||||
completed_event = A2ACompletionBridgeTransformation.create_status_update_event(
|
||||
completed_event: Final = A2ACompletionBridgeTransformation.create_status_update_event(
|
||||
ctx=ctx,
|
||||
state="completed",
|
||||
final=True,
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ A2A Streaming Events:
|
|||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
from uuid import uuid4
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -48,7 +48,7 @@ class A2ACompletionBridgeTransformation:
|
|||
@staticmethod
|
||||
def _extract_text_from_a2a_parts(parts: list[dict[str, Any]]) -> str:
|
||||
"""Extract text from A2A parts (with or without explicit ``kind``)."""
|
||||
content_parts: list[str] = []
|
||||
content_parts: Final[list[str]] = []
|
||||
for part in parts:
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
|
|
@ -71,10 +71,10 @@ class A2ACompletionBridgeTransformation:
|
|||
Forwarded once on the LangGraph run payload (``metadata``), not duplicated on
|
||||
each input message — see ``apply_forward_metadata_to_completion_params``.
|
||||
"""
|
||||
merged: dict[str, Any] = {}
|
||||
merged: Final[dict[str, Any]] = {}
|
||||
if params and isinstance(params.get("metadata"), dict):
|
||||
merged.update(params["metadata"])
|
||||
message_metadata = a2a_message.get("metadata")
|
||||
message_metadata: Final = a2a_message.get("metadata")
|
||||
if isinstance(message_metadata, dict):
|
||||
merged.update(message_metadata)
|
||||
return merged or None
|
||||
|
|
@ -90,7 +90,7 @@ class A2ACompletionBridgeTransformation:
|
|||
|
||||
Uses ``extra_body`` so we do not collide with LiteLLM's spend-log ``metadata`` kwarg.
|
||||
"""
|
||||
forward_metadata = A2ACompletionBridgeTransformation.get_forward_metadata(
|
||||
forward_metadata: Final = A2ACompletionBridgeTransformation.get_forward_metadata(
|
||||
a2a_message=a2a_message,
|
||||
params=params,
|
||||
)
|
||||
|
|
@ -103,9 +103,9 @@ class A2ACompletionBridgeTransformation:
|
|||
# Layer client-supplied A2A metadata under any agent-owner-configured
|
||||
# ``extra_body.metadata`` so the configured keys remain authoritative
|
||||
# and an A2A caller cannot overwrite server-set run metadata.
|
||||
existing_metadata = extra_body.get("metadata")
|
||||
existing_dict: dict[str, Any] = existing_metadata if isinstance(existing_metadata, dict) else {}
|
||||
merged_metadata: dict[str, Any] = {**forward_metadata, **existing_dict}
|
||||
existing_metadata: Final = extra_body.get("metadata")
|
||||
existing_dict: Final[dict[str, Any]] = existing_metadata if isinstance(existing_metadata, dict) else {}
|
||||
merged_metadata: Final[dict[str, Any]] = {**forward_metadata, **existing_dict}
|
||||
extra_body = {**extra_body, "metadata": merged_metadata}
|
||||
completion_params["extra_body"] = extra_body
|
||||
|
||||
|
|
@ -124,7 +124,7 @@ class A2ACompletionBridgeTransformation:
|
|||
Returns:
|
||||
List of OpenAI-format messages
|
||||
"""
|
||||
role = a2a_message.get("role", "user")
|
||||
role: Final = a2a_message.get("role", "user")
|
||||
parts = a2a_message.get("parts", [])
|
||||
|
||||
# Map A2A roles to OpenAI roles
|
||||
|
|
@ -139,11 +139,11 @@ class A2ACompletionBridgeTransformation:
|
|||
if not isinstance(parts, list):
|
||||
parts = []
|
||||
|
||||
content = A2ACompletionBridgeTransformation._extract_text_from_a2a_parts(parts)
|
||||
content: Final = A2ACompletionBridgeTransformation._extract_text_from_a2a_parts(parts)
|
||||
|
||||
# Do not attach A2A message.metadata here — the completion bridge forwards it
|
||||
# once at run level via extra_body.metadata (LangGraph POST /runs/wait shape).
|
||||
openai_message: dict[str, Any] = {"role": openai_role, "content": content}
|
||||
openai_message: Final[dict[str, Any]] = {"role": openai_role, "content": content}
|
||||
|
||||
verbose_logger.debug(
|
||||
"A2A -> OpenAI transform: role=%s -> %s, content_length=%s", role, openai_role, len(content)
|
||||
|
|
@ -169,12 +169,12 @@ class A2ACompletionBridgeTransformation:
|
|||
# Extract content from response
|
||||
content = ""
|
||||
if hasattr(response, "choices") and response.choices:
|
||||
choice = response.choices[0]
|
||||
choice: Final = response.choices[0]
|
||||
if hasattr(choice, "message") and choice.message:
|
||||
content = choice.message.content or ""
|
||||
|
||||
# Build A2A message
|
||||
a2a_message = {
|
||||
a2a_message: Final = {
|
||||
"kind": "message",
|
||||
"role": "agent",
|
||||
"parts": [{"kind": "text", "text": content}],
|
||||
|
|
@ -182,7 +182,7 @@ class A2ACompletionBridgeTransformation:
|
|||
}
|
||||
|
||||
# Build A2A response
|
||||
a2a_response = {
|
||||
a2a_response: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"result": a2a_message,
|
||||
|
|
@ -245,7 +245,7 @@ class A2ACompletionBridgeTransformation:
|
|||
final: Whether this is the final event
|
||||
message_text: Optional message text for 'working' status
|
||||
"""
|
||||
status: dict[str, Any] = {
|
||||
status: Final[dict[str, Any]] = {
|
||||
"state": state,
|
||||
"timestamp": A2ACompletionBridgeTransformation._get_timestamp(),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -13,12 +13,7 @@ import asyncio
|
|||
import datetime
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Coroutine
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Optional,
|
||||
cast,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
|
|
@ -80,7 +75,7 @@ from litellm.a2a_protocol.exception_mapping_utils import (
|
|||
from litellm.a2a_protocol.exceptions import A2ALocalhostURLError
|
||||
|
||||
# Use our custom resolver instead of the default A2A SDK resolver
|
||||
A2ACardResolver = LiteLLMA2ACardResolver
|
||||
A2ACardResolver: Final = LiteLLMA2ACardResolver
|
||||
|
||||
|
||||
def _set_usage_on_logging_obj(
|
||||
|
|
@ -96,9 +91,9 @@ def _set_usage_on_logging_obj(
|
|||
prompt_tokens: Number of input tokens
|
||||
completion_tokens: Number of output tokens
|
||||
"""
|
||||
litellm_logging_obj = kwargs.get("litellm_logging_obj")
|
||||
litellm_logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
if litellm_logging_obj is not None:
|
||||
usage = litellm.Usage(
|
||||
usage: Final = litellm.Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
|
|
@ -120,13 +115,13 @@ def _set_agent_id_on_logging_obj(
|
|||
if agent_id is None:
|
||||
return
|
||||
|
||||
litellm_logging_obj = kwargs.get("litellm_logging_obj")
|
||||
litellm_logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
if litellm_logging_obj is not None:
|
||||
# Set agent_id directly on model_call_details (same pattern as custom_llm_provider)
|
||||
litellm_logging_obj.model_call_details["agent_id"] = agent_id
|
||||
|
||||
|
||||
_A2A_COST_PARAM_KEYS = ("cost_per_query", "input_cost_per_token", "output_cost_per_token")
|
||||
_A2A_COST_PARAM_KEYS: Final = ("cost_per_query", "input_cost_per_token", "output_cost_per_token")
|
||||
|
||||
|
||||
def _set_litellm_params_on_logging_obj(
|
||||
|
|
@ -141,7 +136,7 @@ def _set_litellm_params_on_logging_obj(
|
|||
litellm_params already carries metadata / proxy_server_request / user-key
|
||||
context, so merge the pricing keys in rather than replacing the dict.
|
||||
"""
|
||||
logging_obj = kwargs.get("litellm_logging_obj")
|
||||
logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
if logging_obj is None:
|
||||
return
|
||||
|
||||
|
|
@ -149,7 +144,7 @@ def _set_litellm_params_on_logging_obj(
|
|||
if not cost_params:
|
||||
return
|
||||
|
||||
existing = logging_obj.model_call_details.get("litellm_params") or {}
|
||||
existing: Final = logging_obj.model_call_details.get("litellm_params") or {}
|
||||
logging_obj.model_call_details["litellm_params"] = {**existing, **cost_params}
|
||||
|
||||
|
||||
|
|
@ -162,17 +157,17 @@ def _get_a2a_model_info(a2a_client: Any, kwargs: dict[str, Any]) -> str:
|
|||
"""
|
||||
agent_name = "unknown"
|
||||
|
||||
agent_card = _get_a2a_client_agent_card(a2a_client)
|
||||
agent_card: Final = _get_a2a_client_agent_card(a2a_client)
|
||||
|
||||
if agent_card is not None:
|
||||
agent_name = getattr(agent_card, "name", "unknown") or "unknown"
|
||||
|
||||
# Build model string
|
||||
model = f"a2a_agent/{agent_name}"
|
||||
custom_llm_provider = "a2a_agent"
|
||||
model: Final = f"a2a_agent/{agent_name}"
|
||||
custom_llm_provider: Final = "a2a_agent"
|
||||
|
||||
# Set on litellm_logging_obj if available (for standard logging payload)
|
||||
litellm_logging_obj = kwargs.get("litellm_logging_obj")
|
||||
litellm_logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
if litellm_logging_obj is not None:
|
||||
litellm_logging_obj.model = model
|
||||
litellm_logging_obj.custom_llm_provider = custom_llm_provider
|
||||
|
|
@ -212,7 +207,7 @@ async def _send_message_via_completion_bridge(
|
|||
|
||||
params = request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params)
|
||||
|
||||
response_dict = await A2ACompletionBridgeHandler.handle_non_streaming(
|
||||
response_dict: Final = await A2ACompletionBridgeHandler.handle_non_streaming(
|
||||
request_id=str(request.id),
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
|
|
@ -230,18 +225,18 @@ async def _send_message(a2a_client: "A2AClientType", request: "SendMessageReques
|
|||
"The 'a2a' package is required for A2A agent invocation. Install it with: pip install a2a-sdk"
|
||||
)
|
||||
|
||||
pb_request = _a2a_conversions.to_core_send_message_request(request)
|
||||
pb_request: Final = _a2a_conversions.to_core_send_message_request(request)
|
||||
last_event = None
|
||||
async for event in a2a_client.send_message(pb_request):
|
||||
last_event = event
|
||||
if last_event is None:
|
||||
raise RuntimeError("A2A send_message failed: no response received from agent.")
|
||||
|
||||
stream_compat = _a2a_conversions.to_compat_stream_response(
|
||||
stream_compat: Final = _a2a_conversions.to_compat_stream_response(
|
||||
last_event,
|
||||
request_id=request.id,
|
||||
)
|
||||
result = stream_compat.result
|
||||
result: Final = stream_compat.result
|
||||
if not isinstance(result, (Message, Task)):
|
||||
raise RuntimeError(
|
||||
"A2A send_message failed: non-streaming message/send expects the "
|
||||
|
|
@ -305,7 +300,7 @@ async def _stream_messages(
|
|||
"The 'a2a' package is required for A2A agent invocation. Install it with: pip install a2a-sdk"
|
||||
)
|
||||
|
||||
pb_request = _a2a_conversions.to_core_send_message_request(request)
|
||||
pb_request: Final = _a2a_conversions.to_core_send_message_request(request)
|
||||
async for event in a2a_client.send_message(pb_request):
|
||||
compat_chunk = _a2a_conversions.to_compat_stream_response(
|
||||
event,
|
||||
|
|
@ -425,9 +420,9 @@ async def asend_message(
|
|||
```
|
||||
"""
|
||||
litellm_params = litellm_params or {}
|
||||
logging_obj = kwargs.get("litellm_logging_obj")
|
||||
logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
trace_id = getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
|
||||
|
||||
# Route through completion bridge if custom_llm_provider is set
|
||||
if custom_llm_provider:
|
||||
|
|
@ -450,7 +445,7 @@ async def asend_message(
|
|||
if api_base is None:
|
||||
raise ValueError("Either a2a_client or api_base is required for standard A2A flow")
|
||||
trace_id = trace_id or str(uuid.uuid4())
|
||||
extra_headers: dict[str, str] = {"X-LiteLLM-Trace-Id": trace_id}
|
||||
extra_headers: Final[dict[str, str]] = {"X-LiteLLM-Trace-Id": trace_id}
|
||||
if agent_id:
|
||||
extra_headers["X-LiteLLM-Agent-Id"] = agent_id
|
||||
# Overlay agent-level headers (agent headers take precedence over LiteLLM internal ones)
|
||||
|
|
@ -461,15 +456,15 @@ async def asend_message(
|
|||
# Type assertion: a2a_client is guaranteed to be non-None here
|
||||
assert a2a_client is not None
|
||||
|
||||
agent_name = _get_a2a_model_info(a2a_client, kwargs)
|
||||
agent_name: Final = _get_a2a_model_info(a2a_client, kwargs)
|
||||
|
||||
verbose_logger.info("A2A send_message request_id=%s, agent=%s", request.id, agent_name)
|
||||
|
||||
# Get agent card URL for localhost retry logic
|
||||
agent_card = _get_a2a_client_agent_card(a2a_client)
|
||||
card_url = get_agent_card_url(agent_card) if agent_card else None
|
||||
agent_card: Final = _get_a2a_client_agent_card(a2a_client)
|
||||
card_url: Final = get_agent_card_url(agent_card) if agent_card else None
|
||||
|
||||
a2a_response = await _execute_a2a_send_with_retry(
|
||||
a2a_response: Final = await _execute_a2a_send_with_retry(
|
||||
a2a_client=a2a_client,
|
||||
request=request,
|
||||
agent_card=agent_card,
|
||||
|
|
@ -481,10 +476,10 @@ async def asend_message(
|
|||
verbose_logger.info("A2A send_message completed, request_id=%s", request.id)
|
||||
|
||||
# Wrap in LiteLLM response type for _hidden_params support
|
||||
response = LiteLLMSendMessageResponse.from_a2a_response(a2a_response, request_id=str(request.id))
|
||||
response: Final = LiteLLMSendMessageResponse.from_a2a_response(a2a_response, request_id=str(request.id))
|
||||
|
||||
# Calculate token usage from request and response
|
||||
response_dict = a2a_response.model_dump(mode="json", exclude_none=True)
|
||||
response_dict: Final = a2a_response.model_dump(mode="json", exclude_none=True)
|
||||
(
|
||||
prompt_tokens,
|
||||
completion_tokens,
|
||||
|
|
@ -549,10 +544,10 @@ def _build_streaming_logging_obj(
|
|||
proxy_server_request: dict[str, Any] | None,
|
||||
) -> Logging:
|
||||
"""Build logging object for streaming A2A requests."""
|
||||
start_time = datetime.datetime.now()
|
||||
model = f"a2a_agent/{agent_name}"
|
||||
start_time: Final = datetime.datetime.now()
|
||||
model: Final = f"a2a_agent/{agent_name}"
|
||||
|
||||
logging_obj = Logging(
|
||||
logging_obj: Final = Logging(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "streaming-request"}],
|
||||
stream=False,
|
||||
|
|
@ -569,7 +564,7 @@ def _build_streaming_logging_obj(
|
|||
if agent_id:
|
||||
logging_obj.model_call_details["agent_id"] = agent_id
|
||||
|
||||
_litellm_params = litellm_params.copy() if litellm_params else {}
|
||||
_litellm_params: Final = litellm_params.copy() if litellm_params else {}
|
||||
if metadata:
|
||||
_litellm_params["metadata"] = metadata
|
||||
if proxy_server_request:
|
||||
|
|
@ -632,7 +627,7 @@ async def asend_message_streaming(
|
|||
```
|
||||
"""
|
||||
litellm_params = litellm_params or {}
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
|
||||
|
||||
# Route through completion bridge if custom_llm_provider is set
|
||||
if custom_llm_provider:
|
||||
|
|
@ -647,7 +642,7 @@ async def asend_message_streaming(
|
|||
)
|
||||
|
||||
# Extract params from request
|
||||
params = (
|
||||
params: Final = (
|
||||
request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params)
|
||||
)
|
||||
|
||||
|
|
@ -664,15 +659,15 @@ async def asend_message_streaming(
|
|||
if request is None:
|
||||
raise ValueError("request is required")
|
||||
|
||||
_raw_logging_obj = kwargs.get("litellm_logging_obj")
|
||||
_raw_logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
logging_obj: Logging | None = _raw_logging_obj if isinstance(_raw_logging_obj, Logging) else None
|
||||
|
||||
if a2a_client is None:
|
||||
if api_base is None:
|
||||
raise ValueError("Either a2a_client or api_base is required for standard A2A flow")
|
||||
logging_trace_id = getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None
|
||||
trace_id = logging_trace_id or (str(request.id) if request.id else str(uuid.uuid4()))
|
||||
extra_headers: dict[str, str] = {"X-LiteLLM-Trace-Id": trace_id}
|
||||
logging_trace_id: Final = getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None
|
||||
trace_id: Final = logging_trace_id or (str(request.id) if request.id else str(uuid.uuid4()))
|
||||
extra_headers: Final[dict[str, str]] = {"X-LiteLLM-Trace-Id": trace_id}
|
||||
if agent_id:
|
||||
extra_headers["X-LiteLLM-Agent-Id"] = agent_id
|
||||
if agent_extra_headers:
|
||||
|
|
@ -685,7 +680,7 @@ async def asend_message_streaming(
|
|||
|
||||
assert a2a_client is not None
|
||||
|
||||
agent_name = _get_a2a_model_info(a2a_client, kwargs)
|
||||
agent_name: Final = _get_a2a_model_info(a2a_client, kwargs)
|
||||
|
||||
if logging_obj is None:
|
||||
logging_obj = _build_streaming_logging_obj(
|
||||
|
|
@ -699,10 +694,10 @@ async def asend_message_streaming(
|
|||
|
||||
verbose_logger.info("A2A send_message_streaming request_id=%s, agent=%s", request.id, agent_name)
|
||||
|
||||
agent_card = _get_a2a_client_agent_card(a2a_client)
|
||||
card_url = get_agent_card_url(agent_card) if agent_card else None
|
||||
agent_card: Final = _get_a2a_client_agent_card(a2a_client)
|
||||
card_url: Final = get_agent_card_url(agent_card) if agent_card else None
|
||||
|
||||
stream = _execute_a2a_stream_with_retry(
|
||||
stream: Final = _execute_a2a_stream_with_retry(
|
||||
a2a_client=a2a_client,
|
||||
request=request,
|
||||
agent_card=agent_card,
|
||||
|
|
@ -769,21 +764,21 @@ async def create_a2a_client(
|
|||
# Only pass params that AsyncHTTPHandler.__init__ accepts (e.g. timeout).
|
||||
# Use "disable_aiohttp_transport" key for cache-key-only data (it's
|
||||
# filtered out before reaching the constructor).
|
||||
_client_params: dict = {"timeout": timeout}
|
||||
_client_params: Final[dict] = {"timeout": timeout}
|
||||
if extra_headers:
|
||||
# Encode headers into a cache-key-only param so each unique header
|
||||
# set produces a distinct cache key.
|
||||
_client_params["disable_aiohttp_transport"] = str(sorted(extra_headers.items()))
|
||||
_async_handler = get_async_httpx_client(
|
||||
_async_handler: Final = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.A2AProvider,
|
||||
params=_client_params,
|
||||
)
|
||||
httpx_client = _async_handler.client
|
||||
httpx_client: Final = _async_handler.client
|
||||
if extra_headers:
|
||||
httpx_client.headers.update(extra_headers)
|
||||
verbose_proxy_logger.debug("A2A client created with extra_headers=%s", list(extra_headers.keys()))
|
||||
|
||||
a2a_client = await create_client( # pyright: ignore[reportOptionalCall]
|
||||
a2a_client: Final = await create_client( # pyright: ignore[reportOptionalCall]
|
||||
base_url,
|
||||
client_config=ClientConfig( # pyright: ignore[reportOptionalCall]
|
||||
httpx_client=httpx_client,
|
||||
|
|
@ -794,7 +789,7 @@ async def create_a2a_client(
|
|||
# the configured httpx client (with this agent's trace-id/auth headers) without
|
||||
# excavating a2a-sdk private internals.
|
||||
a2a_client._litellm_httpx_client = httpx_client # type: ignore[attr-defined]
|
||||
agent_card = getattr(a2a_client, "_card", None)
|
||||
agent_card: Final = getattr(a2a_client, "_card", None)
|
||||
if agent_card is not None:
|
||||
a2a_client._litellm_agent_card = agent_card # type: ignore[attr-defined]
|
||||
|
||||
|
|
@ -827,17 +822,17 @@ async def aget_agent_card(
|
|||
verbose_logger.info("Fetching agent card from %s", base_url)
|
||||
|
||||
# Use LiteLLM's cached httpx client
|
||||
http_handler = get_async_httpx_client(
|
||||
http_handler: Final = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.A2A,
|
||||
params={"timeout": timeout},
|
||||
)
|
||||
httpx_client = http_handler.client
|
||||
httpx_client: Final = http_handler.client
|
||||
|
||||
resolver = A2ACardResolver(
|
||||
resolver: Final = A2ACardResolver(
|
||||
httpx_client=httpx_client,
|
||||
base_url=base_url,
|
||||
)
|
||||
agent_card = await resolver.get_agent_card()
|
||||
agent_card: Final = await resolver.get_agent_card()
|
||||
|
||||
verbose_logger.info("Fetched agent card: %s", agent_card.name if hasattr(agent_card, "name") else "unknown")
|
||||
return agent_card
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ Bedrock AgentCore A2A provider configuration.
|
|||
"""
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig
|
||||
from litellm.a2a_protocol.providers.bedrock_agentcore.handler import (
|
||||
|
|
@ -28,7 +28,7 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig):
|
|||
**kwargs,
|
||||
) -> dict[str, Any]:
|
||||
"""Handle non-streaming request to AgentCore A2A agent."""
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
litellm_params: Final = kwargs.get("litellm_params")
|
||||
if not litellm_params:
|
||||
raise ValueError(
|
||||
"litellm_params is required for BedrockAgentCoreA2AConfig (must contain model with AgentCore ARN)"
|
||||
|
|
@ -48,7 +48,7 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig):
|
|||
**kwargs,
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
"""Handle streaming request to AgentCore A2A agent."""
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
litellm_params: Final = kwargs.get("litellm_params")
|
||||
if not litellm_params:
|
||||
raise ValueError(
|
||||
"litellm_params is required for BedrockAgentCoreA2AConfig (must contain model with AgentCore ARN)"
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ completion bridge that would otherwise strip the envelope.
|
|||
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any, cast
|
||||
from typing import Any, Final, cast
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
|
||||
|
|
@ -55,16 +55,16 @@ class BedrockAgentCoreA2AHandler:
|
|||
|
||||
verbose_logger.info("BedrockAgentCore A2A: Sending non-streaming request to %s", url)
|
||||
|
||||
client = get_async_httpx_client(
|
||||
client: Final = get_async_httpx_client(
|
||||
llm_provider=cast(Any, httpxSpecialProvider.A2AProvider),
|
||||
)
|
||||
response = await client.post(
|
||||
response: Final = await client.post(
|
||||
url,
|
||||
headers=headers,
|
||||
data=body,
|
||||
)
|
||||
response.raise_for_status()
|
||||
response_data = response.json()
|
||||
response_data: Final = response.json()
|
||||
|
||||
if "error" in response_data:
|
||||
verbose_logger.warning("BedrockAgentCore A2A: Agent returned error: %s", response_data["error"])
|
||||
|
|
@ -102,10 +102,10 @@ class BedrockAgentCoreA2AHandler:
|
|||
|
||||
verbose_logger.info("BedrockAgentCore A2A: Sending streaming request to %s", url)
|
||||
|
||||
client = get_async_httpx_client(
|
||||
client: Final = get_async_httpx_client(
|
||||
llm_provider=cast(Any, httpxSpecialProvider.A2AProvider),
|
||||
)
|
||||
response = await client.post(
|
||||
response: Final = await client.post(
|
||||
url,
|
||||
headers=headers,
|
||||
data=body,
|
||||
|
|
@ -114,15 +114,15 @@ class BedrockAgentCoreA2AHandler:
|
|||
response.raise_for_status()
|
||||
|
||||
# Check content type — AgentCore may return JSON instead of SSE
|
||||
content_type = response.headers.get("content-type", "").lower()
|
||||
content_type: Final = response.headers.get("content-type", "").lower()
|
||||
|
||||
if "application/json" in content_type:
|
||||
# Single JSON response fallback (not SSE)
|
||||
verbose_logger.debug(
|
||||
"BedrockAgentCore A2A streaming: received JSON instead of SSE, yielding as single event"
|
||||
)
|
||||
response_body = await response.aread()
|
||||
response_data = json.loads(response_body)
|
||||
response_body: Final = await response.aread()
|
||||
response_data: Final = json.loads(response_body)
|
||||
yield response_data
|
||||
else:
|
||||
# SSE stream — parse data: lines
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ and signs requests via AmazonAgentCoreConfig (SigV4 or JWT).
|
|||
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Mapping
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreConfig
|
||||
|
|
@ -23,13 +23,13 @@ from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreCo
|
|||
# ``runtimeSessionId`` / ``runtimeUserId`` in the agent's ``litellm_params``;
|
||||
# ``authorization`` is set by the AgentCore signer (JWT or SigV4); ``host`` and
|
||||
# the ``x-amz-*`` family are owned by SigV4 itself.
|
||||
_RESERVED_EXACT_HEADERS = frozenset(
|
||||
_RESERVED_EXACT_HEADERS: Final = frozenset(
|
||||
{
|
||||
"authorization",
|
||||
"host",
|
||||
}
|
||||
)
|
||||
_RESERVED_PREFIX_HEADERS: tuple[str, ...] = (
|
||||
_RESERVED_PREFIX_HEADERS: Final[tuple[str, ...]] = (
|
||||
"x-amzn-bedrock-agentcore-runtime-",
|
||||
"x-amz-",
|
||||
)
|
||||
|
|
@ -47,8 +47,8 @@ def _filter_reserved_headers(
|
|||
if not agent_extra_headers:
|
||||
return None
|
||||
|
||||
filtered: dict[str, str] = {}
|
||||
dropped: list = []
|
||||
filtered: Final[dict[str, str]] = {}
|
||||
dropped: Final[list] = []
|
||||
for k, v in agent_extra_headers.items():
|
||||
k_lower = k.lower()
|
||||
if k_lower in _RESERVED_EXACT_HEADERS or any(k_lower.startswith(prefix) for prefix in _RESERVED_PREFIX_HEADERS):
|
||||
|
|
@ -107,19 +107,19 @@ class BedrockAgentCoreA2ATransformation:
|
|||
"""
|
||||
# Extract model and strip the "bedrock/" prefix
|
||||
# "bedrock/agentcore/arn:aws:..." → "agentcore/arn:aws:..."
|
||||
model = litellm_params.get("model", "")
|
||||
model: Final = litellm_params.get("model", "")
|
||||
if model.startswith("bedrock/"):
|
||||
agentcore_model = model[len("bedrock/") :]
|
||||
else:
|
||||
agentcore_model = model
|
||||
|
||||
# Build optional_params from litellm_params (everything except model and custom_llm_provider)
|
||||
optional_params = {k: v for k, v in litellm_params.items() if k not in ("model", "custom_llm_provider")}
|
||||
optional_params: Final = {k: v for k, v in litellm_params.items() if k not in ("model", "custom_llm_provider")}
|
||||
|
||||
agentcore_config = AmazonAgentCoreConfig()
|
||||
agentcore_config: Final = AmazonAgentCoreConfig()
|
||||
|
||||
# Derive URL from ARN
|
||||
url = agentcore_config.get_complete_url(
|
||||
url: Final = agentcore_config.get_complete_url(
|
||||
api_base=optional_params.get("api_base"),
|
||||
api_key=optional_params.get("api_key"),
|
||||
model=agentcore_model,
|
||||
|
|
@ -129,7 +129,7 @@ class BedrockAgentCoreA2ATransformation:
|
|||
)
|
||||
|
||||
# Construct JSON-RPC 2.0 envelope
|
||||
json_rpc_body = {
|
||||
json_rpc_body: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"method": method,
|
||||
"id": request_id,
|
||||
|
|
@ -138,17 +138,17 @@ class BedrockAgentCoreA2ATransformation:
|
|||
|
||||
# Set required AgentCore session headers (normally set by transform_request,
|
||||
# which we skip because it also builds {"prompt": "..."})
|
||||
headers: dict = {}
|
||||
session_id = agentcore_config._get_runtime_session_id(optional_params)
|
||||
headers: Final[dict] = {}
|
||||
session_id: Final = agentcore_config._get_runtime_session_id(optional_params)
|
||||
headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] = session_id
|
||||
runtime_user_id = agentcore_config._get_runtime_user_id(optional_params)
|
||||
runtime_user_id: Final = agentcore_config._get_runtime_user_id(optional_params)
|
||||
if runtime_user_id:
|
||||
headers["X-Amzn-Bedrock-AgentCore-Runtime-User-Id"] = runtime_user_id
|
||||
|
||||
# Merge per-request agent headers before signing so SigV4 covers them.
|
||||
# Reserved headers are stripped first to prevent client-controlled values
|
||||
# from spoofing the AgentCore runtime identity / SigV4 metadata.
|
||||
safe_extra_headers = _filter_reserved_headers(agent_extra_headers)
|
||||
safe_extra_headers: Final = _filter_reserved_headers(agent_extra_headers)
|
||||
if safe_extra_headers:
|
||||
headers.update(safe_extra_headers)
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ This handler provides fake streaming by converting non-streaming responses into
|
|||
"""
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.a2a_protocol.providers.pydantic_ai_agents.transformation import (
|
||||
|
|
@ -50,7 +50,7 @@ class PydanticAIHandler:
|
|||
verbose_logger.info("Pydantic AI: Routing to Pydantic AI agent at %s", api_base)
|
||||
|
||||
# Send request directly to Pydantic AI agent
|
||||
response_data = await PydanticAITransformation.send_non_streaming_request(
|
||||
response_data: Final = await PydanticAITransformation.send_non_streaming_request(
|
||||
api_base=api_base,
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
|
|
@ -95,7 +95,7 @@ class PydanticAIHandler:
|
|||
verbose_logger.info("Pydantic AI: Faking streaming for Pydantic AI agent at %s", api_base)
|
||||
|
||||
# Get raw task response first (not the transformed A2A format)
|
||||
raw_response = await PydanticAITransformation.send_and_get_raw_response(
|
||||
raw_response: Final = await PydanticAITransformation.send_and_get_raw_response(
|
||||
api_base=api_base,
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ This module provides fake streaming by converting non-streaming responses into s
|
|||
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any, cast
|
||||
from typing import Any, Final, cast
|
||||
from uuid import uuid4
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -163,7 +163,7 @@ class PydanticAITransformation:
|
|||
params_dict["message"]["kind"] = "message"
|
||||
|
||||
# Build A2A JSON-RPC request using message/send method for FastA2A compatibility
|
||||
a2a_request = {
|
||||
a2a_request: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"method": "message/send",
|
||||
|
|
@ -171,16 +171,16 @@ class PydanticAITransformation:
|
|||
}
|
||||
|
||||
# FastA2A uses root endpoint (/) not /messages
|
||||
endpoint = api_base.rstrip("/")
|
||||
endpoint: Final = api_base.rstrip("/")
|
||||
|
||||
verbose_logger.info("Pydantic AI: Sending non-streaming request to %s", endpoint)
|
||||
|
||||
# Send request to Pydantic AI agent using shared async HTTP client
|
||||
client = get_async_httpx_client(
|
||||
client: Final = get_async_httpx_client(
|
||||
llm_provider=cast(Any, "pydantic_ai_agent"),
|
||||
params={"timeout": timeout},
|
||||
)
|
||||
response = await client.post(
|
||||
response: Final = await client.post(
|
||||
endpoint,
|
||||
json=a2a_request,
|
||||
headers={
|
||||
|
|
@ -192,13 +192,13 @@ class PydanticAITransformation:
|
|||
response_data = response.json()
|
||||
|
||||
# Check if task is already completed
|
||||
result = response_data.get("result", {})
|
||||
status = result.get("status", {})
|
||||
state = status.get("state", "")
|
||||
result: Final = response_data.get("result", {})
|
||||
status: Final = result.get("status", {})
|
||||
state: Final = status.get("state", "")
|
||||
|
||||
if state != "completed":
|
||||
# Need to poll for completion
|
||||
task_id = result.get("id")
|
||||
task_id: Final = result.get("id")
|
||||
if task_id:
|
||||
verbose_logger.info("Pydantic AI: Task %s submitted, polling for completion...", task_id)
|
||||
response_data = await PydanticAITransformation._poll_for_completion(
|
||||
|
|
@ -235,7 +235,7 @@ class PydanticAITransformation:
|
|||
Standard A2A non-streaming response format with message
|
||||
"""
|
||||
# Get raw task response
|
||||
raw_response = await PydanticAITransformation._send_and_poll_raw(
|
||||
raw_response: Final = await PydanticAITransformation._send_and_poll_raw(
|
||||
api_base=api_base,
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
|
|
@ -313,7 +313,7 @@ class PydanticAITransformation:
|
|||
full_text, message_id, parts = PydanticAITransformation._extract_response_text(response_data)
|
||||
|
||||
# Build standard A2A message
|
||||
a2a_message = {
|
||||
a2a_message: Final = {
|
||||
"kind": "message",
|
||||
"role": "agent",
|
||||
"parts": parts if parts else [{"kind": "text", "text": full_text}],
|
||||
|
|
@ -342,10 +342,10 @@ class PydanticAITransformation:
|
|||
Returns:
|
||||
Tuple of (full_text, message_id, parts)
|
||||
"""
|
||||
result = response_data.get("result", {})
|
||||
result: Final = response_data.get("result", {})
|
||||
|
||||
# Try to extract from artifacts first (preferred for results)
|
||||
artifacts = result.get("artifacts", [])
|
||||
artifacts: Final = result.get("artifacts", [])
|
||||
if artifacts:
|
||||
for artifact in artifacts:
|
||||
parts = artifact.get("parts", [])
|
||||
|
|
@ -356,7 +356,7 @@ class PydanticAITransformation:
|
|||
return text, str(uuid4()), parts
|
||||
|
||||
# Fall back to history - get the last agent message
|
||||
history = result.get("history", [])
|
||||
history: Final = result.get("history", [])
|
||||
for msg in reversed(history):
|
||||
if msg.get("role") == "agent":
|
||||
parts = msg.get("parts", [])
|
||||
|
|
@ -369,7 +369,7 @@ class PydanticAITransformation:
|
|||
return full_text, message_id, parts
|
||||
|
||||
# Fall back to message field (original format)
|
||||
message = result.get("message", {})
|
||||
message: Final = result.get("message", {})
|
||||
if message:
|
||||
parts = message.get("parts", [])
|
||||
message_id = message.get("messageId", str(uuid4()))
|
||||
|
|
@ -410,8 +410,8 @@ class PydanticAITransformation:
|
|||
full_text, message_id, parts = PydanticAITransformation._extract_response_text(response_data)
|
||||
|
||||
# Extract input message from raw response for history
|
||||
result = response_data.get("result", {})
|
||||
history = result.get("history", [])
|
||||
result: Final = response_data.get("result", {})
|
||||
history: Final = result.get("history", [])
|
||||
input_message = {}
|
||||
for msg in history:
|
||||
if msg.get("role") == "user":
|
||||
|
|
@ -419,14 +419,14 @@ class PydanticAITransformation:
|
|||
break
|
||||
|
||||
# Generate IDs for streaming events
|
||||
task_id = str(uuid4())
|
||||
context_id = str(uuid4())
|
||||
artifact_id = str(uuid4())
|
||||
input_message_id = input_message.get("messageId", str(uuid4()))
|
||||
task_id: Final = str(uuid4())
|
||||
context_id: Final = str(uuid4())
|
||||
artifact_id: Final = str(uuid4())
|
||||
input_message_id: Final = input_message.get("messageId", str(uuid4()))
|
||||
|
||||
# 1. Emit initial task event (kind: "task", status: "submitted")
|
||||
# Format matches A2ACompletionBridgeTransformation.create_task_event
|
||||
task_event = {
|
||||
task_event: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"result": {
|
||||
|
|
@ -452,7 +452,7 @@ class PydanticAITransformation:
|
|||
|
||||
# 2. Emit status update (kind: "status-update", status: "working")
|
||||
# Format matches A2ACompletionBridgeTransformation.create_status_update_event
|
||||
working_event = {
|
||||
working_event: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"result": {
|
||||
|
|
@ -503,7 +503,7 @@ class PydanticAITransformation:
|
|||
await asyncio.sleep(delay_ms / 1000.0)
|
||||
|
||||
# 4. Emit final status update (kind: "status-update", status: "completed", final: true)
|
||||
completed_event = {
|
||||
completed_event: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"result": {
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ A2A provider configuration for IBM watsonx Orchestrate (WXO).
|
|||
"""
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig
|
||||
from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import (
|
||||
|
|
@ -22,7 +22,7 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
|
|||
**kwargs: Any,
|
||||
) -> dict[str, Any]:
|
||||
"""Handle a non-streaming A2A request via WXO runs API."""
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
litellm_params: Final = kwargs.get("litellm_params")
|
||||
if not litellm_params:
|
||||
raise ValueError(
|
||||
"litellm_params is required for WatsonxOrchestrateA2AConfig "
|
||||
|
|
@ -42,7 +42,7 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
|
|||
**kwargs: Any,
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
"""Handle a streaming A2A request via WXO streaming runs API."""
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
litellm_params: Final = kwargs.get("litellm_params")
|
||||
if not litellm_params:
|
||||
raise ValueError(
|
||||
"litellm_params is required for WatsonxOrchestrateA2AConfig "
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import hashlib
|
|||
import json
|
||||
import time
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any, NamedTuple, cast
|
||||
from typing import Any, Final, NamedTuple, cast
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -21,11 +21,11 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
_IBM_CLOUD_IAM_URL = "https://iam.cloud.ibm.com/identity/token"
|
||||
_POLL_INTERVAL_S = 2.0
|
||||
_MAX_POLL_ATTEMPTS = 90
|
||||
_TOKEN_CACHE_TTL_BUFFER_S = 60
|
||||
_token_cache: dict[str, tuple[str, float]] = {}
|
||||
_IBM_CLOUD_IAM_URL: Final = "https://iam.cloud.ibm.com/identity/token"
|
||||
_POLL_INTERVAL_S: Final = 2.0
|
||||
_MAX_POLL_ATTEMPTS: Final = 90
|
||||
_TOKEN_CACHE_TTL_BUFFER_S: Final = 60
|
||||
_token_cache: Final[dict[str, tuple[str, float]]] = {}
|
||||
|
||||
|
||||
class WXORequestParams(NamedTuple):
|
||||
|
|
@ -53,14 +53,14 @@ class WatsonxOrchestrateHandler:
|
|||
api_key: str,
|
||||
username: str | None,
|
||||
) -> str:
|
||||
material = f"{auth_mode}:{cp4d_host}:{username or ''}:{api_key}"
|
||||
material: Final = f"{auth_mode}:{cp4d_host}:{username or ''}:{api_key}"
|
||||
return hashlib.sha256(material.encode()).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _cp4d_token_ttl_seconds(expiration: Any, now_wall: float | None = None) -> int:
|
||||
# CP4D returns expiration as absolute Unix epoch seconds, not a duration.
|
||||
expires_at = int(expiration)
|
||||
wall = now_wall if now_wall is not None else time.time()
|
||||
expires_at: Final = int(expiration)
|
||||
wall: Final = now_wall if now_wall is not None else time.time()
|
||||
return max(expires_at - int(wall), 0)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -71,9 +71,9 @@ class WatsonxOrchestrateHandler:
|
|||
username: str | None = None,
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
) -> str:
|
||||
cache_key = WatsonxOrchestrateHandler._token_cache_key(auth_mode, cp4d_host, api_key, username)
|
||||
now = time.monotonic()
|
||||
cached = _token_cache.get(cache_key)
|
||||
cache_key: Final = WatsonxOrchestrateHandler._token_cache_key(auth_mode, cp4d_host, api_key, username)
|
||||
now: Final = time.monotonic()
|
||||
cached: Final = _token_cache.get(cache_key)
|
||||
if cached and cached[1] > now:
|
||||
return cached[0]
|
||||
|
||||
|
|
@ -96,7 +96,7 @@ class WatsonxOrchestrateHandler:
|
|||
else:
|
||||
if not username:
|
||||
raise ValueError("'username' is required in litellm_params when auth_mode='cp4d'")
|
||||
token_url = f"{cp4d_host.rstrip('/')}/icp4d-api/v1/authorize"
|
||||
token_url: Final = f"{cp4d_host.rstrip('/')}/icp4d-api/v1/authorize"
|
||||
response = await client.post(
|
||||
token_url,
|
||||
json={"username": username, "api_key": api_key},
|
||||
|
|
@ -105,13 +105,13 @@ class WatsonxOrchestrateHandler:
|
|||
response.raise_for_status()
|
||||
payload = response.json()
|
||||
token = str(payload["token"])
|
||||
expiration = payload.get("expiration")
|
||||
expiration: Final = payload.get("expiration")
|
||||
if expiration is None:
|
||||
ttl_s = 3600
|
||||
else:
|
||||
ttl_s = WatsonxOrchestrateHandler._cp4d_token_ttl_seconds(expiration)
|
||||
|
||||
expires_at = now + max(ttl_s - _TOKEN_CACHE_TTL_BUFFER_S, 0)
|
||||
expires_at: Final = now + max(ttl_s - _TOKEN_CACHE_TTL_BUFFER_S, 0)
|
||||
_token_cache[cache_key] = (token, expires_at)
|
||||
for stale_key, (_, stale_expires_at) in list(_token_cache.items()):
|
||||
if stale_expires_at <= now:
|
||||
|
|
@ -127,7 +127,7 @@ class WatsonxOrchestrateHandler:
|
|||
max_attempts: int = _MAX_POLL_ATTEMPTS,
|
||||
interval_s: float = _POLL_INTERVAL_S,
|
||||
) -> dict[str, Any]:
|
||||
url = f"{base_url}/v1/orchestrate/runs/{run_id}"
|
||||
url: Final = f"{base_url}/v1/orchestrate/runs/{run_id}"
|
||||
|
||||
for attempt in range(max_attempts):
|
||||
await asyncio.sleep(interval_s)
|
||||
|
|
@ -152,7 +152,7 @@ class WatsonxOrchestrateHandler:
|
|||
) -> dict[str, Any]:
|
||||
status = run_data.get("status", "")
|
||||
if status not in WatsonxOrchestrateTransformation.TERMINAL_STATES:
|
||||
run_id = run_data.get("run_id") or run_data.get("id") or ""
|
||||
run_id: Final = run_data.get("run_id") or run_data.get("id") or ""
|
||||
if not run_id:
|
||||
raise ValueError(f"WXO: No run_id in response: {run_data}")
|
||||
run_data = await WatsonxOrchestrateHandler._poll_run(
|
||||
|
|
@ -188,10 +188,10 @@ class WatsonxOrchestrateHandler:
|
|||
|
||||
@staticmethod
|
||||
def _extract_litellm_params(litellm_params: dict[str, Any]) -> WXORequestParams:
|
||||
cp4d_host = litellm_params.get("cp4d_host") or ""
|
||||
instance_id = litellm_params.get("instance_id") or ""
|
||||
wxo_agent_id = litellm_params.get("wxo_agent_id") or ""
|
||||
api_key = litellm_params.get("api_key") or ""
|
||||
cp4d_host: Final = litellm_params.get("cp4d_host") or ""
|
||||
instance_id: Final = litellm_params.get("instance_id") or ""
|
||||
wxo_agent_id: Final = litellm_params.get("wxo_agent_id") or ""
|
||||
api_key: Final = litellm_params.get("api_key") or ""
|
||||
|
||||
if not cp4d_host:
|
||||
raise ValueError("'cp4d_host' is required in litellm_params for WXO agents")
|
||||
|
|
@ -218,29 +218,29 @@ class WatsonxOrchestrateHandler:
|
|||
params: dict[str, Any],
|
||||
litellm_params: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
wxo = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
|
||||
wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
|
||||
|
||||
client = WatsonxOrchestrateHandler._http_client(timeout=90.0)
|
||||
token = await WatsonxOrchestrateHandler._get_bearer_token(
|
||||
client: Final = WatsonxOrchestrateHandler._http_client(timeout=90.0)
|
||||
token: Final = await WatsonxOrchestrateHandler._get_bearer_token(
|
||||
cp4d_host=wxo.cp4d_host,
|
||||
auth_mode=wxo.auth_mode,
|
||||
api_key=wxo.api_key,
|
||||
username=wxo.username,
|
||||
client=client,
|
||||
)
|
||||
base_url = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id)
|
||||
auth_headers = {
|
||||
base_url: Final = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id)
|
||||
auth_headers: Final = {
|
||||
"Authorization": f"Bearer {token}",
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
|
||||
text = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params)
|
||||
body = WatsonxOrchestrateTransformation.build_wxo_run_body(
|
||||
text: Final = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params)
|
||||
body: Final = WatsonxOrchestrateTransformation.build_wxo_run_body(
|
||||
wxo_agent_id=wxo.wxo_agent_id, text=text, thread_id=wxo.thread_id
|
||||
)
|
||||
|
||||
run_response = await client.post(
|
||||
run_response: Final = await client.post(
|
||||
f"{base_url}/v1/orchestrate/runs",
|
||||
json=body,
|
||||
headers=auth_headers,
|
||||
|
|
@ -255,7 +255,7 @@ class WatsonxOrchestrateHandler:
|
|||
client=client,
|
||||
)
|
||||
|
||||
response_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(run_data)
|
||||
response_text: Final = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(run_data)
|
||||
return WatsonxOrchestrateTransformation.build_a2a_message_response(request_id=request_id, text=response_text)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -266,29 +266,29 @@ class WatsonxOrchestrateHandler:
|
|||
chunk_size: int = 50,
|
||||
delay_ms: int = 10,
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
wxo = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
|
||||
wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
|
||||
|
||||
client = WatsonxOrchestrateHandler._http_client(timeout=120.0)
|
||||
token = await WatsonxOrchestrateHandler._get_bearer_token(
|
||||
client: Final = WatsonxOrchestrateHandler._http_client(timeout=120.0)
|
||||
token: Final = await WatsonxOrchestrateHandler._get_bearer_token(
|
||||
cp4d_host=wxo.cp4d_host,
|
||||
auth_mode=wxo.auth_mode,
|
||||
api_key=wxo.api_key,
|
||||
username=wxo.username,
|
||||
client=client,
|
||||
)
|
||||
base_url = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id)
|
||||
auth_headers = {
|
||||
base_url: Final = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id)
|
||||
auth_headers: Final = {
|
||||
"Authorization": f"Bearer {token}",
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "text/event-stream, application/json",
|
||||
}
|
||||
text = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params)
|
||||
body = WatsonxOrchestrateTransformation.build_wxo_run_body(
|
||||
text: Final = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params)
|
||||
body: Final = WatsonxOrchestrateTransformation.build_wxo_run_body(
|
||||
wxo_agent_id=wxo.wxo_agent_id, text=text, thread_id=wxo.thread_id
|
||||
)
|
||||
|
||||
try:
|
||||
response = await client.post(
|
||||
response: Final = await client.post(
|
||||
f"{base_url}/v1/orchestrate/runs/stream",
|
||||
json=body,
|
||||
headers=auth_headers,
|
||||
|
|
@ -306,7 +306,7 @@ class WatsonxOrchestrateHandler:
|
|||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
response_text = WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response(result)
|
||||
response_text: Final = WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response(result)
|
||||
async for chunk in WatsonxOrchestrateTransformation.fake_streaming_from_text(
|
||||
text=response_text,
|
||||
request_id=request_id,
|
||||
|
|
@ -316,9 +316,9 @@ class WatsonxOrchestrateHandler:
|
|||
yield chunk
|
||||
return
|
||||
|
||||
content_type = response.headers.get("content-type", "").lower()
|
||||
content_type: Final = response.headers.get("content-type", "").lower()
|
||||
if "text/event-stream" not in content_type:
|
||||
response_body = await response.aread()
|
||||
response_body: Final = await response.aread()
|
||||
result = json.loads(response_body)
|
||||
result = await WatsonxOrchestrateHandler._get_successful_run_data(
|
||||
run_data=result,
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ WXO uses a REST API (not A2A/JSON-RPC) with an async-poll execution model:
|
|||
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
from uuid import uuid4
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -35,9 +35,9 @@ class WatsonxOrchestrateTransformation:
|
|||
|
||||
A2A format: params.message.parts[*] where part.kind == "text"
|
||||
"""
|
||||
message = params.get("message", {})
|
||||
parts = message.get("parts", [])
|
||||
texts = []
|
||||
message: Final = params.get("message", {})
|
||||
parts: Final = message.get("parts", [])
|
||||
texts: Final = []
|
||||
for part in parts:
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
|
|
@ -53,7 +53,7 @@ class WatsonxOrchestrateTransformation:
|
|||
thread_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Build the WXO POST /v1/orchestrate/runs request body."""
|
||||
body: dict[str, Any] = {
|
||||
body: Final[dict[str, Any]] = {
|
||||
"agent_id": wxo_agent_id,
|
||||
"message": {
|
||||
"role": "user",
|
||||
|
|
@ -96,7 +96,7 @@ class WatsonxOrchestrateTransformation:
|
|||
pass
|
||||
|
||||
# Tertiary: results as a raw string
|
||||
results = result.get("results")
|
||||
results: Final = result.get("results")
|
||||
if results and isinstance(results, str):
|
||||
return results
|
||||
|
||||
|
|
@ -104,11 +104,11 @@ class WatsonxOrchestrateTransformation:
|
|||
|
||||
@staticmethod
|
||||
def extract_text_from_a2a_message_response(a2a_response: dict[str, Any]) -> str:
|
||||
result = a2a_response.get("result")
|
||||
result: Final = a2a_response.get("result")
|
||||
if not isinstance(result, dict):
|
||||
verbose_logger.warning("WXO: A2A response missing result object")
|
||||
return ""
|
||||
parts = result.get("parts")
|
||||
parts: Final = result.get("parts")
|
||||
if not isinstance(parts, list):
|
||||
verbose_logger.warning("WXO: A2A result has no parts list")
|
||||
return ""
|
||||
|
|
@ -150,9 +150,9 @@ class WatsonxOrchestrateTransformation:
|
|||
3. artifact-update chunks
|
||||
4. status-update (kind="status-update", state="completed", final=True)
|
||||
"""
|
||||
task_id = str(uuid4())
|
||||
context_id = str(uuid4())
|
||||
artifact_id = str(uuid4())
|
||||
task_id: Final = str(uuid4())
|
||||
context_id: Final = str(uuid4())
|
||||
artifact_id: Final = str(uuid4())
|
||||
|
||||
# 1. Task submitted
|
||||
yield {
|
||||
|
|
@ -181,7 +181,7 @@ class WatsonxOrchestrateTransformation:
|
|||
await asyncio.sleep(delay_ms / 1000.0)
|
||||
|
||||
# 3. Artifact chunks (always emit at least one chunk, even for empty text)
|
||||
text_to_chunk = text or ""
|
||||
text_to_chunk: Final = text or ""
|
||||
for i in range(0, max(len(text_to_chunk), 1), chunk_size):
|
||||
chunk_text = text_to_chunk[i : i + chunk_size]
|
||||
is_last = (i + chunk_size) >= max(len(text_to_chunk), 1)
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ A2A Streaming Iterator with token tracking and logging support.
|
|||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -47,7 +47,7 @@ class A2AStreamingIterator:
|
|||
|
||||
async def __anext__(self) -> "SendStreamingMessageResponse":
|
||||
try:
|
||||
chunk = await self.stream.__anext__()
|
||||
chunk: Final = await self.stream.__anext__()
|
||||
|
||||
# Store chunk
|
||||
self.chunks.append(chunk)
|
||||
|
|
@ -71,8 +71,8 @@ class A2AStreamingIterator:
|
|||
def _collect_text_from_chunk(self, chunk: Any) -> None:
|
||||
"""Extract text from a streaming chunk and add to collected parts."""
|
||||
try:
|
||||
chunk_dict = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
|
||||
text = A2ARequestUtils.extract_text_from_response(chunk_dict)
|
||||
chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
|
||||
text: Final = A2ARequestUtils.extract_text_from_response(chunk_dict)
|
||||
if text:
|
||||
self.collected_text_parts.append(text)
|
||||
except Exception:
|
||||
|
|
@ -81,10 +81,10 @@ class A2AStreamingIterator:
|
|||
def _is_completed_chunk(self, chunk: Any) -> bool:
|
||||
"""Check if chunk indicates stream completion."""
|
||||
try:
|
||||
chunk_dict = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
|
||||
result = chunk_dict.get("result", {})
|
||||
chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
|
||||
result: Final = chunk_dict.get("result", {})
|
||||
if isinstance(result, dict):
|
||||
status = result.get("status", {})
|
||||
status: Final = result.get("status", {})
|
||||
if isinstance(status, dict):
|
||||
return status.get("state") == "completed"
|
||||
except Exception:
|
||||
|
|
@ -94,21 +94,21 @@ class A2AStreamingIterator:
|
|||
async def _handle_stream_complete(self) -> None:
|
||||
"""Handle logging and token counting when stream completes."""
|
||||
try:
|
||||
end_time = datetime.now()
|
||||
end_time: Final = datetime.now()
|
||||
|
||||
# Calculate tokens from collected text
|
||||
input_message = A2ARequestUtils.get_input_message_from_request(self.request)
|
||||
input_text = A2ARequestUtils.extract_text_from_message(input_message)
|
||||
prompt_tokens = A2ARequestUtils.count_tokens(input_text)
|
||||
input_message: Final = A2ARequestUtils.get_input_message_from_request(self.request)
|
||||
input_text: Final = A2ARequestUtils.extract_text_from_message(input_message)
|
||||
prompt_tokens: Final = A2ARequestUtils.count_tokens(input_text)
|
||||
|
||||
# Use the last (most complete) text from chunks
|
||||
output_text = self.collected_text_parts[-1] if self.collected_text_parts else ""
|
||||
completion_tokens = A2ARequestUtils.count_tokens(output_text)
|
||||
output_text: Final = self.collected_text_parts[-1] if self.collected_text_parts else ""
|
||||
completion_tokens: Final = A2ARequestUtils.count_tokens(output_text)
|
||||
|
||||
total_tokens = prompt_tokens + completion_tokens
|
||||
total_tokens: Final = prompt_tokens + completion_tokens
|
||||
|
||||
# Create usage object
|
||||
usage = litellm.Usage(
|
||||
usage: Final = litellm.Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=total_tokens,
|
||||
|
|
@ -120,11 +120,11 @@ class A2AStreamingIterator:
|
|||
self.logging_obj.model_call_details["stream"] = False
|
||||
|
||||
# Calculate cost using A2ACostCalculator
|
||||
response_cost = A2ACostCalculator.calculate_a2a_cost(self.logging_obj)
|
||||
response_cost: Final = A2ACostCalculator.calculate_a2a_cost(self.logging_obj)
|
||||
self.logging_obj.model_call_details["response_cost"] = response_cost
|
||||
|
||||
# Build result for logging
|
||||
result = self._build_logging_result(usage)
|
||||
result: Final = self._build_logging_result(usage)
|
||||
|
||||
# Call success handlers - they will build standard_logging_object
|
||||
asyncio.create_task(
|
||||
|
|
@ -150,7 +150,7 @@ class A2AStreamingIterator:
|
|||
|
||||
def _build_logging_result(self, usage: litellm.Usage) -> dict[str, Any]:
|
||||
"""Build a result dict for logging."""
|
||||
result: dict[str, Any] = {
|
||||
result: Final[dict[str, Any]] = {
|
||||
"id": getattr(self.request, "id", "unknown"),
|
||||
"jsonrpc": "2.0",
|
||||
"usage": (usage.model_dump() if hasattr(usage, "model_dump") else dict(usage)),
|
||||
|
|
@ -159,7 +159,7 @@ class A2AStreamingIterator:
|
|||
# Add final chunk result if available
|
||||
if self.final_chunk:
|
||||
try:
|
||||
chunk_dict = self.final_chunk.model_dump(mode="json", exclude_none=True)
|
||||
chunk_dict: Final = self.final_chunk.model_dump(mode="json", exclude_none=True)
|
||||
result["result"] = chunk_dict.get("result", {})
|
||||
except Exception:
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
Utility functions for A2A protocol.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -34,7 +34,7 @@ class A2ARequestUtils:
|
|||
else:
|
||||
parts = getattr(message, "parts", []) or []
|
||||
|
||||
text_parts: list[str] = []
|
||||
text_parts: Final[list[str]] = []
|
||||
for part in parts:
|
||||
if isinstance(part, dict):
|
||||
if part.get("kind") == "text":
|
||||
|
|
@ -56,7 +56,7 @@ class A2ARequestUtils:
|
|||
Returns:
|
||||
Text from response message parts
|
||||
"""
|
||||
result = response_dict.get("result", {})
|
||||
result: Final = response_dict.get("result", {})
|
||||
if not isinstance(result, dict):
|
||||
return ""
|
||||
|
||||
|
|
@ -66,7 +66,7 @@ class A2ARequestUtils:
|
|||
if result.get("kind") == "message":
|
||||
return A2ARequestUtils.extract_text_from_message(result)
|
||||
|
||||
message = result.get("message", {})
|
||||
message: Final = result.get("message", {})
|
||||
return A2ARequestUtils.extract_text_from_message(message)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -82,7 +82,7 @@ class A2ARequestUtils:
|
|||
Returns:
|
||||
The message object/dict or None
|
||||
"""
|
||||
params = getattr(request, "params", None)
|
||||
params: Final = getattr(request, "params", None)
|
||||
if params is None:
|
||||
return None
|
||||
return getattr(params, "message", None)
|
||||
|
|
@ -128,14 +128,14 @@ class A2ARequestUtils:
|
|||
input_message = A2ARequestUtils.get_input_message_from_request(request)
|
||||
if input_message is not None and hasattr(input_message, "model_dump"):
|
||||
input_message = input_message.model_dump(mode="json")
|
||||
input_text = A2ARequestUtils.extract_text_from_message(input_message)
|
||||
prompt_tokens = A2ARequestUtils.count_tokens(input_text)
|
||||
input_text: Final = A2ARequestUtils.extract_text_from_message(input_message)
|
||||
prompt_tokens: Final = A2ARequestUtils.count_tokens(input_text)
|
||||
|
||||
# Count output tokens
|
||||
output_text = A2ARequestUtils.extract_text_from_response(response_dict)
|
||||
completion_tokens = A2ARequestUtils.count_tokens(output_text)
|
||||
output_text: Final = A2ARequestUtils.extract_text_from_response(response_dict)
|
||||
completion_tokens: Final = A2ARequestUtils.count_tokens(output_text)
|
||||
|
||||
total_tokens = prompt_tokens + completion_tokens
|
||||
total_tokens: Final = prompt_tokens + completion_tokens
|
||||
|
||||
return prompt_tokens, completion_tokens, total_tokens
|
||||
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ Environment Variables:
|
|||
import json
|
||||
import os
|
||||
from importlib.resources import files
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -46,7 +47,7 @@ class GetAnthropicBetaHeadersConfig:
|
|||
def load_local_beta_headers_config() -> dict:
|
||||
"""Load the local backup beta headers config bundled with the package."""
|
||||
try:
|
||||
content = json.loads(
|
||||
content: Final = json.loads(
|
||||
files("litellm").joinpath("anthropic_beta_headers_config.json").read_text(encoding="utf-8")
|
||||
)
|
||||
return content
|
||||
|
|
@ -79,14 +80,14 @@ class GetAnthropicBetaHeadersConfig:
|
|||
return False
|
||||
|
||||
# Check for at least one provider key
|
||||
provider_keys = [
|
||||
provider_keys: Final = [
|
||||
"anthropic",
|
||||
"azure_ai",
|
||||
"bedrock",
|
||||
"bedrock_converse",
|
||||
"vertex_ai",
|
||||
]
|
||||
has_provider = any(key in fetched_config for key in provider_keys)
|
||||
has_provider: Final = any(key in fetched_config for key in provider_keys)
|
||||
|
||||
if not has_provider:
|
||||
verbose_logger.warning(
|
||||
|
|
@ -113,7 +114,7 @@ class GetAnthropicBetaHeadersConfig:
|
|||
Returns the parsed JSON dict. Raises on network/parse errors
|
||||
(caller is expected to handle).
|
||||
"""
|
||||
response = httpx.get(url, timeout=timeout)
|
||||
response: Final = httpx.get(url, timeout=timeout)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
|
|
@ -138,7 +139,7 @@ def get_beta_headers_config(url: str) -> dict:
|
|||
return GetAnthropicBetaHeadersConfig.load_local_beta_headers_config()
|
||||
|
||||
try:
|
||||
content = GetAnthropicBetaHeadersConfig.fetch_remote_beta_headers_config(url)
|
||||
content: Final = GetAnthropicBetaHeadersConfig.fetch_remote_beta_headers_config(url)
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Failed to fetch remote beta headers config from %s: %s. Falling back to local backup.",
|
||||
|
|
@ -206,8 +207,8 @@ def get_provider_name(provider: str) -> str:
|
|||
Returns:
|
||||
Canonical provider name
|
||||
"""
|
||||
config = _load_beta_headers_config()
|
||||
aliases = config.get("provider_aliases", {})
|
||||
config: Final = _load_beta_headers_config()
|
||||
aliases: Final = config.get("provider_aliases", {})
|
||||
return aliases.get(provider, provider)
|
||||
|
||||
|
||||
|
|
@ -233,13 +234,13 @@ def filter_and_transform_beta_headers(
|
|||
if not beta_headers:
|
||||
return []
|
||||
|
||||
config = _load_beta_headers_config()
|
||||
config: Final = _load_beta_headers_config()
|
||||
provider = get_provider_name(provider)
|
||||
|
||||
# Get the header mapping for this provider
|
||||
provider_mapping = config.get(provider, {})
|
||||
provider_mapping: Final = config.get(provider, {})
|
||||
|
||||
filtered_headers: set[str] = set()
|
||||
filtered_headers: Final[set[str]] = set()
|
||||
|
||||
for header in beta_headers:
|
||||
header = header.strip()
|
||||
|
|
@ -279,9 +280,9 @@ def is_beta_header_supported(
|
|||
Returns:
|
||||
True if the header is in the mapping with a non-null value, False otherwise
|
||||
"""
|
||||
config = _load_beta_headers_config()
|
||||
config: Final = _load_beta_headers_config()
|
||||
provider = get_provider_name(provider)
|
||||
provider_mapping = config.get(provider, {})
|
||||
provider_mapping: Final = config.get(provider, {})
|
||||
|
||||
# Header is supported if it's in the mapping and has a non-null value
|
||||
return beta_header in provider_mapping and provider_mapping[beta_header] is not None
|
||||
|
|
@ -303,11 +304,11 @@ def get_provider_beta_header(
|
|||
Returns:
|
||||
The provider-specific header name if supported, or None if unsupported/unknown
|
||||
"""
|
||||
config = _load_beta_headers_config()
|
||||
config: Final = _load_beta_headers_config()
|
||||
provider = get_provider_name(provider)
|
||||
|
||||
# Get the header mapping for this provider
|
||||
provider_mapping = config.get(provider, {})
|
||||
provider_mapping: Final = config.get(provider, {})
|
||||
|
||||
# Check if header is in the mapping
|
||||
if anthropic_beta_header not in provider_mapping:
|
||||
|
|
@ -332,15 +333,15 @@ def update_headers_with_filtered_beta(
|
|||
Returns:
|
||||
Updated headers dict
|
||||
"""
|
||||
existing_beta = headers.get("anthropic-beta")
|
||||
existing_beta: Final = headers.get("anthropic-beta")
|
||||
if not existing_beta:
|
||||
return headers
|
||||
|
||||
# Parse existing beta headers
|
||||
beta_values = [b.strip() for b in existing_beta.split(",") if b.strip()]
|
||||
beta_values: Final = [b.strip() for b in existing_beta.split(",") if b.strip()]
|
||||
|
||||
# Filter and transform based on provider
|
||||
filtered_beta_values = filter_and_transform_beta_headers(
|
||||
filtered_beta_values: Final = filter_and_transform_beta_headers(
|
||||
beta_headers=beta_values,
|
||||
provider=provider,
|
||||
)
|
||||
|
|
@ -374,11 +375,11 @@ def update_request_with_filtered_beta(
|
|||
"""
|
||||
headers = update_headers_with_filtered_beta(headers=headers, provider=provider)
|
||||
|
||||
existing_body_betas = request_data.get("anthropic_beta")
|
||||
existing_body_betas: Final = request_data.get("anthropic_beta")
|
||||
if not existing_body_betas:
|
||||
return headers, request_data
|
||||
|
||||
filtered_body_betas = filter_and_transform_beta_headers(
|
||||
filtered_body_betas: Final = filter_and_transform_beta_headers(
|
||||
beta_headers=existing_body_betas,
|
||||
provider=provider,
|
||||
)
|
||||
|
|
@ -401,9 +402,9 @@ def get_unsupported_headers(provider: str) -> list[str]:
|
|||
Returns:
|
||||
List of unsupported Anthropic beta header names
|
||||
"""
|
||||
config = _load_beta_headers_config()
|
||||
config: Final = _load_beta_headers_config()
|
||||
provider = get_provider_name(provider)
|
||||
provider_mapping = config.get(provider, {})
|
||||
provider_mapping: Final = config.get(provider, {})
|
||||
|
||||
# Return headers with null values
|
||||
return [header for header, value in provider_mapping.items() if value is None]
|
||||
|
|
|
|||
|
|
@ -4,13 +4,15 @@ Utilities for mapping exceptions to Anthropic error format.
|
|||
Similar to litellm/litellm_core_utils/exception_mapping_utils.py but for Anthropic response format.
|
||||
"""
|
||||
|
||||
from typing import Final
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
|
||||
from .exceptions import AnthropicErrorResponse, AnthropicErrorType
|
||||
|
||||
# HTTP status code -> Anthropic error type
|
||||
# Source: https://docs.anthropic.com/en/api/errors
|
||||
ANTHROPIC_ERROR_TYPE_MAP: dict[int, AnthropicErrorType] = {
|
||||
ANTHROPIC_ERROR_TYPE_MAP: Final[dict[int, AnthropicErrorType]] = {
|
||||
400: "invalid_request_error",
|
||||
401: "authentication_error",
|
||||
403: "permission_error",
|
||||
|
|
@ -50,9 +52,9 @@ class AnthropicExceptionMapping:
|
|||
"request_id": "req_..."
|
||||
}
|
||||
"""
|
||||
error_type = AnthropicExceptionMapping.get_error_type(status_code)
|
||||
error_type: Final = AnthropicExceptionMapping.get_error_type(status_code)
|
||||
|
||||
response: AnthropicErrorResponse = {
|
||||
response: Final[AnthropicErrorResponse] = {
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": error_type,
|
||||
|
|
@ -76,7 +78,7 @@ class AnthropicExceptionMapping:
|
|||
- Generic: {"message": "..."}
|
||||
- Plain strings
|
||||
"""
|
||||
parsed = safe_json_loads(raw_message)
|
||||
parsed: Final = safe_json_loads(raw_message)
|
||||
if isinstance(parsed, dict):
|
||||
# Bedrock format
|
||||
if "detail" in parsed and isinstance(parsed["detail"], dict):
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ import contextvars
|
|||
import os
|
||||
from collections.abc import Coroutine, Iterable
|
||||
from functools import partial
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
import httpx
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
|
|
@ -29,8 +29,8 @@ from ..types.router import *
|
|||
from .utils import get_optional_params_add_message
|
||||
|
||||
####### ENVIRONMENT VARIABLES ###################
|
||||
openai_assistants_api = OpenAIAssistantsAPI()
|
||||
azure_assistants_api = AzureAssistantsAPI()
|
||||
openai_assistants_api: Final = OpenAIAssistantsAPI()
|
||||
azure_assistants_api: Final = AzureAssistantsAPI()
|
||||
|
||||
### ASSISTANTS ###
|
||||
|
||||
|
|
@ -40,23 +40,23 @@ async def aget_assistants(
|
|||
client: AsyncOpenAI | None = None,
|
||||
**kwargs,
|
||||
) -> AsyncCursorPage[Assistant]:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
### PASS ARGS TO GET ASSISTANTS ###
|
||||
kwargs["aget_assistants"] = True
|
||||
try:
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(get_assistants, custom_llm_provider, client, **kwargs)
|
||||
func: Final = partial(get_assistants, custom_llm_provider, client, **kwargs)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
|
||||
# Await normally
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -80,11 +80,11 @@ def get_assistants(
|
|||
api_version: str | None = None,
|
||||
**kwargs,
|
||||
) -> SyncCursorPage[Assistant]:
|
||||
aget_assistants: bool | None = kwargs.pop("aget_assistants", None)
|
||||
aget_assistants: Final[bool | None] = kwargs.pop("aget_assistants", None)
|
||||
if aget_assistants is not None and not isinstance(aget_assistants, bool):
|
||||
raise Exception("Invalid value passed in for aget_assistants. Only bool or None allowed")
|
||||
optional_params = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs)
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
optional_params: Final = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
|
|
@ -95,7 +95,7 @@ def get_assistants(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
|
|
@ -111,7 +111,7 @@ def get_assistants(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -147,7 +147,7 @@ def get_assistants(
|
|||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
|
|
@ -197,25 +197,25 @@ async def acreate_assistants(
|
|||
client: AsyncOpenAI | None = None,
|
||||
**kwargs,
|
||||
) -> Assistant:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
### PASS ARGS TO GET ASSISTANTS ###
|
||||
kwargs["async_create_assistants"] = True
|
||||
model = kwargs.pop("model", None)
|
||||
model: Final = kwargs.pop("model", None)
|
||||
try:
|
||||
kwargs["client"] = client
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(create_assistants, custom_llm_provider, model, **kwargs)
|
||||
func: Final = partial(create_assistants, custom_llm_provider, model, **kwargs)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
|
||||
# Await normally
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -249,11 +249,11 @@ def create_assistants(
|
|||
api_version: str | None = None,
|
||||
**kwargs,
|
||||
) -> Assistant | Coroutine[Any, Any, Assistant]:
|
||||
async_create_assistants: bool | None = kwargs.pop("async_create_assistants", None)
|
||||
async_create_assistants: Final[bool | None] = kwargs.pop("async_create_assistants", None)
|
||||
if async_create_assistants is not None and not isinstance(async_create_assistants, bool):
|
||||
raise ValueError("Invalid value passed in for async_create_assistants. Only bool or None allowed")
|
||||
optional_params = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs)
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
optional_params: Final = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
|
|
@ -264,7 +264,7 @@ def create_assistants(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
|
|
@ -296,7 +296,7 @@ def create_assistants(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -333,7 +333,7 @@ def create_assistants(
|
|||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
|
|
@ -380,24 +380,24 @@ async def adelete_assistant(
|
|||
client: AsyncOpenAI | None = None,
|
||||
**kwargs,
|
||||
) -> AssistantDeleted:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
### PASS ARGS TO GET ASSISTANTS ###
|
||||
kwargs["async_delete_assistants"] = True
|
||||
try:
|
||||
kwargs["client"] = client
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(delete_assistant, custom_llm_provider, **kwargs)
|
||||
func: Final = partial(delete_assistant, custom_llm_provider, **kwargs)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
|
||||
# Await normally
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -422,11 +422,11 @@ def delete_assistant(
|
|||
api_version: str | None = None,
|
||||
**kwargs,
|
||||
) -> AssistantDeleted | Coroutine[Any, Any, AssistantDeleted]:
|
||||
optional_params = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs)
|
||||
optional_params: Final = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs)
|
||||
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
|
||||
async_delete_assistants: bool | None = kwargs.pop("async_delete_assistants", None)
|
||||
async_delete_assistants: Final[bool | None] = kwargs.pop("async_delete_assistants", None)
|
||||
if async_delete_assistants is not None and not isinstance(async_delete_assistants, bool):
|
||||
raise ValueError("Invalid value passed in for async_delete_assistants. Only bool or None allowed")
|
||||
|
||||
|
|
@ -439,7 +439,7 @@ def delete_assistant(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
|
|
@ -455,7 +455,7 @@ def delete_assistant(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) or None
|
||||
)
|
||||
# set API KEY
|
||||
|
|
@ -484,7 +484,7 @@ def delete_assistant(
|
|||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
|
|
@ -530,23 +530,23 @@ def delete_assistant(
|
|||
|
||||
|
||||
async def acreate_thread(custom_llm_provider: Literal["openai", "azure"], **kwargs) -> Thread:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
### PASS ARGS TO GET ASSISTANTS ###
|
||||
kwargs["acreate_thread"] = True
|
||||
try:
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(create_thread, custom_llm_provider, **kwargs)
|
||||
func: Final = partial(create_thread, custom_llm_provider, **kwargs)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
|
||||
# Await normally
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -592,9 +592,9 @@ def create_thread(
|
|||
)
|
||||
```
|
||||
"""
|
||||
acreate_thread = kwargs.get("acreate_thread", None)
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
acreate_thread: Final = kwargs.get("acreate_thread", None)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
|
|
@ -605,7 +605,7 @@ def create_thread(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
|
|
@ -624,7 +624,7 @@ def create_thread(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -661,7 +661,7 @@ def create_thread(
|
|||
|
||||
api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
|
|
@ -704,23 +704,23 @@ async def aget_thread(
|
|||
client: AsyncOpenAI | None = None,
|
||||
**kwargs,
|
||||
) -> Thread:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
### PASS ARGS TO GET ASSISTANTS ###
|
||||
kwargs["aget_thread"] = True
|
||||
try:
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(get_thread, custom_llm_provider, thread_id, client, **kwargs)
|
||||
func: Final = partial(get_thread, custom_llm_provider, thread_id, client, **kwargs)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
|
||||
# Await normally
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -743,9 +743,9 @@ def get_thread(
|
|||
**kwargs,
|
||||
) -> Thread:
|
||||
"""Get the thread object, given a thread_id"""
|
||||
aget_thread = kwargs.pop("aget_thread", None)
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
aget_thread: Final = kwargs.pop("aget_thread", None)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
# set timeout for 10 minutes by default
|
||||
|
|
@ -755,7 +755,7 @@ def get_thread(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
|
|
@ -772,7 +772,7 @@ def get_thread(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -809,7 +809,7 @@ def get_thread(
|
|||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
|
|
@ -858,12 +858,12 @@ async def a_add_message(
|
|||
client=None,
|
||||
**kwargs,
|
||||
) -> OpenAIMessage:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
### PASS ARGS TO GET ASSISTANTS ###
|
||||
kwargs["a_add_message"] = True
|
||||
try:
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
add_message,
|
||||
custom_llm_provider,
|
||||
thread_id,
|
||||
|
|
@ -876,15 +876,15 @@ async def a_add_message(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
|
||||
# Await normally
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -912,12 +912,12 @@ def add_message(
|
|||
**kwargs,
|
||||
) -> OpenAIMessage:
|
||||
### COMMON OBJECTS ###
|
||||
a_add_message = kwargs.pop("a_add_message", None)
|
||||
_message_data = MessageData(role=role, content=content, attachments=attachments, metadata=metadata)
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
a_add_message: Final = kwargs.pop("a_add_message", None)
|
||||
_message_data: Final = MessageData(role=role, content=content, attachments=attachments, metadata=metadata)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
message_data = get_optional_params_add_message(
|
||||
message_data: Final = get_optional_params_add_message(
|
||||
role=_message_data["role"],
|
||||
content=_message_data["content"],
|
||||
attachments=_message_data["attachments"],
|
||||
|
|
@ -934,7 +934,7 @@ def add_message(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
|
|
@ -951,7 +951,7 @@ def add_message(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -988,7 +988,7 @@ def add_message(
|
|||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
|
|
@ -1029,12 +1029,12 @@ async def aget_messages(
|
|||
client: AsyncOpenAI | None = None,
|
||||
**kwargs,
|
||||
) -> AsyncCursorPage[OpenAIMessage]:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
### PASS ARGS TO GET ASSISTANTS ###
|
||||
kwargs["aget_messages"] = True
|
||||
try:
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
get_messages,
|
||||
custom_llm_provider,
|
||||
thread_id,
|
||||
|
|
@ -1043,15 +1043,15 @@ async def aget_messages(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
|
||||
# Await normally
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -1074,9 +1074,9 @@ def get_messages(
|
|||
client: Any | None = None,
|
||||
**kwargs,
|
||||
) -> SyncCursorPage[OpenAIMessage]:
|
||||
aget_messages = kwargs.pop("aget_messages", None)
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
aget_messages: Final = kwargs.pop("aget_messages", None)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
|
|
@ -1087,7 +1087,7 @@ def get_messages(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
|
|
@ -1105,7 +1105,7 @@ def get_messages(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -1141,7 +1141,7 @@ def get_messages(
|
|||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
|
|
@ -1198,12 +1198,12 @@ async def arun_thread(
|
|||
client: Any | None = None,
|
||||
**kwargs,
|
||||
) -> Run:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
### PASS ARGS TO GET ASSISTANTS ###
|
||||
kwargs["arun_thread"] = True
|
||||
try:
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
run_thread,
|
||||
custom_llm_provider,
|
||||
thread_id,
|
||||
|
|
@ -1219,15 +1219,15 @@ async def arun_thread(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
|
||||
# Await normally
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -1267,9 +1267,9 @@ def run_thread(
|
|||
**kwargs,
|
||||
) -> Run:
|
||||
"""Run a given thread + assistant."""
|
||||
arun_thread = kwargs.pop("arun_thread", None)
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
arun_thread: Final = kwargs.pop("arun_thread", None)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
|
|
@ -1280,7 +1280,7 @@ def run_thread(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
|
|
@ -1296,7 +1296,7 @@ def run_thread(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -1341,7 +1341,7 @@ def run_thread(
|
|||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
import litellm
|
||||
|
||||
from ..exceptions import UnsupportedParamsError
|
||||
|
|
@ -17,13 +19,13 @@ def get_optional_params_add_message(
|
|||
|
||||
Reference - https://learn.microsoft.com/en-us/azure/ai-services/openai/assistants-reference-messages?tabs=python#create-message
|
||||
"""
|
||||
passed_params = locals()
|
||||
passed_params: Final = locals()
|
||||
custom_llm_provider = passed_params.pop("custom_llm_provider")
|
||||
special_params = passed_params.pop("kwargs")
|
||||
special_params: Final = passed_params.pop("kwargs")
|
||||
for k, v in special_params.items():
|
||||
passed_params[k] = v
|
||||
|
||||
default_params = {
|
||||
default_params: Final = {
|
||||
"role": None,
|
||||
"content": None,
|
||||
"attachments": None,
|
||||
|
|
@ -36,7 +38,7 @@ def get_optional_params_add_message(
|
|||
## raise exception if non-default value passed for non-openai/azure embedding calls
|
||||
def _check_valid_arg(supported_params):
|
||||
if len(non_default_params.keys()) > 0:
|
||||
keys = list(non_default_params.keys())
|
||||
keys: Final = list(non_default_params.keys())
|
||||
for k in keys:
|
||||
if litellm.drop_params is True and k not in supported_params: # drop the unsupported non-default values
|
||||
non_default_params.pop(k, None)
|
||||
|
|
@ -50,7 +52,7 @@ def get_optional_params_add_message(
|
|||
if custom_llm_provider == "openai":
|
||||
optional_params = non_default_params
|
||||
elif custom_llm_provider == "azure":
|
||||
supported_params = litellm.AzureOpenAIAssistantsAPIConfig().get_supported_openai_create_message_params()
|
||||
supported_params: Final = litellm.AzureOpenAIAssistantsAPIConfig().get_supported_openai_create_message_params()
|
||||
_check_valid_arg(supported_params=supported_params)
|
||||
optional_params = litellm.AzureOpenAIAssistantsAPIConfig().map_openai_params_create_message_params(
|
||||
non_default_params=non_default_params, optional_params=optional_params
|
||||
|
|
@ -72,13 +74,13 @@ def get_optional_params_image_gen(
|
|||
**kwargs,
|
||||
):
|
||||
# retrieve all parameters passed to the function
|
||||
passed_params = locals()
|
||||
passed_params: Final = locals()
|
||||
custom_llm_provider = passed_params.pop("custom_llm_provider")
|
||||
special_params = passed_params.pop("kwargs")
|
||||
special_params: Final = passed_params.pop("kwargs")
|
||||
for k, v in special_params.items():
|
||||
passed_params[k] = v
|
||||
|
||||
default_params = {
|
||||
default_params: Final = {
|
||||
"n": None,
|
||||
"quality": None,
|
||||
"response_format": None,
|
||||
|
|
@ -93,7 +95,7 @@ def get_optional_params_image_gen(
|
|||
## raise exception if non-default value passed for non-openai/azure embedding calls
|
||||
def _check_valid_arg(supported_params):
|
||||
if len(non_default_params.keys()) > 0:
|
||||
keys = list(non_default_params.keys())
|
||||
keys: Final = list(non_default_params.keys())
|
||||
for k in keys:
|
||||
if litellm.drop_params is True and k not in supported_params: # drop the unsupported non-default values
|
||||
non_default_params.pop(k, None)
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose
|
||||
|
|
@ -55,17 +56,17 @@ def batch_completion(
|
|||
Returns:
|
||||
list: A list of completion results.
|
||||
"""
|
||||
args = locals()
|
||||
args: Final = locals()
|
||||
|
||||
batch_messages = messages
|
||||
completions = []
|
||||
batch_messages: Final = messages
|
||||
completions: Final = []
|
||||
model = model
|
||||
custom_llm_provider = None
|
||||
if model.split("/", 1)[0] in litellm.provider_list:
|
||||
custom_llm_provider = model.split("/", 1)[0]
|
||||
model = model.split("/", 1)[1]
|
||||
if custom_llm_provider == "vllm":
|
||||
optional_params = get_optional_params(
|
||||
optional_params: Final = get_optional_params(
|
||||
functions=functions,
|
||||
function_call=function_call,
|
||||
temperature=temperature,
|
||||
|
|
@ -145,7 +146,7 @@ def batch_completion_models(*args, **kwargs):
|
|||
if "model" in kwargs:
|
||||
kwargs.pop("model")
|
||||
if "models" in kwargs:
|
||||
models = kwargs["models"]
|
||||
models: Final = kwargs["models"]
|
||||
kwargs.pop("models")
|
||||
futures = {}
|
||||
with ThreadPoolExecutor(max_workers=len(models)) as executor:
|
||||
|
|
@ -156,10 +157,10 @@ def batch_completion_models(*args, **kwargs):
|
|||
if future.result() is not None:
|
||||
return future.result()
|
||||
elif "deployments" in kwargs:
|
||||
deployments = kwargs["deployments"]
|
||||
deployments: Final = kwargs["deployments"]
|
||||
kwargs.pop("deployments")
|
||||
kwargs.pop("model_list")
|
||||
nested_kwargs = kwargs.pop("kwargs", {})
|
||||
nested_kwargs: Final = kwargs.pop("kwargs", {})
|
||||
futures = {}
|
||||
with ThreadPoolExecutor(max_workers=len(deployments)) as executor:
|
||||
for deployment in deployments:
|
||||
|
|
@ -238,10 +239,10 @@ def batch_completion_models_all_responses(*args, **kwargs):
|
|||
if len(models) == 0:
|
||||
return []
|
||||
|
||||
responses = []
|
||||
responses: Final = []
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=len(models)) as executor:
|
||||
futures = [executor.submit(litellm.completion, *args, model=model, **kwargs) for model in models]
|
||||
futures: Final = [executor.submit(litellm.completion, *args, model=model, **kwargs) for model in models]
|
||||
|
||||
for future in futures:
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import json
|
||||
from collections.abc import Iterable, Iterator
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -141,9 +141,9 @@ def _aggregate_batch_cost_usage_models(
|
|||
) -> tuple[float, Usage, list[str]]:
|
||||
"""Aggregate cost, usage, and models from batch output entries in a single
|
||||
pass, holding one small stats record per line instead of the parsed file."""
|
||||
line_stats = tuple(_iter_successful_output_line_stats(entries, custom_llm_provider, model_name, model_info))
|
||||
line_stats: Final = tuple(_iter_successful_output_line_stats(entries, custom_llm_provider, model_name, model_info))
|
||||
|
||||
cache_token_params = {
|
||||
cache_token_params: Final = {
|
||||
key: tokens
|
||||
for key, tokens in (
|
||||
("cache_read_input_tokens", sum(stats.cache_read_tokens for stats in line_stats)),
|
||||
|
|
@ -151,14 +151,14 @@ def _aggregate_batch_cost_usage_models(
|
|||
)
|
||||
if tokens > 0
|
||||
}
|
||||
batch_usage = Usage(
|
||||
batch_usage: Final = Usage(
|
||||
total_tokens=sum(stats.total_tokens for stats in line_stats),
|
||||
prompt_tokens=sum(stats.prompt_tokens for stats in line_stats),
|
||||
completion_tokens=sum(stats.completion_tokens for stats in line_stats),
|
||||
**cache_token_params,
|
||||
)
|
||||
batch_models = [model_name] if model_name else [stats.model for stats in line_stats if stats.model]
|
||||
total_cost = sum((stats.cost for stats in line_stats), 0.0)
|
||||
batch_models: Final = [model_name] if model_name else [stats.model for stats in line_stats if stats.model]
|
||||
total_cost: Final = sum((stats.cost for stats in line_stats), 0.0)
|
||||
verbose_logger.debug("batch output aggregate: cost=%s usage=%s models=%s", total_cost, batch_usage, batch_models)
|
||||
return total_cost, batch_usage, batch_models
|
||||
|
||||
|
|
@ -184,7 +184,7 @@ def calculate_vertex_ai_batch_cost_and_usage(
|
|||
total_tokens = 0
|
||||
prompt_tokens = 0
|
||||
completion_tokens = 0
|
||||
actual_model_name = model_name or "gemini-2.0-flash-001"
|
||||
actual_model_name: Final = model_name or "gemini-2.0-flash-001"
|
||||
|
||||
for response in vertex_ai_batch_responses:
|
||||
response_body = response.get("response")
|
||||
|
|
@ -254,7 +254,7 @@ async def _fetch_batch_output_file_content(
|
|||
raise ValueError("Output file id is None cannot retrieve file content")
|
||||
|
||||
file_id = batch.output_file_id
|
||||
is_base64_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
|
||||
is_base64_unified_file_id: Final = _is_base64_encoded_unified_file_id(file_id)
|
||||
if is_base64_unified_file_id:
|
||||
try:
|
||||
file_id = is_base64_unified_file_id.split("llm_output_file_id,")[1].split(";")[0]
|
||||
|
|
@ -265,16 +265,16 @@ async def _fetch_batch_output_file_content(
|
|||
)
|
||||
|
||||
# Build kwargs for afile_content with credentials from litellm_params
|
||||
file_content_kwargs = {
|
||||
file_content_kwargs: Final = {
|
||||
"file_id": file_id,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
}
|
||||
|
||||
# Extract and add credentials for file access
|
||||
credentials = _extract_file_access_credentials(litellm_params)
|
||||
credentials: Final = _extract_file_access_credentials(litellm_params)
|
||||
file_content_kwargs.update(credentials)
|
||||
|
||||
_file_content = await afile_content(**file_content_kwargs) # type: ignore[reportArgumentType]
|
||||
_file_content: Final = await afile_content(**file_content_kwargs) # type: ignore[reportArgumentType]
|
||||
return _file_content.content
|
||||
|
||||
|
||||
|
|
@ -291,11 +291,11 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict:
|
|||
Returns:
|
||||
Dictionary containing only the credentials needed for file access
|
||||
"""
|
||||
credentials = {}
|
||||
credentials: Final = {}
|
||||
|
||||
if litellm_params:
|
||||
# List of credential keys that should be passed to file operations
|
||||
credential_keys = [
|
||||
credential_keys: Final = [
|
||||
"api_key",
|
||||
"api_base",
|
||||
"api_version",
|
||||
|
|
@ -355,7 +355,7 @@ def _iter_batch_input_entries(file_content: bytes) -> Iterator[dict]:
|
|||
|
||||
# A batch request's input tokens scale roughly with its serialized size, so this
|
||||
# is a conservative per-row fallback when the token counter cannot measure a row.
|
||||
_BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN = 4
|
||||
_BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN: Final = 4
|
||||
|
||||
|
||||
def _estimate_batch_entry_tokens(raw_line: bytes) -> int:
|
||||
|
|
@ -370,18 +370,18 @@ def _count_entry_tokens(
|
|||
model_name: str | None = None,
|
||||
) -> int:
|
||||
"""Token-count a single batch input entry's body (chat / text / embedding)."""
|
||||
body = entry.get("body", {}) or {}
|
||||
model = body.get("model", model_name or "")
|
||||
body: Final = entry.get("body", {}) or {}
|
||||
model: Final = body.get("model", model_name or "")
|
||||
|
||||
messages = body.get("messages")
|
||||
messages: Final = body.get("messages")
|
||||
if messages:
|
||||
return token_counter(model=model, messages=messages)
|
||||
|
||||
prompt = body.get("prompt")
|
||||
prompt: Final = body.get("prompt")
|
||||
if prompt:
|
||||
return _count_prompt_or_input_tokens(model=model, value=prompt)
|
||||
|
||||
input_data = body.get("input")
|
||||
input_data: Final = body.get("input")
|
||||
if input_data:
|
||||
return _count_prompt_or_input_tokens(model=model, value=input_data)
|
||||
|
||||
|
|
@ -432,8 +432,8 @@ def _get_batch_job_usage_from_response_body(response_body: dict, custom_llm_prov
|
|||
usage_object=response_body.get("usage", None) or {},
|
||||
reasoning_content=None,
|
||||
)
|
||||
_usage_dict = response_body.get("usage", None) or {}
|
||||
usage: Usage = Usage(**_usage_dict)
|
||||
_usage_dict: Final = response_body.get("usage", None) or {}
|
||||
usage: Final[Usage] = Usage(**_usage_dict)
|
||||
return usage
|
||||
|
||||
|
||||
|
|
@ -455,8 +455,8 @@ def _get_response_from_batch_job_output_file(batch_job_output_file: dict, custom
|
|||
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: dict = batch_job_output_file.get("response", None) or {}
|
||||
_response_body = _response.get("body", None) or {}
|
||||
_response: Final[dict] = batch_job_output_file.get("response", None) or {}
|
||||
_response_body: Final = _response.get("body", None) or {}
|
||||
return _response_body
|
||||
|
||||
|
||||
|
|
@ -472,5 +472,5 @@ def _batch_response_was_successful(batch_job_output_file: dict, custom_llm_provi
|
|||
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: dict = batch_job_output_file.get("response", None) or {}
|
||||
_response: Final[dict] = batch_job_output_file.get("response", None) or {}
|
||||
return _response.get("status_code", None) == 200
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ import contextvars
|
|||
import os
|
||||
from collections.abc import Coroutine
|
||||
from functools import partial
|
||||
from typing import Any, Literal, cast
|
||||
from typing import Any, Final, Literal, cast
|
||||
|
||||
import httpx
|
||||
from openai.types.batch import BatchRequestCounts
|
||||
|
|
@ -54,10 +54,10 @@ from litellm.utils import (
|
|||
)
|
||||
|
||||
####### ENVIRONMENT VARIABLES ###################
|
||||
openai_batches_instance = OpenAIBatchesAPI()
|
||||
azure_batches_instance = AzureBatchesAPI()
|
||||
vertex_ai_batches_instance = VertexAIBatchPrediction(gcs_bucket_name="")
|
||||
anthropic_batches_instance = AnthropicBatchesHandler()
|
||||
openai_batches_instance: Final = OpenAIBatchesAPI()
|
||||
azure_batches_instance: Final = AzureBatchesAPI()
|
||||
vertex_ai_batches_instance: Final = VertexAIBatchPrediction(gcs_bucket_name="")
|
||||
anthropic_batches_instance: Final = AnthropicBatchesHandler()
|
||||
base_llm_http_handler = BaseLLMHTTPHandler()
|
||||
#################################################
|
||||
|
||||
|
|
@ -80,13 +80,13 @@ def _resolve_timeout(
|
|||
Returns:
|
||||
Resolved timeout as float
|
||||
"""
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", default_timeout) or default_timeout
|
||||
timeout: Final = optional_params.timeout or kwargs.get("request_timeout", default_timeout) or default_timeout
|
||||
|
||||
# Handle httpx.Timeout objects
|
||||
if isinstance(timeout, httpx.Timeout):
|
||||
if supports_httpx_timeout(custom_llm_provider) is False:
|
||||
# Extract read timeout for providers that don't support httpx.Timeout
|
||||
read_timeout = timeout.read or default_timeout
|
||||
read_timeout: Final = timeout.read or default_timeout
|
||||
return float(read_timeout)
|
||||
else:
|
||||
# For providers that support httpx.Timeout, we still need to return a float
|
||||
|
|
@ -119,11 +119,11 @@ async def acreate_batch(
|
|||
LiteLLM Equivalent of POST: https://api.openai.com/v1/batches
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["acreate_batch"] = True
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
create_batch,
|
||||
completion_window,
|
||||
endpoint,
|
||||
|
|
@ -137,9 +137,9 @@ async def acreate_batch(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -169,10 +169,10 @@ def create_batch(
|
|||
LiteLLM Equivalent of POST: https://api.openai.com/v1/batches
|
||||
"""
|
||||
try:
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_call_id = kwargs.get("litellm_call_id", None)
|
||||
proxy_server_request = kwargs.get("proxy_server_request", None)
|
||||
model_info = kwargs.get("model_info", None)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
litellm_call_id: Final = kwargs.get("litellm_call_id", None)
|
||||
proxy_server_request: Final = kwargs.get("proxy_server_request", None)
|
||||
model_info: Final = kwargs.get("model_info", None)
|
||||
model: str | None = kwargs.get("model", None)
|
||||
try:
|
||||
if model is not None:
|
||||
|
|
@ -185,11 +185,11 @@ def create_batch(
|
|||
"litellm.batches.main.py::create_batch() - Error inferring custom_llm_provider - %s", e
|
||||
)
|
||||
|
||||
_is_async = kwargs.pop("acreate_batch", False) is True
|
||||
litellm_params = dict(GenericLiteLLMParams(**kwargs))
|
||||
litellm_logging_obj: LiteLLMLoggingObj = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj", None))
|
||||
_is_async: Final = kwargs.pop("acreate_batch", False) is True
|
||||
litellm_params: Final = dict(GenericLiteLLMParams(**kwargs))
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj", None))
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = _resolve_timeout(optional_params, kwargs, custom_llm_provider)
|
||||
timeout: Final = _resolve_timeout(optional_params, kwargs, custom_llm_provider)
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
|
|
@ -206,7 +206,7 @@ def create_batch(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
_create_batch_request = CreateBatchRequest(
|
||||
_create_batch_request: Final = CreateBatchRequest(
|
||||
completion_window=completion_window,
|
||||
endpoint=endpoint,
|
||||
input_file_id=input_file_id,
|
||||
|
|
@ -248,7 +248,7 @@ def create_batch(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -301,13 +301,13 @@ def create_batch(
|
|||
)
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or ""
|
||||
vertex_ai_project = (
|
||||
vertex_ai_project: Final = (
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
vertex_ai_location: Final = (
|
||||
optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION")
|
||||
)
|
||||
vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
|
||||
response = vertex_ai_batches_instance.create_batch(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -350,11 +350,11 @@ async def aretrieve_batch(
|
|||
LiteLLM Equivalent of GET https://api.openai.com/v1/batches/{batch_id}
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["aretrieve_batch"] = True
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
retrieve_batch,
|
||||
batch_id,
|
||||
custom_llm_provider,
|
||||
|
|
@ -364,9 +364,9 @@ async def aretrieve_batch(
|
|||
**kwargs,
|
||||
)
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -397,7 +397,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -422,7 +422,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
|
||||
api_version: Final = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -432,7 +432,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
or get_secret_str("AZURE_API_KEY")
|
||||
)
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
if extra_body is not None:
|
||||
extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
|
|
@ -450,13 +450,13 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
)
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or ""
|
||||
vertex_ai_project = (
|
||||
vertex_ai_project: Final = (
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
vertex_ai_location: Final = (
|
||||
optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION")
|
||||
)
|
||||
vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
|
||||
response = vertex_ai_batches_instance.retrieve_batch(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -519,11 +519,11 @@ def retrieve_batch(
|
|||
LiteLLM Equivalent of GET https://api.openai.com/v1/batches/{batch_id}
|
||||
"""
|
||||
try:
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_logging_obj: LiteLLMLoggingObj | None = kwargs.get("litellm_logging_obj", None)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj", None)
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
litellm_params = get_litellm_params(
|
||||
litellm_params: Final = get_litellm_params(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
)
|
||||
|
|
@ -542,21 +542,21 @@ def retrieve_batch(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
_retrieve_batch_request = RetrieveBatchRequest(
|
||||
_retrieve_batch_request: Final = RetrieveBatchRequest(
|
||||
batch_id=batch_id,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
_is_async = kwargs.pop("aretrieve_batch", False) is True
|
||||
client = kwargs.get("client", None)
|
||||
_is_async: Final = kwargs.pop("aretrieve_batch", False) is True
|
||||
client: Final = kwargs.get("client", None)
|
||||
|
||||
# Bedrock has two distinct ARN families that need different APIs:
|
||||
# * async-invoke ARNs (Twelve Labs Marengo embeddings) -> bedrock-runtime data plane
|
||||
|
|
@ -568,7 +568,7 @@ def retrieve_batch(
|
|||
if batch_id.startswith("arn:aws") and ":bedrock:" in batch_id:
|
||||
if ":async-invoke/" in batch_id:
|
||||
# Remove aws_region_name from kwargs to avoid duplicate parameter
|
||||
async_kwargs = kwargs.copy()
|
||||
async_kwargs: Final = kwargs.copy()
|
||||
async_kwargs.pop("aws_region_name", None)
|
||||
|
||||
return BedrockBatchesHandler._handle_async_invoke_status(
|
||||
|
|
@ -578,7 +578,7 @@ def retrieve_batch(
|
|||
**async_kwargs,
|
||||
)
|
||||
if ":model-invocation-job/" in batch_id:
|
||||
mij_kwargs = kwargs.copy()
|
||||
mij_kwargs: Final = kwargs.copy()
|
||||
mij_kwargs.pop("aws_region_name", None)
|
||||
|
||||
return BedrockBatchesHandler._handle_model_invocation_job_status(
|
||||
|
|
@ -589,7 +589,7 @@ def retrieve_batch(
|
|||
)
|
||||
|
||||
# Try to use provider config first (for providers like bedrock)
|
||||
model: str | None = kwargs.get("model", None)
|
||||
model: Final[str | None] = kwargs.get("model", None)
|
||||
if model is not None:
|
||||
provider_config = ProviderConfigManager.get_provider_batches_config(
|
||||
model=model,
|
||||
|
|
@ -599,7 +599,7 @@ def retrieve_batch(
|
|||
provider_config = None
|
||||
|
||||
if provider_config is not None:
|
||||
response = base_llm_http_handler.retrieve_batch(
|
||||
response: Final = base_llm_http_handler.retrieve_batch(
|
||||
batch_id=batch_id,
|
||||
provider_config=provider_config,
|
||||
litellm_params=litellm_params,
|
||||
|
|
@ -656,11 +656,11 @@ async def alist_batches(
|
|||
"""
|
||||
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["alist_batches"] = True
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
list_batches,
|
||||
after,
|
||||
limit,
|
||||
|
|
@ -671,9 +671,9 @@ async def alist_batches(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -700,8 +700,8 @@ def list_batches(
|
|||
"""
|
||||
try:
|
||||
# set API KEY
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params = get_litellm_params(
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = get_litellm_params(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
)
|
||||
|
|
@ -720,14 +720,14 @@ def list_batches(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
_is_async = kwargs.pop("alist_batches", False) is True
|
||||
_is_async: Final = kwargs.pop("alist_batches", False) is True
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
# for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there
|
||||
api_base = (
|
||||
|
|
@ -737,7 +737,7 @@ def list_batches(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -783,13 +783,13 @@ def list_batches(
|
|||
)
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or ""
|
||||
vertex_ai_project = (
|
||||
vertex_ai_project: Final = (
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
vertex_ai_location: Final = (
|
||||
optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION")
|
||||
)
|
||||
vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
|
||||
response = vertex_ai_batches_instance.list_batches(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -836,14 +836,14 @@ async def acancel_batch(
|
|||
LiteLLM Equivalent of POST https://api.openai.com/v1/batches/{batch_id}/cancel
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["acancel_batch"] = True
|
||||
# Preserve model parameter - only pop from kwargs if it exists there
|
||||
# (to avoid passing it twice), otherwise keep the function parameter value
|
||||
model = kwargs.pop("model", None) or model
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
cancel_batch,
|
||||
batch_id,
|
||||
model,
|
||||
|
|
@ -854,9 +854,9 @@ async def acancel_batch(
|
|||
**kwargs,
|
||||
)
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -892,8 +892,8 @@ def cancel_batch(
|
|||
verbose_logger.exception(
|
||||
"litellm.batches.main.py::cancel_batch() - Error inferring custom_llm_provider - %s", e
|
||||
)
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params = get_litellm_params(
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = get_litellm_params(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
)
|
||||
|
|
@ -906,20 +906,20 @@ def cancel_batch(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
_cancel_batch_request = CancelBatchRequest(
|
||||
_cancel_batch_request: Final = CancelBatchRequest(
|
||||
batch_id=batch_id,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
_is_async = kwargs.pop("acancel_batch", False) is True
|
||||
_is_async: Final = kwargs.pop("acancel_batch", False) is True
|
||||
api_base: str | None = None
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
api_base = (
|
||||
|
|
@ -929,7 +929,7 @@ def cancel_batch(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) or None
|
||||
)
|
||||
api_key = optional_params.api_key or litellm.api_key or litellm.openai_key or os.getenv("OPENAI_API_KEY")
|
||||
|
|
@ -973,13 +973,13 @@ def cancel_batch(
|
|||
)
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or None
|
||||
vertex_ai_project = (
|
||||
vertex_ai_project: Final = (
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
vertex_ai_location: Final = (
|
||||
optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION")
|
||||
)
|
||||
vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
|
||||
response = vertex_ai_batches_instance.cancel_batch(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -1025,10 +1025,10 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj
|
|||
|
||||
async def _async_get_status():
|
||||
# Create embedding handler instance
|
||||
embedding_handler = BedrockEmbedding()
|
||||
embedding_handler: Final = BedrockEmbedding()
|
||||
|
||||
# Get the status of the async invoke job
|
||||
status_response = await embedding_handler._get_async_invoke_status(
|
||||
status_response: Final = await embedding_handler._get_async_invoke_status(
|
||||
invocation_arn=batch_id,
|
||||
aws_region_name=aws_region_name,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -1040,16 +1040,16 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj
|
|||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
# Normalize status to lowercase (AWS returns 'Completed', 'Failed', etc.)
|
||||
aws_status_raw = status_response.get("status", "")
|
||||
aws_status_lower = aws_status_raw.lower()
|
||||
aws_status_raw: Final = status_response.get("status", "")
|
||||
aws_status_lower: Final = aws_status_raw.lower()
|
||||
# Map AWS status values to LiteLLM expected values
|
||||
status_mapping: dict[str, BatchJobStatus] = {
|
||||
status_mapping: Final[dict[str, BatchJobStatus]] = {
|
||||
"completed": "completed",
|
||||
"failed": "failed",
|
||||
"inprogress": "in_progress",
|
||||
"in_progress": "in_progress",
|
||||
}
|
||||
normalized_status: BatchJobStatus = status_mapping.get(
|
||||
normalized_status: Final[BatchJobStatus] = status_mapping.get(
|
||||
aws_status_lower, "failed"
|
||||
) # Default to "failed" if unknown status
|
||||
|
||||
|
|
@ -1073,7 +1073,7 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj
|
|||
_,
|
||||
_,
|
||||
) = BedrockBatchesConfig()._parse_timestamps_and_status(status_response, aws_status_raw)
|
||||
result = LiteLLMBatch(
|
||||
result: Final = LiteLLMBatch(
|
||||
id=status_response["invocationArn"],
|
||||
object="batch",
|
||||
status=normalized_status,
|
||||
|
|
@ -1105,7 +1105,7 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj
|
|||
import concurrent.futures
|
||||
|
||||
def run_in_thread():
|
||||
new_loop = asyncio.new_event_loop()
|
||||
new_loop: Final = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(new_loop)
|
||||
try:
|
||||
return new_loop.run_until_complete(_async_get_status())
|
||||
|
|
@ -1113,5 +1113,5 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj
|
|||
new_loop.close()
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor() as executor:
|
||||
future = executor.submit(run_in_thread)
|
||||
future: Final = executor.submit(run_in_thread)
|
||||
return future.result()
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ import json
|
|||
import os
|
||||
import threading
|
||||
import time
|
||||
from typing import Literal
|
||||
from typing import Final, Literal
|
||||
|
||||
import litellm
|
||||
from litellm.constants import (
|
||||
|
|
@ -60,8 +60,8 @@ class BudgetManager:
|
|||
self.print_verbose(f"user dict from local: {self.user_dict}")
|
||||
elif self.client_type == "hosted":
|
||||
# Load the user_dict from hosted db
|
||||
url = self.api_base + "/get_budget"
|
||||
data = {"project_name": self.project_name}
|
||||
url: Final = self.api_base + "/get_budget"
|
||||
data: Final = {"project_name": self.project_name}
|
||||
response = litellm.module_level_client.post(url, headers=self.headers, json=data)
|
||||
response = response.json()
|
||||
if response["status"] == "error":
|
||||
|
|
@ -100,11 +100,11 @@ class BudgetManager:
|
|||
return self.user_dict[user]
|
||||
|
||||
def projected_cost(self, model: str, messages: list, user: str):
|
||||
text = "".join(message["content"] for message in messages)
|
||||
prompt_tokens = litellm.token_counter(model=model, text=text)
|
||||
text: Final = "".join(message["content"] for message in messages)
|
||||
prompt_tokens: Final = litellm.token_counter(model=model, text=text)
|
||||
prompt_cost, _ = litellm.cost_per_token(model=model, prompt_tokens=prompt_tokens, completion_tokens=0)
|
||||
current_cost = self.user_dict[user].get("current_cost", 0)
|
||||
projected_cost = prompt_cost + current_cost
|
||||
current_cost: Final = self.user_dict[user].get("current_cost", 0)
|
||||
projected_cost: Final = prompt_cost + current_cost
|
||||
return projected_cost
|
||||
|
||||
def get_total_budget(self, user: str):
|
||||
|
|
@ -178,11 +178,11 @@ class BudgetManager:
|
|||
|
||||
def reset_on_duration(self, user: str):
|
||||
# Get current and creation time
|
||||
last_updated_at = self.user_dict[user]["last_updated_at"]
|
||||
current_time = time.time()
|
||||
last_updated_at: Final = self.user_dict[user]["last_updated_at"]
|
||||
current_time: Final = time.time()
|
||||
|
||||
# Convert duration from days to seconds
|
||||
duration_in_seconds = self.user_dict[user]["duration"] * HOURS_IN_A_DAY * 60 * 60
|
||||
duration_in_seconds: Final = self.user_dict[user]["duration"] * HOURS_IN_A_DAY * 60 * 60
|
||||
|
||||
# Check if duration has elapsed
|
||||
if current_time - last_updated_at >= duration_in_seconds:
|
||||
|
|
@ -197,7 +197,7 @@ class BudgetManager:
|
|||
self.reset_on_duration(user)
|
||||
|
||||
def _save_data_thread(self):
|
||||
thread = threading.Thread(target=self.save_data) # [Non-Blocking]: saves data without blocking execution
|
||||
thread: Final = threading.Thread(target=self.save_data) # [Non-Blocking]: saves data without blocking execution
|
||||
thread.start()
|
||||
|
||||
def save_data(self):
|
||||
|
|
@ -209,8 +209,8 @@ class BudgetManager:
|
|||
json.dump(self.user_dict, json_file, indent=4) # Indent for pretty formatting
|
||||
return {"status": "success"}
|
||||
elif self.client_type == "hosted":
|
||||
url = self.api_base + "/set_budget"
|
||||
data = {"project_name": self.project_name, "user_dict": self.user_dict}
|
||||
url: Final = self.api_base + "/set_budget"
|
||||
data: Final = {"project_name": self.project_name, "user_dict": self.user_dict}
|
||||
response = litellm.module_level_client.post(url, headers=self.headers, json=data)
|
||||
response = response.json()
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ This module is dependency-injected: callers pass the proxy ``llm_router`` and
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
|
@ -26,7 +26,7 @@ def resolve_embedding_router(
|
|||
"""Return ``llm_router`` iff it serves ``embedding_model`` as a deployment."""
|
||||
if llm_router is None:
|
||||
return None
|
||||
router_model_names: list[str] = (
|
||||
router_model_names: Final[list[str]] = (
|
||||
[m["model_name"] for m in llm_model_list if "model_name" in m] if llm_model_list is not None else []
|
||||
)
|
||||
if embedding_model in router_model_names:
|
||||
|
|
@ -38,6 +38,6 @@ def build_router_embedding_metadata(
|
|||
request_metadata: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
"""Forward the caller's full metadata, flagged as a semantic-cache embedding."""
|
||||
metadata: dict[str, Any] = dict(request_metadata or {})
|
||||
metadata: Final[dict[str, Any]] = dict(request_metadata or {})
|
||||
metadata["semantic-cache-embedding"] = True
|
||||
return metadata
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from collections.abc import Callable
|
||||
from functools import lru_cache
|
||||
from typing import TypeVar
|
||||
from typing import Final, TypeVar
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
|
@ -21,7 +21,7 @@ def lru_cache_wrapper(
|
|||
return ("error", e)
|
||||
|
||||
def wrapped(*args, **kwargs):
|
||||
result = wrapper(*args, **kwargs)
|
||||
result: Final = wrapper(*args, **kwargs)
|
||||
if result[0] == "error":
|
||||
raise result[1]
|
||||
return result[1]
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ Has 4 methods:
|
|||
import asyncio
|
||||
import json
|
||||
from contextlib import suppress
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
|
||||
|
|
@ -41,7 +42,7 @@ class AzureBlobCache(BaseCache):
|
|||
|
||||
def set_cache(self, key, value, **kwargs) -> None:
|
||||
print_verbose(f"LiteLLM SET Cache - Azure Blob. Key={key}. Value={value}")
|
||||
serialized_value = json.dumps(value)
|
||||
serialized_value: Final = json.dumps(value)
|
||||
try:
|
||||
self.container_client.upload_blob(key, serialized_value)
|
||||
except Exception as e:
|
||||
|
|
@ -50,7 +51,7 @@ class AzureBlobCache(BaseCache):
|
|||
|
||||
async def async_set_cache(self, key, value, **kwargs) -> None:
|
||||
print_verbose(f"LiteLLM SET Cache - Azure Blob. Key={key}. Value={value}")
|
||||
serialized_value = json.dumps(value)
|
||||
serialized_value: Final = json.dumps(value)
|
||||
try:
|
||||
await self.async_container_client.upload_blob(key, serialized_value, overwrite=True)
|
||||
except Exception as e:
|
||||
|
|
@ -62,9 +63,9 @@ class AzureBlobCache(BaseCache):
|
|||
|
||||
try:
|
||||
print_verbose(f"Get Azure Blob Cache: key: {key}")
|
||||
as_bytes = self.container_client.download_blob(key).readall()
|
||||
as_str = as_bytes.decode("utf-8")
|
||||
cached_response = json.loads(as_str)
|
||||
as_bytes: Final = self.container_client.download_blob(key).readall()
|
||||
as_str: Final = as_bytes.decode("utf-8")
|
||||
cached_response: Final = json.loads(as_str)
|
||||
|
||||
verbose_logger.debug(
|
||||
"Got Azure Blob Cache: key: %s, cached_response %s. Type Response %s",
|
||||
|
|
@ -82,10 +83,10 @@ class AzureBlobCache(BaseCache):
|
|||
|
||||
try:
|
||||
print_verbose(f"Get Azure Blob Cache: key: {key}")
|
||||
blob = await self.async_container_client.download_blob(key)
|
||||
as_bytes = await blob.readall()
|
||||
as_str = as_bytes.decode("utf-8")
|
||||
cached_response = json.loads(as_str)
|
||||
blob: Final = await self.async_container_client.download_blob(key)
|
||||
as_bytes: Final = await blob.readall()
|
||||
as_str: Final = as_bytes.decode("utf-8")
|
||||
cached_response: Final = json.loads(as_str)
|
||||
verbose_logger.debug(
|
||||
"Got Azure Blob Cache: key: %s, cached_response %s. Type Response %s",
|
||||
key,
|
||||
|
|
@ -105,7 +106,7 @@ class AzureBlobCache(BaseCache):
|
|||
await self.async_container_client.close()
|
||||
|
||||
async def async_set_cache_pipeline(self, cache_list, **kwargs) -> None:
|
||||
tasks = []
|
||||
tasks: Final = []
|
||||
for val in cache_list:
|
||||
tasks.append(self.async_set_cache(val[0], val[1], **kwargs))
|
||||
await asyncio.gather(*tasks)
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ Has 4 methods:
|
|||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Any, Union
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
|
@ -24,7 +24,7 @@ class BaseCache(ABC):
|
|||
self.default_ttl = default_ttl
|
||||
|
||||
def get_ttl(self, **kwargs) -> int | None:
|
||||
kwargs_ttl: int | None = kwargs.get("ttl")
|
||||
kwargs_ttl: Final[int | None] = kwargs.get("ttl")
|
||||
if kwargs_ttl is not None:
|
||||
try:
|
||||
return int(kwargs_ttl)
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ import json
|
|||
import time
|
||||
import traceback
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
|
@ -169,13 +169,13 @@ class Cache:
|
|||
if type == LiteLLMCacheType.REDIS:
|
||||
# Check REDIS_CLUSTER_NODES env var if no explicit startup nodes
|
||||
if not redis_startup_nodes:
|
||||
_env_cluster_nodes = litellm.get_secret("REDIS_CLUSTER_NODES")
|
||||
_env_cluster_nodes: Final = litellm.get_secret("REDIS_CLUSTER_NODES")
|
||||
if _env_cluster_nodes is not None and isinstance(_env_cluster_nodes, str):
|
||||
redis_startup_nodes = json.loads(_env_cluster_nodes)
|
||||
|
||||
if redis_startup_nodes:
|
||||
# Only pass GCP parameters if they are provided
|
||||
cluster_kwargs = {
|
||||
cluster_kwargs: Final = {
|
||||
"host": host,
|
||||
"port": port,
|
||||
"password": password,
|
||||
|
|
@ -312,9 +312,9 @@ class Cache:
|
|||
)
|
||||
|
||||
def _get_semantic_cache_tenant_scope(self, kwargs: dict) -> str:
|
||||
metadata: dict = kwargs.get("metadata") or {}
|
||||
litellm_params: dict = kwargs.get("litellm_params") or {}
|
||||
metadata_in_litellm_params: dict = litellm_params.get("metadata") or {}
|
||||
metadata: Final[dict] = kwargs.get("metadata") or {}
|
||||
litellm_params: Final[dict] = kwargs.get("litellm_params") or {}
|
||||
metadata_in_litellm_params: Final[dict] = litellm_params.get("metadata") or {}
|
||||
|
||||
scope = ""
|
||||
for field in self._SEMANTIC_CACHE_TENANT_SCOPE_FIELDS:
|
||||
|
|
@ -338,15 +338,15 @@ class Cache:
|
|||
cache_key = ""
|
||||
# verbose_logger.debug("\nGetting Cache key. Kwargs: %s", kwargs)
|
||||
|
||||
preset_cache_key = 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 = ModelParamHelper._get_all_llm_api_params()
|
||||
litellm_param_kwargs = all_litellm_params
|
||||
is_semantic_cache = self._is_semantic_cache()
|
||||
scope_excluded_params = self._SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS if is_semantic_cache else frozenset()
|
||||
combined_kwargs: Final = ModelParamHelper._get_all_llm_api_params()
|
||||
litellm_param_kwargs: Final = all_litellm_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:
|
||||
if param in scope_excluded_params:
|
||||
continue
|
||||
|
|
@ -373,7 +373,7 @@ class Cache:
|
|||
)
|
||||
# Remove preset_cache_key from kwargs to avoid "got multiple values" TypeError
|
||||
# when kwargs already contains preset_cache_key from upstream callers
|
||||
kwargs_for_preset = {k: v for k, v in kwargs.items() if k != "preset_cache_key"}
|
||||
kwargs_for_preset: Final = {k: v for k, v in kwargs.items() if k != "preset_cache_key"}
|
||||
self._set_preset_cache_key_in_kwargs(preset_cache_key=hashed_cache_key, **kwargs_for_preset)
|
||||
return hashed_cache_key
|
||||
|
||||
|
|
@ -399,15 +399,15 @@ class Cache:
|
|||
2. Else if a model_group is set, then return the model_group as the model. This is used for all requests sent through the litellm.Router()
|
||||
3. Else use the `model` passed in kwargs
|
||||
"""
|
||||
metadata: dict = kwargs.get("metadata", {}) or {}
|
||||
litellm_params: dict = kwargs.get("litellm_params", {}) or {}
|
||||
metadata_in_litellm_params: dict = litellm_params.get("metadata", {}) or {}
|
||||
model_group: str | None = metadata.get("model_group") or metadata_in_litellm_params.get("model_group")
|
||||
caching_group = self._get_caching_group(metadata, model_group)
|
||||
metadata: Final[dict] = kwargs.get("metadata", {}) or {}
|
||||
litellm_params: Final[dict] = kwargs.get("litellm_params", {}) or {}
|
||||
metadata_in_litellm_params: Final[dict] = litellm_params.get("metadata", {}) or {}
|
||||
model_group: Final[str | None] = metadata.get("model_group") or metadata_in_litellm_params.get("model_group")
|
||||
caching_group: Final = self._get_caching_group(metadata, model_group)
|
||||
return caching_group or model_group or kwargs["model"]
|
||||
|
||||
def _get_caching_group(self, metadata: dict, model_group: str | None) -> str | None:
|
||||
caching_groups: list | None = metadata.get("caching_groups", [])
|
||||
caching_groups: Final[list | None] = metadata.get("caching_groups", [])
|
||||
if caching_groups:
|
||||
for group in caching_groups:
|
||||
if model_group in group:
|
||||
|
|
@ -418,9 +418,9 @@ class Cache:
|
|||
"""
|
||||
Handles getting the value for the 'file' param from kwargs. Used for `transcription` requests
|
||||
"""
|
||||
file = kwargs.get("file")
|
||||
metadata = kwargs.get("metadata", {})
|
||||
litellm_params = kwargs.get("litellm_params", {})
|
||||
file: Final = kwargs.get("file")
|
||||
metadata: Final = kwargs.get("metadata", {})
|
||||
litellm_params: Final = kwargs.get("litellm_params", {})
|
||||
return (
|
||||
metadata.get("file_checksum")
|
||||
or getattr(file, "name", None)
|
||||
|
|
@ -467,9 +467,9 @@ class Cache:
|
|||
Returns:
|
||||
str: The hashed cache key.
|
||||
"""
|
||||
hash_object = hashlib.sha256(cache_key.encode())
|
||||
hash_object: Final = hashlib.sha256(cache_key.encode())
|
||||
# Hexadecimal representation of the hash
|
||||
hash_hex = hash_object.hexdigest()
|
||||
hash_hex: Final = hash_object.hexdigest()
|
||||
verbose_logger.debug("Hashed cache key (SHA-256): %s", hash_hex)
|
||||
return hash_hex
|
||||
|
||||
|
|
@ -484,16 +484,16 @@ class Cache:
|
|||
Returns:
|
||||
str: The final hashed cache key with the redis namespace.
|
||||
"""
|
||||
dynamic_cache_control: DynamicCacheControl = kwargs.get("cache", {})
|
||||
metadata = kwargs.get("metadata") or {}
|
||||
namespace = dynamic_cache_control.get("namespace") or metadata.get("redis_namespace") or self.namespace
|
||||
dynamic_cache_control: Final[DynamicCacheControl] = kwargs.get("cache", {})
|
||||
metadata: Final = kwargs.get("metadata") or {}
|
||||
namespace: Final = dynamic_cache_control.get("namespace") or metadata.get("redis_namespace") or self.namespace
|
||||
if namespace:
|
||||
hash_hex = f"{namespace}:{hash_hex}"
|
||||
verbose_logger.debug("Final hashed key: %s", hash_hex)
|
||||
return hash_hex
|
||||
|
||||
def generate_streaming_content(self, content):
|
||||
chunk_size = 5 # Adjust the chunk size as needed
|
||||
chunk_size: Final = 5 # Adjust the chunk size as needed
|
||||
for i in range(0, len(content), chunk_size):
|
||||
yield {
|
||||
"choices": [
|
||||
|
|
@ -517,11 +517,11 @@ class Cache:
|
|||
"""
|
||||
# Check if a timestamp was stored with the cached response
|
||||
if cached_result is not None and isinstance(cached_result, dict) and "timestamp" in cached_result:
|
||||
timestamp = cached_result["timestamp"]
|
||||
current_time = time.time()
|
||||
timestamp: Final = cached_result["timestamp"]
|
||||
current_time: Final = time.time()
|
||||
|
||||
# Calculate age of the cached response
|
||||
response_age = current_time - timestamp
|
||||
response_age: Final = current_time - timestamp
|
||||
|
||||
# Check if the cached response is older than the max-age
|
||||
if max_age is not None and response_age > max_age:
|
||||
|
|
@ -544,12 +544,12 @@ class Cache:
|
|||
|
||||
@staticmethod
|
||||
def _get_safe_cache_lookup_kwargs(kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
cache_lookup_kwargs: dict[str, Any] = {}
|
||||
cache_lookup_kwargs: Final[dict[str, Any]] = {}
|
||||
for prompt_kwarg in ("messages", "input"):
|
||||
if prompt_kwarg in kwargs:
|
||||
cache_lookup_kwargs[prompt_kwarg] = kwargs[prompt_kwarg]
|
||||
|
||||
metadata = kwargs.get("metadata")
|
||||
metadata: Final = kwargs.get("metadata")
|
||||
if isinstance(metadata, dict):
|
||||
cache_lookup_kwargs["metadata"] = dict(metadata)
|
||||
|
||||
|
|
@ -559,8 +559,8 @@ class Cache:
|
|||
def _update_metadata_from_cache_lookup_kwargs(
|
||||
original_kwargs: dict[str, Any], cache_lookup_kwargs: dict[str, Any]
|
||||
) -> None:
|
||||
original_metadata = original_kwargs.get("metadata")
|
||||
cache_lookup_metadata = cache_lookup_kwargs.get("metadata")
|
||||
original_metadata: Final = original_kwargs.get("metadata")
|
||||
cache_lookup_metadata: Final = cache_lookup_kwargs.get("metadata")
|
||||
if not isinstance(original_metadata, dict) or not isinstance(cache_lookup_metadata, dict):
|
||||
return
|
||||
|
||||
|
|
@ -586,9 +586,9 @@ class Cache:
|
|||
else:
|
||||
cache_key = self.get_cache_key(**kwargs)
|
||||
if cache_key is not None:
|
||||
cache_control_args: DynamicCacheControl = kwargs.get("cache", {})
|
||||
cache_control_args: Final[DynamicCacheControl] = kwargs.get("cache", {})
|
||||
max_age = cache_control_args.get("s-maxage") or cache_control_args.get("s-max-age") or float("inf")
|
||||
cache_lookup_kwargs = self._get_safe_cache_lookup_kwargs(kwargs)
|
||||
cache_lookup_kwargs: Final = self._get_safe_cache_lookup_kwargs(kwargs)
|
||||
if dynamic_cache_object is not None:
|
||||
cached_result = dynamic_cache_object.get_cache(cache_key, **cache_lookup_kwargs)
|
||||
else:
|
||||
|
|
@ -618,8 +618,8 @@ class Cache:
|
|||
else:
|
||||
cache_key = self.get_cache_key(**kwargs)
|
||||
if cache_key is not None:
|
||||
cache_control_args = kwargs.get("cache", {})
|
||||
max_age = cache_control_args.get("s-max-age", cache_control_args.get("s-maxage", float("inf")))
|
||||
cache_control_args: Final = kwargs.get("cache", {})
|
||||
max_age: Final = cache_control_args.get("s-max-age", cache_control_args.get("s-maxage", float("inf")))
|
||||
if dynamic_cache_object is not None:
|
||||
cached_result = await dynamic_cache_object.async_get_cache(cache_key, **kwargs)
|
||||
else:
|
||||
|
|
@ -646,13 +646,13 @@ class Cache:
|
|||
if self.ttl is not None:
|
||||
kwargs["ttl"] = self.ttl
|
||||
## Get Cache-Controls ##
|
||||
_cache_kwargs = kwargs.get("cache", None)
|
||||
_cache_kwargs: Final = kwargs.get("cache", None)
|
||||
if isinstance(_cache_kwargs, dict):
|
||||
for k, v in _cache_kwargs.items():
|
||||
if k == "ttl":
|
||||
kwargs["ttl"] = v
|
||||
|
||||
cached_data = {"timestamp": time.time(), "response": result}
|
||||
cached_data: Final = {"timestamp": time.time(), "response": result}
|
||||
return cache_key, cached_data, kwargs
|
||||
else:
|
||||
raise Exception("cache key is None")
|
||||
|
|
@ -756,7 +756,7 @@ class Cache:
|
|||
if result.usage is None or result.usage.prompt_tokens_details is None:
|
||||
return None
|
||||
|
||||
details = result.usage.prompt_tokens_details
|
||||
details: Final = result.usage.prompt_tokens_details
|
||||
if hasattr(details, "model_dump"):
|
||||
details_dict = details.model_dump(exclude_none=True)
|
||||
elif isinstance(details, dict):
|
||||
|
|
@ -767,12 +767,12 @@ class Cache:
|
|||
if not details_dict:
|
||||
return None
|
||||
|
||||
num_items = len(result.data)
|
||||
num_items: Final = len(result.data)
|
||||
if num_items <= 1:
|
||||
return details_dict
|
||||
|
||||
# Distribute integer/float fields evenly across items
|
||||
per_item: dict = {}
|
||||
per_item: Final[dict] = {}
|
||||
for key, value in details_dict.items():
|
||||
if isinstance(value, int):
|
||||
quotient, remainder = divmod(value, num_items)
|
||||
|
|
@ -798,8 +798,8 @@ class Cache:
|
|||
if result.usage is None or result.usage.prompt_tokens is None:
|
||||
return None
|
||||
|
||||
total = result.usage.prompt_tokens
|
||||
num_items = len(result.data)
|
||||
total: Final = result.usage.prompt_tokens
|
||||
num_items: Final = len(result.data)
|
||||
if num_items <= 1:
|
||||
return total
|
||||
|
||||
|
|
@ -813,23 +813,23 @@ class Cache:
|
|||
kwargs: dict,
|
||||
idx_in_result_data: int = 0,
|
||||
) -> tuple[str, dict, dict]:
|
||||
preset_cache_key = self.get_cache_key(**{**kwargs, "input": input})
|
||||
preset_cache_key: Final = self.get_cache_key(**{**kwargs, "input": input})
|
||||
kwargs["cache_key"] = preset_cache_key
|
||||
embedding_response = result.data[idx_in_result_data]
|
||||
embedding_response: Final = result.data[idx_in_result_data]
|
||||
|
||||
# Extract per-item prompt_tokens + details from response usage
|
||||
prompt_tokens = self._get_per_item_prompt_tokens(
|
||||
prompt_tokens: Final = self._get_per_item_prompt_tokens(
|
||||
result=result,
|
||||
idx_in_result_data=idx_in_result_data,
|
||||
)
|
||||
prompt_tokens_details = self._get_per_item_prompt_tokens_details(
|
||||
prompt_tokens_details: Final = self._get_per_item_prompt_tokens_details(
|
||||
result=result,
|
||||
idx_in_result_data=idx_in_result_data,
|
||||
)
|
||||
|
||||
# Always convert to properly typed CachedEmbedding
|
||||
model_name = result.model
|
||||
embedding_dict: CachedEmbedding = self._convert_to_cached_embedding(
|
||||
model_name: Final = result.model
|
||||
embedding_dict: Final[CachedEmbedding] = self._convert_to_cached_embedding(
|
||||
embedding_response,
|
||||
model_name,
|
||||
prompt_tokens=prompt_tokens,
|
||||
|
|
@ -856,7 +856,7 @@ class Cache:
|
|||
if self.ttl is not None:
|
||||
kwargs["ttl"] = self.ttl
|
||||
|
||||
cache_list = []
|
||||
cache_list: Final = []
|
||||
if isinstance(kwargs["input"], list):
|
||||
for idx, i in enumerate(kwargs["input"]):
|
||||
(
|
||||
|
|
@ -887,7 +887,7 @@ class Cache:
|
|||
return True
|
||||
|
||||
# when mode == default_off -> Cache is opt in only
|
||||
_cache = kwargs.get("cache", None)
|
||||
_cache: Final = kwargs.get("cache", None)
|
||||
verbose_logger.debug("should_use_cache: kwargs: %s; _cache: %s", kwargs, _cache)
|
||||
if _cache and isinstance(_cache, dict):
|
||||
if _cache.get("use-cache", False) is True:
|
||||
|
|
@ -899,13 +899,13 @@ class Cache:
|
|||
await self.cache.batch_cache_write(cache_key, cached_data, **kwargs)
|
||||
|
||||
async def ping(self):
|
||||
cache_ping = getattr(self.cache, "ping")
|
||||
cache_ping: Final = getattr(self.cache, "ping")
|
||||
if cache_ping:
|
||||
return await cache_ping()
|
||||
return None
|
||||
|
||||
async def delete_cache_keys(self, keys):
|
||||
cache_delete_cache_keys = getattr(self.cache, "delete_cache_keys")
|
||||
cache_delete_cache_keys: Final = getattr(self.cache, "delete_cache_keys")
|
||||
if cache_delete_cache_keys:
|
||||
return await cache_delete_cache_keys(keys)
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -19,11 +19,7 @@ import datetime
|
|||
import inspect
|
||||
import time
|
||||
from collections.abc import AsyncGenerator, Callable, Generator
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Optional,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
|
@ -76,7 +72,7 @@ class CachingHandlerResponse(BaseModel):
|
|||
embedding_all_elements_cache_hit: bool = False # this is set to True when all elements in the list have a cache hit in the embedding cache, if true return the final_embedding_cached_response no need to make an API call
|
||||
|
||||
|
||||
in_memory_cache_obj = InMemoryCache()
|
||||
in_memory_cache_obj: Final = InMemoryCache()
|
||||
|
||||
|
||||
def _drop_logging_obj_from_kwargs(request_kwargs: dict[str, object]) -> dict[str, object]:
|
||||
|
|
@ -96,10 +92,10 @@ def _drop_logging_obj_from_kwargs(request_kwargs: dict[str, object]) -> dict[str
|
|||
|
||||
|
||||
def _is_chat_completion_cached_dict(cached_result: dict) -> bool:
|
||||
cached_id = cached_result.get("id")
|
||||
cached_id: Final = cached_result.get("id")
|
||||
if isinstance(cached_id, str) and cached_id.startswith("chatcmpl"):
|
||||
return True
|
||||
obj = cached_result.get("object")
|
||||
obj: Final = cached_result.get("object")
|
||||
if isinstance(obj, str):
|
||||
return obj.startswith("chat.completion")
|
||||
return "choices" in cached_result
|
||||
|
|
@ -184,10 +180,10 @@ class LLMCachingHandler:
|
|||
#########################################################
|
||||
# Init cache timing metrics
|
||||
#########################################################
|
||||
cache_check_start_time = time.perf_counter()
|
||||
cache_check_start_time: Final = time.perf_counter()
|
||||
cache_check_end_time: float | None = None
|
||||
#########################################################
|
||||
parent_otel_span = _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):
|
||||
|
|
@ -201,15 +197,15 @@ class LLMCachingHandler:
|
|||
|
||||
if cached_result is not None and not isinstance(cached_result, list):
|
||||
verbose_logger.debug("Cache Hit!")
|
||||
cache_hit = True
|
||||
end_time = datetime.datetime.now()
|
||||
cache_hit: Final = True
|
||||
end_time: Final = datetime.datetime.now()
|
||||
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model=model,
|
||||
custom_llm_provider=kwargs.get("custom_llm_provider", None),
|
||||
api_base=kwargs.get("api_base", None),
|
||||
api_key=kwargs.get("api_key", None),
|
||||
)
|
||||
cache_duration_ms = (cache_check_end_time - cache_check_start_time) * 1000
|
||||
cache_duration_ms: Final = (cache_check_end_time - cache_check_start_time) * 1000
|
||||
self._update_litellm_logging_obj_environment(
|
||||
logging_obj=logging_obj,
|
||||
model=model,
|
||||
|
|
@ -240,7 +236,7 @@ class LLMCachingHandler:
|
|||
end_time=end_time,
|
||||
cache_hit=cache_hit,
|
||||
)
|
||||
cache_key = (
|
||||
cache_key: Final = (
|
||||
self.preset_cache_key
|
||||
or self.request_kwargs.get("cache_key")
|
||||
or litellm.cache.get_cache_key(**self.request_kwargs)
|
||||
|
|
@ -295,7 +291,7 @@ class LLMCachingHandler:
|
|||
if litellm.cache is not None and self._is_call_type_supported_by_cache(original_function=original_function):
|
||||
args = args or ()
|
||||
# Now that we confirmed caching will happen, prepare kwargs
|
||||
new_kwargs = kwargs.copy()
|
||||
new_kwargs: Final = kwargs.copy()
|
||||
new_kwargs.update(
|
||||
convert_args_to_kwargs(
|
||||
self.original_function,
|
||||
|
|
@ -326,8 +322,8 @@ class LLMCachingHandler:
|
|||
)
|
||||
|
||||
# LOG SUCCESS
|
||||
cache_hit = True
|
||||
end_time = datetime.datetime.now()
|
||||
cache_hit: Final = True
|
||||
end_time: Final = datetime.datetime.now()
|
||||
(
|
||||
model,
|
||||
custom_llm_provider,
|
||||
|
|
@ -354,7 +350,7 @@ class LLMCachingHandler:
|
|||
end_time=end_time,
|
||||
cache_hit=cache_hit,
|
||||
)
|
||||
cache_key = (
|
||||
cache_key: Final = (
|
||||
self.preset_cache_key
|
||||
or self.request_kwargs.get("cache_key")
|
||||
or litellm.cache.get_cache_key(**self.request_kwargs)
|
||||
|
|
@ -420,9 +416,9 @@ class LLMCachingHandler:
|
|||
|
||||
"""
|
||||
embedding_all_elements_cache_hit: bool = False
|
||||
remaining_list = []
|
||||
non_null_list = []
|
||||
kwargs_input_as_list = self.handle_kwargs_input_list_or_str(kwargs)
|
||||
remaining_list: Final = []
|
||||
non_null_list: Final = []
|
||||
kwargs_input_as_list: Final = self.handle_kwargs_input_list_or_str(kwargs)
|
||||
for idx, cr in enumerate(cached_result):
|
||||
if cr is None:
|
||||
remaining_list.append(kwargs_input_as_list[idx])
|
||||
|
|
@ -479,7 +475,7 @@ class LLMCachingHandler:
|
|||
prompt_tokens_details = PromptTokensDetailsWrapper(**aggregated_details)
|
||||
except Exception:
|
||||
prompt_tokens_details = None
|
||||
usage = Usage(
|
||||
usage: Final = Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=0,
|
||||
total_tokens=prompt_tokens,
|
||||
|
|
@ -488,9 +484,9 @@ class LLMCachingHandler:
|
|||
final_embedding_cached_response.usage = usage
|
||||
if len(remaining_list) == 0:
|
||||
# LOG SUCCESS
|
||||
cache_hit = True
|
||||
cache_hit: Final = True
|
||||
embedding_all_elements_cache_hit = True
|
||||
end_time = datetime.datetime.now()
|
||||
end_time: Final = datetime.datetime.now()
|
||||
(
|
||||
model,
|
||||
custom_llm_provider,
|
||||
|
|
@ -546,10 +542,10 @@ class LLMCachingHandler:
|
|||
if details2 is None:
|
||||
return details1
|
||||
|
||||
dict1 = details1.model_dump(exclude_none=True) if hasattr(details1, "model_dump") else {}
|
||||
dict2 = details2.model_dump(exclude_none=True) if hasattr(details2, "model_dump") else {}
|
||||
dict1: Final = details1.model_dump(exclude_none=True) if hasattr(details1, "model_dump") else {}
|
||||
dict2: Final = details2.model_dump(exclude_none=True) if hasattr(details2, "model_dump") else {}
|
||||
|
||||
merged: dict = {}
|
||||
merged: Final[dict] = {}
|
||||
for key in set(dict1.keys()) | set(dict2.keys()):
|
||||
v1 = dict1.get(key, 0)
|
||||
v2 = dict2.get(key, 0)
|
||||
|
|
@ -607,7 +603,7 @@ class LLMCachingHandler:
|
|||
return embedding_response
|
||||
|
||||
idx = 0
|
||||
final_data_list = []
|
||||
final_data_list: Final = []
|
||||
for item in _caching_handler_response.final_embedding_cached_response.data:
|
||||
if item is None and embedding_response.data is not None:
|
||||
final_data_list.append(embedding_response.data[idx])
|
||||
|
|
@ -690,7 +686,7 @@ class LLMCachingHandler:
|
|||
if litellm.cache is None:
|
||||
return None
|
||||
|
||||
new_kwargs = kwargs.copy()
|
||||
new_kwargs: Final = kwargs.copy()
|
||||
new_kwargs.update(
|
||||
convert_args_to_kwargs(
|
||||
self.original_function,
|
||||
|
|
@ -708,7 +704,7 @@ class LLMCachingHandler:
|
|||
new_kwargs["input"] = [new_kwargs["input"]]
|
||||
elif not isinstance(new_kwargs["input"], list):
|
||||
raise ValueError("input must be a string or a list")
|
||||
tasks = []
|
||||
tasks: Final = []
|
||||
for idx, i in enumerate(new_kwargs["input"]):
|
||||
preset_cache_key = litellm.cache.get_cache_key(**{**new_kwargs, "input": i})
|
||||
tasks.append(
|
||||
|
|
@ -724,8 +720,8 @@ class LLMCachingHandler:
|
|||
if all(result is None for result in cached_result):
|
||||
cached_result = None
|
||||
else:
|
||||
request_kwargs = new_kwargs.copy()
|
||||
request_cache_key = request_kwargs.pop("cache_key", None)
|
||||
request_kwargs: Final = new_kwargs.copy()
|
||||
request_cache_key: Final = request_kwargs.pop("cache_key", None)
|
||||
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)
|
||||
|
|
@ -828,7 +824,7 @@ class LLMCachingHandler:
|
|||
elif (call_type == CallTypes.atranscription.value or call_type == CallTypes.transcription.value) and isinstance(
|
||||
cached_result, dict
|
||||
):
|
||||
hidden_params = {
|
||||
hidden_params: Final = {
|
||||
"model": "whisper-1",
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"cache_hit": True,
|
||||
|
|
@ -840,10 +836,10 @@ class LLMCachingHandler:
|
|||
hidden_params=hidden_params,
|
||||
)
|
||||
elif (call_type == "aresponses" or call_type == "responses") and isinstance(cached_result, dict):
|
||||
use_chat_completion_cache = _is_chat_completion_cached_dict(cached_result)
|
||||
use_chat_completion_cache: Final = _is_chat_completion_cached_dict(cached_result)
|
||||
if use_chat_completion_cache:
|
||||
if kwargs.get("stream", False) is True:
|
||||
bridge_call_type = (
|
||||
bridge_call_type: Final = (
|
||||
CallTypes.acompletion.value if call_type == "aresponses" else CallTypes.completion.value
|
||||
)
|
||||
cached_result = self._convert_cached_stream_response(
|
||||
|
|
@ -862,7 +858,7 @@ class LLMCachingHandler:
|
|||
CachedResponsesAPIStreamingIterator,
|
||||
)
|
||||
|
||||
response_obj = ResponsesAPIResponse(**cached_result)
|
||||
response_obj: Final = ResponsesAPIResponse(**cached_result)
|
||||
if (
|
||||
hasattr(response_obj, "_hidden_params")
|
||||
and response_obj._hidden_params is not None
|
||||
|
|
@ -957,14 +953,14 @@ class LLMCachingHandler:
|
|||
if litellm.cache is None:
|
||||
return
|
||||
|
||||
new_kwargs = kwargs.copy()
|
||||
new_kwargs: Final = kwargs.copy()
|
||||
new_kwargs.update(
|
||||
convert_args_to_kwargs(
|
||||
original_function,
|
||||
args,
|
||||
)
|
||||
)
|
||||
parent_otel_span = _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):
|
||||
|
|
@ -1006,7 +1002,7 @@ class LLMCachingHandler:
|
|||
Sync internal method to add the result to the cache
|
||||
"""
|
||||
|
||||
new_kwargs = kwargs.copy()
|
||||
new_kwargs: Final = kwargs.copy()
|
||||
new_kwargs.update(
|
||||
convert_args_to_kwargs(
|
||||
self.original_function,
|
||||
|
|
@ -1067,7 +1063,7 @@ class LLMCachingHandler:
|
|||
|
||||
"""
|
||||
|
||||
complete_streaming_response: ModelResponse | TextCompletionResponse | None = (
|
||||
complete_streaming_response: Final[ModelResponse | TextCompletionResponse | None] = (
|
||||
_assemble_complete_response_from_streaming_chunks(
|
||||
result=processed_chunk,
|
||||
start_time=self.start_time,
|
||||
|
|
@ -1089,7 +1085,7 @@ class LLMCachingHandler:
|
|||
"""
|
||||
Sync internal method to add the streaming response to the cache
|
||||
"""
|
||||
complete_streaming_response: ModelResponse | TextCompletionResponse | None = (
|
||||
complete_streaming_response: Final[ModelResponse | TextCompletionResponse | None] = (
|
||||
_assemble_complete_response_from_streaming_chunks(
|
||||
result=processed_chunk,
|
||||
start_time=self.start_time,
|
||||
|
|
@ -1133,7 +1129,7 @@ class LLMCachingHandler:
|
|||
Returns:
|
||||
None
|
||||
"""
|
||||
litellm_params = {
|
||||
litellm_params: Final = {
|
||||
"logger_fn": kwargs.get("logger_fn", None),
|
||||
"acompletion": is_async,
|
||||
"api_base": kwargs.get("api_base", ""),
|
||||
|
|
@ -1173,13 +1169,13 @@ def convert_args_to_kwargs(
|
|||
args: tuple[Any, ...] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
# Get the signature of the original function
|
||||
signature = inspect.signature(original_function)
|
||||
signature: Final = inspect.signature(original_function)
|
||||
|
||||
# Get parameter names in the order they appear in the original function
|
||||
param_names = list(signature.parameters.keys())
|
||||
param_names: Final = list(signature.parameters.keys())
|
||||
|
||||
# Create a mapping of positional arguments to parameter names
|
||||
args_to_kwargs = {}
|
||||
args_to_kwargs: Final = {}
|
||||
if args:
|
||||
for index, arg in enumerate(args):
|
||||
if index < len(param_names):
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import json
|
||||
from typing import TYPE_CHECKING, Any, Union
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
|
||||
from .base_cache import BaseCache
|
||||
|
||||
|
|
@ -41,7 +41,7 @@ class DiskCache(BaseCache):
|
|||
self.set_cache(key=cache_key, value=cache_value)
|
||||
|
||||
def get_cache(self, key, **kwargs):
|
||||
original_cached_response = self.disk_cache.get(key)
|
||||
original_cached_response: Final = self.disk_cache.get(key)
|
||||
if original_cached_response:
|
||||
try:
|
||||
cached_response = json.loads(original_cached_response) # type: ignore
|
||||
|
|
@ -51,7 +51,7 @@ class DiskCache(BaseCache):
|
|||
return None
|
||||
|
||||
def batch_get_cache(self, keys: list, **kwargs):
|
||||
return_val = []
|
||||
return_val: Final = []
|
||||
for k in keys:
|
||||
val = self.get_cache(key=k, **kwargs)
|
||||
return_val.append(val)
|
||||
|
|
@ -59,9 +59,9 @@ class DiskCache(BaseCache):
|
|||
|
||||
def increment_cache(self, key, value: int, **kwargs) -> int:
|
||||
with self.disk_cache.transact():
|
||||
cached_value = self.get_cache(key=key)
|
||||
init_value = cached_value if isinstance(cached_value, int) else 0
|
||||
new_value = init_value + value
|
||||
cached_value: Final = self.get_cache(key=key)
|
||||
init_value: Final = cached_value if isinstance(cached_value, int) else 0
|
||||
new_value: Final = init_value + value
|
||||
self.set_cache(key, new_value, **kwargs)
|
||||
return new_value
|
||||
|
||||
|
|
@ -69,7 +69,7 @@ class DiskCache(BaseCache):
|
|||
return self.get_cache(key=key, **kwargs)
|
||||
|
||||
async def async_batch_get_cache(self, keys: list, **kwargs):
|
||||
return_val = []
|
||||
return_val: Final = []
|
||||
for k in keys:
|
||||
val = self.get_cache(key=k, **kwargs)
|
||||
return_val.append(val)
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ import time
|
|||
import traceback
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from threading import Lock
|
||||
from typing import TYPE_CHECKING, Any, Union
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
|
|
@ -161,14 +161,14 @@ class DualCache(BaseCache):
|
|||
try:
|
||||
result = None
|
||||
if self.in_memory_cache is not None:
|
||||
in_memory_result = self.in_memory_cache.get_cache(key, **kwargs)
|
||||
in_memory_result: Final = self.in_memory_cache.get_cache(key, **kwargs)
|
||||
|
||||
if in_memory_result is not None:
|
||||
result = in_memory_result
|
||||
|
||||
if result is None and self.redis_cache is not None and local_only is False:
|
||||
# If not found in in-memory cache, try fetching from Redis
|
||||
redis_result = self.redis_cache.get_cache(key, parent_otel_span=parent_otel_span)
|
||||
redis_result: Final = self.redis_cache.get_cache(key, parent_otel_span=parent_otel_span)
|
||||
|
||||
if redis_result is not None:
|
||||
# Update in-memory cache with the value from Redis
|
||||
|
|
@ -188,12 +188,12 @@ class DualCache(BaseCache):
|
|||
local_only: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
received_args = locals()
|
||||
received_args: Final = locals()
|
||||
received_args.pop("self")
|
||||
|
||||
def run_in_new_loop():
|
||||
"""Run the coroutine in a new event loop within this thread."""
|
||||
new_loop = asyncio.new_event_loop()
|
||||
new_loop: Final = asyncio.new_event_loop()
|
||||
try:
|
||||
asyncio.set_event_loop(new_loop)
|
||||
return new_loop.run_until_complete(self.async_batch_get_cache(**received_args))
|
||||
|
|
@ -207,7 +207,7 @@ class DualCache(BaseCache):
|
|||
# If we're already in an event loop, run in a separate thread
|
||||
# to avoid nested event loop issues
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
future = executor.submit(run_in_new_loop)
|
||||
future: Final = executor.submit(run_in_new_loop)
|
||||
return future.result()
|
||||
|
||||
except RuntimeError:
|
||||
|
|
@ -226,7 +226,7 @@ class DualCache(BaseCache):
|
|||
print_verbose(f"async get cache: cache key: {key}; local_only: {local_only}")
|
||||
result = None
|
||||
if self.in_memory_cache is not None:
|
||||
in_memory_result = await self.in_memory_cache.async_get_cache(key, **kwargs)
|
||||
in_memory_result: Final = await self.in_memory_cache.async_get_cache(key, **kwargs)
|
||||
|
||||
print_verbose(f"in_memory_result: {in_memory_result}")
|
||||
if in_memory_result is not None:
|
||||
|
|
@ -234,7 +234,7 @@ class DualCache(BaseCache):
|
|||
|
||||
if result is None and self.redis_cache is not None and local_only is False:
|
||||
# If not found in in-memory cache, try fetching from Redis
|
||||
redis_result = await self.redis_cache.async_get_cache(key, parent_otel_span=parent_otel_span)
|
||||
redis_result: Final = await self.redis_cache.async_get_cache(key, parent_otel_span=parent_otel_span)
|
||||
|
||||
if redis_result is not None:
|
||||
# Update in-memory cache with the value from Redis
|
||||
|
|
@ -257,8 +257,8 @@ class DualCache(BaseCache):
|
|||
Atomically choose keys to fetch from Redis and reserve their access time.
|
||||
This prevents check-then-act races under concurrent async callers.
|
||||
"""
|
||||
sublist_keys: list[str] = []
|
||||
previous_access_times: dict[str, float | None] = {}
|
||||
sublist_keys: Final[list[str]] = []
|
||||
previous_access_times: Final[dict[str, float | None]] = {}
|
||||
|
||||
with self._last_redis_batch_access_time_lock:
|
||||
for key, value in zip(keys, result):
|
||||
|
|
@ -293,7 +293,7 @@ class DualCache(BaseCache):
|
|||
try:
|
||||
result = [None] * len(keys)
|
||||
if self.in_memory_cache is not None:
|
||||
in_memory_result = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs)
|
||||
in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs)
|
||||
|
||||
if in_memory_result is not None:
|
||||
result = in_memory_result
|
||||
|
|
@ -303,14 +303,14 @@ class DualCache(BaseCache):
|
|||
- for the none values in the result
|
||||
- check the redis cache
|
||||
"""
|
||||
current_time = time.time()
|
||||
current_time: Final = time.time()
|
||||
sublist_keys, previous_access_times = self._reserve_redis_batch_keys(current_time, keys, result)
|
||||
|
||||
# Only hit Redis if enough time has passed since last access.
|
||||
if len(sublist_keys) > 0:
|
||||
try:
|
||||
# If not found in in-memory cache, try fetching from Redis
|
||||
redis_result = await self.redis_cache.async_batch_get_cache(
|
||||
redis_result: Final = await self.redis_cache.async_batch_get_cache(
|
||||
sublist_keys, parent_otel_span=parent_otel_span
|
||||
)
|
||||
except Exception:
|
||||
|
|
@ -323,7 +323,7 @@ class DualCache(BaseCache):
|
|||
return result
|
||||
|
||||
# Pre-compute key-to-index mapping for O(1) lookup
|
||||
key_to_index = {key: i for i, key in enumerate(keys)}
|
||||
key_to_index: Final = {key: i for i, key in enumerate(keys)}
|
||||
|
||||
# Update both result and in-memory cache in a single loop
|
||||
for key, value in redis_result.items():
|
||||
|
|
|
|||
|
|
@ -41,14 +41,15 @@ import weakref
|
|||
from collections import deque
|
||||
from collections.abc import Awaitable, Callable, Iterator
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import (
|
||||
EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS,
|
||||
EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING,
|
||||
)
|
||||
|
||||
_CLOSABLE_ANYWHERE = "closable-anywhere"
|
||||
_CLOSABLE_ON_ANY_LOOP = "closable-on-any-loop"
|
||||
_CLOSABLE_ANYWHERE: Final = "closable-anywhere"
|
||||
_CLOSABLE_ON_ANY_LOOP: Final = "closable-on-any-loop"
|
||||
|
||||
_BucketKey = str | int
|
||||
|
||||
|
|
@ -88,7 +89,7 @@ def _running_loop_id() -> int | None:
|
|||
|
||||
|
||||
def _close_function(client: object) -> Callable[[], object] | None:
|
||||
close_fn: Callable[[], object] | None = getattr(client, "aclose", None) or getattr(client, "close", None)
|
||||
close_fn: Final[Callable[[], object] | None] = getattr(client, "aclose", None) or getattr(client, "close", None)
|
||||
return close_fn
|
||||
|
||||
|
||||
|
|
@ -103,7 +104,7 @@ def _transport_of(client: object) -> object:
|
|||
|
||||
def _connection_is_idle(connection: object) -> bool:
|
||||
"""A pooled connection is idle unless it is servicing a request."""
|
||||
is_idle: object = getattr(connection, "is_idle", None)
|
||||
is_idle: Final[object] = getattr(connection, "is_idle", None)
|
||||
return bool(is_idle()) if callable(is_idle) else True
|
||||
|
||||
|
||||
|
|
@ -112,7 +113,7 @@ def _pool_has_busy_connection(transport: object) -> bool | None:
|
|||
|
||||
``None`` when there is no such pool, so the caller can ask the other backend.
|
||||
"""
|
||||
pooled: object = getattr(getattr(transport, "_pool", None), "connections", None)
|
||||
pooled: Final[object] = getattr(getattr(transport, "_pool", None), "connections", None)
|
||||
if not isinstance(pooled, (list, tuple)):
|
||||
return None
|
||||
return any(
|
||||
|
|
@ -134,11 +135,11 @@ def _has_connection_in_flight(client: object) -> bool:
|
|||
window as the only guard, exactly as it was before this check existed.
|
||||
"""
|
||||
try:
|
||||
transport = _transport_of(client)
|
||||
pooled_busy = _pool_has_busy_connection(transport)
|
||||
transport: Final = _transport_of(client)
|
||||
pooled_busy: Final = _pool_has_busy_connection(transport)
|
||||
if pooled_busy is not None:
|
||||
return pooled_busy
|
||||
session: object = getattr(transport, "client", None)
|
||||
session: Final[object] = getattr(transport, "client", None)
|
||||
return bool(getattr(getattr(session, "connector", None), "_acquired", None))
|
||||
except Exception: # noqa: BLE001 - a client that cannot report its state is treated as idle
|
||||
return False
|
||||
|
|
@ -192,7 +193,7 @@ class EvictedClientCloser:
|
|||
"""
|
||||
if client is None or not self._is_owned(client):
|
||||
return
|
||||
close_fn = _close_function(client)
|
||||
close_fn: Final = _close_function(client)
|
||||
if close_fn is None:
|
||||
return
|
||||
if self._pending_count >= self._max_pending:
|
||||
|
|
@ -214,7 +215,7 @@ class EvictedClientCloser:
|
|||
"""
|
||||
if not self._pending_count:
|
||||
return
|
||||
now = self._clock()
|
||||
now: Final = self._clock()
|
||||
for pending in self._take_due(_running_loop_id(), now):
|
||||
client = pending.client_ref()
|
||||
if client is None:
|
||||
|
|
@ -236,7 +237,7 @@ class EvictedClientCloser:
|
|||
the front rather than having to be searched for.
|
||||
"""
|
||||
with self._queue_lock:
|
||||
bucket = self._buckets.setdefault(_bucket_key(pending), deque()) # mutable-ok: FIFO by design
|
||||
bucket: Final = self._buckets.setdefault(_bucket_key(pending), deque()) # mutable-ok: FIFO by design
|
||||
while bucket and bucket[0].client_ref() is None:
|
||||
bucket.popleft()
|
||||
self._pending_count -= 1
|
||||
|
|
@ -249,7 +250,7 @@ class EvictedClientCloser:
|
|||
return tuple(pending for key in buckets for pending in self._drain_locked(key, now))
|
||||
|
||||
def _drain_locked(self, key: _BucketKey, now: float) -> Iterator[_PendingClose]:
|
||||
bucket = self._buckets.get(key)
|
||||
bucket: Final = self._buckets.get(key)
|
||||
if bucket is None:
|
||||
return
|
||||
while bucket and bucket[0].close_after <= now:
|
||||
|
|
@ -259,18 +260,18 @@ class EvictedClientCloser:
|
|||
del self._buckets[key]
|
||||
|
||||
def _close(self, client: object) -> None:
|
||||
close_fn = _close_function(client)
|
||||
close_fn: Final = _close_function(client)
|
||||
if close_fn is None:
|
||||
return
|
||||
try:
|
||||
closing = close_fn()
|
||||
closing: Final = close_fn()
|
||||
except Exception: # noqa: BLE001 - a discarded client's close must never surface to callers
|
||||
return
|
||||
if not inspect.isawaitable(closing):
|
||||
return
|
||||
task = asyncio.get_running_loop().create_task(_close_quietly(closing))
|
||||
task: Final = asyncio.get_running_loop().create_task(_close_quietly(closing))
|
||||
self._close_tasks.add(task)
|
||||
task.add_done_callback(self._close_tasks.discard)
|
||||
|
||||
|
||||
default_evicted_client_closer = EvictedClientCloser()
|
||||
default_evicted_client_closer: Final = EvictedClientCloser()
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ Supports syncing responses to Google Cloud Storage Buckets using HTTP requests.
|
|||
|
||||
import asyncio
|
||||
import json
|
||||
from typing import Final
|
||||
from urllib.parse import quote
|
||||
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
|
|
@ -33,7 +34,7 @@ class GCSCache(BaseCache):
|
|||
self.sync_client = _get_httpx_client()
|
||||
|
||||
def _construct_headers(self) -> dict:
|
||||
base = GCSBucketBase(bucket_name=self.bucket_name)
|
||||
base: Final = GCSBucketBase(bucket_name=self.bucket_name)
|
||||
base.path_service_account_json = self.path_service_account
|
||||
base.BUCKET_NAME = self.bucket_name
|
||||
return base.sync_construct_request_headers()
|
||||
|
|
@ -41,35 +42,35 @@ class GCSCache(BaseCache):
|
|||
def set_cache(self, key, value, **kwargs):
|
||||
try:
|
||||
print_verbose(f"LiteLLM SET Cache - GCS. Key={key}. Value={value}")
|
||||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
headers: Final = self._construct_headers()
|
||||
object_name: Final = self.key_prefix + key
|
||||
bucket_name: Final = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={quote(object_name, safe='')}"
|
||||
data = json.dumps(value)
|
||||
data: Final = json.dumps(value)
|
||||
self.sync_client.post(url=url, data=data, headers=headers)
|
||||
except Exception as e:
|
||||
print_verbose(f"GCS Caching: set_cache() - Got exception from GCS: {e}")
|
||||
|
||||
async def async_set_cache(self, key, value, **kwargs):
|
||||
try:
|
||||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
headers: Final = self._construct_headers()
|
||||
object_name: Final = self.key_prefix + key
|
||||
bucket_name: Final = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={quote(object_name, safe='')}"
|
||||
data = json.dumps(value)
|
||||
data: Final = json.dumps(value)
|
||||
await self.async_client.post(url=url, data=data, headers=headers)
|
||||
except Exception as e:
|
||||
print_verbose(f"GCS Caching: async_set_cache() - Got exception from GCS: {e}")
|
||||
|
||||
def get_cache(self, key, **kwargs):
|
||||
try:
|
||||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
headers: Final = self._construct_headers()
|
||||
object_name: Final = self.key_prefix + key
|
||||
bucket_name: Final = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{quote(object_name, safe='')}?alt=media"
|
||||
response = self.sync_client.get(url=url, headers=headers)
|
||||
response: Final = self.sync_client.get(url=url, headers=headers)
|
||||
if response.status_code == 200:
|
||||
cached_response = json.loads(response.text)
|
||||
cached_response: Final = json.loads(response.text)
|
||||
verbose_logger.debug(
|
||||
"Got GCS Cache: key: %s, cached_response %s. Type Response %s",
|
||||
key,
|
||||
|
|
@ -83,11 +84,11 @@ class GCSCache(BaseCache):
|
|||
|
||||
async def async_get_cache(self, key, **kwargs):
|
||||
try:
|
||||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
headers: Final = self._construct_headers()
|
||||
object_name: Final = self.key_prefix + key
|
||||
bucket_name: Final = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{quote(object_name, safe='')}?alt=media"
|
||||
response = await self.async_client.get(url=url, headers=headers)
|
||||
response: Final = await self.async_client.get(url=url, headers=headers)
|
||||
if response.status_code == 200:
|
||||
return json.loads(response.text)
|
||||
return None
|
||||
|
|
@ -101,7 +102,7 @@ class GCSCache(BaseCache):
|
|||
pass
|
||||
|
||||
async def async_set_cache_pipeline(self, cache_list, **kwargs):
|
||||
tasks = []
|
||||
tasks: Final = []
|
||||
for val in cache_list:
|
||||
tasks.append(self.async_set_cache(val[0], val[1], **kwargs))
|
||||
await asyncio.gather(*tasks)
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ import json
|
|||
import sys
|
||||
import threading
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
|
|
@ -67,7 +67,7 @@ class InMemoryCache(BaseCache):
|
|||
|
||||
# Handle special types without full conversion when possible
|
||||
if hasattr(value, "__sizeof__"): # Use __sizeof__ if available
|
||||
size = value.__sizeof__() / 1024
|
||||
size: Final = value.__sizeof__() / 1024
|
||||
return size <= self.max_size_per_item
|
||||
|
||||
# Fallback for complex types
|
||||
|
|
@ -111,7 +111,7 @@ class InMemoryCache(BaseCache):
|
|||
- 3. the size of in-memory cache is bounded
|
||||
|
||||
"""
|
||||
current_time = time.time()
|
||||
current_time: Final = time.time()
|
||||
|
||||
# Step 1: Remove expired or outdated items
|
||||
while self.expiration_heap:
|
||||
|
|
@ -144,7 +144,7 @@ class InMemoryCache(BaseCache):
|
|||
"""
|
||||
Check if ttl is set for a key
|
||||
"""
|
||||
ttl_time = self.ttl_dict.get(key)
|
||||
ttl_time: Final = self.ttl_dict.get(key)
|
||||
if ttl_time is None or float(ttl_time) < time.time(): # if ttl is not set, allow override
|
||||
return True
|
||||
else:
|
||||
|
|
@ -186,7 +186,7 @@ class InMemoryCache(BaseCache):
|
|||
Add value to set
|
||||
"""
|
||||
# get the value
|
||||
init_value = self.get_cache(key=key) or set()
|
||||
init_value: Final = self.get_cache(key=key) or set()
|
||||
for val in value:
|
||||
init_value.add(val)
|
||||
self.set_cache(key, init_value, ttl=ttl)
|
||||
|
|
@ -207,7 +207,7 @@ class InMemoryCache(BaseCache):
|
|||
if key in self.cache_dict:
|
||||
if self.evict_element_if_expired(key):
|
||||
return None
|
||||
original_cached_response = self.cache_dict[key]
|
||||
original_cached_response: Final = self.cache_dict[key]
|
||||
try:
|
||||
cached_response = json.loads(original_cached_response)
|
||||
except Exception:
|
||||
|
|
@ -216,7 +216,7 @@ class InMemoryCache(BaseCache):
|
|||
return None
|
||||
|
||||
def batch_get_cache(self, keys: list, **kwargs):
|
||||
return_val = []
|
||||
return_val: Final = []
|
||||
for k in keys:
|
||||
val = self.get_cache(key=k, **kwargs)
|
||||
return_val.append(val)
|
||||
|
|
@ -225,7 +225,7 @@ class InMemoryCache(BaseCache):
|
|||
def increment_cache(self, key, value: float, **kwargs) -> float:
|
||||
with self._increment_lock:
|
||||
# keep read-modify-write atomic
|
||||
init_value = self.get_cache(key=key) or 0
|
||||
init_value: Final = self.get_cache(key=key) or 0
|
||||
value = init_value + value
|
||||
self.set_cache(key, value, **kwargs)
|
||||
return value
|
||||
|
|
@ -234,7 +234,7 @@ class InMemoryCache(BaseCache):
|
|||
return self.get_cache(key=key, **kwargs)
|
||||
|
||||
async def async_batch_get_cache(self, keys: list, **kwargs):
|
||||
return_val = []
|
||||
return_val: Final = []
|
||||
for k in keys:
|
||||
val = self.get_cache(key=k, **kwargs)
|
||||
return_val.append(val)
|
||||
|
|
@ -246,7 +246,7 @@ class InMemoryCache(BaseCache):
|
|||
async def async_increment_pipeline(
|
||||
self, increment_list: list["RedisPipelineIncrementOperation"], **kwargs
|
||||
) -> list[float] | None:
|
||||
results = []
|
||||
results: Final = []
|
||||
for increment in increment_list:
|
||||
result = await self.async_increment(increment["key"], increment["increment_value"], **kwargs)
|
||||
results.append(result)
|
||||
|
|
@ -274,5 +274,5 @@ class InMemoryCache(BaseCache):
|
|||
Get the oldest n keys in the cache
|
||||
"""
|
||||
# sorted ttl dict by ttl
|
||||
sorted_ttl_dict = sorted(self.ttl_dict.items(), key=lambda x: x[1])
|
||||
sorted_ttl_dict: Final = sorted(self.ttl_dict.items(), key=lambda x: x[1])
|
||||
return [key for key, _ in sorted_ttl_dict[:n]]
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ Add the event loop to the cache key, to prevent event loop closed errors.
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Final
|
||||
|
||||
from .evicted_client_closer import EvictedClientCloser, default_evicted_client_closer
|
||||
from .in_memory_cache import InMemoryCache
|
||||
|
|
@ -37,7 +38,7 @@ class LLMClientCache(InMemoryCache):
|
|||
self.evicted_client_closer = evicted_client_closer or default_evicted_client_closer
|
||||
|
||||
def _remove_key(self, key: str) -> None:
|
||||
evicted: object = self.cache_dict.get(key)
|
||||
evicted: Final[object] = self.cache_dict.get(key)
|
||||
super()._remove_key(key)
|
||||
self.evicted_client_closer.schedule(evicted)
|
||||
self.evicted_client_closer.reap()
|
||||
|
|
@ -48,8 +49,8 @@ class LLMClientCache(InMemoryCache):
|
|||
If none, use the key as is.
|
||||
"""
|
||||
try:
|
||||
event_loop = asyncio.get_running_loop()
|
||||
stringified_event_loop = str(id(event_loop))
|
||||
event_loop: Final = asyncio.get_running_loop()
|
||||
stringified_event_loop: Final = str(id(event_loop))
|
||||
return f"{key}-{stringified_event_loop}"
|
||||
except RuntimeError: # handle no current running event loop
|
||||
return key
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import ast
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
from typing import Any, cast
|
||||
from typing import Any, Final, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose
|
||||
|
|
@ -88,7 +88,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
|
||||
if quantization_config is None:
|
||||
print_verbose("Quantization config is not provided. Default binary quantization will be used.")
|
||||
collection_exists = self.sync_client.get(
|
||||
collection_exists: Final = self.sync_client.get(
|
||||
url=f"{self.qdrant_api_base}/collections/{self.collection_name}/exists",
|
||||
headers=self.headers,
|
||||
)
|
||||
|
|
@ -124,7 +124,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
else:
|
||||
raise Exception("Quantization config must be one of 'scalar', 'binary' or 'product'")
|
||||
|
||||
new_collection_status = self.sync_client.put(
|
||||
new_collection_status: Final = self.sync_client.put(
|
||||
url=f"{self.qdrant_api_base}/collections/{self.collection_name}",
|
||||
json={
|
||||
"vectors": {"size": self.vector_size, "distance": "Cosine"},
|
||||
|
|
@ -167,7 +167,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
|
||||
def _ensure_cache_key_payload_index(self) -> None:
|
||||
try:
|
||||
response = self.sync_client.put(
|
||||
response: Final = self.sync_client.put(
|
||||
url=f"{self.qdrant_api_base}/collections/{self.collection_name}/index",
|
||||
headers=self.headers,
|
||||
json={
|
||||
|
|
@ -185,7 +185,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
# payload field. Reassigning them to a caller's key would risk
|
||||
# cross-scope hits, so they're treated as misses and re-populated on
|
||||
# the next set_cache.
|
||||
cached_key = payload.get(self.CACHE_KEY_FIELD_NAME)
|
||||
cached_key: Final = payload.get(self.CACHE_KEY_FIELD_NAME)
|
||||
return cached_key is not None and str(cached_key) == str(key)
|
||||
|
||||
def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse:
|
||||
|
|
@ -196,7 +196,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
llm_model_list = None
|
||||
llm_router = None
|
||||
|
||||
router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
if router is not None:
|
||||
return router.embedding(
|
||||
model=self.embedding_model,
|
||||
|
|
@ -217,7 +217,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
llm_model_list = None
|
||||
llm_router = None
|
||||
|
||||
router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
if router is not None:
|
||||
return await router.aembedding(
|
||||
model=self.embedding_model,
|
||||
|
|
@ -237,22 +237,22 @@ class QdrantSemanticCache(BaseCache):
|
|||
from litellm._uuid import uuid
|
||||
|
||||
# get the prompt
|
||||
messages = kwargs["messages"]
|
||||
prompt = get_str_from_messages(messages)
|
||||
messages: Final = kwargs["messages"]
|
||||
prompt: Final = get_str_from_messages(messages)
|
||||
|
||||
# create an embedding for prompt
|
||||
embedding_response = cast(
|
||||
embedding_response: Final = cast(
|
||||
EmbeddingResponse,
|
||||
self._get_embedding(prompt, metadata=kwargs.get("metadata")),
|
||||
)
|
||||
|
||||
# get the embedding
|
||||
embedding = embedding_response["data"][0]["embedding"]
|
||||
embedding: Final = embedding_response["data"][0]["embedding"]
|
||||
|
||||
value = str(value)
|
||||
assert isinstance(value, str)
|
||||
|
||||
data = {
|
||||
data: Final = {
|
||||
"points": [
|
||||
{
|
||||
"id": str(uuid.uuid4()),
|
||||
|
|
@ -275,19 +275,19 @@ class QdrantSemanticCache(BaseCache):
|
|||
print_verbose(f"sync qdrant semantic-cache get_cache, kwargs: {kwargs}")
|
||||
|
||||
# get the messages
|
||||
messages = kwargs["messages"]
|
||||
prompt = get_str_from_messages(messages)
|
||||
messages: Final = kwargs["messages"]
|
||||
prompt: Final = get_str_from_messages(messages)
|
||||
|
||||
# convert to embedding
|
||||
embedding_response = cast(
|
||||
embedding_response: Final = cast(
|
||||
EmbeddingResponse,
|
||||
self._get_embedding(prompt, metadata=kwargs.get("metadata")),
|
||||
)
|
||||
|
||||
# get the embedding
|
||||
embedding = embedding_response["data"][0]["embedding"]
|
||||
embedding: Final = embedding_response["data"][0]["embedding"]
|
||||
|
||||
data = {
|
||||
data: Final = {
|
||||
"vector": embedding,
|
||||
"params": {
|
||||
"quantization": {
|
||||
|
|
@ -301,12 +301,12 @@ class QdrantSemanticCache(BaseCache):
|
|||
}
|
||||
self._add_cache_key_filter_to_search_data(data=data, key=key)
|
||||
|
||||
search_response = self.sync_client.post(
|
||||
search_response: Final = self.sync_client.post(
|
||||
url=f"{self.qdrant_api_base}/collections/{self.collection_name}/points/search",
|
||||
headers=self.headers,
|
||||
json=data,
|
||||
)
|
||||
results = search_response.json()["result"]
|
||||
results: Final = search_response.json()["result"]
|
||||
|
||||
if results is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
|
|
@ -316,14 +316,14 @@ class QdrantSemanticCache(BaseCache):
|
|||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
similarity = results[0]["score"]
|
||||
payload = results[0]["payload"]
|
||||
similarity: Final = results[0]["score"]
|
||||
payload: Final = results[0]["payload"]
|
||||
if not self._payload_matches_cache_key(payload=payload, key=key):
|
||||
print_verbose("Qdrant semantic-cache hit did not match cache key scope")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
cached_prompt = payload["text"]
|
||||
cached_prompt: Final = payload["text"]
|
||||
|
||||
# check similarity, if more than self.similarity_threshold, return results
|
||||
print_verbose(
|
||||
|
|
@ -335,7 +335,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
|
||||
if similarity >= self.similarity_threshold:
|
||||
# cache hit !
|
||||
cached_value = payload["response"]
|
||||
cached_value: Final = payload["response"]
|
||||
print_verbose(
|
||||
f"got a cache hit, similarity: {similarity}, Current prompt: {prompt}, cached_prompt: {cached_prompt}"
|
||||
)
|
||||
|
|
@ -350,17 +350,17 @@ class QdrantSemanticCache(BaseCache):
|
|||
print_verbose(f"async qdrant semantic-cache set_cache, kwargs: {kwargs}")
|
||||
|
||||
# get the prompt
|
||||
messages = kwargs["messages"]
|
||||
prompt = get_str_from_messages(messages)
|
||||
embedding_response = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
messages: Final = kwargs["messages"]
|
||||
prompt: Final = get_str_from_messages(messages)
|
||||
embedding_response: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
# get the embedding
|
||||
embedding = embedding_response["data"][0]["embedding"]
|
||||
embedding: Final = embedding_response["data"][0]["embedding"]
|
||||
|
||||
value = str(value)
|
||||
assert isinstance(value, str)
|
||||
|
||||
data = {
|
||||
data: Final = {
|
||||
"points": [
|
||||
{
|
||||
"id": str(uuid.uuid4()),
|
||||
|
|
@ -384,15 +384,15 @@ class QdrantSemanticCache(BaseCache):
|
|||
print_verbose(f"async qdrant semantic-cache get_cache, kwargs: {kwargs}")
|
||||
|
||||
# get the messages
|
||||
messages = kwargs["messages"]
|
||||
prompt = get_str_from_messages(messages)
|
||||
messages: Final = kwargs["messages"]
|
||||
prompt: Final = get_str_from_messages(messages)
|
||||
|
||||
embedding_response = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
embedding_response: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
# get the embedding
|
||||
embedding = embedding_response["data"][0]["embedding"]
|
||||
embedding: Final = embedding_response["data"][0]["embedding"]
|
||||
|
||||
data = {
|
||||
data: Final = {
|
||||
"vector": embedding,
|
||||
"params": {
|
||||
"quantization": {
|
||||
|
|
@ -406,13 +406,13 @@ class QdrantSemanticCache(BaseCache):
|
|||
}
|
||||
self._add_cache_key_filter_to_search_data(data=data, key=key)
|
||||
|
||||
search_response = await self.async_client.post(
|
||||
search_response: Final = await self.async_client.post(
|
||||
url=f"{self.qdrant_api_base}/collections/{self.collection_name}/points/search",
|
||||
headers=self.headers,
|
||||
json=data,
|
||||
)
|
||||
|
||||
results = search_response.json()["result"]
|
||||
results: Final = search_response.json()["result"]
|
||||
|
||||
if results is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
|
|
@ -422,14 +422,14 @@ class QdrantSemanticCache(BaseCache):
|
|||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
similarity = results[0]["score"]
|
||||
payload = results[0]["payload"]
|
||||
similarity: Final = results[0]["score"]
|
||||
payload: Final = results[0]["payload"]
|
||||
if not self._payload_matches_cache_key(payload=payload, key=key):
|
||||
print_verbose("Qdrant semantic-cache hit did not match cache key scope")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
cached_prompt = payload["text"]
|
||||
cached_prompt: Final = payload["text"]
|
||||
|
||||
# check similarity, if more than self.similarity_threshold, return results
|
||||
print_verbose(
|
||||
|
|
@ -441,7 +441,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
|
||||
if similarity >= self.similarity_threshold:
|
||||
# cache hit !
|
||||
cached_value = payload["response"]
|
||||
cached_value: Final = payload["response"]
|
||||
print_verbose(
|
||||
f"got a cache hit, similarity: {similarity}, Current prompt: {prompt}, cached_prompt: {cached_prompt}"
|
||||
)
|
||||
|
|
@ -454,7 +454,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
return self.collection_info
|
||||
|
||||
async def async_set_cache_pipeline(self, cache_list, **kwargs):
|
||||
tasks = []
|
||||
tasks: Final = []
|
||||
for val in cache_list:
|
||||
tasks.append(self.async_set_cache(val[0], val[1], **kwargs))
|
||||
await asyncio.gather(*tasks)
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ import time
|
|||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from contextvars import ContextVar
|
||||
from datetime import timedelta
|
||||
from typing import TYPE_CHECKING, Any, TypeVar, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, TypeVar, Union, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
|
|
@ -69,18 +69,18 @@ def _get_call_stack_info(num_frames: int = 2) -> str:
|
|||
A string with format "current_function <- caller_function [<- grandparent_function]"
|
||||
"""
|
||||
try:
|
||||
current_frame = inspect.currentframe()
|
||||
current_frame: Final = inspect.currentframe()
|
||||
if current_frame is None:
|
||||
return "unknown"
|
||||
|
||||
# Skip this function and the immediate caller (which sets call_type)
|
||||
f_back = current_frame.f_back
|
||||
f_back: Final = current_frame.f_back
|
||||
if f_back is None:
|
||||
return "unknown"
|
||||
frame = f_back.f_back
|
||||
if frame is None:
|
||||
return "unknown"
|
||||
function_names = []
|
||||
function_names: Final = []
|
||||
|
||||
for _ in range(num_frames):
|
||||
if frame is None:
|
||||
|
|
@ -172,7 +172,7 @@ class RedisCircuitBreaker:
|
|||
_RedisCallResult = TypeVar("_RedisCallResult")
|
||||
|
||||
|
||||
_swallowed_redis_failures: ContextVar[int] = ContextVar("litellm_swallowed_redis_failures", default=0)
|
||||
_swallowed_redis_failures: Final[ContextVar[int]] = ContextVar("litellm_swallowed_redis_failures", default=0)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
|
|
@ -230,9 +230,9 @@ async def _run_under_circuit_breaker(
|
|||
"""
|
||||
if breaker.is_open():
|
||||
raise Exception(f"Redis circuit breaker is open — skipping {name}")
|
||||
swallowed_before = _swallowed_redis_failures.get()
|
||||
swallowed_before: Final = _swallowed_redis_failures.get()
|
||||
try:
|
||||
result = await call()
|
||||
result: Final = await call()
|
||||
except Exception as e:
|
||||
if _is_redis_health_failure(e):
|
||||
breaker.record_failure()
|
||||
|
|
@ -282,7 +282,7 @@ class RedisCache(BaseCache):
|
|||
|
||||
from .._redis import get_redis_client, get_redis_connection_pool
|
||||
|
||||
redis_kwargs = {}
|
||||
redis_kwargs: Final = {}
|
||||
if host is not None:
|
||||
redis_kwargs["host"] = host
|
||||
if port is not None:
|
||||
|
|
@ -363,9 +363,9 @@ class RedisCache(BaseCache):
|
|||
def _handle_async_ping_error(self, e: Exception):
|
||||
"""Handle async ping error with service failure hook."""
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
start_time = time.time()
|
||||
end_time = start_time
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
start_time: Final = time.time()
|
||||
end_time: Final = start_time
|
||||
loop.create_task(
|
||||
self.service_logger_obj.async_service_failure_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
|
|
@ -380,9 +380,9 @@ class RedisCache(BaseCache):
|
|||
def _handle_sync_ping_error(self, e: Exception):
|
||||
"""Handle sync ping error with service failure hook."""
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
start_time = time.time()
|
||||
end_time = start_time
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
start_time: Final = time.time()
|
||||
end_time: Final = start_time
|
||||
loop.create_task(
|
||||
self.service_logger_obj.async_service_failure_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
|
|
@ -401,9 +401,9 @@ class RedisCache(BaseCache):
|
|||
"""
|
||||
# Create a stable representation of redis_kwargs for hashing
|
||||
# Sort keys to ensure consistent hash regardless of parameter order
|
||||
sorted_kwargs = sorted(self.redis_kwargs.items())
|
||||
kwargs_str = json.dumps(sorted_kwargs, sort_keys=True)
|
||||
kwargs_hash = hashlib.sha256(kwargs_str.encode()).hexdigest()[:16]
|
||||
sorted_kwargs: Final = sorted(self.redis_kwargs.items())
|
||||
kwargs_str: Final = json.dumps(sorted_kwargs, sort_keys=True)
|
||||
kwargs_hash: Final = hashlib.sha256(kwargs_str.encode()).hexdigest()[:16]
|
||||
return f"async-redis-client-{kwargs_hash}"
|
||||
|
||||
def init_async_client(
|
||||
|
|
@ -413,8 +413,8 @@ class RedisCache(BaseCache):
|
|||
|
||||
from .._redis import get_redis_async_client, get_redis_connection_pool
|
||||
|
||||
cache_key = self._get_async_client_cache_key()
|
||||
cached_client = in_memory_llm_clients_cache.get_cache(key=cache_key)
|
||||
cache_key: Final = self._get_async_client_cache_key()
|
||||
cached_client: Final = in_memory_llm_clients_cache.get_cache(key=cache_key)
|
||||
if cached_client is not None:
|
||||
redis_async_client = cast(async_redis_client | async_redis_cluster_client, cached_client)
|
||||
else:
|
||||
|
|
@ -454,7 +454,7 @@ class RedisCache(BaseCache):
|
|||
return DEFAULT_REDIS_MAJOR_VERSION
|
||||
|
||||
try:
|
||||
version_str = str(self.redis_version).strip()
|
||||
version_str: Final = str(self.redis_version).strip()
|
||||
# Handle cases where there's no dot (e.g., "7" or 7)
|
||||
if "." in version_str:
|
||||
major_version = int(version_str.split(".")[0])
|
||||
|
|
@ -467,14 +467,14 @@ class RedisCache(BaseCache):
|
|||
return DEFAULT_REDIS_MAJOR_VERSION
|
||||
|
||||
def set_cache(self, key, value, **kwargs):
|
||||
ttl = self.get_ttl(**kwargs)
|
||||
ttl: Final = self.get_ttl(**kwargs)
|
||||
print_verbose(f"Set Redis Cache: key: {key}\nValue {value}\nttl={ttl}, redis_version={self.redis_version}")
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
try:
|
||||
start_time = time.time()
|
||||
start_time: Final = time.time()
|
||||
self.redis_client.set(name=key, value=str(value), ex=ttl)
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
end_time: Final = time.time()
|
||||
_duration: Final = end_time - start_time
|
||||
self.service_logger_obj.service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
|
|
@ -487,13 +487,13 @@ class RedisCache(BaseCache):
|
|||
print_verbose(f"litellm.caching.caching: set() - Got exception from REDIS : {e}")
|
||||
|
||||
def increment_cache(self, key, value: int, ttl: float | None = None, **kwargs) -> int:
|
||||
_redis_client = self.redis_client
|
||||
_redis_client: Final = self.redis_client
|
||||
start_time = time.time()
|
||||
set_ttl = self.get_ttl(ttl=ttl)
|
||||
set_ttl: Final = self.get_ttl(ttl=ttl)
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
try:
|
||||
start_time = time.time()
|
||||
result: int = _redis_client.incr(name=key, amount=value) # type: ignore
|
||||
result: Final[int] = _redis_client.incr(name=key, amount=value) # type: ignore
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
self.service_logger_obj.service_success_hook(
|
||||
|
|
@ -507,7 +507,7 @@ class RedisCache(BaseCache):
|
|||
if set_ttl is not None:
|
||||
# check if key already has ttl, if not -> set ttl
|
||||
start_time = time.time()
|
||||
current_ttl = _redis_client.ttl(key)
|
||||
current_ttl: Final = _redis_client.ttl(key)
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
self.service_logger_obj.service_success_hook(
|
||||
|
|
@ -544,10 +544,10 @@ class RedisCache(BaseCache):
|
|||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_scan_iter(self, pattern: str, count: int = 100) -> list:
|
||||
start_time = time.time()
|
||||
start_time: Final = time.time()
|
||||
try:
|
||||
keys = []
|
||||
_redis_client = self.init_async_client()
|
||||
keys: Final = []
|
||||
_redis_client: Final = self.init_async_client()
|
||||
if not hasattr(_redis_client, "scan_iter"):
|
||||
verbose_logger.debug(
|
||||
"Redis client does not support scan_iter, potentially using Redis Cluster. Returning empty list."
|
||||
|
|
@ -620,7 +620,7 @@ class RedisCache(BaseCache):
|
|||
# different key prefixes never share an executor; in_memory_llm_clients_cache
|
||||
# then adds the running loop, completing the per-(client, namespace, loop)
|
||||
# scoping.
|
||||
script_cache_key = (
|
||||
script_cache_key: Final = (
|
||||
f"redis-registered-script-{self._get_async_client_cache_key()}-"
|
||||
f"{self.namespace}-{hashlib.sha256(script.encode()).hexdigest()[:16]}"
|
||||
)
|
||||
|
|
@ -646,21 +646,21 @@ class RedisCache(BaseCache):
|
|||
Kept separate from async_register_script so each loop caches its own
|
||||
executor; see that method for why the binding must be per loop.
|
||||
"""
|
||||
_redis_client: Any = self.init_async_client()
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
if hasattr(_redis_client, "register_script"):
|
||||
registered_script = _redis_client.register_script(script)
|
||||
registered_script: Final = _redis_client.register_script(script)
|
||||
|
||||
async def standalone_executor(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any:
|
||||
namespaced_keys = tuple(self.check_and_fix_namespace(key=key) for key in keys)
|
||||
namespaced_keys: Final = tuple(self.check_and_fix_namespace(key=key) for key in keys)
|
||||
return await registered_script(keys=namespaced_keys, args=args, client=client)
|
||||
|
||||
return standalone_executor
|
||||
|
||||
if hasattr(_redis_client, "script_load"):
|
||||
script_sha = _redis_client.script_load(script)
|
||||
script_sha: Final = _redis_client.script_load(script)
|
||||
|
||||
async def cluster_executor(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any:
|
||||
namespaced_keys = tuple(self.check_and_fix_namespace(key=key) for key in keys)
|
||||
namespaced_keys: Final = tuple(self.check_and_fix_namespace(key=key) for key in keys)
|
||||
return await _redis_client.evalsha(script_sha, len(namespaced_keys), *namespaced_keys, *args)
|
||||
|
||||
return cluster_executor
|
||||
|
|
@ -678,9 +678,9 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
return None
|
||||
|
||||
start_time = time.time()
|
||||
start_time: Final = time.time()
|
||||
try:
|
||||
_redis_client: Redis = self.init_async_client() # type: ignore
|
||||
_redis_client: Final[Redis] = self.init_async_client() # type: ignore
|
||||
except Exception as e:
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
|
|
@ -704,14 +704,14 @@ class RedisCache(BaseCache):
|
|||
raise e
|
||||
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
ttl = self.get_ttl(**kwargs)
|
||||
nx = kwargs.get("nx", False)
|
||||
ttl: Final = self.get_ttl(**kwargs)
|
||||
nx: Final = kwargs.get("nx", False)
|
||||
print_verbose(f"Set ASYNC Redis Cache: key: {key}\nValue {value}\nttl={ttl}")
|
||||
|
||||
try:
|
||||
if not hasattr(_redis_client, "set"):
|
||||
raise Exception("Redis client cannot set cache. Attribute not found.")
|
||||
result = await _redis_client.set(
|
||||
result: Final = await _redis_client.set(
|
||||
name=key,
|
||||
value=json.dumps(value),
|
||||
nx=nx,
|
||||
|
|
@ -779,7 +779,7 @@ class RedisCache(BaseCache):
|
|||
ex=_td,
|
||||
)
|
||||
# Execute the pipeline and return the results.
|
||||
results = await pipe.execute()
|
||||
results: Final = await pipe.execute()
|
||||
return results
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
|
|
@ -791,14 +791,14 @@ class RedisCache(BaseCache):
|
|||
if len(cache_list) == 0:
|
||||
return
|
||||
|
||||
_redis_client = self.init_async_client()
|
||||
start_time = time.time()
|
||||
_redis_client: Final = self.init_async_client()
|
||||
start_time: Final = time.time()
|
||||
|
||||
print_verbose(f"Set Async Redis Cache: key list: {cache_list}\nttl={ttl}, redis_version={self.redis_version}")
|
||||
cache_value: Any = None
|
||||
cache_value: Final[Any] = None
|
||||
try:
|
||||
async with _redis_client.pipeline(transaction=False) as pipe:
|
||||
results = await self._pipeline_helper(pipe, cache_list, ttl)
|
||||
results: Final = await self._pipeline_helper(pipe, cache_list, ttl)
|
||||
|
||||
print_verbose(f"pipeline results: {results}")
|
||||
# Optionally, you could process 'results' to make sure that all set operations were successful.
|
||||
|
|
@ -851,7 +851,7 @@ class RedisCache(BaseCache):
|
|||
try:
|
||||
await redis_client.sadd(key, *value) # type: ignore
|
||||
if ttl is not None:
|
||||
_td = timedelta(seconds=ttl)
|
||||
_td: Final = timedelta(seconds=ttl)
|
||||
await redis_client.expire(key, _td)
|
||||
except Exception:
|
||||
raise
|
||||
|
|
@ -860,9 +860,9 @@ class RedisCache(BaseCache):
|
|||
async def async_set_cache_sadd(self, key, value: list, ttl: float | None, **kwargs):
|
||||
from redis.asyncio import Redis
|
||||
|
||||
start_time = time.time()
|
||||
start_time: Final = time.time()
|
||||
try:
|
||||
_redis_client: Redis = self.init_async_client() # type: ignore
|
||||
_redis_client: Final[Redis] = self.init_async_client() # type: ignore
|
||||
except Exception as e:
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
|
|
@ -945,17 +945,17 @@ class RedisCache(BaseCache):
|
|||
) -> float:
|
||||
from redis.asyncio import Redis
|
||||
|
||||
_redis_client: Redis = self.init_async_client() # type: ignore
|
||||
start_time = time.time()
|
||||
_used_ttl = self.get_ttl(ttl=ttl)
|
||||
_redis_client: Final[Redis] = self.init_async_client() # type: ignore
|
||||
start_time: Final = time.time()
|
||||
_used_ttl: Final = self.get_ttl(ttl=ttl)
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
try:
|
||||
result = await _redis_client.incrbyfloat(name=key, amount=value)
|
||||
result: Final = await _redis_client.incrbyfloat(name=key, amount=value)
|
||||
if _used_ttl is not None:
|
||||
if refresh_ttl:
|
||||
await _redis_client.expire(key, _used_ttl)
|
||||
else:
|
||||
current_ttl = await _redis_client.ttl(key)
|
||||
current_ttl: Final = await _redis_client.ttl(key)
|
||||
if current_ttl == -1:
|
||||
await _redis_client.expire(key, _used_ttl)
|
||||
|
||||
|
|
@ -1012,10 +1012,10 @@ class RedisCache(BaseCache):
|
|||
GET/compare/SET runs in a single Lua call, so it is also atomic across
|
||||
racing callers and pods. Returns the resulting value.
|
||||
"""
|
||||
_redis_client = self.init_async_client()
|
||||
_used_ttl = self.get_ttl(ttl=ttl)
|
||||
_redis_client: Final = self.init_async_client()
|
||||
_used_ttl: Final = self.get_ttl(ttl=ttl)
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
lua = (
|
||||
lua: Final = (
|
||||
"local cur = redis.call('GET', KEYS[1]) "
|
||||
"if cur == false or tonumber(cur) < tonumber(ARGV[1]) then "
|
||||
"redis.call('SET', KEYS[1], ARGV[1]) "
|
||||
|
|
@ -1056,10 +1056,10 @@ class RedisCache(BaseCache):
|
|||
try:
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
print_verbose(f"Get Redis Cache: key: {key}")
|
||||
start_time = time.time()
|
||||
cached_response = self.redis_client.get(key)
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
start_time: Final = time.time()
|
||||
cached_response: Final = self.redis_client.get(key)
|
||||
end_time: Final = time.time()
|
||||
_duration: Final = end_time - start_time
|
||||
self.service_logger_obj.service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
|
|
@ -1088,7 +1088,7 @@ class RedisCache(BaseCache):
|
|||
|
||||
We use a wrapper so RedisCluster can override this method
|
||||
"""
|
||||
async_redis_client = self.init_async_client()
|
||||
async_redis_client: Final = self.init_async_client()
|
||||
return await async_redis_client.mget(keys=keys) # type: ignore
|
||||
|
||||
def batch_get_cache(
|
||||
|
|
@ -1107,17 +1107,17 @@ class RedisCache(BaseCache):
|
|||
dict: A dictionary mapping keys to their cached values
|
||||
"""
|
||||
key_value_dict = {}
|
||||
_key_list = [key for key in key_list if key is not None]
|
||||
_key_list: Final = [key for key in key_list if key is not None]
|
||||
|
||||
try:
|
||||
_keys = []
|
||||
_keys: Final = []
|
||||
for cache_key in _key_list:
|
||||
cache_key = self.check_and_fix_namespace(key=cache_key or "")
|
||||
_keys.append(cache_key)
|
||||
start_time = time.time()
|
||||
results: list = self._run_redis_mget_operation(keys=_keys)
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
start_time: Final = time.time()
|
||||
results: Final[list] = self._run_redis_mget_operation(keys=_keys)
|
||||
end_time: Final = time.time()
|
||||
_duration: Final = end_time - start_time
|
||||
self.service_logger_obj.service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
|
|
@ -1131,7 +1131,7 @@ class RedisCache(BaseCache):
|
|||
# 'results' is a list of values corresponding to the order of keys in '_key_list'.
|
||||
key_value_dict = dict(zip(_key_list, results))
|
||||
|
||||
decoded_results = {}
|
||||
decoded_results: Final = {}
|
||||
for k, v in key_value_dict.items():
|
||||
if isinstance(k, bytes):
|
||||
k = k.decode("utf-8")
|
||||
|
|
@ -1147,15 +1147,15 @@ class RedisCache(BaseCache):
|
|||
async def async_get_cache(self, key, parent_otel_span: Span | None = None, **kwargs):
|
||||
from redis.asyncio import Redis
|
||||
|
||||
_redis_client: Redis = self.init_async_client() # type: ignore
|
||||
_redis_client: Final[Redis] = self.init_async_client() # type: ignore
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
start_time = time.time()
|
||||
start_time: Final = time.time()
|
||||
|
||||
try:
|
||||
print_verbose(f"Get Async Redis Cache: key: {key}")
|
||||
cached_response = await _redis_client.get(key)
|
||||
cached_response: Final = await _redis_client.get(key)
|
||||
print_verbose(f"Got Async Redis Cache: key: {key}, cached_response {cached_response}")
|
||||
response = self._get_cache_logic(cached_response=cached_response)
|
||||
response: Final = self._get_cache_logic(cached_response=cached_response)
|
||||
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
|
|
@ -1209,14 +1209,14 @@ class RedisCache(BaseCache):
|
|||
"""
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `mget`
|
||||
key_value_dict = {}
|
||||
start_time = time.time()
|
||||
_key_list = [key for key in key_list if key is not None]
|
||||
start_time: Final = time.time()
|
||||
_key_list: Final = [key for key in key_list if key is not None]
|
||||
try:
|
||||
_keys = []
|
||||
_keys: Final = []
|
||||
for cache_key in _key_list:
|
||||
cache_key = self.check_and_fix_namespace(key=cache_key)
|
||||
_keys.append(cache_key)
|
||||
results = await self._async_run_redis_mget_operation(keys=_keys)
|
||||
results: Final = await self._async_run_redis_mget_operation(keys=_keys)
|
||||
## LOGGING ##
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
|
|
@ -1235,7 +1235,7 @@ class RedisCache(BaseCache):
|
|||
# 'results' is a list of values corresponding to the order of keys in 'key_list'.
|
||||
key_value_dict = dict(zip(_key_list, results))
|
||||
|
||||
decoded_results = {}
|
||||
decoded_results: Final = {}
|
||||
for k, v in key_value_dict.items():
|
||||
if isinstance(k, bytes):
|
||||
k = k.decode("utf-8")
|
||||
|
|
@ -1267,9 +1267,9 @@ class RedisCache(BaseCache):
|
|||
Tests if the sync redis client is correctly setup.
|
||||
"""
|
||||
print_verbose("Pinging Sync Redis Cache")
|
||||
start_time = time.time()
|
||||
start_time: Final = time.time()
|
||||
try:
|
||||
response: bool = self.redis_client.ping() # type: ignore
|
||||
response: Final[bool] = self.redis_client.ping() # type: ignore
|
||||
print_verbose(f"Redis Cache PING: {response}")
|
||||
## LOGGING ##
|
||||
end_time = time.time()
|
||||
|
|
@ -1298,11 +1298,11 @@ class RedisCache(BaseCache):
|
|||
|
||||
async def ping(self) -> bool:
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `ping`
|
||||
_redis_client: Any = self.init_async_client()
|
||||
start_time = time.time()
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
start_time: Final = time.time()
|
||||
print_verbose("Pinging Async Redis Cache")
|
||||
try:
|
||||
response = await _redis_client.ping()
|
||||
response: Final = await _redis_client.ping()
|
||||
## LOGGING ##
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
|
|
@ -1333,17 +1333,17 @@ class RedisCache(BaseCache):
|
|||
@_redis_circuit_breaker_guard
|
||||
async def delete_cache_keys(self, keys):
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
|
||||
_redis_client: Any = self.init_async_client()
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
keys = [self.check_and_fix_namespace(key=key) for key in keys]
|
||||
# keys is a list, unpack it so it gets passed as individual elements to delete
|
||||
await _redis_client.delete(*keys)
|
||||
|
||||
def client_list(self) -> list:
|
||||
client_list: list = self.redis_client.client_list() # type: ignore
|
||||
client_list: Final[list] = self.redis_client.client_list() # type: ignore
|
||||
return client_list
|
||||
|
||||
def info(self):
|
||||
info = self.redis_client.info()
|
||||
info: Final = self.redis_client.info()
|
||||
return info
|
||||
|
||||
def flush_cache(self):
|
||||
|
|
@ -1373,10 +1373,10 @@ class RedisCache(BaseCache):
|
|||
import redis.asyncio as redis_async
|
||||
|
||||
# Create a fresh Redis client with current settings
|
||||
redis_client = redis_async.Redis(**self.redis_kwargs)
|
||||
redis_client: Final = redis_async.Redis(**self.redis_kwargs)
|
||||
|
||||
# Test the connection
|
||||
ping_result = await redis_client.ping() # type: ignore[misc]
|
||||
ping_result: Final = await redis_client.ping() # type: ignore[misc]
|
||||
|
||||
# Close the connection
|
||||
await redis_client.aclose() # type: ignore[attr-defined]
|
||||
|
|
@ -1399,7 +1399,7 @@ class RedisCache(BaseCache):
|
|||
@_redis_circuit_breaker_guard
|
||||
async def async_delete_cache(self, key: str):
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
|
||||
_redis_client: Any = self.init_async_client()
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
# keys is str
|
||||
return await _redis_client.delete(key)
|
||||
|
|
@ -1425,7 +1425,7 @@ class RedisCache(BaseCache):
|
|||
_td = timedelta(seconds=increment_op["ttl"])
|
||||
pipe.expire(cache_key, _td)
|
||||
# Execute the pipeline and return results
|
||||
results = await pipe.execute()
|
||||
results: Final = await pipe.execute()
|
||||
# only return float values
|
||||
verbose_logger.debug("Increment ASYNC Redis Cache PIPELINE: results: %s", results)
|
||||
return [r for r in results if isinstance(r, float)]
|
||||
|
|
@ -1448,14 +1448,14 @@ class RedisCache(BaseCache):
|
|||
|
||||
from redis.asyncio import Redis
|
||||
|
||||
_redis_client: Redis = self.init_async_client() # type: ignore
|
||||
start_time = time.time()
|
||||
_redis_client: Final[Redis] = self.init_async_client() # type: ignore
|
||||
start_time: Final = time.time()
|
||||
|
||||
print_verbose(f"Increment Async Redis Cache Pipeline: increment list: {increment_list}")
|
||||
|
||||
try:
|
||||
async with _redis_client.pipeline(transaction=False) as pipe:
|
||||
results = await self._pipeline_increment_helper(pipe, increment_list)
|
||||
results: Final = await self._pipeline_increment_helper(pipe, increment_list)
|
||||
|
||||
## LOGGING ##
|
||||
end_time = time.time()
|
||||
|
|
@ -1507,9 +1507,9 @@ class RedisCache(BaseCache):
|
|||
"""
|
||||
try:
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `ttl`
|
||||
_redis_client: Any = self.init_async_client()
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
ttl = await _redis_client.ttl(key)
|
||||
ttl: Final = await _redis_client.ttl(key)
|
||||
if ttl <= -1: # -1 means the key does not exist, -2 key does not exist
|
||||
return None
|
||||
return ttl
|
||||
|
|
@ -1537,11 +1537,11 @@ class RedisCache(BaseCache):
|
|||
Returns:
|
||||
int: The length of the list after the push operation
|
||||
"""
|
||||
_redis_client: Any = self.init_async_client()
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
start_time = time.time()
|
||||
start_time: Final = time.time()
|
||||
try:
|
||||
response = await _redis_client.rpush(key, *values)
|
||||
response: Final = await _redis_client.rpush(key, *values)
|
||||
## LOGGING ##
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
|
|
@ -1578,7 +1578,7 @@ class RedisCache(BaseCache):
|
|||
for rpush_op in rpush_list:
|
||||
key = self.check_and_fix_namespace(key=rpush_op["key"])
|
||||
pipe.rpush(key, *rpush_op["values"])
|
||||
results = await pipe.execute()
|
||||
results: Final = await pipe.execute()
|
||||
# Preserve positional correspondence — raise on per-command errors
|
||||
for r in results:
|
||||
if isinstance(r, Exception):
|
||||
|
|
@ -1604,12 +1604,12 @@ class RedisCache(BaseCache):
|
|||
if len(rpush_list) == 0:
|
||||
return []
|
||||
|
||||
_redis_client: Any = self.init_async_client()
|
||||
start_time = time.time()
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
start_time: Final = time.time()
|
||||
|
||||
try:
|
||||
async with _redis_client.pipeline(transaction=False) as pipe:
|
||||
results = await self._pipeline_rpush_helper(pipe, rpush_list)
|
||||
results: Final = await self._pipeline_rpush_helper(pipe, rpush_list)
|
||||
|
||||
## LOGGING ##
|
||||
end_time = time.time()
|
||||
|
|
@ -1641,7 +1641,7 @@ class RedisCache(BaseCache):
|
|||
raise e
|
||||
|
||||
async def handle_lpop_count_for_older_redis_versions(self, pipe: pipeline, key: str, count: int) -> list[bytes]:
|
||||
result: list[bytes] = []
|
||||
result: Final[list[bytes]] = []
|
||||
for _ in range(count):
|
||||
pipe.lpop(key)
|
||||
results = await pipe.execute()
|
||||
|
|
@ -1661,12 +1661,12 @@ class RedisCache(BaseCache):
|
|||
parent_otel_span: Span | None = None,
|
||||
**kwargs,
|
||||
) -> Any | list[Any]:
|
||||
_redis_client: Any = self.init_async_client()
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
start_time = time.time()
|
||||
start_time: Final = time.time()
|
||||
print_verbose(f"LPOP from Redis list: key: {key}, count: {count}")
|
||||
try:
|
||||
major_version = self._parse_redis_major_version()
|
||||
major_version: Final = self._parse_redis_major_version()
|
||||
|
||||
if count is not None and major_version < 7:
|
||||
# For Redis < 7.0, use pipeline to execute multiple LPOP commands
|
||||
|
|
@ -1725,7 +1725,7 @@ class RedisCache(BaseCache):
|
|||
For Redis >= 7, queues one LPOP(key, count) per operation.
|
||||
For Redis < 7, queues `count` individual LPOP(key) commands per operation.
|
||||
"""
|
||||
major_version = self._parse_redis_major_version()
|
||||
major_version: Final = self._parse_redis_major_version()
|
||||
|
||||
if major_version >= 7:
|
||||
for lpop_op in lpop_list:
|
||||
|
|
@ -1735,14 +1735,14 @@ class RedisCache(BaseCache):
|
|||
else:
|
||||
# For Redis < 7, LPOP doesn't support count param.
|
||||
# Issue `count` individual LPOP commands per key, all in one pipeline.
|
||||
counts: list[int] = []
|
||||
counts: Final[list[int]] = []
|
||||
for lpop_op in lpop_list:
|
||||
key = self.check_and_fix_namespace(key=lpop_op["key"])
|
||||
count = lpop_op["count"] or 1
|
||||
counts.append(count)
|
||||
for _ in range(count):
|
||||
pipe.lpop(key)
|
||||
flat_results = await pipe.execute()
|
||||
flat_results: Final = await pipe.execute()
|
||||
|
||||
# Re-group the flat results back into per-key lists
|
||||
raw_results = []
|
||||
|
|
@ -1758,7 +1758,7 @@ class RedisCache(BaseCache):
|
|||
raise r
|
||||
|
||||
# Decode bytes -> str for each result set
|
||||
decoded_results: list[list[str] | None] = []
|
||||
decoded_results: Final[list[list[str] | None]] = []
|
||||
for r in raw_results:
|
||||
if r is None:
|
||||
decoded_results.append(None)
|
||||
|
|
@ -1793,12 +1793,12 @@ class RedisCache(BaseCache):
|
|||
if len(lpop_list) == 0:
|
||||
return []
|
||||
|
||||
_redis_client: Any = self.init_async_client()
|
||||
start_time = time.time()
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
start_time: Final = time.time()
|
||||
|
||||
try:
|
||||
async with _redis_client.pipeline(transaction=False) as pipe:
|
||||
results = await self._pipeline_lpop_helper(pipe, lpop_list)
|
||||
results: Final = await self._pipeline_lpop_helper(pipe, lpop_list)
|
||||
|
||||
## LOGGING ##
|
||||
end_time = time.time()
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Key differences:
|
|||
- RedisClient NEEDs to be re-used across requests, adds 3000ms latency if it's re-created
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Union
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
|
|
@ -37,7 +37,7 @@ class RedisClusterCache(RedisCache):
|
|||
if self.redis_async_redis_cluster_client:
|
||||
return self.redis_async_redis_cluster_client
|
||||
|
||||
_redis_client = get_redis_async_client(connection_pool=self.async_redis_conn_pool, **self.redis_kwargs)
|
||||
_redis_client: Final = get_redis_async_client(connection_pool=self.async_redis_conn_pool, **self.redis_kwargs)
|
||||
if isinstance(_redis_client, RedisCluster):
|
||||
self.redis_async_redis_cluster_client = _redis_client
|
||||
|
||||
|
|
@ -53,7 +53,7 @@ class RedisClusterCache(RedisCache):
|
|||
"""
|
||||
Overrides `_async_run_redis_mget_operation` in redis_cache.py
|
||||
"""
|
||||
async_redis_cluster_client = self.init_async_client()
|
||||
async_redis_cluster_client: Final = self.init_async_client()
|
||||
return await async_redis_cluster_client.mget_nonatomic(keys=keys) # type: ignore
|
||||
|
||||
async def test_connection(self) -> dict:
|
||||
|
|
@ -68,21 +68,21 @@ class RedisClusterCache(RedisCache):
|
|||
from redis.cluster import ClusterNode
|
||||
|
||||
# Create ClusterNode objects from startup_nodes
|
||||
cluster_kwargs = self.redis_kwargs.copy()
|
||||
startup_nodes = cluster_kwargs.pop("startup_nodes", [])
|
||||
cluster_kwargs: Final = self.redis_kwargs.copy()
|
||||
startup_nodes: Final = cluster_kwargs.pop("startup_nodes", [])
|
||||
|
||||
new_startup_nodes: list[ClusterNode] = []
|
||||
new_startup_nodes: Final[list[ClusterNode]] = []
|
||||
for item in startup_nodes:
|
||||
new_startup_nodes.append(ClusterNode(**item))
|
||||
|
||||
# Create a fresh Redis Cluster client with current settings
|
||||
redis_client = redis_async.RedisCluster(
|
||||
redis_client: Final = redis_async.RedisCluster(
|
||||
startup_nodes=new_startup_nodes,
|
||||
**cluster_kwargs, # type: ignore
|
||||
)
|
||||
|
||||
# Test the connection
|
||||
ping_result = await redis_client.ping() # type: ignore[attr-defined, misc]
|
||||
ping_result: Final = await redis_client.ping() # type: ignore[attr-defined, misc]
|
||||
|
||||
# Close the connection
|
||||
await redis_client.aclose() # type: ignore[attr-defined]
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ import ast
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
from typing import Any, cast
|
||||
from typing import Any, Final, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
|
|
@ -95,7 +95,7 @@ class RedisSemanticCache(BaseCache):
|
|||
password = password or os.environ["REDIS_PASSWORD"]
|
||||
except KeyError as e:
|
||||
# Raise a more informative exception if any of the required keys are missing
|
||||
missing_var = e.args[0]
|
||||
missing_var: Final = e.args[0]
|
||||
raise ValueError(
|
||||
f"Missing required Redis configuration: {missing_var}. Provide {missing_var} or redis_url."
|
||||
) from e
|
||||
|
|
@ -130,7 +130,7 @@ class RedisSemanticCache(BaseCache):
|
|||
from redisvl.utils.vectorize import CustomTextVectorizer # type: ignore[import-not-found, import-untyped]
|
||||
|
||||
try:
|
||||
cache_vectorizer = CustomTextVectorizer(self._get_embedding)
|
||||
cache_vectorizer: Final = CustomTextVectorizer(self._get_embedding)
|
||||
return self._init_semantic_cache(
|
||||
semantic_cache_cls=SemanticCache,
|
||||
index_name=self._index_name,
|
||||
|
|
@ -156,7 +156,7 @@ class RedisSemanticCache(BaseCache):
|
|||
cache_vectorizer: Any,
|
||||
) -> Any:
|
||||
def _is_schema_mismatch(exc: ValueError) -> bool:
|
||||
error_message = str(exc).lower()
|
||||
error_message: Final = str(exc).lower()
|
||||
return any(phrase in error_message for phrase in ("schema does not match", "index schema"))
|
||||
|
||||
try:
|
||||
|
|
@ -172,7 +172,7 @@ class RedisSemanticCache(BaseCache):
|
|||
if not _is_schema_mismatch(exc):
|
||||
raise
|
||||
|
||||
isolated_index_name = f"{index_name}_isolated"
|
||||
isolated_index_name: Final = f"{index_name}_isolated"
|
||||
print_verbose(
|
||||
"Redis semantic-cache existing index schema is not isolated; "
|
||||
f"using isolated index - {isolated_index_name}"
|
||||
|
|
@ -239,16 +239,16 @@ class RedisSemanticCache(BaseCache):
|
|||
"""
|
||||
Extract a semantic-cache prompt from chat or Responses API request kwargs.
|
||||
"""
|
||||
messages = kwargs.get("messages")
|
||||
messages: Final = kwargs.get("messages")
|
||||
if messages:
|
||||
return get_str_from_messages(messages)
|
||||
|
||||
if "input" not in kwargs:
|
||||
return None
|
||||
|
||||
prompt_parts: list[str] = []
|
||||
prompt_parts: Final[list[str]] = []
|
||||
cls._collect_responses_input_text(kwargs.get("input"), prompt_parts)
|
||||
prompt = "\n".join(prompt_parts).strip()
|
||||
prompt: Final = "\n".join(prompt_parts).strip()
|
||||
return prompt or None
|
||||
|
||||
@classmethod
|
||||
|
|
@ -258,7 +258,7 @@ class RedisSemanticCache(BaseCache):
|
|||
return
|
||||
|
||||
if isinstance(value, str):
|
||||
stripped_value = value.strip()
|
||||
stripped_value: Final = value.strip()
|
||||
if stripped_value:
|
||||
prompt_parts.append(stripped_value)
|
||||
return
|
||||
|
|
@ -298,10 +298,10 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
@staticmethod
|
||||
def _coerce_response_input_value(value: Any) -> Any:
|
||||
model_dump = getattr(value, "model_dump", None)
|
||||
model_dump: Final = getattr(value, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
return model_dump()
|
||||
dict_method = getattr(value, "dict", None)
|
||||
dict_method: Final = getattr(value, "dict", None)
|
||||
if callable(dict_method):
|
||||
return dict_method()
|
||||
return value
|
||||
|
|
@ -318,7 +318,7 @@ class RedisSemanticCache(BaseCache):
|
|||
llm_model_list = None
|
||||
llm_router = None
|
||||
|
||||
router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
if router is not None:
|
||||
embedding_response = cast(
|
||||
EmbeddingResponse,
|
||||
|
|
@ -383,22 +383,22 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
value_str: str | None = None
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
value_str = str(value)
|
||||
|
||||
prompt_embedding = self._get_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
prompt_embedding: Final = self._get_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
store_kwargs: dict[str, Any] = {
|
||||
store_kwargs: Final[dict[str, Any]] = {
|
||||
"vector": prompt_embedding,
|
||||
"filters": self._get_cache_filters(key),
|
||||
}
|
||||
|
||||
# Get TTL and store in Redis semantic cache
|
||||
ttl = self._get_ttl(**kwargs)
|
||||
ttl: Final = self._get_ttl(**kwargs)
|
||||
if ttl is not None:
|
||||
store_kwargs["ttl"] = int(ttl)
|
||||
self.llmcache.store(prompt, value_str, **store_kwargs)
|
||||
|
|
@ -419,7 +419,7 @@ class RedisSemanticCache(BaseCache):
|
|||
print_verbose(f"Redis semantic-cache get_cache, kwargs: {kwargs}")
|
||||
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic cache lookup")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
|
|
@ -427,13 +427,13 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
# Check the cache for semantically similar prompts in this exact
|
||||
# LiteLLM cache-key scope.
|
||||
prompt_embedding = self._get_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
check_kwargs: dict[str, Any] = {
|
||||
prompt_embedding: Final = self._get_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
check_kwargs: Final[dict[str, Any]] = {
|
||||
"prompt": prompt,
|
||||
"vector": prompt_embedding,
|
||||
"filter_expression": self._get_cache_key_filter_expression(key),
|
||||
}
|
||||
results = self.llmcache.check(**check_kwargs)
|
||||
results: Final = self.llmcache.check(**check_kwargs)
|
||||
|
||||
# Return None if no similar prompts found
|
||||
if not results:
|
||||
|
|
@ -441,20 +441,20 @@ class RedisSemanticCache(BaseCache):
|
|||
return None
|
||||
|
||||
# Process the best matching result
|
||||
cache_hit = results[0]
|
||||
cache_hit: Final = results[0]
|
||||
if not self._cache_hit_matches_key(cache_hit=cache_hit, key=key):
|
||||
print_verbose("Redis semantic-cache hit did not match cache key scope")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
vector_distance = float(cache_hit["vector_distance"])
|
||||
vector_distance: Final = float(cache_hit["vector_distance"])
|
||||
|
||||
# Convert vector distance back to similarity score
|
||||
# For cosine distance: 0 = most similar, 2 = least similar
|
||||
# While similarity: 1 = most similar, 0 = least similar
|
||||
similarity = 1 - vector_distance
|
||||
similarity: Final = 1 - vector_distance
|
||||
|
||||
cached_prompt = cache_hit["prompt"]
|
||||
cached_response = cache_hit["response"]
|
||||
cached_prompt: Final = cache_hit["prompt"]
|
||||
cached_response: Final = cache_hit["response"]
|
||||
|
||||
# update kwargs["metadata"] with similarity, don't rewrite the original metadata
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity
|
||||
|
|
@ -488,7 +488,7 @@ class RedisSemanticCache(BaseCache):
|
|||
llm_model_list = None
|
||||
llm_router = None
|
||||
|
||||
router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
try:
|
||||
if router is not None:
|
||||
embedding_response = await router.aembedding(
|
||||
|
|
@ -521,23 +521,23 @@ class RedisSemanticCache(BaseCache):
|
|||
print_verbose(f"Async Redis semantic-cache set_cache, kwargs: {kwargs}")
|
||||
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
value_str = str(value)
|
||||
value_str: Final = str(value)
|
||||
|
||||
# Generate embedding for the value (response) to cache
|
||||
prompt_embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
prompt_embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
store_kwargs: dict[str, Any] = {
|
||||
store_kwargs: Final[dict[str, Any]] = {
|
||||
"vector": prompt_embedding,
|
||||
"filters": self._get_cache_filters(key),
|
||||
}
|
||||
|
||||
# Get TTL and store in Redis semantic cache
|
||||
ttl = self._get_ttl(**kwargs)
|
||||
ttl: Final = self._get_ttl(**kwargs)
|
||||
if ttl is not None:
|
||||
store_kwargs["ttl"] = ttl
|
||||
await self.llmcache.astore(
|
||||
|
|
@ -562,43 +562,43 @@ class RedisSemanticCache(BaseCache):
|
|||
print_verbose(f"Async Redis semantic-cache get_cache, kwargs: {kwargs}")
|
||||
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic cache lookup")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
# Generate embedding for the prompt
|
||||
prompt_embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
prompt_embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
# Check the cache for semantically similar prompts in this exact
|
||||
# LiteLLM cache-key scope.
|
||||
check_kwargs: dict[str, Any] = {
|
||||
check_kwargs: Final[dict[str, Any]] = {
|
||||
"prompt": prompt,
|
||||
"vector": prompt_embedding,
|
||||
"filter_expression": self._get_cache_key_filter_expression(key),
|
||||
}
|
||||
results = await self.llmcache.acheck(**check_kwargs)
|
||||
results: Final = await self.llmcache.acheck(**check_kwargs)
|
||||
|
||||
# handle results / cache hit
|
||||
if not results:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
cache_hit = results[0]
|
||||
cache_hit: Final = results[0]
|
||||
if not self._cache_hit_matches_key(cache_hit=cache_hit, key=key):
|
||||
print_verbose("Redis semantic-cache hit did not match cache key scope")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
vector_distance = float(cache_hit["vector_distance"])
|
||||
vector_distance: Final = float(cache_hit["vector_distance"])
|
||||
|
||||
# Convert vector distance back to similarity
|
||||
# For cosine distance: 0 = most similar, 2 = least similar
|
||||
# While similarity: 1 = most similar, 0 = least similar
|
||||
similarity = 1 - vector_distance
|
||||
similarity: Final = 1 - vector_distance
|
||||
|
||||
cached_prompt = cache_hit["prompt"]
|
||||
cached_response = cache_hit["response"]
|
||||
cached_prompt: Final = cache_hit["prompt"]
|
||||
cached_response: Final = cache_hit["response"]
|
||||
|
||||
# update kwargs["metadata"] with similarity, don't rewrite the original metadata
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity
|
||||
|
|
@ -622,7 +622,7 @@ class RedisSemanticCache(BaseCache):
|
|||
Returns:
|
||||
Dict[str, Any]: Information about the Redis index
|
||||
"""
|
||||
aindex = await self.llmcache._get_async_index()
|
||||
aindex: Final = await self.llmcache._get_async_index()
|
||||
return await aindex.info()
|
||||
|
||||
async def async_set_cache_pipeline(self, cache_list: list[tuple[str, Any]], **kwargs) -> None:
|
||||
|
|
@ -634,7 +634,7 @@ class RedisSemanticCache(BaseCache):
|
|||
**kwargs: Additional arguments
|
||||
"""
|
||||
try:
|
||||
tasks = []
|
||||
tasks: Final = []
|
||||
for val in cache_list:
|
||||
tasks.append(self.async_set_cache(val[0], val[1], **kwargs))
|
||||
await asyncio.gather(*tasks)
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import asyncio
|
|||
import json
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from functools import partial
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
|
||||
|
|
@ -62,16 +63,16 @@ class S3Cache(BaseCache):
|
|||
def set_cache(self, key, value, **kwargs):
|
||||
try:
|
||||
print_verbose(f"LiteLLM SET Cache - S3. Key={key}. Value={value}")
|
||||
ttl = kwargs.get("ttl", None)
|
||||
ttl: Final = kwargs.get("ttl", None)
|
||||
# Convert value to JSON before storing in S3
|
||||
serialized_value = json.dumps(value)
|
||||
serialized_value: Final = json.dumps(value)
|
||||
key = self._to_s3_key(key)
|
||||
|
||||
if ttl is not None:
|
||||
cache_control = f"immutable, max-age={ttl}, s-maxage={ttl}"
|
||||
|
||||
# Calculate expiration time
|
||||
expiration_time = datetime.now(timezone.utc) + timedelta(seconds=ttl)
|
||||
expiration_time: Final = datetime.now(timezone.utc) + timedelta(seconds=ttl)
|
||||
# Upload the data to S3 with the calculated expiration time
|
||||
self.s3_client.put_object(
|
||||
Bucket=self.bucket_name,
|
||||
|
|
@ -105,8 +106,8 @@ class S3Cache(BaseCache):
|
|||
"""
|
||||
try:
|
||||
verbose_logger.debug("Set ASYNC S3 Cache: Key=%s. Value=%s", key, value)
|
||||
loop = asyncio.get_event_loop()
|
||||
func = partial(self.set_cache, key, value, **kwargs)
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
func: Final = partial(self.set_cache, key, value, **kwargs)
|
||||
await loop.run_in_executor(None, func)
|
||||
except Exception as e:
|
||||
verbose_logger.error("S3 Caching: async_set_cache() - Got exception from S3: %s", e)
|
||||
|
|
@ -123,8 +124,8 @@ class S3Cache(BaseCache):
|
|||
|
||||
if cached_response is not None:
|
||||
if "Expires" in cached_response:
|
||||
expires_time = cached_response["Expires"]
|
||||
current_time = datetime.now(expires_time.tzinfo)
|
||||
expires_time: Final = cached_response["Expires"]
|
||||
current_time: Final = datetime.now(expires_time.tzinfo)
|
||||
|
||||
if current_time > expires_time:
|
||||
return None
|
||||
|
|
@ -160,9 +161,9 @@ class S3Cache(BaseCache):
|
|||
"""
|
||||
try:
|
||||
verbose_logger.debug("Get ASYNC S3 Cache: key: %s", key)
|
||||
loop = asyncio.get_event_loop()
|
||||
func = partial(self.get_cache, key, **kwargs)
|
||||
result = await loop.run_in_executor(None, func)
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
func: Final = partial(self.get_cache, key, **kwargs)
|
||||
result: Final = await loop.run_in_executor(None, func)
|
||||
return result
|
||||
except Exception as e:
|
||||
verbose_logger.error("S3 Caching: async_get_cache() - Got exception from S3: %s", e)
|
||||
|
|
@ -175,7 +176,7 @@ class S3Cache(BaseCache):
|
|||
pass
|
||||
|
||||
async def async_set_cache_pipeline(self, cache_list, **kwargs):
|
||||
tasks = []
|
||||
tasks: Final = []
|
||||
for val in cache_list:
|
||||
tasks.append(self.async_set_cache(val[0], val[1], **kwargs))
|
||||
await asyncio.gather(*tasks)
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ import hashlib
|
|||
import os
|
||||
import struct
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from redis import Redis
|
||||
from redis.asyncio import Redis as AsyncRedis
|
||||
|
|
@ -106,8 +106,8 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
"(or VALKEY_HOST/VALKEY_PORT), or pass redis_url."
|
||||
)
|
||||
|
||||
credentials = f":{password}@" if password else ""
|
||||
scheme = "rediss" if ssl else "redis"
|
||||
credentials: Final = f":{password}@" if password else ""
|
||||
scheme: Final = "rediss" if ssl else "redis"
|
||||
return f"{scheme}://{credentials}{host}:{port}"
|
||||
|
||||
@classmethod
|
||||
|
|
@ -154,7 +154,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
return None
|
||||
|
||||
def _assert_dim_matches(self, info: dict, dim: int) -> None:
|
||||
existing_dim = self._extract_index_dim(info)
|
||||
existing_dim: Final = self._extract_index_dim(info)
|
||||
if existing_dim is not None and existing_dim != dim:
|
||||
raise ValueError(
|
||||
f"Valkey semantic-cache index '{self.index_name}' already exists with "
|
||||
|
|
@ -186,7 +186,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
except Exception as exc:
|
||||
if not self._is_index_exists_error(exc):
|
||||
raise
|
||||
info = await self.async_client.ft(self.index_name).info()
|
||||
info: Final = await self.async_client.ft(self.index_name).info()
|
||||
self._assert_dim_matches(info, dim)
|
||||
self._index_dim = dim
|
||||
|
||||
|
|
@ -202,8 +202,8 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
}
|
||||
|
||||
def _knn_query(self, key: str) -> Query:
|
||||
scope = self._scope_tag(key)
|
||||
query_string = (
|
||||
scope: Final = self._scope_tag(key)
|
||||
query_string: Final = (
|
||||
f"(@{self.CACHE_KEY_FIELD_NAME}:{{{scope}}})"
|
||||
f"=>[KNN 1 @{self.EMBEDDING_FIELD_NAME} $vec AS {self.DISTANCE_FIELD_NAME}]"
|
||||
)
|
||||
|
|
@ -211,10 +211,10 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
|
||||
@classmethod
|
||||
def _first_hit(cls, search_result: Any) -> _ValkeyCacheHit | None:
|
||||
docs = getattr(search_result, "docs", [])
|
||||
docs: Final = getattr(search_result, "docs", [])
|
||||
if not docs:
|
||||
return None
|
||||
doc = docs[0]
|
||||
doc: Final = docs[0]
|
||||
return _ValkeyCacheHit(
|
||||
response=str(getattr(doc, cls.RESPONSE_FIELD_NAME)),
|
||||
distance=float(getattr(doc, cls.DISTANCE_FIELD_NAME)),
|
||||
|
|
@ -225,7 +225,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
similarity = 1 - hit.distance
|
||||
similarity: Final = 1 - hit.distance
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity
|
||||
|
||||
if similarity < self.similarity_threshold:
|
||||
|
|
@ -235,17 +235,17 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
def set_cache(self, key: str, value: Any, **kwargs: Any) -> None:
|
||||
print_verbose(f"Valkey semantic-cache set_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
embedding = self._get_embedding(prompt)
|
||||
embedding: Final = self._get_embedding(prompt)
|
||||
self._ensure_index_sync(len(embedding))
|
||||
|
||||
doc_key = self._doc_key(key)
|
||||
doc_key: Final = self._doc_key(key)
|
||||
self.sync_client.hset(doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding))
|
||||
ttl = self._get_ttl(**kwargs)
|
||||
ttl: Final = self._get_ttl(**kwargs)
|
||||
if ttl is not None:
|
||||
self.sync_client.expire(doc_key, ttl)
|
||||
except Exception as e:
|
||||
|
|
@ -254,15 +254,15 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
def get_cache(self, key: str, **kwargs: Any) -> Any:
|
||||
print_verbose(f"Valkey semantic-cache get_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
embedding = self._get_embedding(prompt)
|
||||
embedding: Final = self._get_embedding(prompt)
|
||||
self._ensure_index_sync(len(embedding))
|
||||
|
||||
search_result = self.sync_client.ft(self.index_name).search(
|
||||
search_result: Final = self.sync_client.ft(self.index_name).search(
|
||||
self._knn_query(key),
|
||||
query_params={"vec": self._embedding_to_bytes(embedding)},
|
||||
)
|
||||
|
|
@ -274,17 +274,17 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
async def async_set_cache(self, key: str, value: Any, **kwargs: Any) -> None:
|
||||
print_verbose(f"Async Valkey semantic-cache set_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
await self._ensure_index_async(len(embedding))
|
||||
|
||||
doc_key = self._doc_key(key)
|
||||
doc_key: Final = self._doc_key(key)
|
||||
await self.async_client.hset(doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding))
|
||||
ttl = self._get_ttl(**kwargs)
|
||||
ttl: Final = self._get_ttl(**kwargs)
|
||||
if ttl is not None:
|
||||
await self.async_client.expire(doc_key, ttl)
|
||||
except Exception as e:
|
||||
|
|
@ -293,15 +293,15 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
async def async_get_cache(self, key: str, **kwargs: Any) -> Any:
|
||||
print_verbose(f"Async Valkey semantic-cache get_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
await self._ensure_index_async(len(embedding))
|
||||
|
||||
search_result = await self.async_client.ft(self.index_name).search(
|
||||
search_result: Final = await self.async_client.ft(self.index_name).search(
|
||||
self._knn_query(key),
|
||||
query_params={"vec": self._embedding_to_bytes(embedding)},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ Handler for transforming /chat/completions api requests to litellm.responses req
|
|||
"""
|
||||
|
||||
from collections.abc import Coroutine
|
||||
from typing import TYPE_CHECKING, Any, Union
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
|
|
@ -60,7 +60,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
raise ValueError("Unexpected responses stream payload")
|
||||
|
||||
if hidden_params:
|
||||
existing = getattr(response, "_hidden_params", None)
|
||||
existing: Final = getattr(response, "_hidden_params", None)
|
||||
if not isinstance(existing, dict) or not existing:
|
||||
setattr(response, "_hidden_params", dict(hidden_params))
|
||||
else:
|
||||
|
|
@ -72,13 +72,13 @@ class ResponsesToCompletionBridgeHandler:
|
|||
for _ in stream_iter:
|
||||
pass
|
||||
|
||||
completed = getattr(stream_iter, "completed_response", None)
|
||||
response_obj = getattr(completed, "response", None) if completed else None
|
||||
completed: Final = getattr(stream_iter, "completed_response", None)
|
||||
response_obj: Final = getattr(completed, "response", None) if completed else None
|
||||
if response_obj is None:
|
||||
raise ValueError("Stream ended without a completed response")
|
||||
|
||||
hidden_params = getattr(stream_iter, "_hidden_params", None)
|
||||
response = self._coerce_response_object(response_obj, hidden_params)
|
||||
hidden_params: Final = getattr(stream_iter, "_hidden_params", None)
|
||||
response: Final = self._coerce_response_object(response_obj, hidden_params)
|
||||
if not isinstance(response, ResponsesAPIResponse):
|
||||
raise ValueError("Stream completed response is invalid")
|
||||
return response
|
||||
|
|
@ -87,13 +87,13 @@ class ResponsesToCompletionBridgeHandler:
|
|||
async for _ in stream_iter:
|
||||
pass
|
||||
|
||||
completed = getattr(stream_iter, "completed_response", None)
|
||||
response_obj = getattr(completed, "response", None) if completed else None
|
||||
completed: Final = getattr(stream_iter, "completed_response", None)
|
||||
response_obj: Final = getattr(completed, "response", None) if completed else None
|
||||
if response_obj is None:
|
||||
raise ValueError("Stream ended without a completed response")
|
||||
|
||||
hidden_params = getattr(stream_iter, "_hidden_params", None)
|
||||
response = self._coerce_response_object(response_obj, hidden_params)
|
||||
hidden_params: Final = getattr(stream_iter, "_hidden_params", None)
|
||||
response: Final = self._coerce_response_object(response_obj, hidden_params)
|
||||
if not isinstance(response, ResponsesAPIResponse):
|
||||
raise ValueError("Stream completed response is invalid")
|
||||
return response
|
||||
|
|
@ -102,35 +102,35 @@ class ResponsesToCompletionBridgeHandler:
|
|||
from litellm import LiteLLMLoggingObj
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
model = kwargs.get("model")
|
||||
model: Final = kwargs.get("model")
|
||||
if model is None or not isinstance(model, str):
|
||||
raise ValueError("model is required")
|
||||
|
||||
custom_llm_provider = kwargs.get("custom_llm_provider")
|
||||
custom_llm_provider: Final = kwargs.get("custom_llm_provider")
|
||||
if custom_llm_provider is None or not isinstance(custom_llm_provider, str):
|
||||
raise ValueError("custom_llm_provider is required")
|
||||
|
||||
messages = kwargs.get("messages")
|
||||
messages: Final = kwargs.get("messages")
|
||||
if messages is None or not isinstance(messages, list):
|
||||
raise ValueError("messages is required")
|
||||
|
||||
optional_params = kwargs.get("optional_params")
|
||||
optional_params: Final = kwargs.get("optional_params")
|
||||
if optional_params is None or not isinstance(optional_params, dict):
|
||||
raise ValueError("optional_params is required")
|
||||
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
litellm_params: Final = kwargs.get("litellm_params")
|
||||
if litellm_params is None or not isinstance(litellm_params, dict):
|
||||
raise ValueError("litellm_params is required")
|
||||
|
||||
headers = kwargs.get("headers")
|
||||
headers: Final = kwargs.get("headers")
|
||||
if headers is None or not isinstance(headers, dict):
|
||||
raise ValueError("headers is required")
|
||||
|
||||
model_response = kwargs.get("model_response")
|
||||
model_response: Final = kwargs.get("model_response")
|
||||
if model_response is None or not isinstance(model_response, ModelResponse):
|
||||
raise ValueError("model_response is required")
|
||||
|
||||
logging_obj = kwargs.get("logging_obj")
|
||||
logging_obj: Final = kwargs.get("logging_obj")
|
||||
if logging_obj is None or not isinstance(logging_obj, LiteLLMLoggingObj):
|
||||
raise ValueError("logging_obj is required")
|
||||
|
||||
|
|
@ -158,19 +158,19 @@ class ResponsesToCompletionBridgeHandler:
|
|||
from litellm import responses
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
validated_kwargs = self.validate_input_kwargs(kwargs)
|
||||
model = validated_kwargs["model"]
|
||||
messages = validated_kwargs["messages"]
|
||||
validated_kwargs: Final = self.validate_input_kwargs(kwargs)
|
||||
model: Final = validated_kwargs["model"]
|
||||
messages: Final = validated_kwargs["messages"]
|
||||
optional_params = validated_kwargs["optional_params"]
|
||||
litellm_params = validated_kwargs["litellm_params"]
|
||||
headers = validated_kwargs["headers"]
|
||||
model_response = validated_kwargs["model_response"]
|
||||
logging_obj = validated_kwargs["logging_obj"]
|
||||
custom_llm_provider = validated_kwargs["custom_llm_provider"]
|
||||
litellm_params: Final = validated_kwargs["litellm_params"]
|
||||
headers: Final = validated_kwargs["headers"]
|
||||
model_response: Final = validated_kwargs["model_response"]
|
||||
logging_obj: Final = validated_kwargs["logging_obj"]
|
||||
custom_llm_provider: Final = validated_kwargs["custom_llm_provider"]
|
||||
if kwargs.get("stream") is True and "stream" not in optional_params:
|
||||
optional_params = {**optional_params, "stream": True}
|
||||
|
||||
request_data = self.transformation_handler.transform_request(
|
||||
request_data: Final = self.transformation_handler.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -188,13 +188,13 @@ class ResponsesToCompletionBridgeHandler:
|
|||
# than adding an explicit kwarg) avoids the duplicate-keyword
|
||||
# TypeError that would otherwise fire on the real bridge path.
|
||||
request_data["custom_llm_provider"] = custom_llm_provider
|
||||
result = responses(
|
||||
result: Final = responses(
|
||||
**request_data,
|
||||
)
|
||||
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
stream = self._resolve_stream_flag(optional_params, litellm_params)
|
||||
stream: Final = self._resolve_stream_flag(optional_params, litellm_params)
|
||||
if isinstance(result, ResponsesAPIResponse):
|
||||
return self.transformation_handler.transform_response(
|
||||
model=model,
|
||||
|
|
@ -220,7 +220,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
elif not stream:
|
||||
responses_api_response = self._collect_response_from_stream(result)
|
||||
responses_api_response: Final = self._collect_response_from_stream(result)
|
||||
return self.transformation_handler.transform_response(
|
||||
model=model,
|
||||
raw_response=responses_api_response,
|
||||
|
|
@ -237,12 +237,12 @@ class ResponsesToCompletionBridgeHandler:
|
|||
else:
|
||||
if self._is_preformatted_cached_chat_stream(result):
|
||||
return self._apply_post_stream_processing(result, model, custom_llm_provider)
|
||||
completion_stream = self.transformation_handler.get_model_response_iterator(
|
||||
completion_stream: Final = self.transformation_handler.get_model_response_iterator(
|
||||
streaming_response=result, # type: ignore
|
||||
sync_stream=True,
|
||||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
streamwrapper = CustomStreamWrapper(
|
||||
streamwrapper: Final = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -254,20 +254,20 @@ class ResponsesToCompletionBridgeHandler:
|
|||
from litellm import aresponses
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
validated_kwargs = self.validate_input_kwargs(kwargs)
|
||||
model = validated_kwargs["model"]
|
||||
messages = validated_kwargs["messages"]
|
||||
validated_kwargs: Final = self.validate_input_kwargs(kwargs)
|
||||
model: Final = validated_kwargs["model"]
|
||||
messages: Final = validated_kwargs["messages"]
|
||||
optional_params = validated_kwargs["optional_params"]
|
||||
litellm_params = validated_kwargs["litellm_params"]
|
||||
headers = validated_kwargs["headers"]
|
||||
model_response = validated_kwargs["model_response"]
|
||||
logging_obj = validated_kwargs["logging_obj"]
|
||||
custom_llm_provider = validated_kwargs["custom_llm_provider"]
|
||||
litellm_params: Final = validated_kwargs["litellm_params"]
|
||||
headers: Final = validated_kwargs["headers"]
|
||||
model_response: Final = validated_kwargs["model_response"]
|
||||
logging_obj: Final = validated_kwargs["logging_obj"]
|
||||
custom_llm_provider: Final = validated_kwargs["custom_llm_provider"]
|
||||
if kwargs.get("stream") is True and "stream" not in optional_params:
|
||||
optional_params = {**optional_params, "stream": True}
|
||||
|
||||
try:
|
||||
request_data = self.transformation_handler.transform_request(
|
||||
request_data: Final = self.transformation_handler.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -285,14 +285,14 @@ class ResponsesToCompletionBridgeHandler:
|
|||
# keyword TypeError when `sanitized_litellm_params` already
|
||||
# carries `custom_llm_provider`.
|
||||
request_data["custom_llm_provider"] = custom_llm_provider
|
||||
result = await aresponses(
|
||||
result: Final = await aresponses(
|
||||
**request_data,
|
||||
aresponses=True,
|
||||
)
|
||||
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
stream = self._resolve_stream_flag(optional_params, litellm_params)
|
||||
stream: Final = self._resolve_stream_flag(optional_params, litellm_params)
|
||||
if isinstance(result, ResponsesAPIResponse):
|
||||
return self.transformation_handler.transform_response(
|
||||
model=model,
|
||||
|
|
@ -318,7 +318,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
elif not stream:
|
||||
responses_api_response = await self._collect_response_from_stream_async(result)
|
||||
responses_api_response: Final = await self._collect_response_from_stream_async(result)
|
||||
return self.transformation_handler.transform_response(
|
||||
model=model,
|
||||
raw_response=responses_api_response,
|
||||
|
|
@ -335,12 +335,12 @@ class ResponsesToCompletionBridgeHandler:
|
|||
else:
|
||||
if self._is_preformatted_cached_chat_stream(result):
|
||||
return self._apply_post_stream_processing(result, model, custom_llm_provider)
|
||||
completion_stream = self.transformation_handler.get_model_response_iterator(
|
||||
completion_stream: Final = self.transformation_handler.get_model_response_iterator(
|
||||
streaming_response=result, # type: ignore
|
||||
sync_stream=False,
|
||||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
streamwrapper = CustomStreamWrapper(
|
||||
streamwrapper: Final = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -359,7 +359,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
|
||||
streamwrapper = CustomStreamWrapper(
|
||||
streamwrapper: Final = CustomStreamWrapper(
|
||||
completion_stream=MockResponseIterator(model_response=response, json_mode=json_mode),
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -378,7 +378,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
try:
|
||||
provider_config = ProviderConfigManager.get_provider_chat_config(
|
||||
provider_config: Final = ProviderConfigManager.get_provider_chat_config(
|
||||
model=model, provider=LlmProviders(custom_llm_provider)
|
||||
)
|
||||
except (ValueError, KeyError):
|
||||
|
|
@ -389,4 +389,4 @@ class ResponsesToCompletionBridgeHandler:
|
|||
return stream
|
||||
|
||||
|
||||
responses_api_bridge = ResponsesToCompletionBridgeHandler()
|
||||
responses_api_bridge: Final = ResponsesToCompletionBridgeHandler()
|
||||
|
|
|
|||
|
|
@ -5,13 +5,7 @@ Handler for transforming /chat/completions api requests to litellm.responses req
|
|||
import json
|
||||
import os
|
||||
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Literal,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Union, cast
|
||||
|
||||
from openai.types.responses.custom_tool_param import CustomToolParam
|
||||
from openai.types.responses.response_input_param import (
|
||||
|
|
@ -69,7 +63,7 @@ def _get_reasoning_items(
|
|||
msg: "AllMessageValues",
|
||||
) -> list[ChatCompletionReasoningItem]:
|
||||
"""Extract reasoning_items from a message dict with proper typing."""
|
||||
items = msg.get("reasoning_items") # type: ignore[union-attr]
|
||||
items: Final = msg.get("reasoning_items") # type: ignore[union-attr]
|
||||
if items:
|
||||
return items # type: ignore[return-value]
|
||||
return []
|
||||
|
|
@ -84,7 +78,7 @@ def _build_reasoning_item(
|
|||
|
||||
Handles both pydantic objects (attribute access) and plain dicts.
|
||||
"""
|
||||
summary: list[dict[str, Any]] = []
|
||||
summary: Final[list[dict[str, Any]]] = []
|
||||
for s in summary_raw or []:
|
||||
if isinstance(s, dict):
|
||||
summary.append({"type": s.get("type", "summary_text"), "text": s.get("text", "")})
|
||||
|
|
@ -118,17 +112,17 @@ def _tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _Ch
|
|||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
|
||||
is_custom = item.get("type") == "custom_tool_call"
|
||||
arguments = (item.get("input") if is_custom else item.get("arguments")) or ""
|
||||
name = item.get("name") or ("custom_tool" if is_custom else "")
|
||||
function_chunk = ChatCompletionToolCallFunctionChunk(name=name, arguments=arguments)
|
||||
tool_call_dict = _ChatToolCallDict(
|
||||
is_custom: Final = item.get("type") == "custom_tool_call"
|
||||
arguments: Final = (item.get("input") if is_custom else item.get("arguments")) or ""
|
||||
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")),
|
||||
type="function",
|
||||
function=function_chunk,
|
||||
index=index,
|
||||
)
|
||||
raw_provider_fields = item.get("provider_specific_fields")
|
||||
raw_provider_fields: Final = item.get("provider_specific_fields")
|
||||
if isinstance(raw_provider_fields, dict):
|
||||
provider_specific_fields = raw_provider_fields
|
||||
elif raw_provider_fields and hasattr(raw_provider_fields, "__dict__"):
|
||||
|
|
@ -151,7 +145,7 @@ def _reasoning_item_to_response_input(
|
|||
r_item: ChatCompletionReasoningItem | dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Convert a stored ChatCompletionReasoningItem back to a Responses API input item."""
|
||||
r_input: dict[str, Any] = {
|
||||
r_input: Final[dict[str, Any]] = {
|
||||
"type": "reasoning",
|
||||
"id": r_item.get("id") or f"rs_{id(r_item)}",
|
||||
# summary is always required by the Responses API, even when empty
|
||||
|
|
@ -174,15 +168,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
"""Chat tool_choice nests the name under function/custom; Responses API expects top-level name."""
|
||||
if not isinstance(tool_choice, dict):
|
||||
return tool_choice
|
||||
choice_type = tool_choice.get("type")
|
||||
choice_type: Final = tool_choice.get("type")
|
||||
if choice_type not in ("function", "custom"):
|
||||
return tool_choice
|
||||
if isinstance(tool_choice.get("name"), str) and tool_choice.get("name"):
|
||||
# Return only Responses shape so stray chat ``function``/``custom`` keys are not sent upstream.
|
||||
return _flat_responses_tool_choice(choice_type, tool_choice["name"])
|
||||
nested = tool_choice.get(choice_type)
|
||||
nested: Final = tool_choice.get(choice_type)
|
||||
if isinstance(nested, dict):
|
||||
nested_name = nested.get("name")
|
||||
nested_name: Final = nested.get("name")
|
||||
if isinstance(nested_name, str) and nested_name:
|
||||
return _flat_responses_tool_choice(choice_type, nested_name)
|
||||
return tool_choice
|
||||
|
|
@ -200,7 +194,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
"""
|
||||
from litellm.types.utils import Choices, Message
|
||||
|
||||
item_type = item.get("type")
|
||||
item_type: Final = item.get("type")
|
||||
|
||||
# Ignore reasoning items for now
|
||||
if item_type == "reasoning":
|
||||
|
|
@ -208,7 +202,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
# Handle message items with output_text content
|
||||
if item_type == "message":
|
||||
content_list = item.get("content", [])
|
||||
content_list: Final = item.get("content", [])
|
||||
for content_item in content_list:
|
||||
if isinstance(content_item, dict):
|
||||
content_type = content_item.get("type")
|
||||
|
|
@ -235,9 +229,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
def convert_chat_completion_messages_to_responses_api(
|
||||
self, messages: list["AllMessageValues"]
|
||||
) -> tuple[list[Any], str | None]:
|
||||
input_items: list[Any] = []
|
||||
input_items: Final[list[Any]] = []
|
||||
instructions: str | None = None
|
||||
custom_tool_call_ids = frozenset(
|
||||
custom_tool_call_ids: Final = frozenset(
|
||||
tool_call["id"]
|
||||
for msg in messages
|
||||
if msg.get("role") == "assistant" and isinstance(msg.get("tool_calls"), list)
|
||||
|
|
@ -386,13 +380,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
def _build_sanitized_litellm_params(self, litellm_params: dict) -> dict[str, Any]:
|
||||
"""Build sanitized litellm_params with merged metadata."""
|
||||
responses_optional_param_keys = set(ResponsesAPIOptionalRequestParams.__annotations__.keys())
|
||||
sanitized: dict[str, Any] = {
|
||||
responses_optional_param_keys: Final = set(ResponsesAPIOptionalRequestParams.__annotations__.keys())
|
||||
sanitized: Final[dict[str, Any]] = {
|
||||
key: value for key, value in litellm_params.items() if key not in responses_optional_param_keys
|
||||
}
|
||||
legacy_metadata = litellm_params.get("metadata")
|
||||
existing_litellm_metadata = litellm_params.get("litellm_metadata")
|
||||
merged_litellm_metadata: dict[str, Any] = {}
|
||||
legacy_metadata: Final = litellm_params.get("metadata")
|
||||
existing_litellm_metadata: Final = litellm_params.get("litellm_metadata")
|
||||
merged_litellm_metadata: Final[dict[str, Any]] = {}
|
||||
if isinstance(legacy_metadata, dict):
|
||||
merged_litellm_metadata.update(legacy_metadata)
|
||||
if isinstance(existing_litellm_metadata, dict):
|
||||
|
|
@ -456,7 +450,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
optional_params = self._extract_extra_body_params(optional_params)
|
||||
|
||||
# Build responses API request using the reverse transformation logic
|
||||
responses_api_request = ResponsesAPIOptionalRequestParams()
|
||||
responses_api_request: Final = ResponsesAPIOptionalRequestParams()
|
||||
|
||||
# Set instructions if we found a system message
|
||||
if instructions:
|
||||
|
|
@ -464,7 +458,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
self._map_optional_params_to_responses_api_request(optional_params, responses_api_request)
|
||||
|
||||
stream = optional_params.get("stream") or litellm_params.get("stream", False)
|
||||
stream: Final = optional_params.get("stream") or litellm_params.get("stream", False)
|
||||
verbose_logger.debug("Chat provider: Stream parameter: %s", stream)
|
||||
|
||||
# Ensure stream is properly set in the request
|
||||
|
|
@ -472,22 +466,22 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
responses_api_request["stream"] = True
|
||||
|
||||
# Handle session management if previous_response_id is provided
|
||||
previous_response_id = optional_params.get("previous_response_id")
|
||||
previous_response_id: Final = optional_params.get("previous_response_id")
|
||||
if previous_response_id:
|
||||
# Use the existing session handler for responses API
|
||||
verbose_logger.debug("Chat provider: Warning ignoring previous response ID: %s", previous_response_id)
|
||||
|
||||
# Convert back to responses API format for the actual request
|
||||
|
||||
api_model = model
|
||||
api_model: Final = model
|
||||
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
setattr(litellm_logging_obj, "call_type", CallTypes.responses.value)
|
||||
|
||||
sanitized_litellm_params = self._build_sanitized_litellm_params(litellm_params)
|
||||
sanitized_litellm_params: Final = self._build_sanitized_litellm_params(litellm_params)
|
||||
|
||||
request_data = {
|
||||
request_data: Final = {
|
||||
"model": api_model,
|
||||
"input": input_items,
|
||||
"litellm_logging_obj": litellm_logging_obj,
|
||||
|
|
@ -534,14 +528,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
from litellm.types.utils import Choices, Message
|
||||
|
||||
choices: list[Choices] = []
|
||||
choices: Final[list[Choices]] = []
|
||||
index = 0
|
||||
reasoning_content: str | None = None
|
||||
pending_reasoning_item: dict[str, Any] | None = None
|
||||
|
||||
# Collect all tool calls to put them in a single choice
|
||||
# (Chat Completions API expects all tool calls in one message)
|
||||
accumulated_tool_calls: list[dict[str, Any]] = []
|
||||
accumulated_tool_calls: Final[list[dict[str, Any]]] = []
|
||||
tool_call_index = 0
|
||||
|
||||
for item in output_items:
|
||||
|
|
@ -649,10 +643,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
@classmethod
|
||||
def _extract_output_from_completed_event(cls, parsed_chunk: dict[str, Any]) -> list[dict[str, Any]] | None:
|
||||
response_payload = parsed_chunk.get("response")
|
||||
response_payload: Final = parsed_chunk.get("response")
|
||||
if not isinstance(response_payload, dict):
|
||||
return None
|
||||
response_output = response_payload.get("output")
|
||||
response_output: Final = response_payload.get("output")
|
||||
if not isinstance(response_output, list) or len(response_output) == 0:
|
||||
return None
|
||||
return cast(list[dict[str, Any]], response_output)
|
||||
|
|
@ -662,8 +656,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if not raw_sse or not isinstance(raw_sse, str):
|
||||
return []
|
||||
|
||||
recovered_output_items: dict[int, dict[str, Any]] = {}
|
||||
recovered_text_only_items: dict[int, dict[str, Any]] = {}
|
||||
recovered_output_items: Final[dict[int, dict[str, Any]]] = {}
|
||||
recovered_text_only_items: Final[dict[int, dict[str, Any]]] = {}
|
||||
|
||||
for chunk in raw_sse.splitlines():
|
||||
parsed_chunk = parse_sse_json_chunk(chunk)
|
||||
|
|
@ -698,7 +692,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
# but text-only items at indices without a matching OUTPUT_ITEM_DONE
|
||||
# must still be preserved (e.g. multi-output responses where some
|
||||
# indices only emitted OUTPUT_TEXT_DONE).
|
||||
merged_items: dict[int, dict[str, Any]] = {**recovered_text_only_items}
|
||||
merged_items: Final[dict[int, dict[str, Any]]] = {**recovered_text_only_items}
|
||||
merged_items.update(recovered_output_items)
|
||||
|
||||
if merged_items:
|
||||
|
|
@ -708,8 +702,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
@classmethod
|
||||
def _recover_output_items_from_logging(cls, logging_obj: "LiteLLMLoggingObj") -> list[dict[str, Any]]:
|
||||
model_call_details = getattr(logging_obj, "model_call_details", {}) or {}
|
||||
original_response = model_call_details.get("original_response")
|
||||
model_call_details: Final = getattr(logging_obj, "model_call_details", {}) or {}
|
||||
original_response: Final = model_call_details.get("original_response")
|
||||
return cls._recover_output_items_from_raw_sse(original_response)
|
||||
|
||||
def transform_response(
|
||||
|
|
@ -738,7 +732,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
output_items = raw_response.output
|
||||
if len(output_items) == 0:
|
||||
recovered_output_items = self._recover_output_items_from_logging(logging_obj)
|
||||
recovered_output_items: Final = self._recover_output_items_from_logging(logging_obj)
|
||||
if recovered_output_items:
|
||||
output_items = cast(Any, recovered_output_items)
|
||||
raw_response.output = cast(Any, recovered_output_items)
|
||||
|
|
@ -748,7 +742,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
)
|
||||
|
||||
# Convert response output to choices using the static helper
|
||||
choices = self._convert_response_output_to_choices(
|
||||
choices: Final = self._convert_response_output_to_choices(
|
||||
output_items=output_items,
|
||||
handle_raw_dict_callback=self._handle_raw_dict_response_item,
|
||||
)
|
||||
|
|
@ -771,7 +765,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
# Preserve hidden params from the ResponsesAPIResponse, especially the headers
|
||||
# which contain important provider information like x-request-id
|
||||
raw_response_hidden_params = getattr(raw_response, "_hidden_params", {})
|
||||
raw_response_hidden_params: Final = getattr(raw_response, "_hidden_params", {})
|
||||
if raw_response_hidden_params:
|
||||
if not hasattr(model_response, "_hidden_params") or model_response._hidden_params is None:
|
||||
model_response._hidden_params = {}
|
||||
|
|
@ -807,7 +801,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
) -> "ResponseInputImageParam":
|
||||
from openai.types.responses import ResponseInputImageParam
|
||||
|
||||
content_image_url = content.get("image_url")
|
||||
content_image_url: Final = content.get("image_url")
|
||||
actual_image_url: str | None = None
|
||||
detail: Literal["low", "high", "auto"] | None = None
|
||||
|
||||
|
|
@ -823,7 +817,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if actual_image_url is None:
|
||||
raise ValueError(f"Invalid image URL: {content_image_url}")
|
||||
|
||||
image_param = ResponseInputImageParam(image_url=actual_image_url, detail="auto", type="input_image")
|
||||
image_param: Final = ResponseInputImageParam(image_url=actual_image_url, detail="auto", type="input_image")
|
||||
|
||||
if detail:
|
||||
image_param["detail"] = detail
|
||||
|
|
@ -921,7 +915,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
def _convert_tools_to_responses_format(self, tools: list[dict[str, Any]]) -> list["ALL_RESPONSES_API_TOOL_PARAMS"]:
|
||||
"""Convert chat completion tools to responses API tools format"""
|
||||
responses_tools: list[ALL_RESPONSES_API_TOOL_PARAMS] = []
|
||||
responses_tools: Final[list[ALL_RESPONSES_API_TOOL_PARAMS]] = []
|
||||
for tool in tools:
|
||||
# convert function tool from chat completion to responses API format
|
||||
if tool.get("type") == "function":
|
||||
|
|
@ -959,11 +953,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
unsupported params remain in extra_body.
|
||||
"""
|
||||
# Extract extra_body and separate supported params from unsupported ones
|
||||
extra_body = optional_params.pop("extra_body", None) or {}
|
||||
extra_body: Final = optional_params.pop("extra_body", None) or {}
|
||||
if not extra_body:
|
||||
return optional_params
|
||||
|
||||
supported_responses_api_params = set(ResponsesAPIOptionalRequestParams.__annotations__.keys())
|
||||
supported_responses_api_params: Final = set(ResponsesAPIOptionalRequestParams.__annotations__.keys())
|
||||
# Also include params we handle specially
|
||||
supported_responses_api_params.update(
|
||||
{
|
||||
|
|
@ -973,7 +967,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
)
|
||||
|
||||
# Extract supported params from extra_body and merge into optional_params
|
||||
extra_body_copy = extra_body.copy()
|
||||
extra_body_copy: Final = extra_body.copy()
|
||||
for key, value in extra_body_copy.items():
|
||||
if key in supported_responses_api_params:
|
||||
# Prefer extra_body value if it exists (may have more complete info like summary in reasoning_effort)
|
||||
|
|
@ -988,7 +982,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
# Check if auto-summary is enabled via flag or environment variable
|
||||
# Priority: litellm.reasoning_auto_summary flag > LITELLM_REASONING_AUTO_SUMMARY env var
|
||||
auto_summary_enabled = (
|
||||
auto_summary_enabled: Final = (
|
||||
litellm.reasoning_auto_summary or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true"
|
||||
)
|
||||
|
||||
|
|
@ -1032,7 +1026,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
tools = []
|
||||
responses_api_request["tools"] = tools
|
||||
|
||||
web_search_tool: dict[str, Any] = {"type": "web_search"}
|
||||
web_search_tool: Final[dict[str, Any]] = {"type": "web_search"}
|
||||
if isinstance(web_search_options, dict):
|
||||
web_search_tool.update(web_search_options)
|
||||
|
||||
|
|
@ -1067,10 +1061,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
return None
|
||||
|
||||
if isinstance(response_format, dict):
|
||||
format_type = response_format.get("type")
|
||||
format_type: Final = response_format.get("type")
|
||||
|
||||
if format_type == "json_schema":
|
||||
json_schema = response_format.get("json_schema", {})
|
||||
json_schema: Final = response_format.get("json_schema", {})
|
||||
return {
|
||||
"format": {
|
||||
"type": "json_schema",
|
||||
|
|
@ -1099,7 +1093,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if not annotations:
|
||||
return None
|
||||
|
||||
result: list[ChatCompletionAnnotation] = []
|
||||
result: Final[list[ChatCompletionAnnotation]] = []
|
||||
for annotation in annotations:
|
||||
try:
|
||||
# Convert Pydantic models to dicts (handles both v1 and v2)
|
||||
|
|
@ -1127,7 +1121,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if not status:
|
||||
return "stop"
|
||||
|
||||
status_mapping = {
|
||||
status_mapping: Final = {
|
||||
"completed": "stop",
|
||||
"incomplete": "length",
|
||||
"failed": "stop",
|
||||
|
|
@ -1154,7 +1148,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
if not str_line or str_line.startswith("event:"):
|
||||
# ignore.
|
||||
return GenericStreamingChunk(text="", tool_use=None, is_finished=False, finish_reason="", usage=None)
|
||||
index = str_line.find("data:")
|
||||
index: Final = str_line.find("data:")
|
||||
if index != -1:
|
||||
str_line = str_line[index + 5 :]
|
||||
|
||||
|
|
@ -1240,10 +1234,10 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
# New output item added
|
||||
output_item = parsed_chunk.get("item", {})
|
||||
if output_item.get("type") in ("function_call", "custom_tool_call"):
|
||||
converted = _tool_call_dict_from_output_item(output_item, parsed_chunk.get("output_index", 0))
|
||||
provider_specific_fields = converted.get("provider_specific_fields")
|
||||
converted: Final = _tool_call_dict_from_output_item(output_item, parsed_chunk.get("output_index", 0))
|
||||
provider_specific_fields: Final = converted.get("provider_specific_fields")
|
||||
|
||||
function_chunk = ChatCompletionToolCallFunctionChunk(
|
||||
function_chunk: Final = ChatCompletionToolCallFunctionChunk(
|
||||
name=converted["function"]["name"] or None,
|
||||
arguments=converted["function"]["arguments"] or parsed_chunk.get("arguments") or "",
|
||||
)
|
||||
|
|
@ -1253,7 +1247,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
tool_call_index = OpenAiResponsesToChatCompletionStreamIterator._sequential_tool_call_index(
|
||||
tool_call_index_map, parsed_chunk.get("output_index", 0)
|
||||
)
|
||||
tool_call_chunk = ChatCompletionToolCallChunk(
|
||||
tool_call_chunk: Final = ChatCompletionToolCallChunk(
|
||||
id=converted["id"],
|
||||
index=tool_call_index,
|
||||
type="function",
|
||||
|
|
@ -1383,16 +1377,16 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
|
||||
# Check if response contains function_call items in output
|
||||
# to determine correct finish_reason
|
||||
response_data = parsed_chunk.get("response", {})
|
||||
output_items = response_data.get("output", []) if response_data else []
|
||||
response_data: Final = parsed_chunk.get("response", {})
|
||||
output_items: Final = response_data.get("output", []) if response_data else []
|
||||
|
||||
has_function_calls = any(
|
||||
has_function_calls: Final = any(
|
||||
item.get("type") in ("function_call", "custom_tool_call")
|
||||
for item in output_items
|
||||
if isinstance(item, dict)
|
||||
)
|
||||
|
||||
finish_reason = "tool_calls" if has_function_calls else "stop"
|
||||
finish_reason: Final = "tool_calls" if has_function_calls else "stop"
|
||||
|
||||
# Extract reasoning items with encrypted_content for round-tripping
|
||||
completed_reasoning_items: list[dict[str, Any]] | None = None
|
||||
|
|
@ -1408,7 +1402,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
summary_raw=item.get("summary"),
|
||||
)
|
||||
)
|
||||
completed_reasoning_items_typed = cast(
|
||||
completed_reasoning_items_typed: Final = cast(
|
||||
list[ChatCompletionReasoningItem] | None,
|
||||
completed_reasoning_items,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ scoring, message stubbing, and retrieval tool injection.
|
|||
"""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, cast
|
||||
from typing import Any, Final, cast
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.compression.message_stubbing import (
|
||||
|
|
@ -20,11 +20,11 @@ from litellm.types.utils import CallTypes
|
|||
|
||||
# CallTypes that produce Anthropic-shaped messages (structured content blocks).
|
||||
# Everything else is treated as OpenAI chat-completions shape.
|
||||
_ANTHROPIC_CALL_TYPES = frozenset({CallTypes.anthropic_messages.value})
|
||||
_ANTHROPIC_CALL_TYPES: Final = frozenset({CallTypes.anthropic_messages.value})
|
||||
# CallTypes that are valid targets for compression. Compression operates on
|
||||
# message-shaped inputs, so we only accept call types whose payload is a list
|
||||
# of role/content messages.
|
||||
_SUPPORTED_CALL_TYPES = frozenset(
|
||||
_SUPPORTED_CALL_TYPES: Final = frozenset(
|
||||
{
|
||||
CallTypes.completion.value,
|
||||
CallTypes.acompletion.value,
|
||||
|
|
@ -54,7 +54,7 @@ def _build_retrieval_tools(keys: list[str], call_type: str) -> list[dict]:
|
|||
if not keys:
|
||||
return []
|
||||
|
||||
openai_tools = [build_retrieval_tool(keys)]
|
||||
openai_tools: Final = [build_retrieval_tool(keys)]
|
||||
if not _is_anthropic_call_type(call_type):
|
||||
return openai_tools
|
||||
|
||||
|
|
@ -77,8 +77,8 @@ def _content_to_text(content: Any) -> str:
|
|||
|
||||
Implemented iteratively (stack-based) to avoid unbounded recursion.
|
||||
"""
|
||||
parts: list[str] = []
|
||||
stack: list[Any] = [content]
|
||||
parts: Final[list[str]] = []
|
||||
stack: Final[list[Any]] = [content]
|
||||
while stack:
|
||||
item = stack.pop()
|
||||
if isinstance(item, str):
|
||||
|
|
@ -111,9 +111,9 @@ def _normalize_messages_for_compression(
|
|||
f"Unsupported call_type={call_type!r} for compression. Expected one of: {sorted(_SUPPORTED_CALL_TYPES)}."
|
||||
)
|
||||
|
||||
original_messages: list[dict[str, Any]] = [dict(m) for m in messages]
|
||||
original_messages: Final[list[dict[str, Any]]] = [dict(m) for m in messages]
|
||||
|
||||
normalized_messages: list[dict] = []
|
||||
normalized_messages: Final[list[dict]] = []
|
||||
for msg in original_messages:
|
||||
normalized_messages.append(
|
||||
{
|
||||
|
|
@ -135,7 +135,7 @@ def _extract_last_user_message(messages: list[dict]) -> str:
|
|||
def _extract_tool_use_ids(content: Any) -> list[str]:
|
||||
if not isinstance(content, list):
|
||||
return []
|
||||
tool_use_ids: list[str] = []
|
||||
tool_use_ids: Final[list[str]] = []
|
||||
for part in content:
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
|
|
@ -150,7 +150,7 @@ def _extract_tool_use_ids(content: Any) -> list[str]:
|
|||
def _extract_tool_result_ids(content: Any) -> set[str]:
|
||||
if not isinstance(content, list):
|
||||
return set()
|
||||
tool_result_ids: set[str] = set()
|
||||
tool_result_ids: Final[set[str]] = set()
|
||||
for part in content:
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
|
|
@ -171,7 +171,7 @@ def _extract_anthropic_tool_exchange_spans(
|
|||
Each assistant message containing `tool_use` must be immediately followed by a
|
||||
user message containing matching `tool_result` blocks for all tool_use ids.
|
||||
"""
|
||||
spans: list[set[int]] = []
|
||||
spans: Final[list[set[int]]] = []
|
||||
i = 0
|
||||
while i < len(messages):
|
||||
current = messages[i]
|
||||
|
|
@ -216,8 +216,8 @@ def get_protected_indices(messages: Sequence[Mapping[str, object]]) -> tuple[int
|
|||
so compressing it replaces the live instruction with a marker. Compression
|
||||
guardrails share this policy; see the Headroom guardrail.
|
||||
"""
|
||||
system_indices = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "system")
|
||||
last_user = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "user")[-1:]
|
||||
system_indices: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "system")
|
||||
last_user: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "user")[-1:]
|
||||
last_assistant = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "assistant")[-1:]
|
||||
return system_indices + last_user + last_assistant
|
||||
|
||||
|
|
@ -230,16 +230,16 @@ def _combine_scores(
|
|||
"""Weighted average of BM25 and embedding scores, with min-max normalization."""
|
||||
|
||||
def _normalize(scores: list[float]) -> list[float]:
|
||||
min_s = min(scores) if scores else 0.0
|
||||
max_s = max(scores) if scores else 0.0
|
||||
rng = max_s - min_s
|
||||
min_s: Final = min(scores) if scores else 0.0
|
||||
max_s: Final = max(scores) if scores else 0.0
|
||||
rng: Final = max_s - min_s
|
||||
if rng == 0:
|
||||
return [0.0] * len(scores)
|
||||
return [(s - min_s) / rng for s in scores]
|
||||
|
||||
norm_bm25 = _normalize(bm25_scores)
|
||||
norm_emb = _normalize(emb_scores)
|
||||
emb_weight = 1.0 - bm25_weight
|
||||
norm_bm25: Final = _normalize(bm25_scores)
|
||||
norm_emb: Final = _normalize(emb_scores)
|
||||
emb_weight: Final = 1.0 - bm25_weight
|
||||
|
||||
return [bm25_weight * b + emb_weight * e for b, e in zip(norm_bm25, norm_emb)]
|
||||
|
||||
|
|
@ -253,7 +253,7 @@ def _select_kept_indices_for_budget(
|
|||
initial_kept_indices: set[int],
|
||||
tool_exchange_spans: list[set[int]],
|
||||
) -> tuple[set[int], dict[int, dict]]:
|
||||
kept_indices = set(initial_kept_indices)
|
||||
kept_indices: Final = set(initial_kept_indices)
|
||||
current_tokens = 0
|
||||
for i in kept_indices:
|
||||
current_tokens += token_counter(
|
||||
|
|
@ -265,14 +265,14 @@ def _select_kept_indices_for_budget(
|
|||
# A unit is either:
|
||||
# 1) a single message index, or
|
||||
# 2) an Anthropic tool-exchange span that must be kept/dropped atomically.
|
||||
truncated_overrides: dict[int, dict] = {} # idx -> truncated message dict
|
||||
span_id_by_index: dict[int, int] = {}
|
||||
truncated_overrides: Final[dict[int, dict]] = {} # idx -> truncated message dict
|
||||
span_id_by_index: Final[dict[int, int]] = {}
|
||||
for span_id, span in enumerate(tool_exchange_spans):
|
||||
for idx in span:
|
||||
span_id_by_index[idx] = span_id
|
||||
|
||||
# Build single-message candidate units (non-span messages).
|
||||
candidate_units: list[tuple[float, tuple[int, ...], bool]] = []
|
||||
candidate_units: Final[list[tuple[float, tuple[int, ...], bool]]] = []
|
||||
for idx in range(len(normalized_messages)):
|
||||
if idx in span_id_by_index or idx in kept_indices:
|
||||
continue
|
||||
|
|
@ -323,7 +323,7 @@ def _select_kept_indices_for_budget(
|
|||
|
||||
|
||||
def _get_dropped_tool_span_indices(kept_indices: set[int], tool_exchange_spans: list[set[int]]) -> set[int]:
|
||||
dropped_tool_span_indices: set[int] = set()
|
||||
dropped_tool_span_indices: Final[set[int]] = set()
|
||||
for span in tool_exchange_spans:
|
||||
if not any(idx in kept_indices for idx in span):
|
||||
dropped_tool_span_indices.update(span)
|
||||
|
|
@ -372,7 +372,7 @@ def compress(
|
|||
A ``CompressedResult`` dict containing compressed messages, token
|
||||
counts, a cache of original content, and the retrieval tool definition.
|
||||
"""
|
||||
call_type_str = _normalize_call_type(call_type)
|
||||
call_type_str: Final = _normalize_call_type(call_type)
|
||||
normalized_messages, original_messages = _normalize_messages_for_compression(
|
||||
messages=messages,
|
||||
call_type=call_type_str,
|
||||
|
|
@ -381,7 +381,7 @@ def compress(
|
|||
if compression_target is None:
|
||||
compression_target = compression_trigger * 7 // 10
|
||||
|
||||
original_tokens = token_counter(
|
||||
original_tokens: Final = token_counter(
|
||||
model=model,
|
||||
messages=cast(list[Any], original_messages),
|
||||
)
|
||||
|
|
@ -399,17 +399,17 @@ def compress(
|
|||
)
|
||||
|
||||
# Extract query for relevance scoring
|
||||
query = _extract_last_user_message(normalized_messages)
|
||||
query: Final = _extract_last_user_message(normalized_messages)
|
||||
|
||||
# Score each message
|
||||
bm25_scores = bm25_score_messages(query, normalized_messages)
|
||||
bm25_scores: Final = bm25_score_messages(query, normalized_messages)
|
||||
|
||||
if embedding_model:
|
||||
from litellm.compression.scoring.embedding_scorer import (
|
||||
embedding_score_messages,
|
||||
)
|
||||
|
||||
emb_scores = embedding_score_messages(
|
||||
emb_scores: Final = embedding_score_messages(
|
||||
query,
|
||||
normalized_messages,
|
||||
model=embedding_model,
|
||||
|
|
@ -421,7 +421,7 @@ def compress(
|
|||
combined_scores = bm25_scores
|
||||
|
||||
# Protected messages are never compressed
|
||||
protected_indices = get_protected_indices(normalized_messages)
|
||||
protected_indices: Final = get_protected_indices(normalized_messages)
|
||||
kept_indices: set[int] = set(protected_indices)
|
||||
|
||||
tool_exchange_spans: list[set[int]] = []
|
||||
|
|
@ -454,10 +454,10 @@ def compress(
|
|||
)
|
||||
|
||||
# Build compressed messages and cache
|
||||
compressed_messages: list[dict] = []
|
||||
cache: dict[str, str] = {}
|
||||
used_keys: set[str] = set()
|
||||
dropped_tool_span_indices = _get_dropped_tool_span_indices(
|
||||
compressed_messages: Final[list[dict]] = []
|
||||
cache: Final[dict[str, str]] = {}
|
||||
used_keys: Final[set[str]] = set()
|
||||
dropped_tool_span_indices: Final = _get_dropped_tool_span_indices(
|
||||
kept_indices=kept_indices, tool_exchange_spans=tool_exchange_spans
|
||||
)
|
||||
|
||||
|
|
@ -474,9 +474,9 @@ def compress(
|
|||
compressed_messages.append(stub_message(msg, key))
|
||||
|
||||
# Build retrieval tool in the target request schema
|
||||
tools = _build_retrieval_tools(list(cache.keys()), call_type=call_type_str)
|
||||
tools: Final = _build_retrieval_tools(list(cache.keys()), call_type=call_type_str)
|
||||
|
||||
compressed_tokens = token_counter(
|
||||
compressed_tokens: Final = token_counter(
|
||||
model=model,
|
||||
messages=cast(list[Any], compressed_messages),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,8 +4,9 @@ Auto-detect content type per message: code, JSON, or text.
|
|||
|
||||
import json
|
||||
import re
|
||||
from typing import Final
|
||||
|
||||
_CODE_KEYWORDS = re.compile(
|
||||
_CODE_KEYWORDS: Final = re.compile(
|
||||
r"\b(?:def |function |class |import |from |require\(|#include|fn |func |const |let |var |public |private |static )\b"
|
||||
)
|
||||
|
||||
|
|
@ -16,7 +17,7 @@ def detect_content_type(content: str) -> str:
|
|||
|
||||
Returns one of: "code", "json", "text"
|
||||
"""
|
||||
stripped = content.strip()
|
||||
stripped: Final = content.strip()
|
||||
if not stripped:
|
||||
return "text"
|
||||
|
||||
|
|
@ -30,10 +31,10 @@ def detect_content_type(content: str) -> str:
|
|||
|
||||
# Check code indicators
|
||||
# Sample first 5000 chars for performance
|
||||
sample = stripped[:5000]
|
||||
keyword_matches = len(_CODE_KEYWORDS.findall(sample))
|
||||
lines = sample.split("\n")
|
||||
indented_lines = sum(1 for line in lines if line.startswith((" ", "\t")) and line.strip())
|
||||
sample: Final = stripped[:5000]
|
||||
keyword_matches: Final = len(_CODE_KEYWORDS.findall(sample))
|
||||
lines: Final = sample.split("\n")
|
||||
indented_lines: Final = sum(1 for line in lines if line.startswith((" ", "\t")) and line.strip())
|
||||
|
||||
# If we see multiple code keywords or significant indentation, it's likely code
|
||||
if keyword_matches >= 3 or (indented_lines > len(lines) * 0.3 and len(lines) > 5):
|
||||
|
|
|
|||
|
|
@ -3,11 +3,12 @@ Replace messages with compact stubs and extract human-readable keys.
|
|||
"""
|
||||
|
||||
import re
|
||||
from typing import Final
|
||||
|
||||
from litellm.compression.content_detection import detect_content_type
|
||||
|
||||
# Patterns for extracting file paths from content
|
||||
_FILE_PATH_PATTERNS = [
|
||||
_FILE_PATH_PATTERNS: Final = [
|
||||
re.compile(r"^#\s*(\S+\.\w+)", re.MULTILINE), # # filename.py
|
||||
re.compile(r"^//\s*(\S+\.\w+)", re.MULTILINE), # // filename.js
|
||||
re.compile(r"^File:\s*(\S+)", re.MULTILINE), # File: path/to/file
|
||||
|
|
@ -40,7 +41,7 @@ def extract_key(message: dict, fallback_index: int, used_keys: set[str]) -> str:
|
|||
key = f"message_{fallback_index}"
|
||||
|
||||
# Handle duplicates
|
||||
base_key = key
|
||||
base_key: Final = key
|
||||
counter = 2
|
||||
while key in used_keys:
|
||||
key = f"{base_key}_{counter}"
|
||||
|
|
@ -61,10 +62,10 @@ def stub_message(message: dict, key: str) -> dict:
|
|||
if isinstance(content, list):
|
||||
content = " ".join(p.get("text", "") if isinstance(p, dict) else str(p) for p in content)
|
||||
|
||||
line_count = content.count("\n") + 1
|
||||
content_type = detect_content_type(content)
|
||||
line_count: Final = content.count("\n") + 1
|
||||
content_type: Final = detect_content_type(content)
|
||||
|
||||
stub_content = (
|
||||
stub_content: Final = (
|
||||
f"[Compressed: {key} — {line_count} lines, {content_type}. "
|
||||
f"Use litellm_content_retrieve tool to get full content.]"
|
||||
)
|
||||
|
|
@ -89,23 +90,23 @@ def truncate_message(message: dict, max_tokens: int) -> dict:
|
|||
content = " ".join(p.get("text", "") if isinstance(p, dict) else str(p) for p in content)
|
||||
|
||||
# Rough conversion: 1 token ≈ 3 characters
|
||||
target_chars = max(100, max_tokens * 3)
|
||||
target_chars: Final = max(100, max_tokens * 3)
|
||||
|
||||
if len(content) <= target_chars:
|
||||
return {**message, "content": content}
|
||||
|
||||
lines = content.split("\n")
|
||||
lines: Final = content.split("\n")
|
||||
|
||||
# Estimate target line count from character budget
|
||||
avg_line_len = max(1, len(content) // max(1, len(lines)))
|
||||
target_lines = max(2, target_chars // avg_line_len)
|
||||
avg_line_len: Final = max(1, len(content) // max(1, len(lines)))
|
||||
target_lines: Final = max(2, target_chars // avg_line_len)
|
||||
|
||||
if len(lines) <= target_lines:
|
||||
return {**message, "content": content}
|
||||
|
||||
first_count = (target_lines * 7) // 10
|
||||
last_count = target_lines - first_count
|
||||
truncated = (
|
||||
first_count: Final = (target_lines * 7) // 10
|
||||
last_count: Final = target_lines - first_count
|
||||
truncated: Final = (
|
||||
"\n".join(lines[:first_count]) + "\n...[truncated for context window]...\n" + "\n".join(lines[-last_count:])
|
||||
)
|
||||
return {**message, "content": truncated}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ No external dependencies — uses only stdlib.
|
|||
import math
|
||||
import re
|
||||
from collections import Counter
|
||||
from typing import Final
|
||||
|
||||
|
||||
def _tokenize(text: str) -> list[str]:
|
||||
|
|
@ -16,11 +17,11 @@ def _tokenize(text: str) -> list[str]:
|
|||
|
||||
def _extract_content(message: dict) -> str:
|
||||
"""Extract text content from a message dict."""
|
||||
content = message.get("content", "")
|
||||
content: Final = message.get("content", "")
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
parts = []
|
||||
parts: Final = []
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
parts.append(part.get("text", ""))
|
||||
|
|
@ -48,32 +49,32 @@ def bm25_score_messages(
|
|||
Returns:
|
||||
List of float scores, one per message. Higher = more relevant.
|
||||
"""
|
||||
query_terms = _tokenize(query)
|
||||
query_terms: Final = _tokenize(query)
|
||||
if not query_terms:
|
||||
return [0.0] * len(messages)
|
||||
|
||||
# Tokenize all documents
|
||||
doc_tokens: list[list[str]] = []
|
||||
doc_tokens: Final[list[list[str]]] = []
|
||||
for msg in messages:
|
||||
doc_tokens.append(_tokenize(_extract_content(msg)))
|
||||
|
||||
n = len(doc_tokens)
|
||||
n: Final = len(doc_tokens)
|
||||
if n == 0:
|
||||
return []
|
||||
|
||||
# Average document length
|
||||
doc_lengths = [len(dt) for dt in doc_tokens]
|
||||
avgdl = sum(doc_lengths) / n if n > 0 else 1.0
|
||||
doc_lengths: Final = [len(dt) for dt in doc_tokens]
|
||||
avgdl: Final = sum(doc_lengths) / n if n > 0 else 1.0
|
||||
|
||||
# Document frequency for each term
|
||||
df: dict[str, int] = {}
|
||||
df: Final[dict[str, int]] = {}
|
||||
for dt in doc_tokens:
|
||||
seen = set(dt)
|
||||
for term in seen:
|
||||
df[term] = df.get(term, 0) + 1
|
||||
|
||||
# IDF for query terms
|
||||
idf: dict[str, float] = {}
|
||||
idf: Final[dict[str, float]] = {}
|
||||
for term in set(query_terms):
|
||||
term_df = df.get(term, 0)
|
||||
# Standard BM25 IDF: log((N - df + 0.5) / (df + 0.5) + 1)
|
||||
|
|
@ -85,7 +86,7 @@ def bm25_score_messages(
|
|||
# stemmer dependency.
|
||||
def _expand_tf(query_term: str, tf_counts: Counter) -> int: # type: ignore[type-arg]
|
||||
"""Sum TF across all doc tokens that are prefixed by query_term."""
|
||||
exact = tf_counts.get(query_term, 0)
|
||||
exact: Final = tf_counts.get(query_term, 0)
|
||||
if exact:
|
||||
return exact
|
||||
if len(query_term) < 4:
|
||||
|
|
@ -93,7 +94,7 @@ def bm25_score_messages(
|
|||
return sum(count for token, count in tf_counts.items() if token != query_term and token.startswith(query_term))
|
||||
|
||||
# Score each document
|
||||
scores: list[float] = []
|
||||
scores: Final[list[float]] = []
|
||||
for i, dt in enumerate(doc_tokens):
|
||||
if not dt:
|
||||
scores.append(0.0)
|
||||
|
|
|
|||
|
|
@ -5,18 +5,18 @@ Computes cosine similarity between the query embedding and each message embeddin
|
|||
"""
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
||||
|
||||
def _extract_content(message: dict) -> str:
|
||||
"""Extract text content from a message dict."""
|
||||
content = message.get("content", "")
|
||||
content: Final = message.get("content", "")
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
parts = []
|
||||
parts: Final = []
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
parts.append(part.get("text", ""))
|
||||
|
|
@ -30,15 +30,15 @@ def _truncate_text(text: str, max_chars: int = 30000) -> str:
|
|||
"""Truncate long text, keeping first and last portions."""
|
||||
if len(text) <= max_chars:
|
||||
return text
|
||||
half = max_chars // 2
|
||||
half: Final = max_chars // 2
|
||||
return text[:half] + "\n...\n" + text[-half:]
|
||||
|
||||
|
||||
def _cosine_similarity(a: list[float], b: list[float]) -> float:
|
||||
"""Compute cosine similarity between two vectors."""
|
||||
dot = sum(x * y for x, y in zip(a, b))
|
||||
norm_a = math.sqrt(sum(x * x for x in a))
|
||||
norm_b = math.sqrt(sum(x * x for x in b))
|
||||
dot: Final = sum(x * y for x, y in zip(a, b))
|
||||
norm_a: Final = math.sqrt(sum(x * x for x in a))
|
||||
norm_b: Final = math.sqrt(sum(x * x for x in b))
|
||||
if norm_a == 0 or norm_b == 0:
|
||||
return 0.0
|
||||
return dot / (norm_a * norm_b)
|
||||
|
|
@ -67,12 +67,12 @@ def embedding_score_messages(
|
|||
"""
|
||||
import litellm
|
||||
|
||||
texts = [_truncate_text(query)]
|
||||
texts: Final = [_truncate_text(query)]
|
||||
for msg in messages:
|
||||
texts.append(_truncate_text(_extract_content(msg)))
|
||||
|
||||
# Filter out empty texts — replace with a placeholder to maintain indexing
|
||||
processed_texts = [t if t.strip() else "empty" for t in texts]
|
||||
processed_texts: Final = [t if t.strip() else "empty" for t in texts]
|
||||
|
||||
kwargs: dict[str, Any] = {
|
||||
"model": model,
|
||||
|
|
@ -82,13 +82,13 @@ def embedding_score_messages(
|
|||
if embedding_model_params:
|
||||
kwargs = {**kwargs, **embedding_model_params}
|
||||
|
||||
response = litellm.embedding(**kwargs)
|
||||
response: Final = litellm.embedding(**kwargs)
|
||||
|
||||
# Extract embedding vectors
|
||||
embeddings = [item["embedding"] for item in response.data]
|
||||
embeddings: Final = [item["embedding"] for item in response.data]
|
||||
|
||||
query_embedding = embeddings[0]
|
||||
scores: list[float] = []
|
||||
query_embedding: Final = embeddings[0]
|
||||
scores: Final[list[float]] = []
|
||||
for i in range(1, len(embeddings)):
|
||||
scores.append(_cosine_similarity(query_embedding, embeddings[i]))
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -11,7 +11,7 @@ import json
|
|||
from collections.abc import Callable
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
import litellm
|
||||
from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT
|
||||
|
|
@ -28,7 +28,7 @@ from litellm.types.router import GenericLiteLLMParams
|
|||
from litellm.utils import ProviderConfigManager, client
|
||||
|
||||
# Response type mapping
|
||||
RESPONSE_TYPES: dict[str, type] = {
|
||||
RESPONSE_TYPES: Final[dict[str, type]] = {
|
||||
"ContainerFileListResponse": ContainerFileListResponse,
|
||||
"ContainerFileObject": ContainerFileObject,
|
||||
"DeleteContainerFileResponse": DeleteContainerFileResponse,
|
||||
|
|
@ -37,7 +37,7 @@ RESPONSE_TYPES: dict[str, type] = {
|
|||
|
||||
def _load_endpoints_config() -> dict:
|
||||
"""Load the endpoints configuration from JSON file."""
|
||||
config_path = Path(__file__).parent / "endpoints.json"
|
||||
config_path: Final = Path(__file__).parent / "endpoints.json"
|
||||
with open(config_path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
|
@ -48,9 +48,9 @@ def create_sync_endpoint_function(endpoint_config: dict) -> Callable:
|
|||
|
||||
Uses the generic container handler instead of individual handler methods.
|
||||
"""
|
||||
endpoint_name = endpoint_config["name"]
|
||||
response_type = RESPONSE_TYPES.get(endpoint_config["response_type"])
|
||||
path_params = endpoint_config.get("path_params", [])
|
||||
endpoint_name: Final = endpoint_config["name"]
|
||||
response_type: Final = RESPONSE_TYPES.get(endpoint_config["response_type"])
|
||||
path_params: Final = endpoint_config.get("path_params", [])
|
||||
|
||||
@client
|
||||
def endpoint_func(
|
||||
|
|
@ -61,12 +61,12 @@ def create_sync_endpoint_function(endpoint_config: dict) -> Callable:
|
|||
extra_body: dict[str, Any] | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
resolved_custom_llm_provider: str = custom_llm_provider
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj")
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id")
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
# Check for mock response
|
||||
mock_response = kwargs.get("mock_response")
|
||||
|
|
@ -99,7 +99,7 @@ def create_sync_endpoint_function(endpoint_config: dict) -> Callable:
|
|||
raise ValueError(f"Container provider config not found for: {resolved_custom_llm_provider}")
|
||||
|
||||
# Build optional params for logging
|
||||
optional_params = {k: kwargs.get(k) for k in path_params if k in kwargs}
|
||||
optional_params: Final = {k: kwargs.get(k) for k in path_params if k in kwargs}
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
|
|
@ -150,12 +150,12 @@ def create_async_endpoint_function(
|
|||
extra_body: dict[str, Any] | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
sync_func,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -165,9 +165,9 @@ def create_async_endpoint_function(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -193,8 +193,8 @@ def generate_container_endpoints() -> dict[str, Callable]:
|
|||
|
||||
Returns a dict mapping function names to their implementations.
|
||||
"""
|
||||
config = _load_endpoints_config()
|
||||
endpoints = {}
|
||||
config: Final = _load_endpoints_config()
|
||||
endpoints: Final = {}
|
||||
|
||||
for endpoint_config in config["endpoints"]:
|
||||
# Create sync function
|
||||
|
|
@ -210,8 +210,8 @@ def generate_container_endpoints() -> dict[str, Callable]:
|
|||
|
||||
def get_all_endpoint_names() -> list[str]:
|
||||
"""Get all endpoint names (sync and async) from config."""
|
||||
config = _load_endpoints_config()
|
||||
names = []
|
||||
config: Final = _load_endpoints_config()
|
||||
names: Final = []
|
||||
for endpoint in config["endpoints"]:
|
||||
names.append(endpoint["name"])
|
||||
names.append(endpoint["async_name"])
|
||||
|
|
@ -220,21 +220,21 @@ def get_all_endpoint_names() -> list[str]:
|
|||
|
||||
def get_async_endpoint_names() -> list[str]:
|
||||
"""Get all async endpoint names for router registration."""
|
||||
config = _load_endpoints_config()
|
||||
config: Final = _load_endpoints_config()
|
||||
return [endpoint["async_name"] for endpoint in config["endpoints"]]
|
||||
|
||||
|
||||
# Generate endpoints on module load
|
||||
_generated_endpoints = generate_container_endpoints()
|
||||
_generated_endpoints: Final = generate_container_endpoints()
|
||||
|
||||
# Export generated functions dynamically
|
||||
list_container_files = _generated_endpoints.get("list_container_files")
|
||||
alist_container_files = _generated_endpoints.get("alist_container_files")
|
||||
upload_container_file = _generated_endpoints.get("upload_container_file")
|
||||
aupload_container_file = _generated_endpoints.get("aupload_container_file")
|
||||
retrieve_container_file = _generated_endpoints.get("retrieve_container_file")
|
||||
aretrieve_container_file = _generated_endpoints.get("aretrieve_container_file")
|
||||
delete_container_file = _generated_endpoints.get("delete_container_file")
|
||||
adelete_container_file = _generated_endpoints.get("adelete_container_file")
|
||||
retrieve_container_file_content = _generated_endpoints.get("retrieve_container_file_content")
|
||||
aretrieve_container_file_content = _generated_endpoints.get("aretrieve_container_file_content")
|
||||
list_container_files: Final = _generated_endpoints.get("list_container_files")
|
||||
alist_container_files: Final = _generated_endpoints.get("alist_container_files")
|
||||
upload_container_file: Final = _generated_endpoints.get("upload_container_file")
|
||||
aupload_container_file: Final = _generated_endpoints.get("aupload_container_file")
|
||||
retrieve_container_file: Final = _generated_endpoints.get("retrieve_container_file")
|
||||
aretrieve_container_file: Final = _generated_endpoints.get("aretrieve_container_file")
|
||||
delete_container_file: Final = _generated_endpoints.get("delete_container_file")
|
||||
adelete_container_file: Final = _generated_endpoints.get("adelete_container_file")
|
||||
retrieve_container_file_content: Final = _generated_endpoints.get("retrieve_container_file_content")
|
||||
aretrieve_container_file_content: Final = _generated_endpoints.get("aretrieve_container_file_content")
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import contextvars
|
|||
import json
|
||||
from collections.abc import Coroutine
|
||||
from functools import partial
|
||||
from typing import Any, Literal, overload
|
||||
from typing import Any, Final, Literal, overload
|
||||
|
||||
import litellm
|
||||
from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT
|
||||
|
|
@ -76,12 +76,12 @@ async def acreate_container(
|
|||
Returns:
|
||||
- `response` (ContainerObject): The created container object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
create_container,
|
||||
name=name,
|
||||
expires_after=expires_after,
|
||||
|
|
@ -94,9 +94,9 @@ async def acreate_container(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -185,11 +185,11 @@ def create_container(
|
|||
print(response)
|
||||
```
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id")
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
# Check for mock response first
|
||||
mock_response = kwargs.get("mock_response")
|
||||
|
|
@ -197,12 +197,12 @@ def create_container(
|
|||
if isinstance(mock_response, str):
|
||||
mock_response = json.loads(mock_response)
|
||||
|
||||
response = ContainerObject(**mock_response)
|
||||
response: Final = ContainerObject(**mock_response)
|
||||
return response
|
||||
|
||||
# get llm provider logic
|
||||
# Pass credential params explicitly since they're named args, not in kwargs
|
||||
litellm_params = GenericLiteLLMParams(
|
||||
litellm_params: Final = GenericLiteLLMParams(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
|
|
@ -218,12 +218,12 @@ def create_container(
|
|||
|
||||
local_vars.update(kwargs)
|
||||
# Get ContainerCreateOptionalRequestParams with only valid parameters
|
||||
container_create_optional_params: ContainerCreateOptionalRequestParams = (
|
||||
container_create_optional_params: Final[ContainerCreateOptionalRequestParams] = (
|
||||
ContainerRequestUtils.get_requested_container_create_optional_param(local_vars)
|
||||
)
|
||||
|
||||
# Get optional parameters for the container API
|
||||
container_create_request_params: dict = ContainerRequestUtils.get_optional_params_container_create(
|
||||
container_create_request_params: Final[dict] = ContainerRequestUtils.get_optional_params_container_create(
|
||||
container_provider_config=container_provider_config,
|
||||
container_create_optional_params=container_create_optional_params,
|
||||
)
|
||||
|
|
@ -306,12 +306,12 @@ async def alist_containers(
|
|||
Returns:
|
||||
- `response` (ContainerListResponse): The list of containers
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
list_containers,
|
||||
after=after,
|
||||
limit=limit,
|
||||
|
|
@ -324,9 +324,9 @@ async def alist_containers(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -403,11 +403,11 @@ def list_containers(
|
|||
|
||||
Currently supports OpenAI
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id")
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
# Check for mock response first
|
||||
mock_response = kwargs.get("mock_response")
|
||||
|
|
@ -415,12 +415,12 @@ def list_containers(
|
|||
if isinstance(mock_response, str):
|
||||
mock_response = json.loads(mock_response)
|
||||
|
||||
response = ContainerListResponse(**mock_response)
|
||||
response: Final = ContainerListResponse(**mock_response)
|
||||
return response
|
||||
|
||||
# get llm provider logic
|
||||
# Pass credential params explicitly since they're named args, not in kwargs
|
||||
litellm_params = GenericLiteLLMParams(
|
||||
litellm_params: Final = GenericLiteLLMParams(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
|
|
@ -435,7 +435,7 @@ def list_containers(
|
|||
raise ValueError(f"Container provider config not found for provider: {custom_llm_provider}")
|
||||
|
||||
# Get container list request parameters
|
||||
container_list_optional_params: ContainerListOptionalRequestParams = (
|
||||
container_list_optional_params: Final[ContainerListOptionalRequestParams] = (
|
||||
ContainerRequestUtils.get_requested_container_list_optional_param(local_vars)
|
||||
)
|
||||
|
||||
|
|
@ -504,12 +504,12 @@ async def aretrieve_container(
|
|||
Returns:
|
||||
- `response` (ContainerObject): The container object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
retrieve_container,
|
||||
container_id=container_id,
|
||||
timeout=timeout,
|
||||
|
|
@ -520,9 +520,9 @@ async def aretrieve_container(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -593,12 +593,12 @@ def retrieve_container(
|
|||
|
||||
Currently supports OpenAI
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
resolved_custom_llm_provider: str = custom_llm_provider
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id")
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
# Check for mock response first
|
||||
mock_response = kwargs.get("mock_response")
|
||||
|
|
@ -606,7 +606,7 @@ def retrieve_container(
|
|||
if isinstance(mock_response, str):
|
||||
mock_response = json.loads(mock_response)
|
||||
|
||||
response = ContainerObject(**mock_response)
|
||||
response: Final = ContainerObject(**mock_response)
|
||||
return response
|
||||
|
||||
# get llm provider logic
|
||||
|
|
@ -625,7 +625,7 @@ def retrieve_container(
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
# True when input was a LiteLLM-managed ID (any length); needed to re-encode output for routing affinity
|
||||
was_encoded = original_container_id != container_id
|
||||
was_encoded: Final = original_container_id != container_id
|
||||
|
||||
# get provider config
|
||||
container_provider_config: BaseContainerConfig | None = ProviderConfigManager.get_provider_container_config(
|
||||
|
|
@ -719,12 +719,12 @@ async def adelete_container(
|
|||
Returns:
|
||||
- `response` (DeleteContainerResult): The deletion result
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
delete_container,
|
||||
container_id=container_id,
|
||||
timeout=timeout,
|
||||
|
|
@ -735,9 +735,9 @@ async def adelete_container(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -808,12 +808,12 @@ def delete_container(
|
|||
|
||||
Currently supports OpenAI
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
resolved_custom_llm_provider: str = custom_llm_provider
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id")
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
# Check for mock response first
|
||||
mock_response = kwargs.get("mock_response")
|
||||
|
|
@ -821,7 +821,7 @@ def delete_container(
|
|||
if isinstance(mock_response, str):
|
||||
mock_response = json.loads(mock_response)
|
||||
|
||||
response = DeleteContainerResult(**mock_response)
|
||||
response: Final = DeleteContainerResult(**mock_response)
|
||||
return response
|
||||
|
||||
# get llm provider logic
|
||||
|
|
@ -840,7 +840,7 @@ def delete_container(
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
# True when input was a LiteLLM-managed ID (any length); needed to re-encode output for routing affinity
|
||||
was_encoded = original_container_id != container_id
|
||||
was_encoded: Final = original_container_id != container_id
|
||||
|
||||
# get provider config
|
||||
container_provider_config: BaseContainerConfig | None = ProviderConfigManager.get_provider_container_config(
|
||||
|
|
@ -938,12 +938,12 @@ async def alist_container_files(
|
|||
Returns:
|
||||
- `response` (ContainerFileListResponse): The list of container files
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
list_container_files,
|
||||
container_id=container_id,
|
||||
after=after,
|
||||
|
|
@ -957,9 +957,9 @@ async def alist_container_files(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -1037,12 +1037,12 @@ def list_container_files(
|
|||
|
||||
Currently supports OpenAI
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
resolved_custom_llm_provider: str = custom_llm_provider
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id")
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
# Check for mock response first
|
||||
mock_response = kwargs.get("mock_response")
|
||||
|
|
@ -1050,7 +1050,7 @@ def list_container_files(
|
|||
if isinstance(mock_response, str):
|
||||
mock_response = json.loads(mock_response)
|
||||
|
||||
response = ContainerFileListResponse(**mock_response)
|
||||
response: Final = ContainerFileListResponse(**mock_response)
|
||||
return response
|
||||
|
||||
# get llm provider logic
|
||||
|
|
@ -1168,12 +1168,12 @@ async def aupload_container_file(
|
|||
print(response)
|
||||
```
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
upload_container_file,
|
||||
container_id=container_id,
|
||||
file=file,
|
||||
|
|
@ -1185,9 +1185,9 @@ async def aupload_container_file(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -1288,12 +1288,12 @@ def upload_container_file(
|
|||
"""
|
||||
from litellm.llms.custom_httpx.container_handler import generic_container_handler
|
||||
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
resolved_custom_llm_provider: str = custom_llm_provider
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id")
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
# Check for mock response first
|
||||
mock_response = kwargs.get("mock_response")
|
||||
|
|
@ -1301,7 +1301,7 @@ def upload_container_file(
|
|||
if isinstance(mock_response, str):
|
||||
mock_response = json.loads(mock_response)
|
||||
|
||||
response = ContainerFileObject(**mock_response)
|
||||
response: Final = ContainerFileObject(**mock_response)
|
||||
return response
|
||||
|
||||
# get llm provider logic
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Any, TypeVar
|
||||
from typing import Any, Final, TypeVar
|
||||
|
||||
from litellm.llms.base_llm.containers.transformation import BaseContainerConfig
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
|
@ -19,14 +19,14 @@ def decode_managed_container_id_for_request(
|
|||
Returns:
|
||||
(original_container_id, resolved_provider, updated_litellm_params)
|
||||
"""
|
||||
decoded = ResponsesAPIRequestUtils._decode_container_id(container_id)
|
||||
original_container_id = decoded.get("response_id", container_id)
|
||||
decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id)
|
||||
original_container_id: Final = decoded.get("response_id", container_id)
|
||||
|
||||
decoded_provider = decoded.get("custom_llm_provider")
|
||||
decoded_provider: Final = decoded.get("custom_llm_provider")
|
||||
if decoded_provider and custom_llm_provider == "openai":
|
||||
custom_llm_provider = decoded_provider
|
||||
|
||||
decoded_model_id = decoded.get("model_id")
|
||||
decoded_model_id: Final = decoded.get("model_id")
|
||||
if decoded_model_id and not litellm_params.get("model_id"):
|
||||
litellm_params["model_id"] = decoded_model_id
|
||||
|
||||
|
|
@ -42,9 +42,9 @@ class ContainerRequestUtils:
|
|||
passed_params: dict,
|
||||
) -> ContainerCreateOptionalRequestParams:
|
||||
"""Extract only valid container creation parameters from the passed parameters."""
|
||||
container_create_optional_params = ContainerCreateOptionalRequestParams()
|
||||
container_create_optional_params: Final = ContainerCreateOptionalRequestParams()
|
||||
|
||||
valid_params = [
|
||||
valid_params: Final = [
|
||||
"expires_after",
|
||||
"file_ids",
|
||||
"extra_headers",
|
||||
|
|
@ -63,10 +63,10 @@ class ContainerRequestUtils:
|
|||
container_create_optional_params: ContainerCreateOptionalRequestParams,
|
||||
) -> dict:
|
||||
"""Get the optional parameters for container creation."""
|
||||
supported_params = container_provider_config.get_supported_openai_params()
|
||||
supported_params: Final = container_provider_config.get_supported_openai_params()
|
||||
|
||||
# Filter out unsupported parameters
|
||||
filtered_params = {k: v for k, v in container_create_optional_params.items() if k in supported_params}
|
||||
filtered_params: Final = {k: v for k, v in container_create_optional_params.items() if k in supported_params}
|
||||
|
||||
return container_provider_config.map_openai_params(
|
||||
container_create_optional_params=filtered_params, # type: ignore
|
||||
|
|
@ -78,9 +78,9 @@ class ContainerRequestUtils:
|
|||
passed_params: dict,
|
||||
) -> ContainerListOptionalRequestParams:
|
||||
"""Extract only valid container list parameters from the passed parameters."""
|
||||
container_list_optional_params = ContainerListOptionalRequestParams()
|
||||
container_list_optional_params: Final = ContainerListOptionalRequestParams()
|
||||
|
||||
valid_params = [
|
||||
valid_params: Final = [
|
||||
"after",
|
||||
"limit",
|
||||
"order",
|
||||
|
|
@ -124,7 +124,7 @@ class ContainerRequestUtils:
|
|||
"""
|
||||
# Extract model_id from litellm_metadata
|
||||
litellm_metadata = litellm_metadata or {}
|
||||
model_info: dict[str, Any] = litellm_metadata.get("model_info", {}) or {}
|
||||
model_info: Final[dict[str, Any]] = litellm_metadata.get("model_info", {}) or {}
|
||||
model_id = model_info.get("id")
|
||||
|
||||
# Check if we should encode based on routing metadata
|
||||
|
|
@ -139,7 +139,7 @@ class ContainerRequestUtils:
|
|||
should_encode = True
|
||||
# Extract model_id from target_model_names if not already set
|
||||
if model_id is None:
|
||||
target_models = extra_body["target_model_names"]
|
||||
target_models: Final = extra_body["target_model_names"]
|
||||
# Use first model as model_id for encoding
|
||||
if isinstance(target_models, str):
|
||||
model_id = target_models.split(",")[0].strip()
|
||||
|
|
@ -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 = 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,
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
import logging
|
||||
import time
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING, Any, Literal, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, cast
|
||||
|
||||
from httpx import Response
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -131,14 +131,14 @@ else:
|
|||
LitellmLoggingObject = Any
|
||||
|
||||
# Pre-resolved CallTypes enum values for fast membership checks
|
||||
_A2A_CALL_TYPES = frozenset(
|
||||
_A2A_CALL_TYPES: Final = frozenset(
|
||||
{
|
||||
CallTypes.asend_message.value,
|
||||
CallTypes.send_message.value,
|
||||
}
|
||||
)
|
||||
|
||||
_VIDEO_CALL_TYPES = frozenset(
|
||||
_VIDEO_CALL_TYPES: Final = frozenset(
|
||||
{
|
||||
CallTypes.create_video.value,
|
||||
CallTypes.acreate_video.value,
|
||||
|
|
@ -149,36 +149,36 @@ _VIDEO_CALL_TYPES = frozenset(
|
|||
}
|
||||
)
|
||||
|
||||
_SPEECH_CALL_TYPES = frozenset(
|
||||
_SPEECH_CALL_TYPES: Final = frozenset(
|
||||
{
|
||||
CallTypes.speech.value,
|
||||
CallTypes.aspeech.value,
|
||||
}
|
||||
)
|
||||
|
||||
_TRANSCRIPTION_CALL_TYPES = frozenset(
|
||||
_TRANSCRIPTION_CALL_TYPES: Final = frozenset(
|
||||
{
|
||||
CallTypes.atranscription.value,
|
||||
CallTypes.transcription.value,
|
||||
}
|
||||
)
|
||||
|
||||
_RERANK_CALL_TYPES = frozenset(
|
||||
_RERANK_CALL_TYPES: Final = frozenset(
|
||||
{
|
||||
CallTypes.rerank.value,
|
||||
CallTypes.arerank.value,
|
||||
}
|
||||
)
|
||||
|
||||
_SEARCH_CALL_TYPES = frozenset(
|
||||
_SEARCH_CALL_TYPES: Final = frozenset(
|
||||
{
|
||||
CallTypes.search.value,
|
||||
CallTypes.asearch.value,
|
||||
}
|
||||
)
|
||||
|
||||
_AREALTIME_CALL_TYPE = CallTypes.arealtime.value
|
||||
_MCP_CALL_TYPE = CallTypes.call_mcp_tool.value
|
||||
_AREALTIME_CALL_TYPE: Final = CallTypes.arealtime.value
|
||||
_MCP_CALL_TYPE: Final = CallTypes.call_mcp_tool.value
|
||||
|
||||
|
||||
def _cost_per_token_custom_pricing_helper(
|
||||
|
|
@ -201,24 +201,24 @@ def _cost_per_token_custom_pricing_helper(
|
|||
return None
|
||||
|
||||
if custom_cost_per_token is not None:
|
||||
input_cost_per_token = custom_cost_per_token["input_cost_per_token"]
|
||||
output_cost_per_token = custom_cost_per_token["output_cost_per_token"]
|
||||
input_cost_per_token: Final = custom_cost_per_token["input_cost_per_token"]
|
||||
output_cost_per_token: Final = custom_cost_per_token["output_cost_per_token"]
|
||||
|
||||
cache_read_input_token_cost = custom_cost_per_token.get(
|
||||
cache_read_input_token_cost: Final = custom_cost_per_token.get(
|
||||
"cache_read_input_token_cost",
|
||||
input_cost_per_token,
|
||||
)
|
||||
cache_creation_input_token_cost = custom_cost_per_token.get(
|
||||
cache_creation_input_token_cost: Final = custom_cost_per_token.get(
|
||||
"cache_creation_input_token_cost",
|
||||
input_cost_per_token,
|
||||
)
|
||||
|
||||
regular_prompt_tokens = max(
|
||||
regular_prompt_tokens: Final = max(
|
||||
prompt_tokens - cached_tokens - cache_creation_tokens,
|
||||
0,
|
||||
)
|
||||
|
||||
input_cost = (
|
||||
input_cost: Final = (
|
||||
regular_prompt_tokens * input_cost_per_token
|
||||
+ cached_tokens * cache_read_input_token_cost
|
||||
+ cache_creation_tokens * cache_creation_input_token_cost
|
||||
|
|
@ -284,13 +284,13 @@ def _transcription_usage_has_token_details(
|
|||
if usage_block is None:
|
||||
return False
|
||||
|
||||
prompt_tokens_val = getattr(usage_block, "prompt_tokens", 0) or 0
|
||||
completion_tokens_val = getattr(usage_block, "completion_tokens", 0) or 0
|
||||
prompt_details = getattr(usage_block, "prompt_tokens_details", None)
|
||||
prompt_tokens_val: Final = getattr(usage_block, "prompt_tokens", 0) or 0
|
||||
completion_tokens_val: Final = getattr(usage_block, "completion_tokens", 0) or 0
|
||||
prompt_details: Final = getattr(usage_block, "prompt_tokens_details", None)
|
||||
|
||||
if prompt_details is not None:
|
||||
audio_token_count = getattr(prompt_details, "audio_tokens", 0) or 0
|
||||
text_token_count = getattr(prompt_details, "text_tokens", 0) or 0
|
||||
audio_token_count: Final = getattr(prompt_details, "audio_tokens", 0) or 0
|
||||
text_token_count: Final = getattr(prompt_details, "text_tokens", 0) or 0
|
||||
if audio_token_count > 0 or text_token_count > 0:
|
||||
return True
|
||||
|
||||
|
|
@ -375,7 +375,7 @@ def cost_per_token(
|
|||
_is_anthropic_style = False
|
||||
|
||||
if usage_object is not None:
|
||||
_pt_details = getattr(usage_object, "prompt_tokens_details", None)
|
||||
_pt_details: Final = getattr(usage_object, "prompt_tokens_details", None)
|
||||
if _pt_details is not None:
|
||||
_cache_read_tokens = float(getattr(_pt_details, "cached_tokens", 0) or 0)
|
||||
# OpenAI-compatible providers report cache-write tokens under
|
||||
|
|
@ -385,8 +385,8 @@ def cost_per_token(
|
|||
getattr(_pt_details, "cache_write_tokens", 0) or getattr(_pt_details, "cache_creation_tokens", 0) or 0
|
||||
)
|
||||
|
||||
_anthropic_read = getattr(usage_object, "cache_read_input_tokens", None)
|
||||
_anthropic_create = getattr(usage_object, "cache_creation_input_tokens", None)
|
||||
_anthropic_read: Final = getattr(usage_object, "cache_read_input_tokens", None)
|
||||
_anthropic_create: Final = getattr(usage_object, "cache_creation_input_tokens", None)
|
||||
if _anthropic_read is not None or _anthropic_create is not None:
|
||||
_is_anthropic_style = True
|
||||
if _anthropic_read is not None:
|
||||
|
|
@ -407,7 +407,7 @@ def cost_per_token(
|
|||
if _is_anthropic_style:
|
||||
_normalized_prompt_tokens += _cache_read_tokens + _cache_creation_tokens
|
||||
|
||||
response_cost = _cost_per_token_custom_pricing_helper(
|
||||
response_cost: Final = _cost_per_token_custom_pricing_helper(
|
||||
prompt_tokens=_normalized_prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
response_time_ms=response_time_ms,
|
||||
|
|
@ -423,23 +423,23 @@ def cost_per_token(
|
|||
# given
|
||||
prompt_tokens_cost_usd_dollar: float = 0
|
||||
completion_tokens_cost_usd_dollar: float = 0
|
||||
model_cost_ref = litellm.model_cost
|
||||
model_cost_ref: Final = litellm.model_cost
|
||||
# Only callers that explicitly pass `custom_llm_provider` get the
|
||||
# dedup/prefix-join treatment. When provider is omitted, preserve legacy
|
||||
# behavior: `model_with_provider` stays equal to the raw `model` string
|
||||
# (provider is detected below for downstream use only).
|
||||
caller_supplied_provider = custom_llm_provider is not None
|
||||
caller_supplied_provider: Final = custom_llm_provider is not None
|
||||
|
||||
# `model` is normally a string, but callers that mock the transport can pass
|
||||
# non-string objects. Only run the string-based dedup/prefix-join when it is
|
||||
# actually a string — e.g. a MagicMock's `.startswith()` is always truthy and
|
||||
# its slices return new mocks, which would spin the dedup loop forever.
|
||||
model_is_str = isinstance(model, str)
|
||||
model_is_str: Final = isinstance(model, str)
|
||||
|
||||
# Router/proxy deployments may repeat the provider segment (e.g. model_name
|
||||
# "openai/openai/gpt-5.5"). Strip duplicated `{provider}/` chains before joining.
|
||||
if caller_supplied_provider and model_is_str:
|
||||
_dup_prefix = f"{custom_llm_provider}/"
|
||||
_dup_prefix: Final = f"{custom_llm_provider}/"
|
||||
while model.startswith(_dup_prefix):
|
||||
_remainder = model[len(_dup_prefix) :]
|
||||
if _remainder.startswith(_dup_prefix):
|
||||
|
|
@ -449,13 +449,13 @@ def cost_per_token(
|
|||
|
||||
model_with_provider = model
|
||||
if caller_supplied_provider:
|
||||
_prov_prefix = f"{custom_llm_provider}/"
|
||||
_prov_prefix: Final = f"{custom_llm_provider}/"
|
||||
if model_is_str and model.startswith(_prov_prefix):
|
||||
model_with_provider = model
|
||||
else:
|
||||
model_with_provider = f"{custom_llm_provider}/{model}"
|
||||
if region_name is not None:
|
||||
model_with_provider_and_region = f"{custom_llm_provider}/{region_name}/{model}"
|
||||
model_with_provider_and_region: Final = f"{custom_llm_provider}/{region_name}/{model}"
|
||||
if model_with_provider_and_region in model_cost_ref: # use region based pricing, if it's available
|
||||
model_with_provider = model_with_provider_and_region
|
||||
else:
|
||||
|
|
@ -464,7 +464,7 @@ def cost_per_token(
|
|||
assert custom_llm_provider is not None # caller-supplied or get_llm_provider
|
||||
|
||||
model_without_prefix = model
|
||||
model_parts = model.split("/", 1)
|
||||
model_parts: Final = model.split("/", 1)
|
||||
if len(model_parts) > 1:
|
||||
model_without_prefix = model_parts[1]
|
||||
else:
|
||||
|
|
@ -487,7 +487,7 @@ def cost_per_token(
|
|||
# see this https://learn.microsoft.com/en-us/azure/ai-services/openai/concepts/models
|
||||
if call_type == "speech" or call_type == "aspeech":
|
||||
speech_model_info = litellm.get_model_info(model=model_without_prefix, custom_llm_provider=custom_llm_provider)
|
||||
cost_metric = select_cost_metric_for_model(speech_model_info)
|
||||
cost_metric: Final = select_cost_metric_for_model(speech_model_info)
|
||||
prompt_cost: float = 0.0
|
||||
completion_cost: float = 0.0
|
||||
if cost_metric == "cost_per_character":
|
||||
|
|
@ -574,7 +574,7 @@ def cost_per_token(
|
|||
optional_params=(response._hidden_params if response and hasattr(response, "_hidden_params") else None),
|
||||
)
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
cost_router = google_cost_router(
|
||||
cost_router: Final = google_cost_router(
|
||||
model=model_without_prefix,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
call_type=call_type,
|
||||
|
|
@ -643,7 +643,7 @@ def cost_per_token(
|
|||
service_tier=service_tier,
|
||||
)
|
||||
else:
|
||||
model_info = _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 (model_info.get("input_cost_per_token") or 0.0) > 0 or (model_info.get("output_cost_per_token") or 0.0) > 0:
|
||||
return generic_cost_per_token(
|
||||
|
|
@ -654,7 +654,7 @@ def cost_per_token(
|
|||
data_residency=data_residency,
|
||||
)
|
||||
|
||||
input_cost_per_second = model_info.get("input_cost_per_second")
|
||||
input_cost_per_second: Final = model_info.get("input_cost_per_second")
|
||||
if input_cost_per_second is not None and response_time_ms is not None:
|
||||
verbose_logger.debug(
|
||||
"For model=%s - input_cost_per_second: %s; response time: %s",
|
||||
|
|
@ -665,7 +665,7 @@ def cost_per_token(
|
|||
## COST PER SECOND ##
|
||||
prompt_tokens_cost_usd_dollar = input_cost_per_second * response_time_ms / 1000
|
||||
|
||||
output_cost_per_second = model_info.get("output_cost_per_second")
|
||||
output_cost_per_second: Final = model_info.get("output_cost_per_second")
|
||||
if output_cost_per_second is not None and response_time_ms is not None:
|
||||
verbose_logger.debug(
|
||||
"For model=%s - output_cost_per_second: %s; response time: %s",
|
||||
|
|
@ -688,12 +688,12 @@ def cost_per_token(
|
|||
def get_replicate_completion_pricing(completion_response: dict, total_time=0.0):
|
||||
# see https://replicate.com/pricing
|
||||
# for all litellm currently supported LLMs, almost all requests go to a100_80gb
|
||||
a100_80gb_price_per_second_public = (
|
||||
a100_80gb_price_per_second_public: Final = (
|
||||
DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND # assume all calls sent to A100 80GB for now
|
||||
)
|
||||
if total_time == 0.0: # total time is in ms
|
||||
start_time = completion_response.get("created", time.time())
|
||||
end_time = getattr(completion_response, "ended", time.time())
|
||||
start_time: Final = completion_response.get("created", time.time())
|
||||
end_time: Final = getattr(completion_response, "ended", time.time())
|
||||
total_time = end_time - start_time
|
||||
|
||||
return a100_80gb_price_per_second_public * total_time / 1000
|
||||
|
|
@ -747,11 +747,11 @@ def _select_model_name_for_cost_calc(
|
|||
completion_response_model = getattr(completion_response, "model", None)
|
||||
elif isinstance(completion_response, dict):
|
||||
completion_response_model = completion_response.get("model", None)
|
||||
hidden_params: dict | None = getattr(completion_response, "_hidden_params", None)
|
||||
hidden_params: Final[dict | None] = getattr(completion_response, "_hidden_params", None)
|
||||
|
||||
if custom_pricing is True:
|
||||
if router_model_id is not None and router_model_id in litellm.model_cost:
|
||||
entry = litellm.model_cost[router_model_id]
|
||||
entry: Final = litellm.model_cost[router_model_id]
|
||||
if (
|
||||
entry.get("input_cost_per_token") is not None
|
||||
or entry.get("input_cost_per_second") is not None
|
||||
|
|
@ -796,7 +796,7 @@ def _model_contains_known_llm_provider(model: str) -> bool:
|
|||
"""
|
||||
Check if the model contains a known llm provider
|
||||
"""
|
||||
_provider_prefix = model.split("/")[0]
|
||||
_provider_prefix: Final = model.split("/")[0]
|
||||
return _provider_prefix in LlmProvidersSet
|
||||
|
||||
|
||||
|
|
@ -818,7 +818,7 @@ def _get_response_model(completion_response: Any) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
_GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER: dict = {
|
||||
_GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER: Final[dict] = {
|
||||
# ON_DEMAND_PRIORITY maps to "priority" — selects input_cost_per_token_priority, etc.
|
||||
"ON_DEMAND_PRIORITY": "priority",
|
||||
# FLEX / BATCH maps to "flex" — selects input_cost_per_token_flex, etc.
|
||||
|
|
@ -844,7 +844,7 @@ def _map_traffic_type_to_service_tier(traffic_type: str | None) -> str | None:
|
|||
"""
|
||||
if traffic_type is None:
|
||||
return None
|
||||
service_tier = _GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER.get(str(traffic_type).upper())
|
||||
service_tier: Final = _GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER.get(str(traffic_type).upper())
|
||||
return service_tier
|
||||
|
||||
|
||||
|
|
@ -865,7 +865,7 @@ def _normalize_service_tier(service_tier: object) -> str | None:
|
|||
def _get_usage_object(
|
||||
completion_response: Any,
|
||||
) -> Usage | None:
|
||||
usage_obj = cast(
|
||||
usage_obj: Final = cast(
|
||||
Usage | ResponseAPIUsage | dict | BaseModel,
|
||||
(
|
||||
completion_response.get("usage")
|
||||
|
|
@ -950,14 +950,14 @@ def _apply_cost_discount(
|
|||
Returns:
|
||||
Tuple of (final_cost, discount_percent, discount_amount)
|
||||
"""
|
||||
original_cost = base_cost
|
||||
original_cost: Final = base_cost
|
||||
discount_percent = 0.0
|
||||
discount_amount = 0.0
|
||||
|
||||
if custom_llm_provider and custom_llm_provider in litellm.cost_discount_config:
|
||||
discount_percent = litellm.cost_discount_config[custom_llm_provider]
|
||||
discount_amount = original_cost * discount_percent
|
||||
final_cost = original_cost - discount_amount
|
||||
final_cost: Final = original_cost - discount_amount
|
||||
|
||||
if verbose_logger.isEnabledFor(logging.DEBUG):
|
||||
verbose_logger.debug(
|
||||
|
|
@ -984,7 +984,7 @@ def _apply_cost_margin(
|
|||
Returns:
|
||||
Tuple of (final_cost, margin_percent, margin_fixed_amount, margin_total_amount)
|
||||
"""
|
||||
original_cost = base_cost
|
||||
original_cost: Final = base_cost
|
||||
margin_percent = 0.0
|
||||
margin_fixed_amount = 0.0
|
||||
margin_total_amount = 0.0
|
||||
|
|
@ -1022,7 +1022,7 @@ def _apply_cost_margin(
|
|||
margin_fixed_amount = float(margin_config["fixed_amount"])
|
||||
margin_total_amount += margin_fixed_amount
|
||||
|
||||
final_cost = original_cost + margin_total_amount
|
||||
final_cost: Final = original_cost + margin_total_amount
|
||||
|
||||
if verbose_logger.isEnabledFor(logging.DEBUG):
|
||||
verbose_logger.debug(
|
||||
|
|
@ -1180,7 +1180,7 @@ def completion_cost(
|
|||
cache_creation_input_tokens: int | None = None
|
||||
cache_read_input_tokens: int | None = None
|
||||
audio_transcription_file_duration: float = 0.0
|
||||
cost_per_token_usage_object: Usage | None = _get_usage_object(completion_response=completion_response)
|
||||
cost_per_token_usage_object: Final[Usage | None] = _get_usage_object(completion_response=completion_response)
|
||||
rerank_billed_units: RerankBilledUnits | None = None
|
||||
|
||||
# Extract service_tier from optional_params if not provided directly
|
||||
|
|
@ -1207,7 +1207,7 @@ def completion_cost(
|
|||
|
||||
service_tier = _normalize_service_tier(service_tier)
|
||||
|
||||
selected_model = _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,
|
||||
|
|
@ -1216,7 +1216,7 @@ def completion_cost(
|
|||
router_model_id=router_model_id,
|
||||
)
|
||||
|
||||
potential_model_names = [
|
||||
potential_model_names: Final = [
|
||||
selected_model,
|
||||
_get_response_model(completion_response),
|
||||
]
|
||||
|
|
@ -1691,9 +1691,9 @@ def get_response_cost_from_hidden_params(
|
|||
else:
|
||||
_hidden_params_dict = hidden_params
|
||||
|
||||
additional_headers = _hidden_params_dict.get("additional_headers", {})
|
||||
additional_headers: Final = _hidden_params_dict.get("additional_headers", {})
|
||||
if additional_headers and "llm_provider-x-litellm-response-cost" in additional_headers:
|
||||
response_cost = additional_headers["llm_provider-x-litellm-response-cost"]
|
||||
response_cost: Final = additional_headers["llm_provider-x-litellm-response-cost"]
|
||||
if response_cost is None:
|
||||
return None
|
||||
return float(additional_headers["llm_provider-x-litellm-response-cost"])
|
||||
|
|
@ -1761,7 +1761,7 @@ def response_cost_calculator(
|
|||
if isinstance(response_object, BaseModel):
|
||||
if hasattr(response_object, "_hidden_params"):
|
||||
response_object._hidden_params["optional_params"] = optional_params
|
||||
provider_response_cost = get_response_cost_from_hidden_params(response_object._hidden_params)
|
||||
provider_response_cost: Final = get_response_cost_from_hidden_params(response_object._hidden_params)
|
||||
if provider_response_cost is not None:
|
||||
return provider_response_cost
|
||||
|
||||
|
|
@ -1818,7 +1818,7 @@ def ocr_cost(
|
|||
except Exception:
|
||||
model_info = None
|
||||
|
||||
credits = getattr(response.usage_info, "credits", None)
|
||||
credits: Final = getattr(response.usage_info, "credits", None)
|
||||
cost_per_credit = None
|
||||
if model_info is not None:
|
||||
cost_per_credit = model_info.get("ocr_cost_per_credit")
|
||||
|
|
@ -1829,7 +1829,7 @@ def ocr_cost(
|
|||
if model_info is not None:
|
||||
ocr_cost_per_page = model_info.get("ocr_cost_per_page")
|
||||
|
||||
pages_processed = response.usage_info.pages_processed
|
||||
pages_processed: Final = response.usage_info.pages_processed
|
||||
if pages_processed is None:
|
||||
if cost_per_credit is not None or ocr_cost_per_page is None:
|
||||
# Surface missing usage data instead of silently under-reporting
|
||||
|
|
@ -1862,7 +1862,7 @@ def ocr_cost(
|
|||
)
|
||||
return 0.0, 0.0
|
||||
|
||||
total_ocr_processing_cost: float = ocr_cost_per_page * pages_processed
|
||||
total_ocr_processing_cost: Final[float] = ocr_cost_per_page * pages_processed
|
||||
return total_ocr_processing_cost, 0.0
|
||||
|
||||
|
||||
|
|
@ -1884,7 +1884,7 @@ def vector_store_search_cost(
|
|||
model=model,
|
||||
)
|
||||
|
||||
config = ProviderConfigManager.get_provider_vector_stores_config(
|
||||
config: Final = ProviderConfigManager.get_provider_vector_stores_config(
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
api_type=api_type,
|
||||
)
|
||||
|
|
@ -1910,7 +1910,7 @@ def rerank_cost(
|
|||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
try:
|
||||
config = ProviderConfigManager.get_provider_rerank_config(
|
||||
config: Final = ProviderConfigManager.get_provider_rerank_config(
|
||||
model=model,
|
||||
api_base=None,
|
||||
present_version_params=[],
|
||||
|
|
@ -1973,19 +1973,19 @@ def default_image_cost_calculator(
|
|||
if custom_llm_provider and model.startswith(f"{custom_llm_provider}/"):
|
||||
model_name_without_custom_llm_provider = model.replace(f"{custom_llm_provider}/", "")
|
||||
base_model_name = f"{custom_llm_provider}/{size_str}/{model_name_without_custom_llm_provider}"
|
||||
model_name_with_quality = f"{quality}/{base_model_name}" if quality else base_model_name
|
||||
model_name_with_quality: Final = f"{quality}/{base_model_name}" if quality else base_model_name
|
||||
|
||||
# gpt-image-1 models use low, medium, high quality. If user did not specify quality, use medium fot gpt-image-1 model family
|
||||
model_name_with_v2_quality = f"{ImageGenerationRequestQuality.HIGH.value}/{base_model_name}"
|
||||
model_name_with_v2_quality: Final = f"{ImageGenerationRequestQuality.HIGH.value}/{base_model_name}"
|
||||
|
||||
verbose_logger.debug("Looking up cost for models: %s, %s", model_name_with_quality, base_model_name)
|
||||
|
||||
model_without_provider = f"{size_str}/{model.split('/')[-1]}"
|
||||
model_without_provider: Final = f"{size_str}/{model.split('/')[-1]}"
|
||||
model_with_quality_without_provider = f"{quality}/{model_without_provider}" if quality else model_without_provider
|
||||
|
||||
# Try model with quality first, fall back to base model name
|
||||
cost_info: dict | None = None
|
||||
models_to_check: list[str | None] = [
|
||||
models_to_check: Final[list[str | None]] = [
|
||||
model_name_with_quality,
|
||||
base_model_name,
|
||||
model_name_with_v2_quality,
|
||||
|
|
@ -2050,10 +2050,10 @@ def default_video_cost_calculator(
|
|||
|
||||
verbose_logger.debug("Looking up cost for video model: %s", base_model_name)
|
||||
|
||||
model_without_provider = model.split("/")[-1]
|
||||
model_without_provider: Final = model.split("/")[-1]
|
||||
|
||||
# Try model with provider first, fall back to base model name
|
||||
models_to_check: list[str | None] = [
|
||||
models_to_check: Final[list[str | None]] = [
|
||||
base_model_name,
|
||||
model,
|
||||
model_without_provider,
|
||||
|
|
@ -2066,7 +2066,7 @@ def default_video_cost_calculator(
|
|||
|
||||
# If still not found, try with custom_llm_provider prefix
|
||||
if cost_info is None and custom_llm_provider:
|
||||
prefixed_model = f"{custom_llm_provider}/{model}"
|
||||
prefixed_model: Final = f"{custom_llm_provider}/{model}"
|
||||
if prefixed_model in litellm.model_cost:
|
||||
cost_info = litellm.model_cost[prefixed_model]
|
||||
|
||||
|
|
@ -2074,11 +2074,11 @@ def default_video_cost_calculator(
|
|||
raise Exception(f"Model not found in cost map for model={model}")
|
||||
|
||||
# Check for video-specific cost per second first
|
||||
video_cost_per_second = cost_info.get("output_cost_per_video_per_second")
|
||||
video_cost_per_second: Final = cost_info.get("output_cost_per_video_per_second")
|
||||
if video_cost_per_second is not None:
|
||||
return video_cost_per_second * duration_seconds
|
||||
|
||||
output_cost_per_second = _video_output_cost_per_second(cost_info, video_resolution)
|
||||
output_cost_per_second: Final = _video_output_cost_per_second(cost_info, video_resolution)
|
||||
if output_cost_per_second is not None:
|
||||
return output_cost_per_second * duration_seconds
|
||||
|
||||
|
|
@ -2133,7 +2133,7 @@ def batch_cost_calculator(
|
|||
# but carries no pricing fields. Fall back to the global pricing table so
|
||||
# that standard model pricing is used instead of silently returning $0.
|
||||
try:
|
||||
global_info = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
global_info: Final = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
if global_info:
|
||||
model_info = global_info
|
||||
except Exception:
|
||||
|
|
@ -2142,31 +2142,31 @@ def batch_cost_calculator(
|
|||
if not model_info:
|
||||
return 0.0, 0.0
|
||||
|
||||
input_cost_per_token_batches = model_info.get("input_cost_per_token_batches")
|
||||
input_cost_per_token = model_info.get("input_cost_per_token")
|
||||
output_cost_per_token_batches = model_info.get("output_cost_per_token_batches")
|
||||
output_cost_per_token = model_info.get("output_cost_per_token")
|
||||
input_cost_per_token_batches: Final = model_info.get("input_cost_per_token_batches")
|
||||
input_cost_per_token: Final = model_info.get("input_cost_per_token")
|
||||
output_cost_per_token_batches: Final = model_info.get("output_cost_per_token_batches")
|
||||
output_cost_per_token: Final = model_info.get("output_cost_per_token")
|
||||
total_prompt_cost = 0.0
|
||||
total_completion_cost = 0.0
|
||||
if input_cost_per_token_batches:
|
||||
total_prompt_cost = usage.prompt_tokens * input_cost_per_token_batches
|
||||
elif input_cost_per_token:
|
||||
details = _parse_prompt_tokens_details(usage)
|
||||
cache_read_tokens = details["cache_hit_tokens"]
|
||||
cache_creation_tokens = details["cache_creation_tokens"]
|
||||
details: Final = _parse_prompt_tokens_details(usage)
|
||||
cache_read_tokens: Final = details["cache_hit_tokens"]
|
||||
cache_creation_tokens: Final = details["cache_creation_tokens"]
|
||||
|
||||
# Subtract cached tokens from prompt_tokens before calculating cost
|
||||
# Fixes issue where cached tokens are being charged again
|
||||
base_input_tokens = get_billable_input_tokens(usage) - cache_creation_tokens
|
||||
base_input_tokens: Final = get_billable_input_tokens(usage) - cache_creation_tokens
|
||||
total_prompt_cost = (
|
||||
base_input_tokens * (input_cost_per_token) / 2
|
||||
) # batch cost is usually half of the regular token cost
|
||||
|
||||
# Add cache read cost if applicable
|
||||
cache_read_cost_key = _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 = model_info.get("cache_creation_input_token_cost") or input_cost_per_token
|
||||
cache_creation_cost: Final = model_info.get("cache_creation_input_token_cost") or input_cost_per_token
|
||||
total_prompt_cost += cache_creation_tokens * cache_creation_cost / 2
|
||||
if output_cost_per_token_batches:
|
||||
total_completion_cost = usage.completion_tokens * output_cost_per_token_batches
|
||||
|
|
@ -2175,7 +2175,7 @@ def batch_cost_calculator(
|
|||
usage.completion_tokens * (output_cost_per_token) / 2
|
||||
) # batch cost is usually half of the regular token cost
|
||||
|
||||
uplift = _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
|
||||
|
|
@ -2184,7 +2184,7 @@ def batch_cost_calculator(
|
|||
|
||||
|
||||
def _summable_prompt_token_fields(prompt_tokens_details: BaseModel) -> list[str]:
|
||||
field_names = list(type(prompt_tokens_details).model_fields)
|
||||
field_names: Final = list(type(prompt_tokens_details).model_fields)
|
||||
if getattr(prompt_tokens_details, "cache_write_tokens", None) is None:
|
||||
return field_names
|
||||
return [attr for attr in field_names if attr != "cache_creation_tokens"]
|
||||
|
|
@ -2202,7 +2202,7 @@ class BaseTokenUsageProcessor:
|
|||
Usage,
|
||||
)
|
||||
|
||||
combined = Usage()
|
||||
combined: Final = Usage()
|
||||
|
||||
# Sum basic token counts
|
||||
for usage in usage_objects:
|
||||
|
|
@ -2268,11 +2268,11 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor):
|
|||
"""
|
||||
Collect usage from realtime stream results
|
||||
"""
|
||||
response_done_events: list[OpenAIRealtimeStreamResponseBaseObject] = cast(
|
||||
response_done_events: Final[list[OpenAIRealtimeStreamResponseBaseObject]] = cast(
|
||||
list[OpenAIRealtimeStreamResponseBaseObject],
|
||||
[result for result in results if result["type"] == "response.done"],
|
||||
)
|
||||
usage_objects: list[Usage] = []
|
||||
usage_objects: Final[list[Usage]] = []
|
||||
for result in response_done_events:
|
||||
usage_object = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
result["response"].get("usage", {})
|
||||
|
|
@ -2288,7 +2288,7 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor):
|
|||
Collect and combine usage from realtime stream results
|
||||
"""
|
||||
collected_usage_objects = RealtimeAPITokenUsageProcessor.collect_usage_from_realtime_stream_results(results)
|
||||
combined_usage_object = RealtimeAPITokenUsageProcessor.combine_usage_objects(collected_usage_objects)
|
||||
combined_usage_object: Final = RealtimeAPITokenUsageProcessor.combine_usage_objects(collected_usage_objects)
|
||||
return combined_usage_object
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -2301,7 +2301,7 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor):
|
|||
)
|
||||
|
||||
|
||||
_TRANSCRIPTION_COMPLETED_EVENT_TYPE = "conversation.item.input_audio_transcription.completed"
|
||||
_TRANSCRIPTION_COMPLETED_EVENT_TYPE: Final = "conversation.item.input_audio_transcription.completed"
|
||||
|
||||
|
||||
def handle_realtime_stream_cost_calculation(
|
||||
|
|
@ -2321,7 +2321,7 @@ def handle_realtime_stream_cost_calculation(
|
|||
results: A list of OpenAIRealtimeStreamBaseObject objects
|
||||
"""
|
||||
received_model = None
|
||||
potential_model_names = []
|
||||
potential_model_names: Final = []
|
||||
for result in results:
|
||||
if result["type"] == "session.created":
|
||||
received_model = cast(OpenAIRealtimeStreamSessionEvents, result)["session"].get("model", None)
|
||||
|
|
@ -2346,7 +2346,7 @@ def handle_realtime_stream_cost_calculation(
|
|||
input_cost_per_token += _input_cost_per_token
|
||||
output_cost_per_token += _output_cost_per_token
|
||||
break # exit if we find a valid model
|
||||
transcription_cost = (
|
||||
transcription_cost: Final = (
|
||||
handle_realtime_transcription_cost_calculation(
|
||||
results=results,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -2355,7 +2355,7 @@ def handle_realtime_stream_cost_calculation(
|
|||
if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results)
|
||||
else 0.0
|
||||
)
|
||||
total_cost = input_cost_per_token + output_cost_per_token + transcription_cost
|
||||
total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost
|
||||
|
||||
_store_cost_breakdown_in_logging_obj(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
|
|
@ -2384,13 +2384,13 @@ def handle_realtime_transcription_cost_calculation(
|
|||
- {"type": "duration", "seconds": <float>} → priced via input_cost_per_second
|
||||
- {"type": "tokens", "input_tokens": ...} → priced via input/audio token cost
|
||||
"""
|
||||
completed_events = [
|
||||
completed_events: Final = [
|
||||
cast(dict, result) for result in results if result.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE
|
||||
]
|
||||
if not completed_events:
|
||||
return 0.0
|
||||
|
||||
model_name = _get_transcription_model_name_from_results(results) or litellm_model_name
|
||||
model_name: Final = _get_transcription_model_name_from_results(results) or litellm_model_name
|
||||
try:
|
||||
model_info = litellm.get_model_info(model=model_name, custom_llm_provider=custom_llm_provider)
|
||||
except Exception:
|
||||
|
|
@ -2427,20 +2427,20 @@ def _get_transcription_model_name_from_results(
|
|||
def _transcription_usage_cost(usage: dict, model_info: ModelInfo | None) -> float:
|
||||
if model_info is None:
|
||||
return 0.0
|
||||
usage_type = usage.get("type")
|
||||
usage_type: Final = usage.get("type")
|
||||
if usage_type == "duration":
|
||||
seconds = usage.get("seconds") or 0.0
|
||||
per_second = model_info.get("input_cost_per_second") or 0.0
|
||||
seconds: Final = usage.get("seconds") or 0.0
|
||||
per_second: Final = model_info.get("input_cost_per_second") or 0.0
|
||||
return float(seconds) * float(per_second)
|
||||
if usage_type == "tokens":
|
||||
input_token_details = usage.get("input_token_details") or {}
|
||||
audio_tokens = input_token_details.get("audio_tokens") or 0
|
||||
text_tokens = input_token_details.get("text_tokens") or 0
|
||||
output_tokens = usage.get("output_tokens") or 0
|
||||
audio_cost = float(audio_tokens) * float(
|
||||
input_token_details: Final = usage.get("input_token_details") or {}
|
||||
audio_tokens: Final = input_token_details.get("audio_tokens") or 0
|
||||
text_tokens: Final = input_token_details.get("text_tokens") or 0
|
||||
output_tokens: Final = usage.get("output_tokens") or 0
|
||||
audio_cost: Final = float(audio_tokens) * float(
|
||||
model_info.get("input_cost_per_audio_token") or model_info.get("input_cost_per_token") or 0.0
|
||||
)
|
||||
text_cost = float(text_tokens) * float(model_info.get("input_cost_per_token") or 0.0)
|
||||
output_cost = float(output_tokens) * float(model_info.get("output_cost_per_token") or 0.0)
|
||||
text_cost: Final = float(text_tokens) * float(model_info.get("input_cost_per_token") or 0.0)
|
||||
output_cost: Final = float(output_tokens) * float(model_info.get("output_cost_per_token") or 0.0)
|
||||
return audio_cost + text_cost + output_cost
|
||||
return 0.0
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
Handler for transforming /chat/completions api requests to litellm.responses requests
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
|
|
@ -32,23 +32,23 @@ class SpeechToCompletionBridgeHandler:
|
|||
def validate_input_kwargs(self, kwargs: dict) -> SpeechToCompletionBridgeHandlerInputKwargs:
|
||||
from litellm import LiteLLMLoggingObj
|
||||
|
||||
model = kwargs.get("model")
|
||||
model: Final = kwargs.get("model")
|
||||
if model is None or not isinstance(model, str):
|
||||
raise ValueError("model is required")
|
||||
|
||||
custom_llm_provider = kwargs.get("custom_llm_provider")
|
||||
custom_llm_provider: Final = kwargs.get("custom_llm_provider")
|
||||
if custom_llm_provider is None or not isinstance(custom_llm_provider, str):
|
||||
raise ValueError("custom_llm_provider is required")
|
||||
|
||||
input = kwargs.get("input")
|
||||
input: Final = kwargs.get("input")
|
||||
if input is None or not isinstance(input, str):
|
||||
raise ValueError("input is required")
|
||||
|
||||
optional_params = kwargs.get("optional_params")
|
||||
optional_params: Final = kwargs.get("optional_params")
|
||||
if optional_params is None or not isinstance(optional_params, dict):
|
||||
raise ValueError("optional_params is required")
|
||||
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
litellm_params: Final = kwargs.get("litellm_params")
|
||||
if litellm_params is None or not isinstance(litellm_params, dict):
|
||||
raise ValueError("litellm_params is required")
|
||||
|
||||
|
|
@ -60,7 +60,7 @@ class SpeechToCompletionBridgeHandler:
|
|||
if headers is None or not isinstance(headers, dict):
|
||||
raise ValueError("headers is required")
|
||||
|
||||
logging_obj = kwargs.get("logging_obj")
|
||||
logging_obj: Final = kwargs.get("logging_obj")
|
||||
if logging_obj is None or not isinstance(logging_obj, LiteLLMLoggingObj):
|
||||
raise ValueError("logging_obj is required")
|
||||
|
||||
|
|
@ -86,11 +86,11 @@ class SpeechToCompletionBridgeHandler:
|
|||
logging_obj: "LiteLLMLoggingObj",
|
||||
custom_llm_provider: str,
|
||||
) -> "HttpxBinaryResponseContent":
|
||||
received_args = locals()
|
||||
received_args: Final = locals()
|
||||
from litellm import completion
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
validated_kwargs = self.validate_input_kwargs(received_args)
|
||||
validated_kwargs: Final = self.validate_input_kwargs(received_args)
|
||||
model = validated_kwargs["model"]
|
||||
input = validated_kwargs["input"]
|
||||
optional_params = validated_kwargs["optional_params"]
|
||||
|
|
@ -100,7 +100,7 @@ class SpeechToCompletionBridgeHandler:
|
|||
custom_llm_provider = validated_kwargs["custom_llm_provider"]
|
||||
voice = validated_kwargs["voice"]
|
||||
|
||||
request_data = self.transformation_handler.transform_request(
|
||||
request_data: Final = self.transformation_handler.transform_request(
|
||||
model=model,
|
||||
input=input,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -111,7 +111,7 @@ class SpeechToCompletionBridgeHandler:
|
|||
voice=voice,
|
||||
)
|
||||
|
||||
result = completion(
|
||||
result: Final = completion(
|
||||
**request_data,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import TYPE_CHECKING, cast
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
|
||||
from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS
|
||||
|
||||
|
|
@ -20,7 +20,7 @@ class SpeechToCompletionBridgeTransformationHandler:
|
|||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
custom_llm_provider: str,
|
||||
) -> dict:
|
||||
passed_optional_params = {}
|
||||
passed_optional_params: Final = {}
|
||||
for op in optional_params:
|
||||
if op in OPENAI_CHAT_COMPLETION_PARAMS:
|
||||
passed_optional_params[op] = optional_params[op]
|
||||
|
|
@ -66,13 +66,13 @@ class SpeechToCompletionBridgeTransformationHandler:
|
|||
import struct
|
||||
|
||||
# WAV header parameters
|
||||
byte_rate = sample_rate * channels * 2 # 2 bytes per sample (16-bit)
|
||||
block_align = channels * 2
|
||||
data_size = len(pcm_data)
|
||||
file_size = 36 + data_size
|
||||
byte_rate: Final = sample_rate * channels * 2 # 2 bytes per sample (16-bit)
|
||||
block_align: Final = channels * 2
|
||||
data_size: Final = len(pcm_data)
|
||||
file_size: Final = 36 + data_size
|
||||
|
||||
# Create WAV header
|
||||
wav_header = struct.pack(
|
||||
wav_header: Final = struct.pack(
|
||||
"<4sI4s4sIHHIIHH4sI",
|
||||
b"RIFF", # Chunk ID
|
||||
file_size, # Chunk Size
|
||||
|
|
@ -103,17 +103,17 @@ class SpeechToCompletionBridgeTransformationHandler:
|
|||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
from litellm.types.utils import Choices
|
||||
|
||||
audio_part = cast(Choices, model_response.choices[0]).message.audio
|
||||
audio_part: Final = cast(Choices, model_response.choices[0]).message.audio
|
||||
if audio_part is None:
|
||||
raise ValueError("No audio part found in the response")
|
||||
audio_content = audio_part.data
|
||||
audio_content: Final = audio_part.data
|
||||
|
||||
# Decode base64 to get binary content
|
||||
binary_data = base64.b64decode(audio_content)
|
||||
|
||||
# Check if this is a Gemini TTS model that returns raw PCM16 data
|
||||
model = getattr(model_response, "model", "")
|
||||
headers = {}
|
||||
model: Final = getattr(model_response, "model", "")
|
||||
headers: Final = {}
|
||||
if self._is_gemini_tts_model(model):
|
||||
# Convert PCM16 to WAV format for proper audio file playback
|
||||
binary_data = self._convert_pcm16_to_wav(binary_data)
|
||||
|
|
@ -122,5 +122,5 @@ class SpeechToCompletionBridgeTransformationHandler:
|
|||
headers["Content-Type"] = "audio/mpeg"
|
||||
|
||||
# Create an httpx.Response object
|
||||
response = httpx.Response(status_code=200, content=binary_data, headers=headers)
|
||||
response: Final = httpx.Response(status_code=200, content=binary_data, headers=headers)
|
||||
return HttpxBinaryResponseContent(response)
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import asyncio
|
|||
import contextvars
|
||||
from collections.abc import Coroutine
|
||||
from functools import partial
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -36,7 +36,7 @@ from litellm.utils import ProviderConfigManager, client
|
|||
|
||||
# Initialize HTTP handler
|
||||
base_llm_http_handler = BaseLLMHTTPHandler()
|
||||
DEFAULT_OPENAI_API_BASE = "https://api.openai.com"
|
||||
DEFAULT_OPENAI_API_BASE: Final = "https://api.openai.com"
|
||||
|
||||
|
||||
@client
|
||||
|
|
@ -70,12 +70,12 @@ async def acreate_eval(
|
|||
Returns:
|
||||
Eval object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["acreate_eval"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
create_eval,
|
||||
data_source_config=data_source_config,
|
||||
testing_criteria=testing_criteria,
|
||||
|
|
@ -89,9 +89,9 @@ async def acreate_eval(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -139,14 +139,14 @@ def create_eval(
|
|||
Returns:
|
||||
Eval object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("acreate_eval", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("acreate_eval", False) is True
|
||||
|
||||
# Get LiteLLM parameters
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# Determine provider
|
||||
if custom_llm_provider is None:
|
||||
|
|
@ -161,7 +161,7 @@ def create_eval(
|
|||
raise ValueError(f"CREATE eval is not supported for {custom_llm_provider}")
|
||||
|
||||
# Build create request
|
||||
create_request: CreateEvalRequest = {
|
||||
create_request: Final[CreateEvalRequest] = {
|
||||
"data_source_config": data_source_config, # type: ignore
|
||||
"testing_criteria": testing_criteria, # type: ignore
|
||||
}
|
||||
|
|
@ -177,15 +177,15 @@ def create_eval(
|
|||
headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
# Transform request
|
||||
request_body = evals_api_provider_config.transform_create_eval_request(
|
||||
request_body: Final = evals_api_provider_config.transform_create_eval_request(
|
||||
create_request=create_request,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
# Get API base and URL
|
||||
api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
url = evals_api_provider_config.get_complete_url(api_base=api_base, endpoint="evals")
|
||||
api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
url: Final = evals_api_provider_config.get_complete_url(api_base=api_base, endpoint="evals")
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
|
|
@ -199,7 +199,7 @@ def create_eval(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response = base_llm_http_handler.create_eval_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.create_eval_handler( # type: ignore
|
||||
url=url,
|
||||
request_body=request_body,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
|
|
@ -255,12 +255,12 @@ async def alist_evals(
|
|||
Returns:
|
||||
ListEvalsResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["alist_evals"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
list_evals,
|
||||
limit=limit,
|
||||
after=after,
|
||||
|
|
@ -274,9 +274,9 @@ async def alist_evals(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -324,14 +324,14 @@ def list_evals(
|
|||
Returns:
|
||||
ListEvalsResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("alist_evals", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("alist_evals", False) is True
|
||||
|
||||
# Get LiteLLM parameters
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# Determine provider
|
||||
if custom_llm_provider is None:
|
||||
|
|
@ -346,7 +346,7 @@ def list_evals(
|
|||
raise ValueError(f"LIST evals is not supported for {custom_llm_provider}")
|
||||
|
||||
# Build list parameters
|
||||
list_params: ListEvalsParams = {}
|
||||
list_params: Final[ListEvalsParams] = {}
|
||||
if limit is not None:
|
||||
list_params["limit"] = limit
|
||||
if after is not None:
|
||||
|
|
@ -385,7 +385,7 @@ def list_evals(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response = base_llm_http_handler.list_evals_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.list_evals_handler( # type: ignore
|
||||
url=url,
|
||||
query_params=query_params,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
|
|
@ -433,12 +433,12 @@ async def aget_eval(
|
|||
Returns:
|
||||
Eval object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["aget_eval"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
get_eval,
|
||||
eval_id=eval_id,
|
||||
extra_headers=extra_headers,
|
||||
|
|
@ -448,9 +448,9 @@ async def aget_eval(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -490,14 +490,14 @@ def get_eval(
|
|||
Returns:
|
||||
Eval object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("aget_eval", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("aget_eval", False) is True
|
||||
|
||||
# Get LiteLLM parameters
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# Determine provider
|
||||
if custom_llm_provider is None:
|
||||
|
|
@ -516,7 +516,7 @@ def get_eval(
|
|||
headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
# Transform request
|
||||
api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
url, headers = evals_api_provider_config.transform_get_eval_request(
|
||||
eval_id=eval_id,
|
||||
api_base=api_base,
|
||||
|
|
@ -536,7 +536,7 @@ def get_eval(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response = base_llm_http_handler.get_eval_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.get_eval_handler( # type: ignore
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -589,12 +589,12 @@ async def aupdate_eval(
|
|||
Returns:
|
||||
Eval object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["aupdate_eval"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
update_eval,
|
||||
eval_id=eval_id,
|
||||
name=name,
|
||||
|
|
@ -607,9 +607,9 @@ async def aupdate_eval(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -655,14 +655,14 @@ def update_eval(
|
|||
Returns:
|
||||
Eval object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("aupdate_eval", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("aupdate_eval", False) is True
|
||||
|
||||
# Get LiteLLM parameters
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# Determine provider
|
||||
if custom_llm_provider is None:
|
||||
|
|
@ -677,14 +677,14 @@ def update_eval(
|
|||
raise ValueError(f"UPDATE eval is not supported for {custom_llm_provider}")
|
||||
|
||||
# Build update request
|
||||
update_request: UpdateEvalRequest = {}
|
||||
update_request: Final[UpdateEvalRequest] = {}
|
||||
if name is not None:
|
||||
update_request["name"] = name
|
||||
|
||||
# Filter metadata to exclude internal LiteLLM fields
|
||||
if metadata is not None:
|
||||
# List of internal LiteLLM metadata keys that should NOT be sent to OpenAI
|
||||
internal_keys = {
|
||||
internal_keys: Final = {
|
||||
"headers",
|
||||
"requester_metadata",
|
||||
"user_api_key_hash",
|
||||
|
|
@ -717,7 +717,7 @@ def update_eval(
|
|||
"user_agent",
|
||||
}
|
||||
# Only include user-provided metadata keys
|
||||
filtered_metadata = {k: v for k, v in metadata.items() if k not in internal_keys}
|
||||
filtered_metadata: Final = {k: v for k, v in metadata.items() if k not in internal_keys}
|
||||
if filtered_metadata: # Only add if there's user metadata
|
||||
update_request["metadata"] = filtered_metadata
|
||||
|
||||
|
|
@ -730,7 +730,7 @@ def update_eval(
|
|||
headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
# Transform request
|
||||
api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
(
|
||||
url,
|
||||
headers,
|
||||
|
|
@ -755,7 +755,7 @@ def update_eval(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response = base_llm_http_handler.update_eval_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.update_eval_handler( # type: ignore
|
||||
url=url,
|
||||
request_body=request_body,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
|
|
@ -803,12 +803,12 @@ async def adelete_eval(
|
|||
Returns:
|
||||
DeleteEvalResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["adelete_eval"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
delete_eval,
|
||||
eval_id=eval_id,
|
||||
extra_headers=extra_headers,
|
||||
|
|
@ -818,9 +818,9 @@ async def adelete_eval(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -860,14 +860,14 @@ def delete_eval(
|
|||
Returns:
|
||||
DeleteEvalResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("adelete_eval", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("adelete_eval", False) is True
|
||||
|
||||
# Get LiteLLM parameters
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# Determine provider
|
||||
if custom_llm_provider is None:
|
||||
|
|
@ -886,7 +886,7 @@ def delete_eval(
|
|||
headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
# Transform request
|
||||
api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
url, headers = evals_api_provider_config.transform_delete_eval_request(
|
||||
eval_id=eval_id,
|
||||
api_base=api_base,
|
||||
|
|
@ -906,7 +906,7 @@ def delete_eval(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response = base_llm_http_handler.delete_eval_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.delete_eval_handler( # type: ignore
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -953,12 +953,12 @@ async def acancel_eval(
|
|||
Returns:
|
||||
CancelEvalResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["acancel_eval"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
cancel_eval,
|
||||
eval_id=eval_id,
|
||||
extra_headers=extra_headers,
|
||||
|
|
@ -968,9 +968,9 @@ async def acancel_eval(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -1010,14 +1010,14 @@ def cancel_eval(
|
|||
Returns:
|
||||
CancelEvalResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("acancel_eval", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("acancel_eval", False) is True
|
||||
|
||||
# Get LiteLLM parameters
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# Determine provider
|
||||
if custom_llm_provider is None:
|
||||
|
|
@ -1036,7 +1036,7 @@ def cancel_eval(
|
|||
headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
# Transform request
|
||||
api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
(
|
||||
url,
|
||||
headers,
|
||||
|
|
@ -1060,7 +1060,7 @@ def cancel_eval(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response = base_llm_http_handler.cancel_eval_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.cancel_eval_handler( # type: ignore
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -1120,12 +1120,12 @@ async def acreate_run(
|
|||
Returns:
|
||||
Run object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["acreate_run"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
create_run,
|
||||
eval_id=eval_id,
|
||||
data_source=data_source,
|
||||
|
|
@ -1139,9 +1139,9 @@ async def acreate_run(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -1189,14 +1189,14 @@ def create_run(
|
|||
Returns:
|
||||
Run object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("acreate_run", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("acreate_run", False) is True
|
||||
|
||||
# Get LiteLLM parameters
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# Determine provider
|
||||
if custom_llm_provider is None:
|
||||
|
|
@ -1211,7 +1211,7 @@ def create_run(
|
|||
raise ValueError(f"CREATE run is not supported for {custom_llm_provider}")
|
||||
|
||||
# Build create request
|
||||
create_request: CreateRunRequest = {
|
||||
create_request: Final[CreateRunRequest] = {
|
||||
"data_source": data_source, # type: ignore
|
||||
}
|
||||
if name is not None:
|
||||
|
|
@ -1228,7 +1228,7 @@ def create_run(
|
|||
headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
# Transform request
|
||||
api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
url, request_body = evals_api_provider_config.transform_create_run_request(
|
||||
eval_id=eval_id,
|
||||
create_request=create_request,
|
||||
|
|
@ -1248,7 +1248,7 @@ def create_run(
|
|||
)
|
||||
|
||||
# Make HTTP request (default 600s timeout for long-running operations)
|
||||
response = base_llm_http_handler.create_run_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.create_run_handler( # type: ignore
|
||||
url=url,
|
||||
request_body=request_body,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
|
|
@ -1304,12 +1304,12 @@ async def alist_runs(
|
|||
Returns:
|
||||
ListRunsResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["alist_runs"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
list_runs,
|
||||
eval_id=eval_id,
|
||||
limit=limit,
|
||||
|
|
@ -1323,9 +1323,9 @@ async def alist_runs(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -1373,14 +1373,14 @@ def list_runs(
|
|||
Returns:
|
||||
ListRunsResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("alist_runs", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("alist_runs", False) is True
|
||||
|
||||
# Get LiteLLM parameters
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# Determine provider
|
||||
if custom_llm_provider is None:
|
||||
|
|
@ -1395,7 +1395,7 @@ def list_runs(
|
|||
raise ValueError(f"LIST runs is not supported for {custom_llm_provider}")
|
||||
|
||||
# Build list parameters
|
||||
list_params: ListRunsParams = {}
|
||||
list_params: Final[ListRunsParams] = {}
|
||||
if limit is not None:
|
||||
list_params["limit"] = limit
|
||||
if after is not None:
|
||||
|
|
@ -1433,7 +1433,7 @@ def list_runs(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response = base_llm_http_handler.list_runs_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.list_runs_handler( # type: ignore
|
||||
url=url,
|
||||
query_params=query_params,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
|
|
@ -1483,12 +1483,12 @@ async def aget_run(
|
|||
Returns:
|
||||
Run object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["aget_run"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
get_run,
|
||||
eval_id=eval_id,
|
||||
run_id=run_id,
|
||||
|
|
@ -1499,9 +1499,9 @@ async def aget_run(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -1543,14 +1543,14 @@ def get_run(
|
|||
Returns:
|
||||
Run object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("aget_run", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("aget_run", False) is True
|
||||
|
||||
# Get LiteLLM parameters
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# Determine provider
|
||||
if custom_llm_provider is None:
|
||||
|
|
@ -1569,7 +1569,7 @@ def get_run(
|
|||
headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
# Transform request
|
||||
api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
url, headers = evals_api_provider_config.transform_get_run_request(
|
||||
eval_id=eval_id,
|
||||
run_id=run_id,
|
||||
|
|
@ -1590,7 +1590,7 @@ def get_run(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response = base_llm_http_handler.get_run_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.get_run_handler( # type: ignore
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -1639,12 +1639,12 @@ async def acancel_run(
|
|||
Returns:
|
||||
CancelRunResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["acancel_run"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
cancel_run,
|
||||
eval_id=eval_id,
|
||||
run_id=run_id,
|
||||
|
|
@ -1655,9 +1655,9 @@ async def acancel_run(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -1699,14 +1699,14 @@ def cancel_run(
|
|||
Returns:
|
||||
CancelRunResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("acancel_run", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("acancel_run", False) is True
|
||||
|
||||
# Get LiteLLM parameters
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# Determine provider
|
||||
if custom_llm_provider is None:
|
||||
|
|
@ -1725,7 +1725,7 @@ def cancel_run(
|
|||
headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
# Transform request
|
||||
api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
(
|
||||
url,
|
||||
headers,
|
||||
|
|
@ -1750,7 +1750,7 @@ def cancel_run(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response = base_llm_http_handler.cancel_run_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.cancel_run_handler( # type: ignore
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -1804,12 +1804,12 @@ async def adelete_run(
|
|||
Returns:
|
||||
RunDeleteResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["adelete_run"] = True
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
delete_run,
|
||||
eval_id=eval_id,
|
||||
run_id=run_id,
|
||||
|
|
@ -1820,9 +1820,9 @@ async def adelete_run(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -1864,14 +1864,14 @@ def delete_run(
|
|||
Returns:
|
||||
RunDeleteResponse object
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("adelete_run", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("adelete_run", False) is True
|
||||
|
||||
# Get LiteLLM parameters
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# Determine provider
|
||||
if custom_llm_provider is None:
|
||||
|
|
@ -1890,7 +1890,7 @@ def delete_run(
|
|||
headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
# Transform request
|
||||
api_base = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
api_base: Final = litellm_params.api_base or DEFAULT_OPENAI_API_BASE
|
||||
(
|
||||
url,
|
||||
headers,
|
||||
|
|
@ -1915,7 +1915,7 @@ def delete_run(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response = base_llm_http_handler.delete_run_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.delete_run_handler( # type: ignore
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@
|
|||
## LiteLLM versions of the OpenAI Exception Types
|
||||
|
||||
import enum
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
import httpx
|
||||
import openai
|
||||
|
|
@ -81,8 +81,8 @@ class RateLimitType(str, enum.Enum):
|
|||
"""Per-session max-iterations cap reached (agent-style flows)."""
|
||||
|
||||
|
||||
_RATE_LIMIT_CATEGORY_VALUES = frozenset(c.value for c in RateLimitErrorCategory)
|
||||
_RATE_LIMIT_TYPE_VALUES = frozenset(t.value for t in RateLimitType)
|
||||
_RATE_LIMIT_CATEGORY_VALUES: Final = frozenset(c.value for c in RateLimitErrorCategory)
|
||||
_RATE_LIMIT_TYPE_VALUES: Final = frozenset(t.value for t in RateLimitType)
|
||||
|
||||
|
||||
def validate_rate_limit_category(value: Any) -> str | None:
|
||||
|
|
@ -339,7 +339,7 @@ class Timeout(openai.APITimeoutError): # type: ignore
|
|||
headers: dict | None = None,
|
||||
exception_status_code: int | None = None,
|
||||
):
|
||||
request = httpx.Request(
|
||||
request: Final = httpx.Request(
|
||||
method="POST",
|
||||
url="https://api.openai.com/v1",
|
||||
)
|
||||
|
|
@ -464,7 +464,7 @@ class RateLimitError(openai.RateLimitError): # type: ignore
|
|||
# headers stay reachable on `e.response.headers` for callers that
|
||||
# explicitly want them; only the proxy-supplied `headers=` kwarg
|
||||
# makes it onto `self.headers`.
|
||||
_response_headers = getattr(response, "headers", None) if response is not None else None
|
||||
_response_headers: Final = getattr(response, "headers", None) if response is not None else None
|
||||
self.headers: dict[str, str] | None = {k: str(v) for k, v in headers.items()} if headers else None
|
||||
# Mirrors FastAPI HTTPException.detail so the same instance can be
|
||||
# serialized through both the ProxyException and HTTPException paths.
|
||||
|
|
@ -558,8 +558,8 @@ class RejectedRequestError(BadRequestError): # type: ignore
|
|||
self.llm_provider = llm_provider
|
||||
self.litellm_debug_info = litellm_debug_info
|
||||
self.request_data = request_data
|
||||
request = httpx.Request(method="POST", url="https://api.openai.com/v1")
|
||||
response = httpx.Response(status_code=400, request=request)
|
||||
request: Final = httpx.Request(method="POST", url="https://api.openai.com/v1")
|
||||
response: Final = httpx.Response(status_code=400, request=request)
|
||||
super().__init__(
|
||||
message=self.message,
|
||||
model=self.model, # type: ignore
|
||||
|
|
@ -648,7 +648,7 @@ class ServiceUnavailableError(openai.APIStatusError): # type: ignore
|
|||
self.litellm_debug_info = litellm_debug_info
|
||||
self.max_retries = max_retries
|
||||
self.num_retries = num_retries
|
||||
_response_headers = getattr(response, "headers", None) if response is not None else None
|
||||
_response_headers: Final = getattr(response, "headers", None) if response is not None else None
|
||||
self.response = httpx.Response(
|
||||
status_code=self.status_code,
|
||||
headers=_response_headers,
|
||||
|
|
@ -696,7 +696,7 @@ class BadGatewayError(openai.APIStatusError): # type: ignore
|
|||
self.litellm_debug_info = litellm_debug_info
|
||||
self.max_retries = max_retries
|
||||
self.num_retries = num_retries
|
||||
_response_headers = getattr(response, "headers", None) if response is not None else None
|
||||
_response_headers: Final = getattr(response, "headers", None) if response is not None else None
|
||||
self.response = httpx.Response(
|
||||
status_code=self.status_code,
|
||||
headers=_response_headers,
|
||||
|
|
@ -744,7 +744,7 @@ class InternalServerError(openai.InternalServerError): # type: ignore
|
|||
self.litellm_debug_info = litellm_debug_info
|
||||
self.max_retries = max_retries
|
||||
self.num_retries = num_retries
|
||||
_response_headers = getattr(response, "headers", None) if response is not None else None
|
||||
_response_headers: Final = getattr(response, "headers", None) if response is not None else None
|
||||
self.response = httpx.Response(
|
||||
status_code=self.status_code,
|
||||
headers=_response_headers,
|
||||
|
|
@ -868,8 +868,8 @@ class APIResponseValidationError(openai.APIResponseValidationError): # type: ig
|
|||
self.message = f"litellm.APIResponseValidationError: {message}"
|
||||
self.llm_provider = llm_provider
|
||||
self.model = model
|
||||
request = httpx.Request(method="POST", url="https://api.openai.com/v1")
|
||||
response = httpx.Response(status_code=500, request=request)
|
||||
request: Final = httpx.Request(method="POST", url="https://api.openai.com/v1")
|
||||
response: Final = httpx.Response(status_code=500, request=request)
|
||||
self.litellm_debug_info = litellm_debug_info
|
||||
self.max_retries = max_retries
|
||||
self.num_retries = num_retries
|
||||
|
|
@ -933,7 +933,7 @@ class UnsupportedParamsError(BadRequestError):
|
|||
self.num_retries = num_retries
|
||||
|
||||
|
||||
LITELLM_EXCEPTION_TYPES = [
|
||||
LITELLM_EXCEPTION_TYPES: Final = [
|
||||
AuthenticationError,
|
||||
NotFoundError,
|
||||
BadRequestError,
|
||||
|
|
@ -1046,7 +1046,7 @@ class GuardrailRaisedException(Exception):
|
|||
should_wrap_with_default_message: bool = True,
|
||||
status_code: int = 400,
|
||||
):
|
||||
default_message = f"Guardrail raised an exception, Guardrail: {guardrail_name}, Message: {message}"
|
||||
default_message: Final = f"Guardrail raised an exception, Guardrail: {guardrail_name}, Message: {message}"
|
||||
self.guardrail_name = guardrail_name
|
||||
self.status_code = status_code
|
||||
self.message = default_message if should_wrap_with_default_message else message
|
||||
|
|
@ -1084,7 +1084,7 @@ class MidStreamFallbackError(ServiceUnavailableError): # type: ignore
|
|||
generated_content: str = "",
|
||||
is_pre_first_chunk: bool = False,
|
||||
):
|
||||
original_status = getattr(original_exception, "status_code", None)
|
||||
original_status: Final = getattr(original_exception, "status_code", None)
|
||||
self.status_code = int(original_status) if original_status is not None else 503
|
||||
self.message = f"litellm.MidStreamFallbackError: {message}"
|
||||
self.model = model
|
||||
|
|
@ -1109,11 +1109,11 @@ class MidStreamFallbackError(ServiceUnavailableError): # type: ignore
|
|||
self.response = response
|
||||
|
||||
# Save the original attributes before they are overridden by ServiceUnavailableError
|
||||
_saved_response = self.response
|
||||
_saved_request = getattr(self.response, "request", None) or httpx.Request(
|
||||
_saved_response: Final = self.response
|
||||
_saved_request: Final = getattr(self.response, "request", None) or httpx.Request(
|
||||
method="POST", url=f"https://{llm_provider}.com/v1/"
|
||||
)
|
||||
_saved_message = self.message
|
||||
_saved_message: Final = self.message
|
||||
|
||||
# Call the parent constructor (which hardcodes status_code=503 and modifies the response object)
|
||||
super().__init__(
|
||||
|
|
|
|||
|
|
@ -6,10 +6,7 @@ import asyncio
|
|||
import base64
|
||||
import os
|
||||
from collections.abc import Awaitable, Callable, Generator
|
||||
from typing import (
|
||||
Any,
|
||||
TypeVar,
|
||||
)
|
||||
from typing import Any, Final, TypeVar
|
||||
|
||||
import httpx
|
||||
from mcp import ClientSession, ReadResourceResult, Resource, StdioServerParameters
|
||||
|
|
@ -61,7 +58,7 @@ def _strip_header_whitespace(headers: dict[str, str]) -> dict[str, str]:
|
|||
|
||||
|
||||
def _first_non_cancelled_cause(exc: BaseException) -> BaseException | None:
|
||||
queue: list[BaseException] = [exc]
|
||||
queue: Final[list[BaseException]] = [exc]
|
||||
while queue:
|
||||
current = queue.pop(0)
|
||||
nested = getattr(current, "exceptions", None)
|
||||
|
|
@ -123,7 +120,7 @@ class MCPSigV4Auth(httpx.Auth):
|
|||
# Fall back to default boto3 credential chain
|
||||
import botocore.session
|
||||
|
||||
session = botocore.session.get_session()
|
||||
session: Final = botocore.session.get_session()
|
||||
self.credentials = session.get_credentials()
|
||||
if self.credentials is None:
|
||||
raise ValueError(
|
||||
|
|
@ -145,19 +142,19 @@ class MCPSigV4Auth(httpx.Auth):
|
|||
import boto3
|
||||
from botocore.credentials import Credentials
|
||||
|
||||
session_name = aws_session_name or f"litellm-mcp-{int(__import__('time').time())}"
|
||||
sts_kwargs: dict = {"region_name": aws_region_name}
|
||||
session_name: Final = aws_session_name or f"litellm-mcp-{int(__import__('time').time())}"
|
||||
sts_kwargs: Final[dict] = {"region_name": aws_region_name}
|
||||
if aws_access_key_id and aws_secret_access_key:
|
||||
sts_kwargs["aws_access_key_id"] = aws_access_key_id
|
||||
sts_kwargs["aws_secret_access_key"] = aws_secret_access_key
|
||||
if aws_session_token:
|
||||
sts_kwargs["aws_session_token"] = aws_session_token
|
||||
sts_client = boto3.client("sts", **sts_kwargs)
|
||||
sts_response = sts_client.assume_role(
|
||||
sts_client: Final = boto3.client("sts", **sts_kwargs)
|
||||
sts_response: Final = sts_client.assume_role(
|
||||
RoleArn=aws_role_name,
|
||||
RoleSessionName=session_name,
|
||||
)
|
||||
sts_creds = sts_response["Credentials"]
|
||||
sts_creds: Final = sts_response["Credentials"]
|
||||
return Credentials(
|
||||
access_key=sts_creds["AccessKeyId"],
|
||||
secret_key=sts_creds["SecretAccessKey"],
|
||||
|
|
@ -170,7 +167,7 @@ class MCPSigV4Auth(httpx.Auth):
|
|||
|
||||
# Build AWSRequest from the httpx Request.
|
||||
# Pass all request headers so the canonical SigV4 signature covers them.
|
||||
aws_request = AWSRequest(
|
||||
aws_request: Final = AWSRequest(
|
||||
method=request.method,
|
||||
url=str(request.url),
|
||||
data=request.content,
|
||||
|
|
@ -179,7 +176,7 @@ class MCPSigV4Auth(httpx.Auth):
|
|||
# Sign the request — SigV4Auth.add_auth() adds Authorization,
|
||||
# X-Amz-Date, and X-Amz-Security-Token (if session token present).
|
||||
# Host header is derived automatically from the URL.
|
||||
sigv4 = SigV4Auth(self.credentials, self.service_name, self.region_name)
|
||||
sigv4: Final = SigV4Auth(self.credentials, self.service_name, self.region_name)
|
||||
sigv4.add_auth(aws_request)
|
||||
# Copy SigV4 headers back to the httpx request
|
||||
for header_name, header_value in aws_request.headers.items():
|
||||
|
|
@ -246,7 +243,7 @@ class MCPClient:
|
|||
if self.transport_type == MCPTransport.stdio:
|
||||
if not self.stdio_config:
|
||||
raise ValueError("stdio_config is required for stdio transport")
|
||||
server_params = StdioServerParameters(
|
||||
server_params: Final = StdioServerParameters(
|
||||
command=self.stdio_config.get("command", ""),
|
||||
args=self.stdio_config.get("args", []),
|
||||
env=self._get_safe_stdio_env(self.stdio_config.get("env")),
|
||||
|
|
@ -274,7 +271,7 @@ class MCPClient:
|
|||
headers=headers,
|
||||
timeout=httpx.Timeout(self.timeout),
|
||||
)
|
||||
transport_ctx = streamable_http_client(
|
||||
transport_ctx: Final = streamable_http_client(
|
||||
url=self.server_url,
|
||||
http_client=http_client,
|
||||
)
|
||||
|
|
@ -292,7 +289,7 @@ class MCPClient:
|
|||
return provided_env
|
||||
|
||||
# Minimal allowlist of safe/standard environment variables
|
||||
safe_keys = {
|
||||
safe_keys: Final = {
|
||||
"PATH",
|
||||
"HOME",
|
||||
"USER",
|
||||
|
|
@ -316,7 +313,7 @@ class MCPClient:
|
|||
"WINDIR",
|
||||
}
|
||||
|
||||
safe_env = {}
|
||||
safe_env: Final = {}
|
||||
for key in safe_keys:
|
||||
if key in os.environ:
|
||||
safe_env[key] = os.environ[key]
|
||||
|
|
@ -338,25 +335,25 @@ class MCPClient:
|
|||
so that upstream MCP servers can request LLM inference (sampling),
|
||||
user input (elicitation), or send log messages.
|
||||
"""
|
||||
transport = await transport_ctx.__aenter__()
|
||||
transport: Final = await transport_ctx.__aenter__()
|
||||
in_flight_error: BaseException | None = None
|
||||
try:
|
||||
read_stream, write_stream = transport[0], transport[1]
|
||||
# Build session kwargs with optional callbacks
|
||||
session_kwargs: dict[str, Any] = {}
|
||||
session_kwargs: Final[dict[str, Any]] = {}
|
||||
if self._sampling_callback is not None:
|
||||
session_kwargs["sampling_callback"] = self._sampling_callback
|
||||
if self._elicitation_callback is not None:
|
||||
session_kwargs["elicitation_callback"] = self._elicitation_callback
|
||||
if self._logging_callback is not None:
|
||||
session_kwargs["logging_callback"] = self._logging_callback
|
||||
session_ctx = ClientSession(read_stream, write_stream, **session_kwargs)
|
||||
session = await session_ctx.__aenter__()
|
||||
session_ctx: Final = ClientSession(read_stream, write_stream, **session_kwargs)
|
||||
session: Final = await session_ctx.__aenter__()
|
||||
try:
|
||||
init_result = await session.initialize()
|
||||
init_result: Final = await session.initialize()
|
||||
self._last_initialize_instructions = None
|
||||
if init_result is not None:
|
||||
ins = getattr(init_result, "instructions", None)
|
||||
ins: Final = getattr(init_result, "instructions", None)
|
||||
if isinstance(ins, str) and ins.strip():
|
||||
self._last_initialize_instructions = ins.strip()
|
||||
return await operation(session)
|
||||
|
|
@ -373,7 +370,7 @@ class MCPClient:
|
|||
await transport_ctx.__aexit__(None, None, None)
|
||||
except BaseException as exit_error:
|
||||
verbose_logger.debug("Error during transport context exit: %s", exit_error)
|
||||
root_cause = _first_non_cancelled_cause(exit_error)
|
||||
root_cause: Final = _first_non_cancelled_cause(exit_error)
|
||||
if root_cause is not None and isinstance(in_flight_error, asyncio.CancelledError):
|
||||
raise root_cause from in_flight_error
|
||||
|
||||
|
|
@ -394,7 +391,7 @@ class MCPClient:
|
|||
transport_ctx, http_client = self._create_transport_context()
|
||||
return await self._execute_session_operation(transport_ctx, operation)
|
||||
except Exception:
|
||||
_log = verbose_logger.debug if quiet_on_error else verbose_logger.warning
|
||||
_log: Final = verbose_logger.debug if quiet_on_error else verbose_logger.warning
|
||||
_log("MCP client run_with_session failed for %s", self.server_url or "stdio")
|
||||
raise
|
||||
finally:
|
||||
|
|
@ -418,7 +415,7 @@ class MCPClient:
|
|||
|
||||
def _get_auth_headers(self) -> dict:
|
||||
"""Generate authentication headers based on auth type."""
|
||||
headers = {}
|
||||
headers: Final = {}
|
||||
if self._mcp_auth_value:
|
||||
if isinstance(self._mcp_auth_value, str):
|
||||
if self.auth_type == MCPAuth.bearer_token:
|
||||
|
|
@ -463,13 +460,13 @@ class MCPClient:
|
|||
) -> httpx.AsyncClient:
|
||||
"""Create an httpx.AsyncClient with LiteLLM's SSL configuration."""
|
||||
# Get unified SSL configuration using the same logic as http_handler.py
|
||||
ssl_config = get_ssl_configuration(self.ssl_verify)
|
||||
ssl_config: Final = get_ssl_configuration(self.ssl_verify)
|
||||
verbose_logger.debug("MCP client using SSL configuration: %s", type(ssl_config).__name__)
|
||||
# The MCP SDK's sse_client and streamable_http_client call this factory without
|
||||
# passing auth=, so the fallback is used: a v2-resolved auth if present, else the
|
||||
# SigV4 aws_auth. Both are None for the common case — no behavior change.
|
||||
fallback_auth = self._resolved_auth if self._resolved_auth is not None else self._aws_auth
|
||||
effective_auth = auth if auth is not None else fallback_auth
|
||||
fallback_auth: Final = self._resolved_auth if self._resolved_auth is not None else self._aws_auth
|
||||
effective_auth: Final = auth if auth is not None else fallback_auth
|
||||
return httpx.AsyncClient(
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
|
|
@ -496,9 +493,9 @@ class MCPClient:
|
|||
return await session.list_tools()
|
||||
|
||||
try:
|
||||
result = await self.run_with_session(_list_tools_operation, quiet_on_error=raise_on_error)
|
||||
tool_count = len(result.tools)
|
||||
tool_names = [tool.name for tool in result.tools]
|
||||
result: Final = await self.run_with_session(_list_tools_operation, quiet_on_error=raise_on_error)
|
||||
tool_count: Final = len(result.tools)
|
||||
tool_names: Final = [tool.name for tool in result.tools]
|
||||
verbose_logger.info(
|
||||
"MCP client listed %s tools from %s: %s", tool_count, self.server_url or "stdio", tool_names
|
||||
)
|
||||
|
|
@ -507,13 +504,13 @@ class MCPClient:
|
|||
verbose_logger.warning("MCP client list_tools was cancelled")
|
||||
raise
|
||||
except Exception as e:
|
||||
error_type = type(e).__name__
|
||||
error_type: Final = type(e).__name__
|
||||
# Mirror call_tool: when the caller opted into raise_on_error it owns the exception and
|
||||
# logs it at the fitting level (an expected pass-through re-auth 401 is info, not an
|
||||
# error), so log at debug here to avoid an error-level line + traceback that would trip
|
||||
# error-rate alerts on that expected signal. The swallow path still logs the full
|
||||
# exception because nothing downstream will surface the failure.
|
||||
_log = verbose_logger.debug if raise_on_error else verbose_logger.exception
|
||||
_log: Final = verbose_logger.debug if raise_on_error else verbose_logger.exception
|
||||
_log(
|
||||
f"MCP client list_tools failed - "
|
||||
f"Error Type: {error_type}, "
|
||||
|
|
@ -523,7 +520,7 @@ class MCPClient:
|
|||
)
|
||||
# Check if it's a stream/connection error
|
||||
if "BrokenResourceError" in error_type or "Broken" in error_type:
|
||||
_log_broken = verbose_logger.debug if raise_on_error else verbose_logger.error
|
||||
_log_broken: Final = verbose_logger.debug if raise_on_error else verbose_logger.error
|
||||
_log_broken(
|
||||
"MCP client detected broken connection/stream during list_tools - "
|
||||
"the MCP server may have crashed, disconnected, or timed out"
|
||||
|
|
@ -560,7 +557,7 @@ class MCPClient:
|
|||
verbose_logger.info("MCP client calling tool '%s'", call_tool_request_params.name)
|
||||
|
||||
async def on_progress(progress: float, total: float | None, message: str | None):
|
||||
percentage = (progress / total * 100) if total else 0
|
||||
percentage: Final = (progress / total * 100) if total else 0
|
||||
verbose_logger.info(
|
||||
f"MCP Tool '{call_tool_request_params.name}' progress: "
|
||||
f"{progress}/{total} ({percentage:.0f}%) - {message or ''}"
|
||||
|
|
@ -581,7 +578,7 @@ class MCPClient:
|
|||
)
|
||||
|
||||
try:
|
||||
tool_result = await self.run_with_session(_call_tool_operation, quiet_on_error=raise_on_error)
|
||||
tool_result: Final = await self.run_with_session(_call_tool_operation, quiet_on_error=raise_on_error)
|
||||
verbose_logger.info("MCP client tool call '%s' completed successfully", call_tool_request_params.name)
|
||||
return tool_result
|
||||
except asyncio.CancelledError:
|
||||
|
|
@ -590,16 +587,16 @@ class MCPClient:
|
|||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
error_trace = traceback.format_exc()
|
||||
error_trace: Final = traceback.format_exc()
|
||||
verbose_logger.debug("MCP client tool call traceback:\n%s", error_trace)
|
||||
# Log detailed error information
|
||||
error_type = type(e).__name__
|
||||
error_type: Final = type(e).__name__
|
||||
# When the caller opted into raise_on_error it owns the exception and logs it at the
|
||||
# level that fits (an expected pass-through re-auth 401 is info, not an operator-actionable
|
||||
# error), so log at debug here to avoid an error-level line that would trip error-rate
|
||||
# alerts on that expected signal. The swallow path (raise_on_error=False) still logs at
|
||||
# error because nothing downstream will surface the failure.
|
||||
_log = verbose_logger.debug if raise_on_error else verbose_logger.error
|
||||
_log: Final = verbose_logger.debug if raise_on_error else verbose_logger.error
|
||||
_log(
|
||||
f"MCP client call_tool failed - "
|
||||
f"Error Type: {error_type}, "
|
||||
|
|
@ -627,9 +624,9 @@ class MCPClient:
|
|||
return await session.list_prompts()
|
||||
|
||||
try:
|
||||
result = await self.run_with_session(_list_prompts_operation)
|
||||
prompt_count = len(result.prompts)
|
||||
prompt_names = [prompt.name for prompt in result.prompts]
|
||||
result: Final = await self.run_with_session(_list_prompts_operation)
|
||||
prompt_count: Final = len(result.prompts)
|
||||
prompt_names: Final = [prompt.name for prompt in result.prompts]
|
||||
verbose_logger.info(
|
||||
"MCP client listed %s tools from %s: %s", prompt_count, self.server_url or "stdio", prompt_names
|
||||
)
|
||||
|
|
@ -638,7 +635,7 @@ class MCPClient:
|
|||
verbose_logger.warning("MCP client list_prompts was cancelled")
|
||||
raise
|
||||
except Exception as e:
|
||||
error_type = type(e).__name__
|
||||
error_type: Final = type(e).__name__
|
||||
verbose_logger.error(
|
||||
"MCP client list_prompts failed - Error Type: %s, Error: %s, Server: %s, Transport: %s",
|
||||
error_type,
|
||||
|
|
@ -667,7 +664,7 @@ class MCPClient:
|
|||
)
|
||||
|
||||
try:
|
||||
get_prompt_result = await self.run_with_session(_get_prompt_operation)
|
||||
get_prompt_result: Final = await self.run_with_session(_get_prompt_operation)
|
||||
verbose_logger.info("MCP client get_prompt '%s' completed successfully", get_prompt_request_params.name)
|
||||
return get_prompt_result
|
||||
except asyncio.CancelledError:
|
||||
|
|
@ -676,10 +673,10 @@ class MCPClient:
|
|||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
error_trace = traceback.format_exc()
|
||||
error_trace: Final = traceback.format_exc()
|
||||
verbose_logger.debug("MCP client get_prompt traceback:\n%s", error_trace)
|
||||
# Log detailed error information
|
||||
error_type = type(e).__name__
|
||||
error_type: Final = type(e).__name__
|
||||
verbose_logger.error(
|
||||
"MCP client get_prompt failed - Error Type: %s, Error: %s, Prompt: %s, Server: %s, Transport: %s",
|
||||
error_type,
|
||||
|
|
@ -704,9 +701,9 @@ class MCPClient:
|
|||
return await session.list_resources()
|
||||
|
||||
try:
|
||||
result = await self.run_with_session(_list_resources_operation)
|
||||
resource_count = len(result.resources)
|
||||
resource_names = [resource.name for resource in result.resources]
|
||||
result: Final = await self.run_with_session(_list_resources_operation)
|
||||
resource_count: Final = len(result.resources)
|
||||
resource_names: Final = [resource.name for resource in result.resources]
|
||||
verbose_logger.info(
|
||||
"MCP client listed %s resources from %s: %s", resource_count, self.server_url or "stdio", resource_names
|
||||
)
|
||||
|
|
@ -715,7 +712,7 @@ class MCPClient:
|
|||
verbose_logger.warning("MCP client list_resources was cancelled")
|
||||
raise
|
||||
except Exception as e:
|
||||
error_type = type(e).__name__
|
||||
error_type: Final = type(e).__name__
|
||||
verbose_logger.error(
|
||||
"MCP client list_resources failed - Error Type: %s, Error: %s, Server: %s, Transport: %s",
|
||||
error_type,
|
||||
|
|
@ -740,9 +737,9 @@ class MCPClient:
|
|||
return await session.list_resource_templates()
|
||||
|
||||
try:
|
||||
result = await self.run_with_session(_list_resource_templates_operation)
|
||||
resource_template_count = len(result.resourceTemplates)
|
||||
resource_template_names = [resourceTemplate.name for resourceTemplate in result.resourceTemplates]
|
||||
result: Final = await self.run_with_session(_list_resource_templates_operation)
|
||||
resource_template_count: Final = len(result.resourceTemplates)
|
||||
resource_template_names: Final = [resourceTemplate.name for resourceTemplate in result.resourceTemplates]
|
||||
verbose_logger.info(
|
||||
"MCP client listed %s resource templates from %s: %s",
|
||||
resource_template_count,
|
||||
|
|
@ -754,7 +751,7 @@ class MCPClient:
|
|||
verbose_logger.warning("MCP client list_resource_templates was cancelled")
|
||||
raise
|
||||
except Exception as e:
|
||||
error_type = type(e).__name__
|
||||
error_type: Final = type(e).__name__
|
||||
verbose_logger.error(
|
||||
"MCP client list_resource_templates failed - Error Type: %s, Error: %s, Server: %s, Transport: %s",
|
||||
error_type,
|
||||
|
|
@ -780,7 +777,7 @@ class MCPClient:
|
|||
return await session.read_resource(url)
|
||||
|
||||
try:
|
||||
read_resource_result = await self.run_with_session(_read_resource_operation)
|
||||
read_resource_result: Final = await self.run_with_session(_read_resource_operation)
|
||||
verbose_logger.info("MCP client read_resource '%s' completed successfully", url)
|
||||
return read_resource_result
|
||||
except asyncio.CancelledError:
|
||||
|
|
@ -789,10 +786,10 @@ class MCPClient:
|
|||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
error_trace = traceback.format_exc()
|
||||
error_trace: Final = traceback.format_exc()
|
||||
verbose_logger.debug("MCP client read_resource traceback:\n%s", error_trace)
|
||||
# Log detailed error information
|
||||
error_type = type(e).__name__
|
||||
error_type: Final = type(e).__name__
|
||||
verbose_logger.error(
|
||||
"MCP client read_resource failed - Error Type: %s, Error: %s, Url: %s, Server: %s, Transport: %s",
|
||||
error_type,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import json
|
||||
from typing import Literal
|
||||
from typing import Final, Literal
|
||||
|
||||
from mcp import ClientSession
|
||||
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
|
||||
|
|
@ -18,7 +18,7 @@ from litellm.types.utils import ChatCompletionMessageToolCall
|
|||
########################################################
|
||||
def transform_mcp_tool_to_openai_tool(mcp_tool: MCPTool) -> ChatCompletionToolParam:
|
||||
"""Convert an MCP tool to an OpenAI tool."""
|
||||
normalized_parameters = _normalize_mcp_input_schema(mcp_tool.inputSchema)
|
||||
normalized_parameters: Final = _normalize_mcp_input_schema(mcp_tool.inputSchema)
|
||||
|
||||
return ChatCompletionToolParam(
|
||||
type="function",
|
||||
|
|
@ -44,7 +44,7 @@ def _normalize_mcp_input_schema(input_schema: dict) -> dict:
|
|||
return {"type": "object", "properties": {}, "additionalProperties": False}
|
||||
|
||||
# Make a copy to avoid modifying the original
|
||||
normalized_schema = dict(input_schema)
|
||||
normalized_schema: Final = dict(input_schema)
|
||||
|
||||
# Ensure type is 'object'
|
||||
if "type" not in normalized_schema:
|
||||
|
|
@ -65,7 +65,7 @@ def transform_mcp_tool_to_openai_responses_api_tool(
|
|||
mcp_tool: MCPTool,
|
||||
) -> FunctionToolParam:
|
||||
"""Convert an MCP tool to an OpenAI Responses API tool."""
|
||||
normalized_parameters = _normalize_mcp_input_schema(mcp_tool.inputSchema)
|
||||
normalized_parameters: Final = _normalize_mcp_input_schema(mcp_tool.inputSchema)
|
||||
|
||||
return FunctionToolParam(
|
||||
name=mcp_tool.name,
|
||||
|
|
@ -103,7 +103,7 @@ async def load_mcp_tools(
|
|||
|
||||
If format is set to "openai", the tools are converted to OpenAI API compatible tools.
|
||||
"""
|
||||
tools = await session.list_tools()
|
||||
tools: Final = await session.list_tools()
|
||||
if format == "openai":
|
||||
return [transform_mcp_tool_to_openai_tool(mcp_tool=tool) for tool in tools.tools]
|
||||
return tools.tools
|
||||
|
|
@ -119,7 +119,7 @@ async def call_mcp_tool(
|
|||
call_tool_request_params: MCPCallToolRequestParams,
|
||||
) -> MCPCallToolResult:
|
||||
"""Call an MCP tool."""
|
||||
tool_result = await session.call_tool(
|
||||
tool_result: Final = await session.call_tool(
|
||||
name=call_tool_request_params.name,
|
||||
arguments=call_tool_request_params.arguments,
|
||||
)
|
||||
|
|
@ -141,7 +141,7 @@ def transform_openai_tool_call_request_to_mcp_tool_call_request(
|
|||
openai_tool: ChatCompletionMessageToolCall | dict,
|
||||
) -> MCPCallToolRequestParams:
|
||||
"""Convert an OpenAI ChatCompletionMessageToolCall to an MCP CallToolRequestParams."""
|
||||
function = openai_tool["function"]
|
||||
function: Final = openai_tool["function"]
|
||||
return MCPCallToolRequestParams(
|
||||
name=function["name"],
|
||||
arguments=_get_function_arguments(function),
|
||||
|
|
@ -161,7 +161,7 @@ async def call_openai_tool(
|
|||
Returns:
|
||||
The result of the MCP tool call.
|
||||
"""
|
||||
mcp_tool_call_request_params = transform_openai_tool_call_request_to_mcp_tool_call_request(
|
||||
mcp_tool_call_request_params: Final = transform_openai_tool_call_request_to_mcp_tool_call_request(
|
||||
openai_tool=openai_tool,
|
||||
)
|
||||
return await call_mcp_tool(
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import uuid as uuid_module
|
|||
from collections.abc import Coroutine
|
||||
from functools import partial
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Literal, cast
|
||||
from typing import Any, Final, Literal, cast
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -78,17 +78,17 @@ def _should_sdk_support_streaming(
|
|||
return custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS
|
||||
|
||||
|
||||
openai_files_instance = OpenAIFilesAPI()
|
||||
azure_files_instance = AzureOpenAIFilesAPI()
|
||||
vertex_ai_files_instance = VertexAIFilesHandler()
|
||||
bedrock_files_instance = BedrockFilesHandler()
|
||||
openai_files_instance: Final = OpenAIFilesAPI()
|
||||
azure_files_instance: Final = AzureOpenAIFilesAPI()
|
||||
vertex_ai_files_instance: Final = VertexAIFilesHandler()
|
||||
bedrock_files_instance: Final = BedrockFilesHandler()
|
||||
#################################################
|
||||
|
||||
|
||||
def _add_trusted_model_credentials_to_litellm_params(
|
||||
litellm_params_dict: dict[str, Any], kwargs: dict[str, Any]
|
||||
) -> None:
|
||||
trusted_model_credentials = kwargs.get("_litellm_internal_model_credentials")
|
||||
trusted_model_credentials: Final = kwargs.get("_litellm_internal_model_credentials")
|
||||
if isinstance(trusted_model_credentials, type(MappingProxyType({}))):
|
||||
litellm_params_dict["_litellm_internal_model_credentials"] = trusted_model_credentials
|
||||
|
||||
|
|
@ -109,10 +109,10 @@ async def acreate_file(
|
|||
LiteLLM Equivalent of POST: POST https://api.openai.com/v1/files
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["acreate_file"] = True
|
||||
|
||||
call_args = {
|
||||
call_args: Final = {
|
||||
"file": file,
|
||||
"purpose": purpose,
|
||||
"expires_after": expires_after,
|
||||
|
|
@ -123,11 +123,11 @@ async def acreate_file(
|
|||
}
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(create_file, **call_args)
|
||||
func: Final = partial(create_file, **call_args)
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -156,13 +156,13 @@ def create_file(
|
|||
Specify either provider_list or custom_llm_provider.
|
||||
"""
|
||||
try:
|
||||
_is_async = kwargs.pop("acreate_file", False) is True
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict = dict(**kwargs)
|
||||
logging_obj = cast(LiteLLMLoggingObj | None, kwargs.get("litellm_logging_obj"))
|
||||
_is_async: Final = kwargs.pop("acreate_file", False) is True
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict: Final = dict(**kwargs)
|
||||
logging_obj: Final = cast(LiteLLMLoggingObj | None, kwargs.get("litellm_logging_obj"))
|
||||
if logging_obj is None:
|
||||
raise ValueError("logging_obj is required")
|
||||
client = kwargs.get("client")
|
||||
client: Final = kwargs.get("client")
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
|
|
@ -173,7 +173,7 @@ def create_file(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(cast(str, custom_llm_provider)) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
|
|
@ -196,7 +196,7 @@ def create_file(
|
|||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
provider_config = ProviderConfigManager.get_provider_files_config(
|
||||
provider_config: Final = ProviderConfigManager.get_provider_files_config(
|
||||
model="",
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
|
@ -214,7 +214,7 @@ def create_file(
|
|||
timeout=timeout,
|
||||
)
|
||||
elif custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
openai_creds = get_openai_credentials(
|
||||
openai_creds: Final = get_openai_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
|
|
@ -229,7 +229,7 @@ def create_file(
|
|||
create_file_data=_create_file_request,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
azure_creds = get_azure_credentials(
|
||||
azure_creds: Final = get_azure_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
api_version=optional_params.api_version,
|
||||
|
|
@ -274,11 +274,11 @@ async def afile_retrieve(
|
|||
LiteLLM Equivalent of GET https://api.openai.com/v1/files
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["is_async"] = True
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
file_retrieve,
|
||||
file_id,
|
||||
custom_llm_provider,
|
||||
|
|
@ -288,9 +288,9 @@ async def afile_retrieve(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -315,7 +315,7 @@ def file_retrieve(
|
|||
LiteLLM Equivalent of POST: POST https://api.openai.com/v1/files
|
||||
"""
|
||||
try:
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
# set timeout for 10 minutes by default
|
||||
|
|
@ -325,17 +325,17 @@ def file_retrieve(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
_is_async = kwargs.pop("is_async", False) is True
|
||||
_is_async: Final = kwargs.pop("is_async", False) is True
|
||||
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
openai_creds = get_openai_credentials(
|
||||
openai_creds: Final = get_openai_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
|
|
@ -350,7 +350,7 @@ def file_retrieve(
|
|||
organization=openai_creds.organization,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
azure_creds = get_azure_credentials(
|
||||
azure_creds: Final = get_azure_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
api_version=optional_params.api_version,
|
||||
|
|
@ -366,12 +366,12 @@ def file_retrieve(
|
|||
)
|
||||
else:
|
||||
# Try using provider config pattern (for Manus, Bedrock, etc.)
|
||||
provider_config = ProviderConfigManager.get_provider_files_config(
|
||||
provider_config: Final = ProviderConfigManager.get_provider_files_config(
|
||||
model="",
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
if provider_config is not None:
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
_add_trusted_model_credentials_to_litellm_params(
|
||||
litellm_params_dict=litellm_params_dict,
|
||||
kwargs=kwargs,
|
||||
|
|
@ -395,7 +395,7 @@ def file_retrieve(
|
|||
function_id=str(kwargs.get("id") or ""),
|
||||
)
|
||||
|
||||
client = kwargs.get("client")
|
||||
client: Final = kwargs.get("client")
|
||||
response = base_llm_http_handler.retrieve_file(
|
||||
file_id=file_id,
|
||||
provider_config=provider_config,
|
||||
|
|
@ -443,12 +443,12 @@ async def afile_delete(
|
|||
LiteLLM Equivalent of DELETE https://api.openai.com/v1/files
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
model = kwargs.pop("model", None)
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
model: Final = kwargs.pop("model", None)
|
||||
kwargs["is_async"] = True
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
file_delete,
|
||||
file_id,
|
||||
model,
|
||||
|
|
@ -459,9 +459,9 @@ async def afile_delete(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -492,8 +492,8 @@ def file_delete(
|
|||
_, custom_llm_provider, _, _ = get_llm_provider(model, custom_llm_provider)
|
||||
except Exception:
|
||||
pass
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
_add_trusted_model_credentials_to_litellm_params(
|
||||
litellm_params_dict=litellm_params_dict,
|
||||
kwargs=kwargs,
|
||||
|
|
@ -501,22 +501,22 @@ def file_delete(
|
|||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
# set timeout for 10 minutes by default
|
||||
client = kwargs.get("client")
|
||||
client: Final = kwargs.get("client")
|
||||
|
||||
if (
|
||||
timeout is not None
|
||||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
_is_async = kwargs.pop("is_async", False) is True
|
||||
_is_async: Final = kwargs.pop("is_async", False) is True
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
openai_creds = get_openai_credentials(
|
||||
openai_creds: Final = get_openai_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
|
|
@ -531,7 +531,7 @@ def file_delete(
|
|||
organization=openai_creds.organization,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
azure_creds = get_azure_credentials(
|
||||
azure_creds: Final = get_azure_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
api_version=optional_params.api_version,
|
||||
|
|
@ -549,7 +549,7 @@ def file_delete(
|
|||
)
|
||||
else:
|
||||
# Try using provider config pattern (for Manus, Bedrock, etc.)
|
||||
provider_config = ProviderConfigManager.get_provider_files_config(
|
||||
provider_config: Final = ProviderConfigManager.get_provider_files_config(
|
||||
model="",
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
|
@ -619,11 +619,11 @@ async def afile_list(
|
|||
LiteLLM Equivalent of GET https://api.openai.com/v1/files
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["is_async"] = True
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
file_list,
|
||||
custom_llm_provider,
|
||||
purpose,
|
||||
|
|
@ -633,9 +633,9 @@ async def afile_list(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -660,7 +660,7 @@ def file_list(
|
|||
LiteLLM Equivalent of GET https://api.openai.com/v1/files
|
||||
"""
|
||||
try:
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
# set timeout for 10 minutes by default
|
||||
|
|
@ -670,22 +670,22 @@ def file_list(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
_is_async = kwargs.pop("is_async", False) is True
|
||||
_is_async: Final = kwargs.pop("is_async", False) is True
|
||||
|
||||
# Check if provider has a custom files config (e.g., Manus, Bedrock, Vertex AI)
|
||||
provider_config = ProviderConfigManager.get_provider_files_config(
|
||||
provider_config: Final = ProviderConfigManager.get_provider_files_config(
|
||||
model="",
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
if provider_config is not None:
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
litellm_params_dict["api_key"] = optional_params.api_key
|
||||
litellm_params_dict["api_base"] = optional_params.api_base
|
||||
|
||||
|
|
@ -705,7 +705,7 @@ def file_list(
|
|||
function_id=str(kwargs.get("id", "")),
|
||||
)
|
||||
|
||||
client = kwargs.get("client")
|
||||
client: Final = kwargs.get("client")
|
||||
response = base_llm_http_handler.list_files(
|
||||
purpose=purpose,
|
||||
provider_config=provider_config,
|
||||
|
|
@ -718,7 +718,7 @@ def file_list(
|
|||
)
|
||||
return response
|
||||
elif custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
openai_creds = get_openai_credentials(
|
||||
openai_creds: Final = get_openai_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
|
|
@ -733,7 +733,7 @@ def file_list(
|
|||
organization=openai_creds.organization,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
azure_creds = get_azure_credentials(
|
||||
azure_creds: Final = get_azure_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
api_version=optional_params.api_version,
|
||||
|
|
@ -779,12 +779,12 @@ async def afile_content(
|
|||
LiteLLM Equivalent of GET https://api.openai.com/v1/files
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["afile_content"] = True
|
||||
model = kwargs.pop("model", None)
|
||||
model: Final = kwargs.pop("model", None)
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
file_content,
|
||||
file_id=file_id,
|
||||
model=model,
|
||||
|
|
@ -797,9 +797,9 @@ async def afile_content(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -832,15 +832,15 @@ def file_content(
|
|||
LiteLLM Equivalent of POST: POST https://api.openai.com/v1/files
|
||||
"""
|
||||
try:
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
_add_trusted_model_credentials_to_litellm_params(
|
||||
litellm_params_dict=litellm_params_dict,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
client = kwargs.get("client")
|
||||
client: Final = kwargs.get("client")
|
||||
# set timeout for 10 minutes by default
|
||||
|
||||
try:
|
||||
|
|
@ -854,20 +854,20 @@ def file_content(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(cast(str, custom_llm_provider)) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
_file_content_request = FileContentRequest(
|
||||
_file_content_request: Final = FileContentRequest(
|
||||
file_id=file_id,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
_is_async = kwargs.pop("afile_content", False) is True
|
||||
_is_async: Final = kwargs.pop("afile_content", False) is True
|
||||
|
||||
if stream and _should_sdk_support_streaming(custom_llm_provider):
|
||||
return file_content_streaming(
|
||||
|
|
@ -885,7 +885,7 @@ def file_content(
|
|||
)
|
||||
|
||||
# Check if provider has a custom files config (e.g., Anthropic, Manus)
|
||||
provider_config = ProviderConfigManager.get_provider_files_config(
|
||||
provider_config: Final = ProviderConfigManager.get_provider_files_config(
|
||||
model="",
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
|
@ -918,7 +918,7 @@ def file_content(
|
|||
return response
|
||||
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
openai_creds = get_openai_credentials(
|
||||
openai_creds: Final = get_openai_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
|
|
@ -933,7 +933,7 @@ def file_content(
|
|||
organization=openai_creds.organization,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
azure_creds = get_azure_credentials(
|
||||
azure_creds: Final = get_azure_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
api_version=optional_params.api_version,
|
||||
|
|
@ -950,14 +950,14 @@ def file_content(
|
|||
litellm_params=litellm_params_dict,
|
||||
)
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or ""
|
||||
vertex_ai_project = (
|
||||
api_base: Final = optional_params.api_base or ""
|
||||
vertex_ai_project: Final = (
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
vertex_ai_location: Final = (
|
||||
optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION")
|
||||
)
|
||||
vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
|
||||
response = vertex_ai_files_instance.file_content(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -1014,7 +1014,7 @@ def file_content_streaming(
|
|||
logging_obj.model_call_details["model"] = model or ""
|
||||
logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
litellm_params = logging_obj.model_call_details.get("litellm_params", {}) or {}
|
||||
litellm_params: Final = logging_obj.model_call_details.get("litellm_params", {}) or {}
|
||||
if optional_params.api_base is not None:
|
||||
litellm_params["api_base"] = optional_params.api_base
|
||||
logging_obj.model_call_details["litellm_params"] = litellm_params
|
||||
|
|
@ -1037,7 +1037,7 @@ def file_content_streaming(
|
|||
stream_iterator=iter(()), headers={}
|
||||
)
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
openai_creds = get_openai_credentials(
|
||||
openai_creds: Final = get_openai_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
|
|
|
|||
|
|
@ -1,12 +1,7 @@
|
|||
import datetime
|
||||
import traceback
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Optional,
|
||||
cast,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional, cast
|
||||
|
||||
import anyio
|
||||
|
||||
|
|
@ -91,7 +86,7 @@ class FileContentStreamingResponse:
|
|||
|
||||
self._close_completed = True
|
||||
self._logging_completed = True
|
||||
stream_to_close = self.stream_iterator
|
||||
stream_to_close: Final = self.stream_iterator
|
||||
self.stream_iterator = cast(Iterator[bytes] | AsyncIterator[bytes], iter(()))
|
||||
|
||||
# Shield cleanup from request cancellation so upstream HTTP connections
|
||||
|
|
@ -100,7 +95,7 @@ class FileContentStreamingResponse:
|
|||
if hasattr(stream_to_close, "aclose"):
|
||||
await cast(AsyncIterator[bytes], stream_to_close).aclose() # type: ignore[attr-defined]
|
||||
elif hasattr(stream_to_close, "close"):
|
||||
result = cast(Iterator[bytes], stream_to_close).close() # type: ignore[attr-defined]
|
||||
result: Final = cast(Iterator[bytes], stream_to_close).close() # type: ignore[attr-defined]
|
||||
if result is not None:
|
||||
await result
|
||||
|
||||
|
|
@ -110,14 +105,14 @@ class FileContentStreamingResponse:
|
|||
|
||||
self._close_completed = True
|
||||
self._logging_completed = True
|
||||
stream_to_close = self.stream_iterator
|
||||
stream_to_close: Final = self.stream_iterator
|
||||
self.stream_iterator = cast(Iterator[bytes] | AsyncIterator[bytes], iter(()))
|
||||
|
||||
if hasattr(stream_to_close, "close"):
|
||||
cast(Iterator[bytes], stream_to_close).close() # type: ignore[attr-defined]
|
||||
|
||||
def _build_logging_response(self) -> dict[str, str]:
|
||||
response = {
|
||||
response: Final = {
|
||||
"id": self.file_id,
|
||||
"object": "file.content",
|
||||
}
|
||||
|
|
@ -154,7 +149,7 @@ class FileContentStreamingResponse:
|
|||
)
|
||||
|
||||
self._sync_hidden_params()
|
||||
payload = get_standard_logging_object_payload(
|
||||
payload: Final = get_standard_logging_object_payload(
|
||||
kwargs=self.logging_obj.model_call_details,
|
||||
init_response_obj=self._build_logging_response(),
|
||||
start_time=self._start_time,
|
||||
|
|
@ -165,7 +160,7 @@ class FileContentStreamingResponse:
|
|||
if payload is None:
|
||||
return None
|
||||
|
||||
merged_hidden_params = cast(
|
||||
merged_hidden_params: Final = cast(
|
||||
"StandardLoggingHiddenParams",
|
||||
{
|
||||
**cast(dict[str, Any], payload.get("hidden_params") or {}),
|
||||
|
|
@ -189,8 +184,8 @@ class FileContentStreamingResponse:
|
|||
return
|
||||
|
||||
self._logging_completed = True
|
||||
end_time = datetime.datetime.now()
|
||||
standard_logging_object = self._build_standard_logging_object(end_time=end_time)
|
||||
end_time: Final = datetime.datetime.now()
|
||||
standard_logging_object: Final = self._build_standard_logging_object(end_time=end_time)
|
||||
await self.logging_obj.async_success_handler(
|
||||
result=self._build_logging_response(),
|
||||
start_time=self._start_time,
|
||||
|
|
@ -208,8 +203,8 @@ class FileContentStreamingResponse:
|
|||
return
|
||||
|
||||
self._logging_completed = True
|
||||
end_time = datetime.datetime.now()
|
||||
standard_logging_object = self._build_standard_logging_object(end_time=end_time)
|
||||
end_time: Final = datetime.datetime.now()
|
||||
standard_logging_object: Final = self._build_standard_logging_object(end_time=end_time)
|
||||
self.logging_obj.success_handler(
|
||||
result=self._build_logging_response(),
|
||||
start_time=self._start_time,
|
||||
|
|
@ -222,8 +217,8 @@ class FileContentStreamingResponse:
|
|||
return
|
||||
|
||||
self._logging_completed = True
|
||||
end_time = datetime.datetime.now()
|
||||
traceback_str = traceback.format_exc()
|
||||
end_time: Final = datetime.datetime.now()
|
||||
traceback_str: Final = traceback.format_exc()
|
||||
self.logging_obj.failure_handler(error, traceback_str, self._start_time, end_time)
|
||||
await self.logging_obj.async_failure_handler(error, traceback_str, self._start_time, end_time)
|
||||
|
||||
|
|
@ -232,5 +227,5 @@ class FileContentStreamingResponse:
|
|||
return
|
||||
|
||||
self._logging_completed = True
|
||||
end_time = datetime.datetime.now()
|
||||
end_time: Final = datetime.datetime.now()
|
||||
self.logging_obj.failure_handler(error, traceback.format_exc(), self._start_time, end_time)
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.types.llms.openai import CreateFileRequest
|
||||
from litellm.types.utils import ExtractedFileData
|
||||
|
||||
|
|
@ -6,7 +8,7 @@ from litellm.types.utils import ExtractedFileData
|
|||
# batch file must not silently bypass the streaming path just because of its
|
||||
# declared type. ``purpose == "batch"`` is the authoritative signal; non-JSONL
|
||||
# content still fails loudly when the rows are parsed.
|
||||
_BATCH_JSONL_CONTENT_TYPES = frozenset(
|
||||
_BATCH_JSONL_CONTENT_TYPES: Final = frozenset(
|
||||
{
|
||||
"application/jsonl",
|
||||
"application/json",
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ import contextvars
|
|||
import os
|
||||
from collections.abc import Coroutine
|
||||
from functools import partial
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -29,9 +29,9 @@ from litellm.types.utils import LiteLLMFineTuningJob
|
|||
from litellm.utils import client, supports_httpx_timeout
|
||||
|
||||
####### ENVIRONMENT VARIABLES ###################
|
||||
openai_fine_tuning_apis_instance = OpenAIFineTuningAPI()
|
||||
azure_fine_tuning_apis_instance = AzureOpenAIFineTuningAPI()
|
||||
vertex_fine_tuning_apis_instance = VertexFineTuningAPI()
|
||||
openai_fine_tuning_apis_instance: Final = OpenAIFineTuningAPI()
|
||||
azure_fine_tuning_apis_instance: Final = AzureOpenAIFineTuningAPI()
|
||||
vertex_fine_tuning_apis_instance: Final = VertexFineTuningAPI()
|
||||
#################################################
|
||||
|
||||
|
||||
|
|
@ -61,7 +61,7 @@ def _prepare_azure_extra_body(
|
|||
extra_body = {}
|
||||
|
||||
# Azure-specific root-level parameters
|
||||
azure_specific_params = ["trainingType"]
|
||||
azure_specific_params: Final = ["trainingType"]
|
||||
for param in azure_specific_params:
|
||||
if param in kwargs:
|
||||
extra_body[param] = kwargs[param]
|
||||
|
|
@ -93,11 +93,11 @@ async def acreate_fine_tuning_job(
|
|||
"""
|
||||
verbose_logger.debug("inside acreate_fine_tuning_job model=%s and kwargs=%s", model, kwargs)
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["acreate_fine_tuning_job"] = True
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
create_fine_tuning_job,
|
||||
model,
|
||||
training_file,
|
||||
|
|
@ -113,9 +113,9 @@ async def acreate_fine_tuning_job(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -171,24 +171,24 @@ def create_fine_tuning_job(
|
|||
|
||||
"""
|
||||
try:
|
||||
_is_async = kwargs.pop("acreate_fine_tuning_job", False) is True
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
_is_async: Final = kwargs.pop("acreate_fine_tuning_job", False) is True
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
# handle hyperparameters
|
||||
hyperparameters = hyperparameters or {} # original hyperparameters
|
||||
|
||||
# For Azure, extract Azure-specific hyperparameters before creating OpenAI-spec hyperparameters
|
||||
azure_specific_hyperparams = {}
|
||||
azure_specific_hyperparams: Final = {}
|
||||
if custom_llm_provider == "azure":
|
||||
azure_hyperparameter_keys = ["prompt_loss_weight"]
|
||||
azure_hyperparameter_keys: Final = ["prompt_loss_weight"]
|
||||
for key in azure_hyperparameter_keys:
|
||||
if key in hyperparameters:
|
||||
azure_specific_hyperparams[key] = hyperparameters.pop(key)
|
||||
|
||||
_oai_hyperparameters: Hyperparameters = Hyperparameters(
|
||||
_oai_hyperparameters: Final[Hyperparameters] = Hyperparameters(
|
||||
**hyperparameters
|
||||
) # Typed Hyperparameters for OpenAI Spec
|
||||
timeout = _resolve_fine_tuning_timeout(
|
||||
timeout: Final = _resolve_fine_tuning_timeout(
|
||||
optional_params.timeout or kwargs.get("request_timeout", 600),
|
||||
custom_llm_provider,
|
||||
)
|
||||
|
|
@ -203,7 +203,7 @@ def create_fine_tuning_job(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -287,13 +287,13 @@ def create_fine_tuning_job(
|
|||
)
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or ""
|
||||
vertex_ai_project = (
|
||||
vertex_ai_project: Final = (
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
vertex_ai_location: Final = (
|
||||
optional_params.vertex_location or litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION")
|
||||
)
|
||||
vertex_credentials = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
vertex_credentials: Final = optional_params.vertex_credentials or get_secret_str("VERTEXAI_CREDENTIALS")
|
||||
response = vertex_fine_tuning_apis_instance.create_fine_tuning_job(
|
||||
_is_async=_is_async,
|
||||
create_fine_tuning_job_data=_build_fine_tuning_job_data(
|
||||
|
|
@ -342,11 +342,11 @@ async def acancel_fine_tuning_job(
|
|||
Async: Immediately cancel a fine-tune job.
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["acancel_fine_tuning_job"] = True
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
cancel_fine_tuning_job,
|
||||
fine_tuning_job_id,
|
||||
custom_llm_provider,
|
||||
|
|
@ -356,9 +356,9 @@ async def acancel_fine_tuning_job(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -383,7 +383,7 @@ def cancel_fine_tuning_job(
|
|||
|
||||
"""
|
||||
try:
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
# set timeout for 10 minutes by default
|
||||
|
|
@ -393,14 +393,14 @@ def cancel_fine_tuning_job(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
_is_async = kwargs.pop("acancel_fine_tuning_job", False) is True
|
||||
_is_async: Final = kwargs.pop("acancel_fine_tuning_job", False) is True
|
||||
|
||||
# OpenAI
|
||||
if custom_llm_provider == "openai":
|
||||
|
|
@ -412,7 +412,7 @@ def cancel_fine_tuning_job(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -493,11 +493,11 @@ async def alist_fine_tuning_jobs(
|
|||
Async: List your organization's fine-tuning jobs
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["alist_fine_tuning_jobs"] = True
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
list_fine_tuning_jobs,
|
||||
after,
|
||||
limit,
|
||||
|
|
@ -508,9 +508,9 @@ async def alist_fine_tuning_jobs(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -537,7 +537,7 @@ def list_fine_tuning_jobs(
|
|||
- limit: Optional[int] = None, Number of fine-tuning jobs to retrieve. Defaults to 20
|
||||
"""
|
||||
try:
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
# set timeout for 10 minutes by default
|
||||
|
|
@ -547,14 +547,14 @@ def list_fine_tuning_jobs(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
_is_async = kwargs.pop("alist_fine_tuning_jobs", False) is True
|
||||
_is_async: Final = kwargs.pop("alist_fine_tuning_jobs", False) is True
|
||||
|
||||
# OpenAI
|
||||
if custom_llm_provider == "openai":
|
||||
|
|
@ -566,7 +566,7 @@ def list_fine_tuning_jobs(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
|
|
@ -649,11 +649,11 @@ async def aretrieve_fine_tuning_job(
|
|||
Async: Get info about a fine-tuning job.
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["aretrieve_fine_tuning_job"] = True
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
retrieve_fine_tuning_job,
|
||||
fine_tuning_job_id,
|
||||
custom_llm_provider,
|
||||
|
|
@ -663,9 +663,9 @@ async def aretrieve_fine_tuning_job(
|
|||
)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
|
|
@ -687,7 +687,7 @@ def retrieve_fine_tuning_job(
|
|||
Get info about a fine-tuning job.
|
||||
"""
|
||||
try:
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
optional_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
# set timeout for 10 minutes by default
|
||||
|
|
@ -697,14 +697,14 @@ def retrieve_fine_tuning_job(
|
|||
and isinstance(timeout, httpx.Timeout)
|
||||
and supports_httpx_timeout(custom_llm_provider) is False
|
||||
):
|
||||
read_timeout = timeout.read or 600
|
||||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
_is_async = kwargs.pop("aretrieve_fine_tuning_job", False) is True
|
||||
_is_async: Final = kwargs.pop("aretrieve_fine_tuning_job", False) is True
|
||||
|
||||
# OpenAI
|
||||
if custom_llm_provider == "openai":
|
||||
|
|
@ -715,7 +715,7 @@ def retrieve_fine_tuning_job(
|
|||
or os.getenv("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
organization: Final = (
|
||||
optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) or None
|
||||
)
|
||||
api_key = optional_params.api_key or litellm.api_key or litellm.openai_key or os.getenv("OPENAI_API_KEY")
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from collections.abc import AsyncIterator, Coroutine
|
||||
from typing import Any, cast
|
||||
from typing import Any, Final, cast
|
||||
|
||||
import litellm
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -8,7 +8,7 @@ from litellm.types.utils import ModelResponse
|
|||
from .transformation import GoogleGenAIAdapter
|
||||
|
||||
# Initialize adapter
|
||||
GOOGLE_GENAI_ADAPTER = GoogleGenAIAdapter()
|
||||
GOOGLE_GENAI_ADAPTER: Final = GoogleGenAIAdapter()
|
||||
|
||||
|
||||
class GenerateContentToCompletionHandler:
|
||||
|
|
@ -26,7 +26,7 @@ class GenerateContentToCompletionHandler:
|
|||
"""Prepare kwargs for litellm.completion/acompletion"""
|
||||
|
||||
# Transform generate_content request to completion format
|
||||
completion_request = GOOGLE_GENAI_ADAPTER.translate_generate_content_to_completion(
|
||||
completion_request: Final = GOOGLE_GENAI_ADAPTER.translate_generate_content_to_completion(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
|
|
@ -34,7 +34,7 @@ class GenerateContentToCompletionHandler:
|
|||
**(extra_kwargs or {}),
|
||||
)
|
||||
|
||||
completion_kwargs: dict[str, Any] = dict(completion_request)
|
||||
completion_kwargs: Final[dict[str, Any]] = dict(completion_request)
|
||||
|
||||
# Forward extra_kwargs that should be passed to completion call
|
||||
if extra_kwargs is not None:
|
||||
|
|
@ -61,7 +61,7 @@ class GenerateContentToCompletionHandler:
|
|||
) -> dict[str, Any] | AsyncIterator[bytes]:
|
||||
"""Handle generate_content call asynchronously using completion adapter"""
|
||||
|
||||
completion_kwargs = GenerateContentToCompletionHandler._prepare_completion_kwargs(
|
||||
completion_kwargs: Final = GenerateContentToCompletionHandler._prepare_completion_kwargs(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
|
|
@ -71,7 +71,7 @@ class GenerateContentToCompletionHandler:
|
|||
)
|
||||
|
||||
try:
|
||||
completion_response = await litellm.acompletion(**completion_kwargs)
|
||||
completion_response: Final = await litellm.acompletion(**completion_kwargs)
|
||||
|
||||
if stream:
|
||||
# Check if completion_response is actually a stream or a ModelResponse
|
||||
|
|
@ -84,7 +84,7 @@ class GenerateContentToCompletionHandler:
|
|||
return generate_content_response
|
||||
else:
|
||||
# Transform streaming completion response to generate_content format
|
||||
transformed_stream = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
|
||||
transformed_stream: Final = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
|
||||
completion_response
|
||||
)
|
||||
if transformed_stream is not None:
|
||||
|
|
@ -122,7 +122,7 @@ class GenerateContentToCompletionHandler:
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
completion_kwargs = GenerateContentToCompletionHandler._prepare_completion_kwargs(
|
||||
completion_kwargs: Final = GenerateContentToCompletionHandler._prepare_completion_kwargs(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
|
|
@ -132,7 +132,7 @@ class GenerateContentToCompletionHandler:
|
|||
)
|
||||
|
||||
try:
|
||||
completion_response = litellm.completion(**completion_kwargs)
|
||||
completion_response: Final = litellm.completion(**completion_kwargs)
|
||||
|
||||
if stream:
|
||||
# Check if completion_response is actually a stream or a ModelResponse
|
||||
|
|
@ -145,7 +145,7 @@ class GenerateContentToCompletionHandler:
|
|||
return generate_content_response
|
||||
else:
|
||||
# Transform streaming completion response to generate_content format
|
||||
transformed_stream = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
|
||||
transformed_stream: Final = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
|
||||
completion_response
|
||||
)
|
||||
if transformed_stream is not None:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import json
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, cast
|
||||
from typing import Any, Final, cast
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm.litellm_core_utils.json_validation_rule import normalize_tool_schema
|
||||
|
|
@ -85,7 +85,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
# After the stream is exhausted, check for any remaining accumulated tool calls
|
||||
if self.accumulated_tool_calls:
|
||||
try:
|
||||
parts = []
|
||||
parts: Final = []
|
||||
for (
|
||||
tool_call_index,
|
||||
tool_call_data,
|
||||
|
|
@ -110,7 +110,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
tool_call_data["arguments"],
|
||||
)
|
||||
if parts:
|
||||
final_chunk = {
|
||||
final_chunk: Final = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": parts, "role": "model"},
|
||||
|
|
@ -197,9 +197,9 @@ class GoogleGenAIAdapter:
|
|||
"""
|
||||
|
||||
# Extract top-level fields from kwargs
|
||||
system_instruction = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
|
||||
tools = kwargs.get("tools")
|
||||
tool_config = kwargs.get("toolConfig") or kwargs.get("tool_config")
|
||||
system_instruction: Final = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
|
||||
tools: Final = kwargs.get("tools")
|
||||
tool_config: Final = kwargs.get("toolConfig") or kwargs.get("tool_config")
|
||||
|
||||
# Normalize contents to list format
|
||||
if isinstance(contents, dict):
|
||||
|
|
@ -208,10 +208,10 @@ class GoogleGenAIAdapter:
|
|||
contents_list = contents
|
||||
|
||||
# Transform contents to OpenAI messages format
|
||||
messages = self._transform_contents_to_messages(contents_list, system_instruction=system_instruction)
|
||||
messages: Final = self._transform_contents_to_messages(contents_list, system_instruction=system_instruction)
|
||||
|
||||
# Create base request as dict (which is compatible with ChatCompletionRequest)
|
||||
completion_request: ChatCompletionRequest = {
|
||||
completion_request: Final[ChatCompletionRequest] = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
}
|
||||
|
|
@ -248,14 +248,14 @@ class GoogleGenAIAdapter:
|
|||
# Check if tools are already in OpenAI format or Google GenAI format
|
||||
if isinstance(tools, list) and len(tools) > 0:
|
||||
# Tools are in Google GenAI format, transform them
|
||||
openai_tools = self._transform_google_genai_tools_to_openai(tools)
|
||||
openai_tools: Final = self._transform_google_genai_tools_to_openai(tools)
|
||||
|
||||
if openai_tools:
|
||||
completion_request["tools"] = openai_tools
|
||||
|
||||
# Handle tool_config (tool choice)
|
||||
if tool_config:
|
||||
tool_choice = self._transform_google_genai_tool_config_to_openai(tool_config)
|
||||
tool_choice: Final = self._transform_google_genai_tool_config_to_openai(tool_config)
|
||||
if tool_choice:
|
||||
completion_request["tool_choice"] = tool_choice
|
||||
|
||||
|
|
@ -285,9 +285,9 @@ class GoogleGenAIAdapter:
|
|||
Returns:
|
||||
Dict[str, Any]
|
||||
"""
|
||||
allowed_fields = GenericLiteLLMParams.model_fields.keys()
|
||||
allowed_fields: Final = GenericLiteLLMParams.model_fields.keys()
|
||||
if litellm_params:
|
||||
litellm_dict = litellm_params.model_dump(exclude_none=True)
|
||||
litellm_dict: Final = litellm_params.model_dump(exclude_none=True)
|
||||
for key, value in litellm_dict.items():
|
||||
if key in allowed_fields:
|
||||
completion_request_dict[key] = value
|
||||
|
|
@ -298,7 +298,7 @@ class GoogleGenAIAdapter:
|
|||
completion_stream: Any,
|
||||
) -> AsyncIterator[bytes] | None:
|
||||
"""Transform streaming completion output to Google GenAI format"""
|
||||
google_genai_wrapper = GoogleGenAIStreamWrapper(completion_stream=completion_stream)
|
||||
google_genai_wrapper: Final = GoogleGenAIStreamWrapper(completion_stream=completion_stream)
|
||||
# Return the SSE-wrapped version for proper event formatting
|
||||
return google_genai_wrapper.async_google_genai_sse_wrapper()
|
||||
|
||||
|
|
@ -307,7 +307,7 @@ class GoogleGenAIAdapter:
|
|||
tools: list[dict[str, Any]],
|
||||
) -> list[ChatCompletionToolParam]:
|
||||
"""Transform Google GenAI tools to OpenAI tools format"""
|
||||
openai_tools: list[dict[str, Any]] = []
|
||||
openai_tools: Final[list[dict[str, Any]]] = []
|
||||
|
||||
for tool in tools:
|
||||
if "functionDeclarations" in tool:
|
||||
|
|
@ -325,7 +325,7 @@ class GoogleGenAIAdapter:
|
|||
openai_tools.append(openai_tool)
|
||||
|
||||
# normalize the tool schemas
|
||||
normalized_tools = [normalize_tool_schema(tool) for tool in openai_tools]
|
||||
normalized_tools: Final = [normalize_tool_schema(tool) for tool in openai_tools]
|
||||
|
||||
return cast(list[ChatCompletionToolParam], normalized_tools)
|
||||
|
||||
|
|
@ -334,12 +334,12 @@ class GoogleGenAIAdapter:
|
|||
tool_config: dict[str, Any],
|
||||
) -> ChatCompletionToolChoiceValues | None:
|
||||
"""Transform Google GenAI tool_config to OpenAI tool_choice"""
|
||||
function_calling_config = tool_config.get("functionCallingConfig", {})
|
||||
mode = function_calling_config.get("mode", "AUTO")
|
||||
function_calling_config: Final = tool_config.get("functionCallingConfig", {})
|
||||
mode: Final = function_calling_config.get("mode", "AUTO")
|
||||
|
||||
mode_mapping = {"AUTO": "auto", "ANY": "required", "NONE": "none"}
|
||||
mode_mapping: Final = {"AUTO": "auto", "ANY": "required", "NONE": "none"}
|
||||
|
||||
tool_choice = mode_mapping.get(mode, "auto")
|
||||
tool_choice: Final = mode_mapping.get(mode, "auto")
|
||||
return cast(ChatCompletionToolChoiceValues, tool_choice)
|
||||
|
||||
def _transform_contents_to_messages(
|
||||
|
|
@ -348,11 +348,11 @@ class GoogleGenAIAdapter:
|
|||
system_instruction: dict[str, Any] | None = None,
|
||||
) -> list[AllMessageValues]:
|
||||
"""Transform Google GenAI contents to OpenAI messages format"""
|
||||
messages: list[AllMessageValues] = []
|
||||
messages: Final[list[AllMessageValues]] = []
|
||||
|
||||
# Handle system instruction
|
||||
if system_instruction:
|
||||
system_parts = system_instruction.get("parts", [])
|
||||
system_parts: Final = system_instruction.get("parts", [])
|
||||
if system_parts and "text" in system_parts[0]:
|
||||
messages.append(ChatCompletionSystemMessage(role="system", content=system_parts[0]["text"]))
|
||||
|
||||
|
|
@ -473,7 +473,7 @@ class GoogleGenAIAdapter:
|
|||
"""
|
||||
|
||||
# Extract the main response content
|
||||
choice = response.choices[0] if response.choices else None
|
||||
choice: Final = response.choices[0] if response.choices else None
|
||||
if not choice:
|
||||
raise ValueError("Invalid completion response: no choices found")
|
||||
|
||||
|
|
@ -490,7 +490,7 @@ class GoogleGenAIAdapter:
|
|||
parts = [{"text": message_content}] if message_content else []
|
||||
|
||||
# Create Google GenAI format response
|
||||
generate_content_response: dict[str, Any] = {
|
||||
generate_content_response: Final[dict[str, Any]] = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": parts, "role": "model"},
|
||||
|
|
@ -537,7 +537,7 @@ class GoogleGenAIAdapter:
|
|||
"""
|
||||
|
||||
# Extract the main response content from streaming chunk
|
||||
choice = response.choices[0] if response.choices else None
|
||||
choice: Final = response.choices[0] if response.choices else None
|
||||
if not choice:
|
||||
# Return empty chunk if no choices
|
||||
return None
|
||||
|
|
@ -551,7 +551,7 @@ class GoogleGenAIAdapter:
|
|||
finish_reason = getattr(choice, "finish_reason", None)
|
||||
else:
|
||||
# Fallback for generic choice objects
|
||||
message_content = getattr(choice, "delta", {}).get("content", "")
|
||||
message_content: Final = getattr(choice, "delta", {}).get("content", "")
|
||||
parts = [{"text": message_content}] if message_content else []
|
||||
finish_reason = getattr(choice, "finish_reason", None)
|
||||
|
||||
|
|
@ -560,7 +560,7 @@ class GoogleGenAIAdapter:
|
|||
return None
|
||||
|
||||
# Create Google GenAI streaming format response
|
||||
streaming_chunk: dict[str, Any] = {
|
||||
streaming_chunk: Final[dict[str, Any]] = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": parts, "role": "model"},
|
||||
|
|
@ -573,7 +573,7 @@ class GoogleGenAIAdapter:
|
|||
|
||||
# Add usage metadata only in the final chunk (when finish_reason is present)
|
||||
if finish_reason:
|
||||
usage_metadata = (
|
||||
usage_metadata: Final = (
|
||||
self._map_usage(getattr(response, "usage", None))
|
||||
if hasattr(response, "usage") and getattr(response, "usage", None)
|
||||
else {
|
||||
|
|
@ -599,7 +599,7 @@ class GoogleGenAIAdapter:
|
|||
message: Any,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Transform OpenAI message to Google GenAI parts format"""
|
||||
parts: list[dict[str, Any]] = []
|
||||
parts: Final[list[dict[str, Any]]] = []
|
||||
|
||||
# Add text content if present
|
||||
if hasattr(message, "content") and message.content:
|
||||
|
|
@ -633,13 +633,13 @@ class GoogleGenAIAdapter:
|
|||
if not hasattr(wrapper, "accumulated_tool_calls"):
|
||||
wrapper.accumulated_tool_calls = {}
|
||||
|
||||
parts: list[dict[str, Any]] = []
|
||||
parts: Final[list[dict[str, Any]]] = []
|
||||
|
||||
if hasattr(delta, "content") and delta.content:
|
||||
parts.append({"text": delta.content})
|
||||
|
||||
# 2. Ensure tool_calls is iterable
|
||||
tool_calls = delta.tool_calls or []
|
||||
tool_calls: Final = delta.tool_calls or []
|
||||
|
||||
for tool_call in tool_calls:
|
||||
if not hasattr(tool_call, "function"):
|
||||
|
|
@ -704,7 +704,7 @@ class GoogleGenAIAdapter:
|
|||
if not finish_reason:
|
||||
return "STOP"
|
||||
|
||||
mapping = {
|
||||
mapping: Final = {
|
||||
"stop": "STOP",
|
||||
"length": "MAX_TOKENS",
|
||||
"content_filter": "SAFETY",
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import asyncio
|
|||
import contextvars
|
||||
from collections.abc import Iterator
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
|
@ -110,11 +110,11 @@ class GenerateContentHelper:
|
|||
Returns:
|
||||
GenerateContentSetupResult containing all setup information
|
||||
"""
|
||||
litellm_logging_obj: LiteLLMLoggingObj | None = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
|
||||
# get llm provider logic
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
## MOCK RESPONSE LOGIC (only for non-streaming)
|
||||
if (
|
||||
|
|
@ -140,7 +140,7 @@ class GenerateContentHelper:
|
|||
litellm_params.custom_llm_provider = custom_llm_provider
|
||||
|
||||
# get provider config
|
||||
generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig | None = (
|
||||
generate_content_provider_config: Final[BaseGoogleGenAIGenerateContentConfig | None] = (
|
||||
ProviderConfigManager.get_provider_google_genai_generate_content_config(
|
||||
model=model,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
|
|
@ -168,19 +168,19 @@ class GenerateContentHelper:
|
|||
# Construct request body
|
||||
#########################################################################################
|
||||
# Create Google Optional Params Config
|
||||
generate_content_config_dict = generate_content_provider_config.map_generate_content_optional_params(
|
||||
generate_content_config_dict: Final = generate_content_provider_config.map_generate_content_optional_params(
|
||||
generate_content_config_dict=config or {},
|
||||
model=model,
|
||||
)
|
||||
# Extract systemInstruction from kwargs to pass to transform
|
||||
system_instruction = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
|
||||
system_instruction: Final = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
|
||||
# Native top-level REST fields arrive as loose kwargs and are otherwise dropped.
|
||||
native_request_fields: dict[str, object] = {
|
||||
native_request_fields: Final[dict[str, object]] = {
|
||||
field: kwargs[field]
|
||||
for field in generate_content_provider_config.get_generate_content_request_top_level_fields()
|
||||
if field in kwargs
|
||||
}
|
||||
request_body = generate_content_provider_config.transform_generate_content_request(
|
||||
request_body: Final = generate_content_provider_config.transform_generate_content_request(
|
||||
model=model,
|
||||
contents=contents,
|
||||
tools=tools,
|
||||
|
|
@ -250,9 +250,9 @@ async def agenerate_content(
|
|||
"""
|
||||
Async: Generate content using Google GenAI
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["agenerate_content"] = True
|
||||
|
||||
# Handle generationConfig parameter from kwargs for backward compatibility
|
||||
|
|
@ -265,7 +265,7 @@ async def agenerate_content(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
generate_content,
|
||||
model=model,
|
||||
contents=contents,
|
||||
|
|
@ -279,9 +279,9 @@ async def agenerate_content(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -318,9 +318,9 @@ def generate_content(
|
|||
"""
|
||||
Generate content using Google GenAI
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
_is_async = kwargs.pop("agenerate_content", False)
|
||||
_is_async: Final = kwargs.pop("agenerate_content", False)
|
||||
|
||||
_mark_async_entrypoint(kwargs.get("litellm_logging_obj"), CallTypes.agenerate_content.value, _is_async)
|
||||
|
||||
|
|
@ -328,12 +328,12 @@ def generate_content(
|
|||
if "generationConfig" in kwargs and config is None:
|
||||
config = kwargs.pop("generationConfig")
|
||||
# Check for mock response first
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
if litellm_params.mock_response and isinstance(litellm_params.mock_response, str):
|
||||
return GenerateContentHelper.mock_generate_content_response(mock_response=litellm_params.mock_response)
|
||||
|
||||
# Setup the call
|
||||
setup_result = GenerateContentHelper.setup_generate_content_call(
|
||||
setup_result: Final = GenerateContentHelper.setup_generate_content_call(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
|
|
@ -343,7 +343,7 @@ def generate_content(
|
|||
)
|
||||
|
||||
# Extract systemInstruction from kwargs to pass to handler
|
||||
system_instruction = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
|
||||
system_instruction: Final = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
|
||||
|
||||
# Check if we should use the adapter (when provider config is None)
|
||||
if setup_result.generate_content_provider_config is None:
|
||||
|
|
@ -360,7 +360,7 @@ def generate_content(
|
|||
)
|
||||
|
||||
# Call the standard handler
|
||||
response = base_llm_http_handler.generate_content_handler(
|
||||
response: Final = base_llm_http_handler.generate_content_handler(
|
||||
model=setup_result.model,
|
||||
contents=contents,
|
||||
tools=tools,
|
||||
|
|
@ -408,7 +408,7 @@ async def agenerate_content_stream(
|
|||
"""
|
||||
Async: Generate content using Google GenAI with streaming response
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
kwargs["agenerate_content_stream"] = True
|
||||
|
||||
|
|
@ -424,7 +424,7 @@ async def agenerate_content_stream(
|
|||
)
|
||||
|
||||
# Setup the call
|
||||
setup_result = GenerateContentHelper.setup_generate_content_call(
|
||||
setup_result: Final = GenerateContentHelper.setup_generate_content_call(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
|
|
@ -434,7 +434,7 @@ async def agenerate_content_stream(
|
|||
)
|
||||
|
||||
# Extract systemInstruction from kwargs to pass to handler
|
||||
system_instruction = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
|
||||
system_instruction: Final = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
|
||||
|
||||
# Check if we should use the adapter (when provider config is None)
|
||||
if setup_result.generate_content_provider_config is None:
|
||||
|
|
@ -503,10 +503,10 @@ def generate_content_stream(
|
|||
"""
|
||||
Generate content using Google GenAI with streaming response
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
# Remove any async-related flags since this is the sync function
|
||||
_is_async = kwargs.pop("agenerate_content_stream", False)
|
||||
_is_async: Final = kwargs.pop("agenerate_content_stream", False)
|
||||
|
||||
_mark_async_entrypoint(kwargs.get("litellm_logging_obj"), CallTypes.agenerate_content_stream.value, _is_async)
|
||||
|
||||
|
|
@ -514,7 +514,7 @@ def generate_content_stream(
|
|||
if "generationConfig" in kwargs and config is None:
|
||||
config = kwargs.pop("generationConfig")
|
||||
# Setup the call
|
||||
setup_result = GenerateContentHelper.setup_generate_content_call(
|
||||
setup_result: Final = GenerateContentHelper.setup_generate_content_call(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
|
|
@ -524,7 +524,7 @@ def generate_content_stream(
|
|||
)
|
||||
|
||||
# Extract systemInstruction from kwargs to pass to handler
|
||||
system_instruction = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
|
||||
system_instruction: Final = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
|
||||
|
||||
# Check if we should use the adapter (when provider config is None)
|
||||
if setup_result.generate_content_provider_config is None:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import asyncio
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy.pass_through_endpoints.success_handler import (
|
||||
|
|
@ -15,7 +15,7 @@ if TYPE_CHECKING:
|
|||
else:
|
||||
BaseGoogleGenAIGenerateContentConfig = Any
|
||||
|
||||
GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ = PassThroughEndpointLogging()
|
||||
GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ: Final = PassThroughEndpointLogging()
|
||||
|
||||
|
||||
def _encode_google_genai_sse_event(event_lines: list[str]) -> bytes:
|
||||
|
|
@ -23,7 +23,7 @@ def _encode_google_genai_sse_event(event_lines: list[str]) -> bytes:
|
|||
|
||||
|
||||
def _next_google_genai_sse_chunk(line_iter) -> bytes:
|
||||
event_lines: list[str] = []
|
||||
event_lines: Final[list[str]] = []
|
||||
while True:
|
||||
try:
|
||||
line = next(line_iter)
|
||||
|
|
@ -39,7 +39,7 @@ def _next_google_genai_sse_chunk(line_iter) -> bytes:
|
|||
|
||||
|
||||
async def _anext_google_genai_sse_chunk(line_iter) -> bytes:
|
||||
event_lines: list[str] = []
|
||||
event_lines: Final[list[str]] = []
|
||||
while True:
|
||||
try:
|
||||
line = await line_iter.__anext__()
|
||||
|
|
@ -82,7 +82,7 @@ class BaseGoogleGenAIGenerateContentStreamingIterator:
|
|||
PassThroughStreamingHandler,
|
||||
)
|
||||
|
||||
end_time = datetime.now()
|
||||
end_time: Final = datetime.now()
|
||||
asyncio.create_task(
|
||||
PassThroughStreamingHandler._route_streaming_logging_to_handler(
|
||||
litellm_logging_obj=self.litellm_logging_obj,
|
||||
|
|
@ -134,7 +134,7 @@ class GoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateContent
|
|||
|
||||
def __next__(self):
|
||||
try:
|
||||
chunk = _next_google_genai_sse_chunk(self.stream_iterator)
|
||||
chunk: Final = _next_google_genai_sse_chunk(self.stream_iterator)
|
||||
self.collected_chunks.append(chunk)
|
||||
return chunk
|
||||
except StopIteration:
|
||||
|
|
@ -185,7 +185,7 @@ class AsyncGoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateCo
|
|||
|
||||
async def __anext__(self):
|
||||
try:
|
||||
chunk = await _anext_google_genai_sse_chunk(self.stream_iterator)
|
||||
chunk: Final = await _anext_google_genai_sse_chunk(self.stream_iterator)
|
||||
self.collected_chunks.append(chunk)
|
||||
return chunk
|
||||
except StopAsyncIteration:
|
||||
|
|
|
|||
|
|
@ -3,14 +3,7 @@ import contextvars
|
|||
import importlib
|
||||
from collections.abc import Coroutine
|
||||
from functools import partial
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Literal,
|
||||
Optional,
|
||||
cast,
|
||||
overload,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast, overload
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.images.utils import ImageEditRequestUtils
|
||||
|
|
@ -76,7 +69,7 @@ def _get_ImageEditRequestUtils() -> "ImageEditRequestUtils":
|
|||
global _ImageEditRequestUtils_cache
|
||||
if _ImageEditRequestUtils_cache is None:
|
||||
# Access via module to trigger __getattr__ if not cached
|
||||
module = importlib.import_module(__name__)
|
||||
module: Final = importlib.import_module(__name__)
|
||||
_ImageEditRequestUtils_cache = module.ImageEditRequestUtils
|
||||
assert _ImageEditRequestUtils_cache is not None # Type narrowing for type checker
|
||||
return _ImageEditRequestUtils_cache
|
||||
|
|
@ -95,23 +88,23 @@ async def aimage_generation(*args, **kwargs) -> ImageResponse:
|
|||
Returns:
|
||||
- `response` (Any): The response returned by the `image_generation` function.
|
||||
"""
|
||||
loop = asyncio.get_event_loop()
|
||||
model = args[0] if len(args) > 0 else kwargs["model"]
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
model: Final = args[0] if len(args) > 0 else kwargs["model"]
|
||||
### PASS ARGS TO Image Generation ###
|
||||
kwargs["aimg_generation"] = True
|
||||
custom_llm_provider = None
|
||||
try:
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(image_generation, *args, **kwargs)
|
||||
func: Final = partial(image_generation, *args, **kwargs)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model=model, api_base=kwargs.get("api_base", None))
|
||||
|
||||
# Await normally
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
response: ImageResponse | None = None
|
||||
if isinstance(init_response, dict):
|
||||
|
|
@ -210,20 +203,20 @@ def image_generation(
|
|||
Currently supports just Azure + OpenAI.
|
||||
"""
|
||||
try:
|
||||
args = locals()
|
||||
aimg_generation = kwargs.get("aimg_generation", False)
|
||||
litellm_call_id = kwargs.get("litellm_call_id", None)
|
||||
logger_fn = kwargs.get("logger_fn", None)
|
||||
mock_response: str | None = kwargs.get("mock_response", None) # type: ignore
|
||||
proxy_server_request = kwargs.get("proxy_server_request", None)
|
||||
args: Final = locals()
|
||||
aimg_generation: Final = kwargs.get("aimg_generation", False)
|
||||
litellm_call_id: Final = kwargs.get("litellm_call_id", None)
|
||||
logger_fn: Final = kwargs.get("logger_fn", None)
|
||||
mock_response: Final[str | None] = kwargs.get("mock_response", None) # type: ignore
|
||||
proxy_server_request: Final = kwargs.get("proxy_server_request", None)
|
||||
azure_ad_token_provider = kwargs.get("azure_ad_token_provider", None)
|
||||
model_info = kwargs.get("model_info", None)
|
||||
metadata = kwargs.get("metadata", {})
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
client = kwargs.get("client", None)
|
||||
extra_headers = kwargs.get("extra_headers", None)
|
||||
headers: dict = kwargs.get("headers", None) or {}
|
||||
base_model = kwargs.get("base_model", None)
|
||||
model_info: Final = kwargs.get("model_info", None)
|
||||
metadata: Final = kwargs.get("metadata", {})
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
client: Final = kwargs.get("client", None)
|
||||
extra_headers: Final = kwargs.get("extra_headers", None)
|
||||
headers: Final[dict] = kwargs.get("headers", None) or {}
|
||||
base_model: Final = kwargs.get("base_model", None)
|
||||
if extra_headers is not None:
|
||||
headers.update(extra_headers)
|
||||
model_response: ImageResponse = litellm.utils.ImageResponse()
|
||||
|
|
@ -238,7 +231,7 @@ def image_generation(
|
|||
model = "dall-e-2"
|
||||
custom_llm_provider = "openai" # default to dall-e-2 on openai
|
||||
model_response._hidden_params["model"] = model
|
||||
openai_params = [
|
||||
openai_params: Final = [
|
||||
"user",
|
||||
"request_timeout",
|
||||
"api_base",
|
||||
|
|
@ -255,9 +248,9 @@ def image_generation(
|
|||
"size",
|
||||
"style",
|
||||
]
|
||||
litellm_params = all_litellm_params
|
||||
default_params = openai_params + litellm_params
|
||||
non_default_params = {
|
||||
litellm_params: Final = all_litellm_params
|
||||
default_params: Final = openai_params + litellm_params
|
||||
non_default_params: Final = {
|
||||
k: v for k, v in kwargs.items() if k not in default_params
|
||||
} # model-specific params - pass them straight to the model/provider
|
||||
|
||||
|
|
@ -268,7 +261,7 @@ def image_generation(
|
|||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
optional_params = get_optional_params_image_gen(
|
||||
optional_params: Final = get_optional_params_image_gen(
|
||||
model=base_model or model,
|
||||
n=n,
|
||||
quality=quality,
|
||||
|
|
@ -281,9 +274,9 @@ def image_generation(
|
|||
**non_default_params,
|
||||
)
|
||||
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
|
||||
logging: Logging = litellm_logging_obj
|
||||
logging: Final[Logging] = litellm_logging_obj
|
||||
logging.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
|
|
@ -308,7 +301,7 @@ def image_generation(
|
|||
|
||||
if custom_llm_provider == "azure":
|
||||
# azure configs
|
||||
api_type = get_secret_str("AZURE_API_TYPE") or "azure"
|
||||
api_type: Final = get_secret_str("AZURE_API_TYPE") or "azure"
|
||||
|
||||
api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
|
||||
|
|
@ -322,7 +315,7 @@ def image_generation(
|
|||
or get_secret_str("AZURE_API_KEY")
|
||||
)
|
||||
|
||||
azure_ad_token = optional_params.pop("azure_ad_token", None) or get_secret_str("AZURE_AD_TOKEN")
|
||||
azure_ad_token: Final = optional_params.pop("azure_ad_token", None) or get_secret_str("AZURE_AD_TOKEN")
|
||||
|
||||
# Create azure_ad_token_provider from tenant_id, client_id, client_secret if not already provided
|
||||
if azure_ad_token_provider is None:
|
||||
|
|
@ -331,9 +324,9 @@ def image_generation(
|
|||
)
|
||||
|
||||
# Extract Azure AD credentials from litellm_params
|
||||
tenant_id = litellm_params_dict.get("tenant_id")
|
||||
client_id = litellm_params_dict.get("client_id")
|
||||
client_secret = litellm_params_dict.get("client_secret")
|
||||
tenant_id: Final = litellm_params_dict.get("tenant_id")
|
||||
client_id: Final = litellm_params_dict.get("client_id")
|
||||
client_secret: Final = litellm_params_dict.get("client_secret")
|
||||
azure_scope = litellm_params_dict.get("azure_scope") or "https://cognitiveservices.azure.com/.default"
|
||||
|
||||
# Create token provider if credentials are available
|
||||
|
|
@ -392,7 +385,7 @@ def image_generation(
|
|||
raise ValueError(f"image generation config is not supported for {custom_llm_provider}")
|
||||
|
||||
# Resolve api_base from litellm.api_base if not explicitly provided
|
||||
_api_base = api_base or litellm.api_base
|
||||
_api_base: Final = api_base or litellm.api_base
|
||||
litellm_params_dict["api_base"] = _api_base
|
||||
|
||||
return llm_http_handler.image_generation_handler(
|
||||
|
|
@ -468,7 +461,7 @@ def image_generation(
|
|||
if extra_headers is not None:
|
||||
optional_params["extra_headers"] = extra_headers
|
||||
# Forward OpenAI organization if present (set by proxy pre-call utils)
|
||||
organization: str | None = kwargs.get("organization", None)
|
||||
organization: Final[str | None] = kwargs.get("organization", None)
|
||||
model_response = openai_chat_completions.image_generation(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
|
|
@ -568,18 +561,18 @@ async def aimage_variation(*args, **kwargs) -> ImageResponse:
|
|||
Returns:
|
||||
- `response` (Any): The response returned by the `image_variation` function.
|
||||
"""
|
||||
loop = asyncio.get_event_loop()
|
||||
model = kwargs.get("model", None)
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
model: Final = kwargs.get("model", None)
|
||||
custom_llm_provider = kwargs.get("custom_llm_provider", None)
|
||||
### PASS ARGS TO Image Generation ###
|
||||
kwargs["async_call"] = True
|
||||
try:
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(image_variation, *args, **kwargs)
|
||||
func: Final = partial(image_variation, *args, **kwargs)
|
||||
|
||||
# Add the context to the function
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
if custom_llm_provider is None and model is not None:
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model=model, api_base=kwargs.get("api_base", None))
|
||||
|
|
@ -618,12 +611,12 @@ def image_variation(
|
|||
**kwargs,
|
||||
) -> ImageResponse:
|
||||
# get non-default params
|
||||
client = kwargs.get("client", None)
|
||||
client: Final = kwargs.get("client", None)
|
||||
# get logging object
|
||||
litellm_logging_obj = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj"))
|
||||
litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj"))
|
||||
|
||||
# get the litellm params
|
||||
litellm_params = get_litellm_params(**kwargs)
|
||||
litellm_params: Final = get_litellm_params(**kwargs)
|
||||
# get the custom llm provider
|
||||
model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider(
|
||||
model=model,
|
||||
|
|
@ -634,17 +627,17 @@ def image_variation(
|
|||
|
||||
# route to the correct provider w/ the params
|
||||
try:
|
||||
llm_provider = LlmProviders(custom_llm_provider)
|
||||
image_variation_provider = LITELLM_IMAGE_VARIATION_PROVIDERS(llm_provider)
|
||||
llm_provider: Final = LlmProviders(custom_llm_provider)
|
||||
image_variation_provider: Final = LITELLM_IMAGE_VARIATION_PROVIDERS(llm_provider)
|
||||
except ValueError:
|
||||
raise ValueError(
|
||||
f"Invalid image variation provider: {custom_llm_provider}. Supported providers are: {LITELLM_IMAGE_VARIATION_PROVIDERS}"
|
||||
)
|
||||
model_response = ImageResponse()
|
||||
model_response: Final = ImageResponse()
|
||||
|
||||
response: ImageResponse | None = None
|
||||
|
||||
provider_config = ProviderConfigManager.get_provider_model_info(
|
||||
provider_config: Final = ProviderConfigManager.get_provider_model_info(
|
||||
model=model or "", # openai defaults to dall-e-2
|
||||
provider=llm_provider,
|
||||
)
|
||||
|
|
@ -654,7 +647,7 @@ def image_variation(
|
|||
f"image variation provider has no known model info config - required for getting api keys, etc.: {custom_llm_provider}. Supported providers are: {LITELLM_IMAGE_VARIATION_PROVIDERS}"
|
||||
)
|
||||
|
||||
api_key = provider_config.get_api_key(litellm_params.get("api_key", None))
|
||||
api_key: Final = provider_config.get_api_key(litellm_params.get("api_key", None))
|
||||
api_base = provider_config.get_api_base(litellm_params.get("api_base", None))
|
||||
|
||||
if image_variation_provider == LITELLM_IMAGE_VARIATION_PROVIDERS.OPENAI:
|
||||
|
|
@ -727,9 +720,9 @@ def image_edit(
|
|||
"""
|
||||
Maps the image edit functionality, similar to OpenAI's images/edits endpoint.
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
openai_params = [
|
||||
openai_params: Final = [
|
||||
"user",
|
||||
"request_timeout",
|
||||
"api_base",
|
||||
|
|
@ -747,22 +740,22 @@ def image_edit(
|
|||
"style",
|
||||
"async_call",
|
||||
]
|
||||
litellm_params_list = all_litellm_params
|
||||
default_params = openai_params + litellm_params_list
|
||||
non_default_params = {
|
||||
litellm_params_list: Final = all_litellm_params
|
||||
default_params: Final = openai_params + litellm_params_list
|
||||
non_default_params: Final = {
|
||||
k: v for k, v in kwargs.items() if k not in default_params
|
||||
} # model-specific params - pass them straight to the model/provider
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
model_info = kwargs.get("model_info", None)
|
||||
metadata = kwargs.get("metadata", {})
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
model_info: Final = kwargs.get("model_info", None)
|
||||
metadata: Final = kwargs.get("metadata", {})
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
# add images / or return a single image
|
||||
images = image if isinstance(image, list) else ([image] if image is not None else [])
|
||||
images: Final = image if isinstance(image, list) else ([image] if image is not None else [])
|
||||
|
||||
headers_from_kwargs = kwargs.get("headers")
|
||||
merged_extra_headers: dict[str, Any] = {}
|
||||
headers_from_kwargs: Final = kwargs.get("headers")
|
||||
merged_extra_headers: Final[dict[str, Any]] = {}
|
||||
if isinstance(headers_from_kwargs, dict):
|
||||
merged_extra_headers.update(headers_from_kwargs)
|
||||
if isinstance(extra_headers, dict):
|
||||
|
|
@ -772,7 +765,7 @@ def image_edit(
|
|||
extra_headers = dict(merged_extra_headers)
|
||||
|
||||
# get llm provider logic
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
model, custom_llm_provider, _, _ = get_llm_provider(
|
||||
model=model or DEFAULT_IMAGE_ENDPOINT_MODEL,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -788,7 +781,7 @@ def image_edit(
|
|||
if custom_handler is None:
|
||||
raise LiteLLMUnknownProvider(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
model_response = ImageResponse()
|
||||
model_response: Final = ImageResponse()
|
||||
|
||||
if _is_async:
|
||||
async_custom_client: AsyncHTTPHandler | None = None
|
||||
|
|
@ -836,11 +829,11 @@ def image_edit(
|
|||
|
||||
local_vars.update(kwargs)
|
||||
# Get ImageEditOptionalRequestParams with only valid parameters
|
||||
image_edit_optional_params: ImageEditOptionalRequestParams = (
|
||||
image_edit_optional_params: Final[ImageEditOptionalRequestParams] = (
|
||||
_get_ImageEditRequestUtils().get_requested_image_edit_optional_param(local_vars)
|
||||
)
|
||||
# Get optional parameters for the responses API
|
||||
image_edit_request_params: dict = _get_ImageEditRequestUtils().get_optional_params_image_edit(
|
||||
image_edit_request_params: Final[dict] = _get_ImageEditRequestUtils().get_optional_params_image_edit(
|
||||
model=model,
|
||||
image_edit_provider_config=image_edit_provider_config,
|
||||
image_edit_optional_params=image_edit_optional_params,
|
||||
|
|
@ -973,9 +966,9 @@ async def aimage_edit(
|
|||
Returns:
|
||||
- `response` (Any): The response returned by the `image_edit` function.
|
||||
"""
|
||||
local_vars = locals()
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
# get custom llm provider so we can use this for mapping exceptions
|
||||
|
|
@ -984,9 +977,9 @@ async def aimage_edit(
|
|||
model=model, api_base=local_vars.get("base_url", None)
|
||||
)
|
||||
|
||||
images = image if isinstance(image, list) else [image]
|
||||
images: Final = image if isinstance(image, list) else [image]
|
||||
|
||||
func = partial(
|
||||
func: Final = partial(
|
||||
image_edit,
|
||||
image=images,
|
||||
prompt=prompt,
|
||||
|
|
@ -1002,9 +995,9 @@ async def aimage_edit(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
|
|
@ -1029,7 +1022,7 @@ def __getattr__(name: str) -> Any:
|
|||
from .utils import ImageEditRequestUtils as _ImageEditRequestUtils
|
||||
|
||||
# Cache it in the module's __dict__ for subsequent accesses
|
||||
module = importlib.import_module(__name__)
|
||||
module: Final = importlib.import_module(__name__)
|
||||
module.__dict__["ImageEditRequestUtils"] = _ImageEditRequestUtils
|
||||
return _ImageEditRequestUtils
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from io import BufferedReader, BytesIO
|
||||
from typing import Any, cast, get_type_hints
|
||||
from typing import Any, Final, cast, get_type_hints
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.token_counter import get_image_type
|
||||
|
|
@ -30,16 +30,16 @@ class ImageEditRequestUtils:
|
|||
Returns:
|
||||
A dictionary of supported parameters for the image edit API
|
||||
"""
|
||||
supported_params = image_edit_provider_config.get_supported_openai_params(model)
|
||||
supported_params: Final = image_edit_provider_config.get_supported_openai_params(model)
|
||||
|
||||
should_drop = litellm.drop_params is True or drop_params is True
|
||||
should_drop: Final = litellm.drop_params is True or drop_params is True
|
||||
|
||||
filtered_optional_params = dict(image_edit_optional_params)
|
||||
filtered_optional_params: Final = dict(image_edit_optional_params)
|
||||
if additional_drop_params:
|
||||
for param in additional_drop_params:
|
||||
filtered_optional_params.pop(param, None)
|
||||
|
||||
unsupported_params = [param for param in filtered_optional_params if param not in supported_params]
|
||||
unsupported_params: Final = [param for param in filtered_optional_params if param not in supported_params]
|
||||
|
||||
if unsupported_params:
|
||||
if should_drop:
|
||||
|
|
@ -51,7 +51,7 @@ class ImageEditRequestUtils:
|
|||
message=f"The following parameters are not supported for model {model}: {', '.join(unsupported_params)}",
|
||||
)
|
||||
|
||||
mapped_params = image_edit_provider_config.map_openai_params(
|
||||
mapped_params: Final = image_edit_provider_config.map_openai_params(
|
||||
image_edit_optional_params=cast(ImageEditOptionalRequestParams, filtered_optional_params),
|
||||
model=model,
|
||||
drop_params=should_drop,
|
||||
|
|
@ -72,8 +72,8 @@ class ImageEditRequestUtils:
|
|||
Returns:
|
||||
ImageEditOptionalRequestParams instance with only the valid parameters
|
||||
"""
|
||||
valid_keys = get_type_hints(ImageEditOptionalRequestParams).keys()
|
||||
filtered_params = {k: v for k, v in params.items() if k in valid_keys and v is not None}
|
||||
valid_keys: Final = get_type_hints(ImageEditOptionalRequestParams).keys()
|
||||
filtered_params: Final = {k: v for k, v in params.items() if k in valid_keys and v is not None}
|
||||
return cast(ImageEditOptionalRequestParams, filtered_params)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -118,13 +118,13 @@ class ImageEditRequestUtils:
|
|||
return FILE_MIME_TYPES[FileType.PNG] # Default fallback
|
||||
|
||||
# Use the existing get_image_type function to detect image type
|
||||
image_type_str = get_image_type(bytes_data)
|
||||
image_type_str: Final = get_image_type(bytes_data)
|
||||
|
||||
if image_type_str is None:
|
||||
return FILE_MIME_TYPES[FileType.PNG] # Default if detection fails
|
||||
|
||||
# Map detected type string to FileType enum and get MIME type
|
||||
type_mapping = {
|
||||
type_mapping: Final = {
|
||||
"png": FileType.PNG,
|
||||
"jpeg": FileType.JPEG,
|
||||
"gif": FileType.GIF,
|
||||
|
|
@ -132,7 +132,7 @@ class ImageEditRequestUtils:
|
|||
"heic": FileType.HEIC,
|
||||
}
|
||||
|
||||
file_type = type_mapping.get(image_type_str)
|
||||
file_type: Final = type_mapping.get(image_type_str)
|
||||
if file_type is None:
|
||||
return FILE_MIME_TYPES[FileType.PNG] # Default to PNG if unknown
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ Slack alerts are sent every 10s or when events are greater than X events
|
|||
see custom_batch_logger.py for more details / defaults
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
|
|
@ -19,7 +19,7 @@ else:
|
|||
|
||||
|
||||
def squash_payloads(queue):
|
||||
squashed = {}
|
||||
squashed: Final = {}
|
||||
if len(queue) == 0:
|
||||
return squashed
|
||||
if len(queue) == 1:
|
||||
|
|
@ -57,12 +57,12 @@ async def send_to_webhook(slackAlertingInstance: SlackAlertingType, item, count)
|
|||
"""
|
||||
import json
|
||||
|
||||
payload = item.get("payload", {})
|
||||
payload: Final = item.get("payload", {})
|
||||
try:
|
||||
if count > 1:
|
||||
payload["text"] = f"[Num Alerts: {count}]\n\n{payload['text']}"
|
||||
|
||||
response = await slackAlertingInstance.async_http_handler.post(
|
||||
response: Final = await slackAlertingInstance.async_http_handler.post(
|
||||
url=item["url"],
|
||||
headers=item["headers"],
|
||||
data=json.dumps(payload),
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from abc import ABC, abstractmethod
|
||||
from typing import Literal
|
||||
from typing import Final, Literal
|
||||
|
||||
from litellm.proxy._types import CallInfo, Litellm_EntityType
|
||||
|
||||
|
|
@ -97,7 +97,7 @@ def get_budget_alert_type(
|
|||
) -> BaseBudgetAlertType:
|
||||
"""Factory function to get the appropriate budget alert type class"""
|
||||
|
||||
alert_types = {
|
||||
alert_types: Final = {
|
||||
"proxy_budget": ProxyBudgetAlert(),
|
||||
"soft_budget": SoftBudgetAlert(),
|
||||
"user_budget": UserBudgetAlert(),
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ Notes:
|
|||
|
||||
import asyncio
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -57,8 +57,8 @@ class AlertingHangingRequestCheck:
|
|||
if request_data is None:
|
||||
return
|
||||
|
||||
request_metadata = get_litellm_metadata_from_kwargs(kwargs=request_data)
|
||||
model = request_data.get("model", "")
|
||||
request_metadata: Final = get_litellm_metadata_from_kwargs(kwargs=request_data)
|
||||
model: Final = request_data.get("model", "")
|
||||
api_base: str | None = None
|
||||
|
||||
if request_data.get("deployment", None) is not None and isinstance(request_data["deployment"], dict):
|
||||
|
|
@ -67,7 +67,7 @@ class AlertingHangingRequestCheck:
|
|||
optional_params=request_data["deployment"].get("litellm_params", {}),
|
||||
)
|
||||
|
||||
hanging_request_data = HangingRequestData(
|
||||
hanging_request_data: Final = HangingRequestData(
|
||||
request_id=request_data.get("litellm_call_id", ""),
|
||||
model=model,
|
||||
api_base=api_base,
|
||||
|
|
@ -96,7 +96,7 @@ class AlertingHangingRequestCheck:
|
|||
if proxy_logging_obj.internal_usage_cache is None:
|
||||
return
|
||||
|
||||
hanging_requests = await self.hanging_request_cache.async_get_oldest_n_keys(
|
||||
hanging_requests: Final = await self.hanging_request_cache.async_get_oldest_n_keys(
|
||||
n=MAX_OLDEST_HANGING_REQUESTS_TO_CHECK,
|
||||
)
|
||||
|
||||
|
|
@ -166,7 +166,7 @@ class AlertingHangingRequestCheck:
|
|||
################
|
||||
# Send the Alert on Slack
|
||||
################
|
||||
request_info = f"""Request Model: `{hanging_request_data.model}`
|
||||
request_info: Final = f"""Request Model: `{hanging_request_data.model}`
|
||||
API Base: `{hanging_request_data.api_base}`
|
||||
Key Alias: `{hanging_request_data.key_alias}`
|
||||
Team Alias: `{hanging_request_data.team_alias}`"""
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ import os
|
|||
import random
|
||||
import time
|
||||
from datetime import timedelta
|
||||
from typing import TYPE_CHECKING, Any, Literal
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
|
||||
from openai import APIError
|
||||
|
||||
|
|
@ -128,7 +128,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
if self.alert_to_webhook_url is None:
|
||||
self.alert_to_webhook_url = process_slack_alerting_variables(alert_to_webhook_url=alert_to_webhook_url)
|
||||
else:
|
||||
_new_values = process_slack_alerting_variables(alert_to_webhook_url=alert_to_webhook_url) or {}
|
||||
_new_values: Final = process_slack_alerting_variables(alert_to_webhook_url=alert_to_webhook_url) or {}
|
||||
self.alert_to_webhook_url.update(_new_values)
|
||||
if llm_router is not None:
|
||||
self.llm_router = llm_router
|
||||
|
|
@ -139,7 +139,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
Converts set objects to lists for JSON serialization.
|
||||
"""
|
||||
# Convert to dict for processing
|
||||
cache_value = dict(outage_value)
|
||||
cache_value: Final = dict(outage_value)
|
||||
|
||||
if "deployment_ids" in cache_value and isinstance(cache_value["deployment_ids"], set):
|
||||
cache_value["deployment_ids"] = list(cache_value["deployment_ids"])
|
||||
|
|
@ -173,19 +173,19 @@ class SlackAlerting(CustomBatchLogger):
|
|||
end_time, # start/end time
|
||||
):
|
||||
try:
|
||||
time_difference = end_time - start_time
|
||||
time_difference: Final = end_time - start_time
|
||||
# Convert the timedelta to float (in seconds)
|
||||
time_difference_float = time_difference.total_seconds()
|
||||
litellm_params = kwargs.get("litellm_params", {})
|
||||
model = kwargs.get("model", "")
|
||||
api_base = litellm.get_api_base(model=model, optional_params=litellm_params)
|
||||
time_difference_float: Final = time_difference.total_seconds()
|
||||
litellm_params: Final = kwargs.get("litellm_params", {})
|
||||
model: Final = kwargs.get("model", "")
|
||||
api_base: Final = litellm.get_api_base(model=model, optional_params=litellm_params)
|
||||
messages = kwargs.get("messages", None)
|
||||
# if messages does not exist fallback to "input"
|
||||
if messages is None:
|
||||
messages = kwargs.get("input", None)
|
||||
|
||||
# only use first 100 chars for alerting
|
||||
_messages = str(messages)[:100]
|
||||
_messages: Final = str(messages)[:100]
|
||||
|
||||
return time_difference_float, model, api_base, _messages
|
||||
except Exception as e:
|
||||
|
|
@ -251,10 +251,10 @@ class SlackAlerting(CustomBatchLogger):
|
|||
if time_difference_float > self.alerting_threshold:
|
||||
# add deployment latencies to alert
|
||||
if kwargs is not None and "litellm_params" in kwargs and "metadata" in kwargs["litellm_params"]:
|
||||
_metadata: dict = kwargs["litellm_params"]["metadata"]
|
||||
_metadata: Final[dict] = kwargs["litellm_params"]["metadata"]
|
||||
request_info = _add_key_name_and_team_to_alert(request_info=request_info, metadata=_metadata)
|
||||
|
||||
_deployment_latency_map = self._get_deployment_latencies_to_alert(metadata=_metadata)
|
||||
_deployment_latency_map: Final = self._get_deployment_latencies_to_alert(metadata=_metadata)
|
||||
if _deployment_latency_map is not None:
|
||||
request_info += f"\nAvailable Deployment Latencies\n{_deployment_latency_map}"
|
||||
|
||||
|
|
@ -324,15 +324,15 @@ class SlackAlerting(CustomBatchLogger):
|
|||
False -> if not sent
|
||||
"""
|
||||
|
||||
ids = router.get_model_ids()
|
||||
ids: Final = router.get_model_ids()
|
||||
|
||||
# get keys
|
||||
failed_request_keys = [f"{id}:{SlackAlertingCacheKeys.failed_requests_key.value}" for id in ids]
|
||||
latency_keys = [f"{id}:{SlackAlertingCacheKeys.latency_key.value}" for id in ids]
|
||||
failed_request_keys: Final = [f"{id}:{SlackAlertingCacheKeys.failed_requests_key.value}" for id in ids]
|
||||
latency_keys: Final = [f"{id}:{SlackAlertingCacheKeys.latency_key.value}" for id in ids]
|
||||
|
||||
combined_metrics_keys = failed_request_keys + latency_keys # reduce cache calls
|
||||
combined_metrics_keys: Final = failed_request_keys + latency_keys # reduce cache calls
|
||||
|
||||
combined_metrics_values = await self.internal_usage_cache.async_batch_get_cache(
|
||||
combined_metrics_values: Final = await self.internal_usage_cache.async_batch_get_cache(
|
||||
keys=combined_metrics_keys
|
||||
) # [1, 2, None, ..]
|
||||
|
||||
|
|
@ -348,8 +348,8 @@ class SlackAlerting(CustomBatchLogger):
|
|||
if all_none:
|
||||
return False
|
||||
|
||||
failed_request_values = combined_metrics_values[: len(failed_request_keys)] # # [1, 2, None, ..]
|
||||
latency_values = combined_metrics_values[len(failed_request_keys) :]
|
||||
failed_request_values: Final = combined_metrics_values[: len(failed_request_keys)] # # [1, 2, None, ..]
|
||||
latency_values: Final = combined_metrics_values[len(failed_request_keys) :]
|
||||
|
||||
# find top 5 failed
|
||||
## Replace None values with a placeholder value (-1 in this case)
|
||||
|
|
@ -367,7 +367,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
# find top 5 slowest
|
||||
# Replace None values with a placeholder value (-1 in this case)
|
||||
placeholder_value = 0
|
||||
replaced_slowest_values = [value if value is not None else placeholder_value for value in latency_values]
|
||||
replaced_slowest_values: Final = [value if value is not None else placeholder_value for value in latency_values]
|
||||
|
||||
# Get the indices of top 5 values with the highest numerical values (ignoring None and 0 values)
|
||||
top_5_slowest = sorted(
|
||||
|
|
@ -420,9 +420,9 @@ class SlackAlerting(CustomBatchLogger):
|
|||
message += f"\t{i + 1}. Deployment: `{deployment_name}`, Latency per output token: `{value}s/token`, API Base: `{api_base}`\n\n"
|
||||
|
||||
# cache cleanup -> reset values to 0
|
||||
latency_cache_keys = [(key, 0) for key in latency_keys]
|
||||
failed_request_cache_keys = [(key, 0) for key in failed_request_keys]
|
||||
combined_metrics_cache_keys = latency_cache_keys + failed_request_cache_keys
|
||||
latency_cache_keys: Final = [(key, 0) for key in latency_keys]
|
||||
failed_request_cache_keys: Final = [(key, 0) for key in failed_request_keys]
|
||||
combined_metrics_cache_keys: Final = latency_cache_keys + failed_request_cache_keys
|
||||
await self.internal_usage_cache.async_set_cache_pipeline(cache_list=combined_metrics_cache_keys)
|
||||
|
||||
message += f"\n\nNext Run is at: `{time.time() + self.alerting_args.daily_report_frequency}`s"
|
||||
|
|
@ -463,10 +463,10 @@ class SlackAlerting(CustomBatchLogger):
|
|||
if "failed_tracking_spend" not in self.alert_types:
|
||||
return
|
||||
|
||||
_cache: DualCache = self.internal_usage_cache
|
||||
message = "Failed Tracking Cost for " + error_message
|
||||
_cache_key = f"budget_alerts:failed_tracking:{failing_model}"
|
||||
result = await _cache.async_get_cache(key=_cache_key)
|
||||
_cache: Final[DualCache] = self.internal_usage_cache
|
||||
message: Final = "Failed Tracking Cost for " + error_message
|
||||
_cache_key: Final = f"budget_alerts:failed_tracking:{failing_model}"
|
||||
result: Final = await _cache.async_get_cache(key=_cache_key)
|
||||
if result is None:
|
||||
await self.send_alert(
|
||||
message=message,
|
||||
|
|
@ -506,7 +506,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
# - Alert once within 24hr period
|
||||
# - Cache this information
|
||||
# - Don't re-alert, if alert already sent
|
||||
_cache: DualCache = self.internal_usage_cache
|
||||
_cache: Final[DualCache] = self.internal_usage_cache
|
||||
|
||||
if self.alerting is None or self.alert_types is None:
|
||||
# do nothing if alerting is not switched on
|
||||
|
|
@ -515,10 +515,10 @@ class SlackAlerting(CustomBatchLogger):
|
|||
return
|
||||
|
||||
# Get the appropriate budget alert type handler
|
||||
budget_alert_class = get_budget_alert_type(type)
|
||||
_id = budget_alert_class.get_id(user_info)
|
||||
user_info_json = user_info.model_dump(exclude_none=True)
|
||||
user_info_str = self._get_user_info_str(user_info)
|
||||
budget_alert_class: Final = get_budget_alert_type(type)
|
||||
_id: Final = budget_alert_class.get_id(user_info)
|
||||
user_info_json: Final = user_info.model_dump(exclude_none=True)
|
||||
user_info_str: Final = self._get_user_info_str(user_info)
|
||||
event_message = budget_alert_class.get_event_message()
|
||||
|
||||
# Set default event unless we're in projected_limit_exceeded
|
||||
|
|
@ -541,8 +541,8 @@ class SlackAlerting(CustomBatchLogger):
|
|||
|
||||
# send alert
|
||||
if event is not None and user_info.event_group is not None:
|
||||
_cache_key = f"budget_alerts:{event}:{_id}"
|
||||
result = await _cache.async_get_cache(key=_cache_key)
|
||||
_cache_key: Final = f"budget_alerts:{event}:{_id}"
|
||||
result: Final = await _cache.async_get_cache(key=_cache_key)
|
||||
if result is None:
|
||||
webhook_event = WebhookEvent(
|
||||
event=event,
|
||||
|
|
@ -581,7 +581,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
|
||||
Handles Max Budget and Soft Budget Alerts
|
||||
"""
|
||||
percent_left: float = self._get_percent_of_max_budget_left(user_info=user_info)
|
||||
percent_left: Final[float] = self._get_percent_of_max_budget_left(user_info=user_info)
|
||||
|
||||
#####################################################################
|
||||
# SOFT BUDGET CHECK
|
||||
|
|
@ -616,8 +616,8 @@ class SlackAlerting(CustomBatchLogger):
|
|||
Get the percent of the max budget that is left
|
||||
"""
|
||||
percent_left: float = 0.0
|
||||
current_spend: float = user_info.spend
|
||||
max_budget: float | None = user_info.max_budget
|
||||
current_spend: Final[float] = user_info.spend
|
||||
max_budget: Final[float | None] = user_info.max_budget
|
||||
if max_budget is None:
|
||||
return percent_left
|
||||
if max_budget <= 0:
|
||||
|
|
@ -629,7 +629,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
"""
|
||||
Create a standard message for a budget alert
|
||||
"""
|
||||
_all_fields_as_dict = user_info.model_dump(exclude_none=True)
|
||||
_all_fields_as_dict: Final = user_info.model_dump(exclude_none=True)
|
||||
_all_fields_as_dict.pop("token")
|
||||
msg = ""
|
||||
for k, v in _all_fields_as_dict.items():
|
||||
|
|
@ -655,7 +655,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
and response_cost is not None
|
||||
):
|
||||
# log customer spend
|
||||
event = WebhookEvent(
|
||||
event: Final = WebhookEvent(
|
||||
spend=response_cost,
|
||||
max_budget=max_budget,
|
||||
token=token,
|
||||
|
|
@ -681,7 +681,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
Returns:
|
||||
- str -> formatted string. This is an alert message, giving a human-friendly description of the errors.
|
||||
"""
|
||||
error_breakdown = {"Timeout Errors": 0, "API Errors": 0, "Unknown Errors": 0}
|
||||
error_breakdown: Final = {"Timeout Errors": 0, "API Errors": 0, "Unknown Errors": 0}
|
||||
for alert in alerts:
|
||||
if alert == 408:
|
||||
error_breakdown["Timeout Errors"] += 1
|
||||
|
|
@ -707,7 +707,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
outage_value: BaseOutageModel,
|
||||
) -> str:
|
||||
"""Format an alert message for slack"""
|
||||
headers = {f"{key} Name": key_val, "Provider": provider}
|
||||
headers: Final = {f"{key} Name": key_val, "Provider": provider}
|
||||
if api_base is not None:
|
||||
headers["API Base"] = api_base # type: ignore
|
||||
|
||||
|
|
@ -739,7 +739,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
if self.llm_router is None:
|
||||
return
|
||||
|
||||
deployment = self.llm_router.get_deployment(model_id=deployment_id)
|
||||
deployment: Final = self.llm_router.get_deployment(model_id=deployment_id)
|
||||
|
||||
if deployment is None:
|
||||
return
|
||||
|
|
@ -761,7 +761,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
return
|
||||
|
||||
### UNIQUE CACHE KEY ###
|
||||
cache_key = provider + region_name
|
||||
cache_key: Final = provider + region_name
|
||||
|
||||
outage_value: ProviderRegionOutageModel | None = await self.internal_usage_cache.async_get_cache(key=cache_key)
|
||||
|
||||
|
|
@ -896,7 +896,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
return
|
||||
|
||||
### EXTRACT MODEL DETAILS ###
|
||||
deployment = self.llm_router.get_deployment(model_id=deployment_id)
|
||||
deployment: Final = self.llm_router.get_deployment(model_id=deployment_id)
|
||||
if deployment is None:
|
||||
return
|
||||
|
||||
|
|
@ -907,7 +907,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
model, provider, _, _ = litellm.get_llm_provider(model=model)
|
||||
except Exception:
|
||||
provider = ""
|
||||
api_base = litellm.get_api_base(model=model, optional_params=deployment.litellm_params)
|
||||
api_base: Final = litellm.get_api_base(model=model, optional_params=deployment.litellm_params)
|
||||
|
||||
if outage_value is None:
|
||||
outage_value = OutageModel(
|
||||
|
|
@ -979,13 +979,13 @@ class SlackAlerting(CustomBatchLogger):
|
|||
|
||||
## update cache ##
|
||||
# Convert set to list for JSON serialization
|
||||
cache_value = self._prepare_outage_value_for_cache(outage_value)
|
||||
cache_value: Final = self._prepare_outage_value_for_cache(outage_value)
|
||||
await self.internal_usage_cache.async_set_cache(key=deployment_id, value=cache_value)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def model_added_alert(self, model_name: str, litellm_model_name: str, passed_model_info: Any):
|
||||
base_model_from_user = getattr(passed_model_info, "base_model", None)
|
||||
base_model_from_user: Final = getattr(passed_model_info, "base_model", None)
|
||||
model_info = {}
|
||||
base_model = ""
|
||||
if base_model_from_user is not None:
|
||||
|
|
@ -1001,7 +1001,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
|
||||
model_info_str += f"{k}: {v}\n"
|
||||
|
||||
message = f"""
|
||||
message: Final = f"""
|
||||
*🚅 New Model Added*
|
||||
Model Name: `{model_name}`
|
||||
{base_model}
|
||||
|
|
@ -1031,7 +1031,7 @@ Model Info:
|
|||
```
|
||||
"""
|
||||
|
||||
alert_val = self.send_alert(
|
||||
alert_val: Final = self.send_alert(
|
||||
message=message,
|
||||
level="Low",
|
||||
alert_type=AlertType.new_model_added,
|
||||
|
|
@ -1056,14 +1056,14 @@ Model Info:
|
|||
- if WEBHOOK_URL is not set
|
||||
"""
|
||||
|
||||
webhook_url = os.getenv("WEBHOOK_URL", None)
|
||||
webhook_url: Final = os.getenv("WEBHOOK_URL", None)
|
||||
if webhook_url is None:
|
||||
raise Exception("Missing webhook_url from environment")
|
||||
|
||||
payload = webhook_event.model_dump_json()
|
||||
headers = {"Content-type": "application/json"}
|
||||
payload: Final = webhook_event.model_dump_json()
|
||||
headers: Final = {"Content-type": "application/json"}
|
||||
|
||||
response = await self.async_http_handler.post(
|
||||
response: Final = await self.async_http_handler.post(
|
||||
url=webhook_url,
|
||||
headers=headers,
|
||||
data=payload,
|
||||
|
|
@ -1108,18 +1108,18 @@ Model Info:
|
|||
if email_support_contact is None:
|
||||
email_support_contact = LITELLM_SUPPORT_CONTACT
|
||||
|
||||
event_name = webhook_event.event_message
|
||||
event_name: Final = webhook_event.event_message
|
||||
recipient_email = webhook_event.user_email
|
||||
recipient_user_id = webhook_event.user_id
|
||||
recipient_user_id: Final = webhook_event.user_id
|
||||
if recipient_email is None and recipient_user_id is not None and prisma_client is not None:
|
||||
user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": recipient_user_id})
|
||||
|
||||
if user_row is not None:
|
||||
recipient_email = user_row.user_email
|
||||
|
||||
key_token = webhook_event.token
|
||||
key_budget = webhook_event.max_budget
|
||||
base_url = os.getenv("PROXY_BASE_URL", "http://0.0.0.0:4000")
|
||||
key_token: Final = webhook_event.token
|
||||
key_budget: Final = webhook_event.max_budget
|
||||
base_url: Final = os.getenv("PROXY_BASE_URL", "http://0.0.0.0:4000")
|
||||
|
||||
email_html_content = "Alert from LiteLLM Server"
|
||||
if recipient_email is None:
|
||||
|
|
@ -1139,10 +1139,10 @@ Model Info:
|
|||
)
|
||||
elif webhook_event.event == "internal_user_created":
|
||||
# GET TEAM NAME
|
||||
team_id = webhook_event.team_id
|
||||
team_id: Final = webhook_event.team_id
|
||||
team_name = "Default Team"
|
||||
if team_id is not None and prisma_client is not None:
|
||||
team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
|
||||
team_row: Final = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
|
||||
if team_row is not None:
|
||||
team_name = team_row.team_alias or "-"
|
||||
email_html_content = USER_INVITED_EMAIL_TEMPLATE.format(
|
||||
|
|
@ -1159,7 +1159,7 @@ Model Info:
|
|||
)
|
||||
|
||||
webhook_event.model_dump_json()
|
||||
email_event = {
|
||||
email_event: Final = {
|
||||
"to": recipient_email,
|
||||
"subject": f"LiteLLM: {event_name}",
|
||||
"html": email_html_content,
|
||||
|
|
@ -1197,10 +1197,10 @@ Model Info:
|
|||
if email_support_contact is None:
|
||||
email_support_contact = LITELLM_SUPPORT_CONTACT
|
||||
|
||||
event_name = webhook_event.event_message
|
||||
recipient_email = webhook_event.user_email
|
||||
user_name = webhook_event.user_id
|
||||
max_budget = webhook_event.max_budget
|
||||
event_name: Final = webhook_event.event_message
|
||||
recipient_email: Final = webhook_event.user_email
|
||||
user_name: Final = webhook_event.user_id
|
||||
max_budget: Final = webhook_event.max_budget
|
||||
email_html_content = "Alert from LiteLLM Server"
|
||||
if recipient_email is None:
|
||||
verbose_proxy_logger.error("Trying to send email alert to no recipient", extra=webhook_event.dict())
|
||||
|
|
@ -1222,7 +1222,7 @@ Model Info:
|
|||
"""
|
||||
|
||||
webhook_event.model_dump_json()
|
||||
email_event = {
|
||||
email_event: Final = {
|
||||
"to": recipient_email,
|
||||
"subject": f"LiteLLM: {event_name}",
|
||||
"html": email_html_content,
|
||||
|
|
@ -1290,8 +1290,8 @@ Model Info:
|
|||
from datetime import datetime
|
||||
|
||||
# Check if digest mode is enabled for this alert type
|
||||
alert_type_name_str = getattr(alert_type, "value", str(alert_type))
|
||||
_atc = self.alert_type_config.get(alert_type_name_str)
|
||||
alert_type_name_str: Final = getattr(alert_type, "value", str(alert_type))
|
||||
_atc: Final = self.alert_type_config.get(alert_type_name_str)
|
||||
if _atc is not None and _atc.digest:
|
||||
# Resolve webhook URL for this alert type (needed for digest entry)
|
||||
if self.alert_to_webhook_url is not None and alert_type in self.alert_to_webhook_url:
|
||||
|
|
@ -1303,10 +1303,10 @@ Model Info:
|
|||
if _digest_webhook is None:
|
||||
raise ValueError("Missing SLACK_WEBHOOK_URL from environment")
|
||||
|
||||
digest_key = f"{alert_type_name_str}:{request_model or ''}:{api_base or ''}"
|
||||
digest_key: Final = f"{alert_type_name_str}:{request_model or ''}:{api_base or ''}"
|
||||
|
||||
async with self.digest_lock:
|
||||
now = datetime.now()
|
||||
now: Final = datetime.now()
|
||||
if digest_key in self.digest_buckets:
|
||||
self.digest_buckets[digest_key]["count"] += 1
|
||||
self.digest_buckets[digest_key]["last_time"] = now
|
||||
|
|
@ -1325,11 +1325,11 @@ Model Info:
|
|||
return # Suppress immediate alert; will be emitted by _flush_digest_buckets
|
||||
|
||||
# Get the current timestamp
|
||||
current_time = datetime.now().strftime("%H:%M:%S")
|
||||
_proxy_base_url = os.getenv("PROXY_BASE_URL", None)
|
||||
current_time: Final = datetime.now().strftime("%H:%M:%S")
|
||||
_proxy_base_url: Final = os.getenv("PROXY_BASE_URL", None)
|
||||
# Use .name if it's an enum, otherwise use as is
|
||||
alert_type_name = getattr(alert_type, "name", alert_type)
|
||||
alert_type_formatted = f"Alert type: `{alert_type_name}`"
|
||||
alert_type_name: Final = getattr(alert_type, "name", alert_type)
|
||||
alert_type_formatted: Final = f"Alert type: `{alert_type_name}`"
|
||||
if alert_type == "daily_reports" or alert_type == "new_model_added":
|
||||
formatted_message = alert_type_formatted + message
|
||||
else:
|
||||
|
|
@ -1356,8 +1356,8 @@ Model Info:
|
|||
|
||||
if slack_webhook_url is None:
|
||||
raise ValueError("Missing SLACK_WEBHOOK_URL from environment")
|
||||
payload = {"text": formatted_message}
|
||||
headers = {"Content-type": "application/json"}
|
||||
payload: Final = {"text": formatted_message}
|
||||
headers: Final = {"Content-type": "application/json"}
|
||||
|
||||
if isinstance(slack_webhook_url, list):
|
||||
for url in slack_webhook_url:
|
||||
|
|
@ -1386,8 +1386,8 @@ Model Info:
|
|||
if not self.log_queue:
|
||||
return
|
||||
|
||||
squashed_queue = squash_payloads(self.log_queue)
|
||||
tasks = [
|
||||
squashed_queue: Final = squash_payloads(self.log_queue)
|
||||
tasks: Final = [
|
||||
send_to_webhook(slackAlertingInstance=self, item=item["item"], count=item["count"])
|
||||
for item in squashed_queue.values()
|
||||
]
|
||||
|
|
@ -1402,8 +1402,8 @@ Model Info:
|
|||
"""
|
||||
from datetime import datetime
|
||||
|
||||
now = datetime.now()
|
||||
flushed_keys: list[str] = []
|
||||
now: Final = datetime.now()
|
||||
flushed_keys: Final[list[str]] = []
|
||||
|
||||
async with self.digest_lock:
|
||||
for key, entry in self.digest_buckets.items():
|
||||
|
|
@ -1474,10 +1474,10 @@ Model Info:
|
|||
"""Log deployment latency"""
|
||||
try:
|
||||
if "daily_reports" in self.alert_types:
|
||||
litellm_params = kwargs.get("litellm_params", {}) or {}
|
||||
model_info = litellm_params.get("model_info", {}) or {}
|
||||
model_id = model_info.get("id", "") or ""
|
||||
response_s: timedelta = end_time - start_time
|
||||
litellm_params: Final = kwargs.get("litellm_params", {}) or {}
|
||||
model_info: Final = litellm_params.get("model_info", {}) or {}
|
||||
model_id: Final = model_info.get("id", "") or ""
|
||||
response_s: Final[timedelta] = end_time - start_time
|
||||
|
||||
final_value = response_s
|
||||
|
||||
|
|
@ -1486,7 +1486,7 @@ Model Info:
|
|||
and response_obj.usage is not None # type: ignore
|
||||
and hasattr(response_obj.usage, "completion_tokens") # type: ignore
|
||||
):
|
||||
completion_tokens = response_obj.usage.completion_tokens # type: ignore
|
||||
completion_tokens: Final = response_obj.usage.completion_tokens # type: ignore
|
||||
if completion_tokens is not None and completion_tokens > 0:
|
||||
final_value = float(response_s.total_seconds() / completion_tokens)
|
||||
if isinstance(final_value, timedelta):
|
||||
|
|
@ -1507,9 +1507,9 @@ Model Info:
|
|||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
"""Log failure + deployment latency"""
|
||||
_litellm_params = kwargs.get("litellm_params", {})
|
||||
_model_info = _litellm_params.get("model_info", {}) or {}
|
||||
model_id = _model_info.get("id", "")
|
||||
_litellm_params: Final = kwargs.get("litellm_params", {})
|
||||
_model_info: Final = _litellm_params.get("model_info", {}) or {}
|
||||
model_id: Final = _model_info.get("id", "")
|
||||
try:
|
||||
if "daily_reports" in self.alert_types:
|
||||
try:
|
||||
|
|
@ -1544,12 +1544,12 @@ Model Info:
|
|||
"""
|
||||
report_sent_bool = False
|
||||
|
||||
report_sent = await self.internal_usage_cache.async_get_cache(
|
||||
report_sent: Final = await self.internal_usage_cache.async_get_cache(
|
||||
key=SlackAlertingCacheKeys.report_sent_key.value,
|
||||
parent_otel_span=None,
|
||||
) # None | float
|
||||
|
||||
current_time = time.time()
|
||||
current_time: Final = time.time()
|
||||
|
||||
if report_sent is None:
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
|
|
@ -1558,7 +1558,7 @@ Model Info:
|
|||
)
|
||||
elif isinstance(report_sent, float):
|
||||
# Check if current time - interval >= time last sent
|
||||
interval_seconds = self.alerting_args.daily_report_frequency
|
||||
interval_seconds: Final = self.alerting_args.daily_report_frequency
|
||||
|
||||
if current_time - report_sent >= interval_seconds:
|
||||
# Sneak in the reporting logic here
|
||||
|
|
@ -1612,20 +1612,20 @@ Model Info:
|
|||
)
|
||||
|
||||
# Parse the time range
|
||||
days = int(time_range[:-1])
|
||||
days: Final = int(time_range[:-1])
|
||||
if time_range[-1].lower() != "d":
|
||||
raise ValueError("Time range must be specified in days, e.g., '7d'")
|
||||
|
||||
todays_date = datetime.datetime.now().date()
|
||||
start_date = todays_date - datetime.timedelta(days=days)
|
||||
todays_date: Final = datetime.datetime.now().date()
|
||||
start_date: Final = todays_date - datetime.timedelta(days=days)
|
||||
|
||||
_event_cache_key = (
|
||||
_event_cache_key: Final = (
|
||||
f"weekly_spend_report_sent_{start_date.strftime('%Y-%m-%d')}_{todays_date.strftime('%Y-%m-%d')}"
|
||||
)
|
||||
if await self.internal_usage_cache.async_get_cache(key=_event_cache_key):
|
||||
return
|
||||
|
||||
_resp = await _get_spend_report_for_time_range(
|
||||
_resp: Final = await _get_spend_report_for_time_range(
|
||||
start_date=start_date.strftime("%Y-%m-%d"),
|
||||
end_date=todays_date.strftime("%Y-%m-%d"),
|
||||
)
|
||||
|
|
@ -1675,8 +1675,8 @@ Model Info:
|
|||
_get_spend_report_for_time_range,
|
||||
)
|
||||
|
||||
todays_date = datetime.datetime.now().date()
|
||||
first_day_of_month = todays_date.replace(day=1)
|
||||
todays_date: Final = datetime.datetime.now().date()
|
||||
first_day_of_month: Final = todays_date.replace(day=1)
|
||||
_, last_day_of_month = monthrange(todays_date.year, todays_date.month)
|
||||
last_day_of_month = first_day_of_month + datetime.timedelta(days=last_day_of_month - 1)
|
||||
|
||||
|
|
@ -1684,7 +1684,7 @@ Model Info:
|
|||
if await self.internal_usage_cache.async_get_cache(key=_event_cache_key):
|
||||
return
|
||||
|
||||
_resp = await _get_spend_report_for_time_range(
|
||||
_resp: Final = await _get_spend_report_for_time_range(
|
||||
start_date=first_day_of_month.strftime("%Y-%m-%d"),
|
||||
end_date=last_day_of_month.strftime("%Y-%m-%d"),
|
||||
)
|
||||
|
|
@ -1742,9 +1742,9 @@ Model Info:
|
|||
)
|
||||
|
||||
# call prometheuslogger.
|
||||
falllback_success_info_prometheus = await get_fallback_metric_from_prometheus()
|
||||
falllback_success_info_prometheus: Final = await get_fallback_metric_from_prometheus()
|
||||
|
||||
fallback_message = f"*Fallback Statistics:*\n{falllback_success_info_prometheus}"
|
||||
fallback_message: Final = f"*Fallback Statistics:*\n{falllback_success_info_prometheus}"
|
||||
|
||||
await self.send_alert(
|
||||
message=fallback_message,
|
||||
|
|
@ -1773,7 +1773,7 @@ Model Info:
|
|||
try:
|
||||
message = f"`{event_name}`\n"
|
||||
|
||||
key_event_dict = key_event.model_dump()
|
||||
key_event_dict: Final = key_event.model_dump()
|
||||
|
||||
# Add Created by information first
|
||||
message += "*Action Done by:*\n"
|
||||
|
|
@ -1783,7 +1783,7 @@ Model Info:
|
|||
|
||||
# Add args sent to function in the alert
|
||||
message += "\n*Arguments passed:*\n"
|
||||
request_kwargs = key_event.request_kwargs
|
||||
request_kwargs: Final = key_event.request_kwargs
|
||||
for key, value in request_kwargs.items():
|
||||
if key == "user_api_key_dict":
|
||||
continue
|
||||
|
|
@ -1808,8 +1808,8 @@ Model Info:
|
|||
|
||||
if request_data.get("litellm_status", "") != "success" and request_data.get("litellm_status", "") != "fail":
|
||||
## CHECK IF CACHE IS UPDATED
|
||||
litellm_call_id = request_data.get("litellm_call_id", "")
|
||||
status: str | None = await self.internal_usage_cache.async_get_cache(
|
||||
litellm_call_id: Final = request_data.get("litellm_call_id", "")
|
||||
status: Final[str | None] = await self.internal_usage_cache.async_get_cache(
|
||||
key=f"request_status:{litellm_call_id}", local_only=True
|
||||
)
|
||||
if status is not None and (status == "success" or status == "fail"):
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ Utils used for slack alerting
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm.proxy._types import AlertType
|
||||
|
|
@ -74,7 +74,7 @@ async def _add_langfuse_trace_id_to_alert(
|
|||
|
||||
if request_data is not None and request_data.get("litellm_logging_obj", None) is not None:
|
||||
trace_id: str | None = None
|
||||
litellm_logging_obj: Logging = request_data["litellm_logging_obj"]
|
||||
litellm_logging_obj: Final[Logging] = request_data["litellm_logging_obj"]
|
||||
|
||||
for _ in range(3):
|
||||
trace_id = litellm_logging_obj._get_trace_id(service_name="langfuse")
|
||||
|
|
@ -82,9 +82,9 @@ async def _add_langfuse_trace_id_to_alert(
|
|||
break
|
||||
await asyncio.sleep(3) # wait 3s before retrying for trace id
|
||||
#########################################################
|
||||
langfuse_object = litellm_logging_obj._get_callback_object(service_name="langfuse")
|
||||
langfuse_object: Final = litellm_logging_obj._get_callback_object(service_name="langfuse")
|
||||
if langfuse_object is not None:
|
||||
base_url = langfuse_object.Langfuse.base_url
|
||||
base_url: Final = langfuse_object.Langfuse.base_url
|
||||
return f"{base_url}/trace/{trace_id}"
|
||||
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ AgentOps integration for LiteLLM - Provides OpenTelemetry tracing for LLM calls
|
|||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
|
||||
|
|
@ -58,21 +58,21 @@ class AgentOps(OpenTelemetry):
|
|||
project_id = None
|
||||
if config.api_key:
|
||||
try:
|
||||
response = self._fetch_auth_token(config.api_key, config.auth_endpoint)
|
||||
response: Final = self._fetch_auth_token(config.api_key, config.auth_endpoint)
|
||||
jwt_token = response.get("token")
|
||||
project_id = response.get("project_id")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
headers = f"Authorization=Bearer {jwt_token}" if jwt_token else None
|
||||
headers: Final = f"Authorization=Bearer {jwt_token}" if jwt_token else None
|
||||
|
||||
otel_config = OpenTelemetryConfig(exporter="otlp_http", endpoint=config.endpoint, headers=headers)
|
||||
otel_config: Final = OpenTelemetryConfig(exporter="otlp_http", endpoint=config.endpoint, headers=headers)
|
||||
|
||||
# Initialize OpenTelemetry with our config
|
||||
super().__init__(config=otel_config, callback_name="agentops")
|
||||
|
||||
# Set AgentOps-specific resource attributes
|
||||
resource_attrs = {
|
||||
resource_attrs: Final = {
|
||||
"service.name": config.service_name or "litellm",
|
||||
"deployment.environment": config.deployment_environment or "production",
|
||||
"telemetry.sdk.name": "agentops",
|
||||
|
|
@ -94,14 +94,14 @@ class AgentOps(OpenTelemetry):
|
|||
Returns:
|
||||
Dict containing JWT token and project ID
|
||||
"""
|
||||
headers = {
|
||||
headers: Final = {
|
||||
"Content-Type": "application/json",
|
||||
"Connection": "keep-alive",
|
||||
}
|
||||
|
||||
client = _get_httpx_client()
|
||||
client: Final = _get_httpx_client()
|
||||
try:
|
||||
response = client.post(
|
||||
response: Final = client.post(
|
||||
url=auth_endpoint,
|
||||
headers=headers,
|
||||
json={"api_key": api_key},
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ Supported for both `v1/chat/completions` (via the prompt-management hook) and
|
|||
"""
|
||||
|
||||
import copy
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -32,7 +32,7 @@ else:
|
|||
|
||||
# Anthropic (and Bedrock Claude) reject requests with more than 4 cache_control
|
||||
# breakpoints: "A maximum of 4 blocks with cache_control may be provided."
|
||||
MAX_CACHE_CONTROL_BLOCKS = 4
|
||||
MAX_CACHE_CONTROL_BLOCKS: Final = 4
|
||||
|
||||
|
||||
class AnthropicCacheControlHook(CustomPromptManagement):
|
||||
|
|
@ -59,7 +59,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
- non_default_params: dict - params with any global cache controls
|
||||
"""
|
||||
# Extract cache control injection points
|
||||
injection_points: list[CacheControlInjectionPoint] = non_default_params.pop(
|
||||
injection_points: Final[list[CacheControlInjectionPoint]] = non_default_params.pop(
|
||||
"cache_control_injection_points", []
|
||||
)
|
||||
if not injection_points:
|
||||
|
|
@ -69,8 +69,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
processed_messages = copy.deepcopy(messages)
|
||||
|
||||
# Separate message-level and non-message-level injection points
|
||||
message_points: list[CacheControlMessageInjectionPoint] = []
|
||||
remaining_points: list[CacheControlInjectionPoint] = []
|
||||
message_points: Final[list[CacheControlMessageInjectionPoint]] = []
|
||||
remaining_points: Final[list[CacheControlInjectionPoint]] = []
|
||||
for point in injection_points:
|
||||
if point.get("location") == "message":
|
||||
message_points.append(cast(CacheControlMessageInjectionPoint, point))
|
||||
|
|
@ -81,7 +81,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
# provider transform, where each tool_config point appends at most one
|
||||
# cachePoint to the tools. That block also counts toward Anthropic's
|
||||
# limit, so reserve a slot for it here to leave room.
|
||||
reserved_blocks = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0
|
||||
reserved_blocks: Final = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0
|
||||
|
||||
processed_messages = self._apply_message_injections(
|
||||
points=message_points,
|
||||
|
|
@ -154,7 +154,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
point: CacheControlMessageInjectionPoint, messages: list[AllMessageValues]
|
||||
) -> list[int]:
|
||||
"""Resolve which message indices an injection point targets."""
|
||||
_targetted_index: int | str | None = point.get("index", None)
|
||||
_targetted_index: Final[int | str | None] = point.get("index", None)
|
||||
targetted_index: int | None = None
|
||||
if isinstance(_targetted_index, str):
|
||||
try:
|
||||
|
|
@ -166,7 +166,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
# Case 1: Target by specific index
|
||||
if targetted_index is not None:
|
||||
original_index = targetted_index
|
||||
original_index: Final = targetted_index
|
||||
if targetted_index < 0:
|
||||
targetted_index += len(messages)
|
||||
|
||||
|
|
@ -182,7 +182,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
return []
|
||||
|
||||
# Case 2: Target by role
|
||||
targetted_role = point.get("role", None)
|
||||
targetted_role: Final = point.get("role", None)
|
||||
if targetted_role is not None:
|
||||
return [idx for idx, msg in enumerate(messages) if msg.get("role") == targetted_role]
|
||||
|
||||
|
|
@ -194,7 +194,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
count = 0
|
||||
if message.get("cache_control") is not None:
|
||||
count += 1
|
||||
content = message.get("content")
|
||||
content: Final = message.get("content")
|
||||
if isinstance(content, list):
|
||||
for block in content:
|
||||
if isinstance(block, dict) and block.get("cache_control") is not None:
|
||||
|
|
@ -221,7 +221,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
Per Anthropic's API specification, when using multiple content blocks,
|
||||
only the last content block can have cache_control.
|
||||
"""
|
||||
message_content = message.get("content", None)
|
||||
message_content: Final = message.get("content", None)
|
||||
|
||||
# 1. if string, insert cache control in the message
|
||||
if isinstance(message_content, str):
|
||||
|
|
@ -248,9 +248,9 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
processed_messages: list[dict] = copy.deepcopy(messages)
|
||||
processed_system = copy.deepcopy(system) if system is not None else None
|
||||
|
||||
message_points: list[CacheControlMessageInjectionPoint] = []
|
||||
system_points: list[CacheControlMessageInjectionPoint] = []
|
||||
remaining_points: list[CacheControlInjectionPoint] = []
|
||||
message_points: Final[list[CacheControlMessageInjectionPoint]] = []
|
||||
system_points: Final[list[CacheControlMessageInjectionPoint]] = []
|
||||
remaining_points: Final[list[CacheControlInjectionPoint]] = []
|
||||
|
||||
for point in injection_points:
|
||||
if point.get("location") == "message":
|
||||
|
|
@ -262,8 +262,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
else:
|
||||
remaining_points.append(point)
|
||||
|
||||
reserved_blocks = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0
|
||||
max_blocks = MAX_CACHE_CONTROL_BLOCKS - reserved_blocks
|
||||
reserved_blocks: Final = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0
|
||||
max_blocks: Final = MAX_CACHE_CONTROL_BLOCKS - reserved_blocks
|
||||
|
||||
used_blocks = sum(
|
||||
AnthropicCacheControlHook._count_cache_control_blocks(cast(AllMessageValues, msg))
|
||||
|
|
@ -275,11 +275,11 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
)
|
||||
|
||||
if system_points and processed_system is not None and used_blocks < max_blocks:
|
||||
system_already_has_cc = isinstance(processed_system, list) and any(
|
||||
system_already_has_cc: Final = isinstance(processed_system, list) and any(
|
||||
isinstance(b, dict) and b.get("cache_control") is not None for b in processed_system
|
||||
)
|
||||
if not system_already_has_cc:
|
||||
control = system_points[0].get("control") or ChatCompletionCachedContent(type="ephemeral")
|
||||
control: Final = system_points[0].get("control") or ChatCompletionCachedContent(type="ephemeral")
|
||||
if isinstance(processed_system, str):
|
||||
processed_system = [{"type": "text", "text": processed_system, "cache_control": control}]
|
||||
used_blocks += 1
|
||||
|
|
@ -309,7 +309,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
"""
|
||||
import litellm
|
||||
|
||||
ttl = litellm.anthropic_prompt_caching_ttl
|
||||
ttl: Final = litellm.anthropic_prompt_caching_ttl
|
||||
if ttl == "5m" or ttl == "1h":
|
||||
return ChatCompletionCachedContent(type="ephemeral", ttl=ttl)
|
||||
return ChatCompletionCachedContent(type="ephemeral")
|
||||
|
|
@ -419,8 +419,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
if AnthropicCacheControlHook._request_has_cache_control(messages, system, tools):
|
||||
return []
|
||||
|
||||
control = AnthropicCacheControlHook._default_control()
|
||||
points: list[CacheControlInjectionPoint] = [
|
||||
control: Final = AnthropicCacheControlHook._default_control()
|
||||
points: Final[list[CacheControlInjectionPoint]] = [
|
||||
CacheControlMessageInjectionPoint(location="message", role="system", index=None, control=control),
|
||||
CacheControlMessageInjectionPoint(location="message", role=None, index=-1, control=control),
|
||||
]
|
||||
|
|
@ -452,7 +452,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
):
|
||||
non_default_params.pop("cache_control_injection_points")
|
||||
return
|
||||
points = AnthropicCacheControlHook.get_default_injection_points(
|
||||
points: Final = AnthropicCacheControlHook.get_default_injection_points(
|
||||
messages=messages,
|
||||
system=None,
|
||||
model=model,
|
||||
|
|
@ -484,7 +484,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
downstream transforms can handle them.
|
||||
"""
|
||||
typed_messages = cast(list[AllMessageValues], messages) # cast-ok: Anthropic-shaped dicts from v1/messages
|
||||
configured = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list
|
||||
configured: Final = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list
|
||||
list[CacheControlInjectionPoint] | None, kwargs.pop("cache_control_injection_points", None)
|
||||
)
|
||||
if configured and AnthropicCacheControlHook._should_stand_down(configured, typed_messages, system, tools):
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import json
|
|||
import os
|
||||
import random
|
||||
import types
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel # type: ignore
|
||||
|
|
@ -29,7 +29,7 @@ from litellm.types.utils import StandardLoggingPayload
|
|||
|
||||
|
||||
def is_serializable(value):
|
||||
non_serializable_types = (
|
||||
non_serializable_types: Final = (
|
||||
types.CoroutineType,
|
||||
types.FunctionType,
|
||||
types.GeneratorType,
|
||||
|
|
@ -62,7 +62,7 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
)
|
||||
|
||||
self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
|
||||
_batch_size = os.getenv("ARGILLA_BATCH_SIZE", None) or litellm.argilla_batch_size
|
||||
_batch_size: Final = os.getenv("ARGILLA_BATCH_SIZE", None) or litellm.argilla_batch_size
|
||||
if _batch_size:
|
||||
self.batch_size = int(_batch_size)
|
||||
asyncio.create_task(self.periodic_flush())
|
||||
|
|
@ -85,11 +85,11 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
argilla_dataset_name: str | None,
|
||||
argilla_base_url: str | None,
|
||||
) -> ArgillaCredentialsObject:
|
||||
_credentials_api_key = argilla_api_key or os.getenv("ARGILLA_API_KEY")
|
||||
_credentials_api_key: Final = argilla_api_key or os.getenv("ARGILLA_API_KEY")
|
||||
if _credentials_api_key is None:
|
||||
raise Exception("Invalid Argilla API Key given. _credentials_api_key=None.")
|
||||
|
||||
_credentials_base_url = argilla_base_url or os.getenv("ARGILLA_BASE_URL") or "http://localhost:6900/"
|
||||
_credentials_base_url: Final = argilla_base_url or os.getenv("ARGILLA_BASE_URL") or "http://localhost:6900/"
|
||||
if _credentials_base_url is None:
|
||||
raise Exception("Invalid Argilla Base URL given. _credentials_base_url=None.")
|
||||
|
||||
|
|
@ -97,11 +97,11 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
if _credentials_dataset_name is None:
|
||||
raise Exception("Invalid Argilla Dataset give. Value=None.")
|
||||
else:
|
||||
dataset_response = litellm.module_level_client.get(
|
||||
dataset_response: Final = litellm.module_level_client.get(
|
||||
url=f"{_credentials_base_url}/api/v1/me/datasets?name={_credentials_dataset_name}",
|
||||
headers={"X-Argilla-Api-Key": _credentials_api_key},
|
||||
)
|
||||
json_response = dataset_response.json()
|
||||
json_response: Final = dataset_response.json()
|
||||
if (
|
||||
"items" in json_response
|
||||
and isinstance(json_response["items"], list)
|
||||
|
|
@ -116,7 +116,7 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
)
|
||||
|
||||
def get_chat_messages(self, payload: StandardLoggingPayload) -> list[dict[str, Any]]:
|
||||
payload_messages = payload.get("messages", None)
|
||||
payload_messages: Final = payload.get("messages", None)
|
||||
|
||||
if payload_messages is None:
|
||||
raise Exception("No chat messages found in payload.")
|
||||
|
|
@ -129,7 +129,7 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
raise Exception(f"Invalid chat messages format: {payload_messages}")
|
||||
|
||||
def get_str_response(self, payload: StandardLoggingPayload) -> str:
|
||||
response = payload["response"]
|
||||
response: Final = payload["response"]
|
||||
|
||||
if response is None:
|
||||
raise Exception("No response found in payload.")
|
||||
|
|
@ -144,14 +144,14 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
def _prepare_log_data(self, kwargs, response_obj, start_time, end_time) -> ArgillaItem | None:
|
||||
try:
|
||||
# Ensure everything in the payload is converted to str
|
||||
payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None)
|
||||
payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None)
|
||||
|
||||
if payload is None:
|
||||
raise Exception("Error logging request payload. Payload=none.")
|
||||
|
||||
argilla_message = self.get_chat_messages(payload)
|
||||
argilla_response = self.get_str_response(payload)
|
||||
argilla_item: ArgillaItem = {"fields": {}}
|
||||
argilla_message: Final = self.get_chat_messages(payload)
|
||||
argilla_response: Final = self.get_str_response(payload)
|
||||
argilla_item: Final[ArgillaItem] = {"fields": {}}
|
||||
for k, v in self.argilla_transformation_object.items():
|
||||
if v == "messages":
|
||||
argilla_item["fields"][k] = argilla_message
|
||||
|
|
@ -168,17 +168,17 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
if not self.log_queue:
|
||||
return
|
||||
|
||||
argilla_api_base = self.default_credentials["ARGILLA_BASE_URL"]
|
||||
argilla_dataset_name = self.default_credentials["ARGILLA_DATASET_NAME"]
|
||||
argilla_api_base: Final = self.default_credentials["ARGILLA_BASE_URL"]
|
||||
argilla_dataset_name: Final = self.default_credentials["ARGILLA_DATASET_NAME"]
|
||||
|
||||
url = f"{argilla_api_base}/api/v1/datasets/{argilla_dataset_name}/records/bulk"
|
||||
url: Final = f"{argilla_api_base}/api/v1/datasets/{argilla_dataset_name}/records/bulk"
|
||||
|
||||
argilla_api_key = self.default_credentials["ARGILLA_API_KEY"]
|
||||
argilla_api_key: Final = self.default_credentials["ARGILLA_API_KEY"]
|
||||
|
||||
headers = {"X-Argilla-Api-Key": argilla_api_key}
|
||||
headers: Final = {"X-Argilla-Api-Key": argilla_api_key}
|
||||
|
||||
try:
|
||||
response = litellm.module_level_client.post(
|
||||
response: Final = litellm.module_level_client.post(
|
||||
url=url,
|
||||
json=self.log_queue,
|
||||
headers=headers,
|
||||
|
|
@ -195,13 +195,13 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
sampling_rate = (
|
||||
sampling_rate: Final = (
|
||||
float(os.getenv("LANGSMITH_SAMPLING_RATE")) # type: ignore
|
||||
if os.getenv("LANGSMITH_SAMPLING_RATE") is not None
|
||||
and os.getenv("LANGSMITH_SAMPLING_RATE").strip().isdigit() # type: ignore
|
||||
else 1.0
|
||||
)
|
||||
random_sample = random.random()
|
||||
random_sample: Final = random.random()
|
||||
if random_sample > sampling_rate:
|
||||
verbose_logger.info(
|
||||
"Skipping Langsmith logging. Sampling rate=%s, random_sample=%s", sampling_rate, random_sample
|
||||
|
|
@ -212,7 +212,7 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
kwargs,
|
||||
response_obj,
|
||||
)
|
||||
data = self._prepare_log_data(kwargs, response_obj, start_time, end_time)
|
||||
data: Final = self._prepare_log_data(kwargs, response_obj, start_time, end_time)
|
||||
if data is None:
|
||||
return
|
||||
|
||||
|
|
@ -227,8 +227,8 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
sampling_rate = self.sampling_rate
|
||||
random_sample = random.random()
|
||||
sampling_rate: Final = self.sampling_rate
|
||||
random_sample: Final = random.random()
|
||||
if random_sample > sampling_rate:
|
||||
verbose_logger.info(
|
||||
"Skipping Langsmith logging. Sampling rate=%s, random_sample=%s", sampling_rate, random_sample
|
||||
|
|
@ -239,7 +239,7 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
kwargs,
|
||||
response_obj,
|
||||
)
|
||||
payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object", None)
|
||||
payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None)
|
||||
|
||||
data = self._prepare_log_data(kwargs, response_obj, start_time, end_time)
|
||||
|
||||
|
|
@ -268,8 +268,8 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
verbose_logger.exception("Argilla Layer Error - error logging async success event.")
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
sampling_rate = self.sampling_rate
|
||||
random_sample = random.random()
|
||||
sampling_rate: Final = self.sampling_rate
|
||||
random_sample: Final = random.random()
|
||||
if random_sample > sampling_rate:
|
||||
verbose_logger.info(
|
||||
"Skipping Langsmith logging. Sampling rate=%s, random_sample=%s", sampling_rate, random_sample
|
||||
|
|
@ -277,7 +277,7 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
return # Skip logging
|
||||
verbose_logger.info("Langsmith Failure Event Logging!")
|
||||
try:
|
||||
data = self._prepare_log_data(kwargs, response_obj, start_time, end_time)
|
||||
data: Final = self._prepare_log_data(kwargs, response_obj, start_time, end_time)
|
||||
self.log_queue.append(data)
|
||||
verbose_logger.debug(
|
||||
"Langsmith logging: queue length %s, batch size %s",
|
||||
|
|
@ -302,17 +302,17 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
if not self.log_queue:
|
||||
return
|
||||
|
||||
argilla_api_base = self.default_credentials["ARGILLA_BASE_URL"]
|
||||
argilla_dataset_name = self.default_credentials["ARGILLA_DATASET_NAME"]
|
||||
argilla_api_base: Final = self.default_credentials["ARGILLA_BASE_URL"]
|
||||
argilla_dataset_name: Final = self.default_credentials["ARGILLA_DATASET_NAME"]
|
||||
|
||||
url = f"{argilla_api_base}/api/v1/datasets/{argilla_dataset_name}/records/bulk"
|
||||
url: Final = f"{argilla_api_base}/api/v1/datasets/{argilla_dataset_name}/records/bulk"
|
||||
|
||||
argilla_api_key = self.default_credentials["ARGILLA_API_KEY"]
|
||||
argilla_api_key: Final = self.default_credentials["ARGILLA_API_KEY"]
|
||||
|
||||
headers = {"X-Argilla-Api-Key": argilla_api_key}
|
||||
headers: Final = {"X-Argilla-Api-Key": argilla_api_key}
|
||||
|
||||
try:
|
||||
response = await self.async_httpx_client.put(
|
||||
response: Final = await self.async_httpx_client.put(
|
||||
url=url,
|
||||
data=json.dumps(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import os
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.custom_prompt_management import CustomPromptManagement
|
||||
|
|
@ -10,22 +10,22 @@ from litellm.types.prompts.init_prompts import SupportedPromptIntegrations
|
|||
from .arize_phoenix_prompt_manager import ArizePhoenixPromptManager
|
||||
|
||||
# Global instances
|
||||
global_arize_config: dict | None = None
|
||||
global_arize_config: Final[dict | None] = None
|
||||
|
||||
|
||||
def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec") -> "CustomPromptManagement":
|
||||
"""
|
||||
Initialize a prompt from Arize Phoenix.
|
||||
"""
|
||||
api_key = getattr(litellm_params, "api_key", None) or os.environ.get("PHOENIX_API_KEY")
|
||||
api_base = getattr(litellm_params, "api_base", None)
|
||||
prompt_id = getattr(litellm_params, "prompt_id", None)
|
||||
api_key: Final = getattr(litellm_params, "api_key", None) or os.environ.get("PHOENIX_API_KEY")
|
||||
api_base: Final = getattr(litellm_params, "api_base", None)
|
||||
prompt_id: Final = getattr(litellm_params, "prompt_id", None)
|
||||
|
||||
if not api_key or not api_base:
|
||||
raise ValueError("api_key and api_base are required for Arize Phoenix prompt integration")
|
||||
|
||||
try:
|
||||
arize_prompt_manager = ArizePhoenixPromptManager(
|
||||
arize_prompt_manager: Final = ArizePhoenixPromptManager(
|
||||
**{
|
||||
"api_key": api_key,
|
||||
"api_base": api_base,
|
||||
|
|
@ -39,6 +39,6 @@ def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "Prom
|
|||
raise e
|
||||
|
||||
|
||||
prompt_initializer_registry = {
|
||||
prompt_initializer_registry: Final = {
|
||||
SupportedPromptIntegrations.ARIZE_PHOENIX.value: prompt_initializer,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import json
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from typing_extensions import override
|
||||
|
||||
|
|
@ -32,12 +32,12 @@ class ArizeOTELAttributes(BaseLLMObsOTELAttributes):
|
|||
@staticmethod
|
||||
@override
|
||||
def set_messages(span: "Span", kwargs: dict[str, Any]):
|
||||
messages = kwargs.get("messages")
|
||||
messages: Final = kwargs.get("messages")
|
||||
|
||||
# for /chat/completions
|
||||
# https://docs.arize.com/arize/large-language-models/tracing/semantic-conventions
|
||||
if messages:
|
||||
last_message = messages[-1]
|
||||
last_message: Final = messages[-1]
|
||||
safe_set_attribute(
|
||||
span,
|
||||
SpanAttributes.INPUT_VALUE,
|
||||
|
|
@ -129,7 +129,7 @@ def _set_choice_outputs(span: "Span", response_obj, msg_attrs, span_attrs):
|
|||
|
||||
|
||||
def _set_image_outputs(span: "Span", response_obj, image_attrs, span_attrs):
|
||||
images = response_obj.get("data", [])
|
||||
images: Final = response_obj.get("data", [])
|
||||
for i, image in enumerate(images):
|
||||
img_url = image.get("url")
|
||||
if img_url is None and image.get("b64_json"):
|
||||
|
|
@ -145,7 +145,7 @@ def _set_image_outputs(span: "Span", response_obj, image_attrs, span_attrs):
|
|||
|
||||
|
||||
def _set_audio_outputs(span: "Span", response_obj, audio_attrs, span_attrs):
|
||||
audio = response_obj.get("audio", [])
|
||||
audio: Final = response_obj.get("audio", [])
|
||||
for i, audio_item in enumerate(audio):
|
||||
audio_url = audio_item.get("url")
|
||||
if audio_url is None and audio_item.get("b64_json"):
|
||||
|
|
@ -166,7 +166,7 @@ def _set_audio_outputs(span: "Span", response_obj, audio_attrs, span_attrs):
|
|||
|
||||
|
||||
def _set_embedding_outputs(span: "Span", response_obj, embedding_attrs, span_attrs):
|
||||
embeddings = response_obj.get("data", [])
|
||||
embeddings: Final = response_obj.get("data", [])
|
||||
for i, embedding_item in enumerate(embeddings):
|
||||
embedding_vector = embedding_item.get("embedding")
|
||||
if embedding_vector:
|
||||
|
|
@ -193,7 +193,7 @@ def _set_embedding_outputs(span: "Span", response_obj, embedding_attrs, span_att
|
|||
|
||||
|
||||
def _set_structured_outputs(span: "Span", response_obj, msg_attrs, span_attrs):
|
||||
output_items = response_obj.get("output", [])
|
||||
output_items: Final = response_obj.get("output", [])
|
||||
for i, item in enumerate(output_items):
|
||||
prefix = f"{span_attrs.LLM_OUTPUT_MESSAGES}.{i}"
|
||||
if not hasattr(item, "type"):
|
||||
|
|
@ -232,7 +232,7 @@ def _safe_get(obj, key, default=None):
|
|||
"""
|
||||
if obj is None:
|
||||
return default
|
||||
getter = getattr(obj, "get", None)
|
||||
getter: Final = getattr(obj, "get", None)
|
||||
if callable(getter):
|
||||
try:
|
||||
return getter(key, default)
|
||||
|
|
@ -243,15 +243,15 @@ def _safe_get(obj, key, default=None):
|
|||
|
||||
|
||||
def _set_usage_outputs(span: "Span", response_obj, span_attrs):
|
||||
usage = response_obj and response_obj.get("usage")
|
||||
usage: Final = response_obj and response_obj.get("usage")
|
||||
if not usage:
|
||||
return
|
||||
|
||||
safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_TOTAL, _safe_get(usage, "total_tokens"))
|
||||
completion_tokens = _safe_get(usage, "completion_tokens") or _safe_get(usage, "output_tokens")
|
||||
completion_tokens: Final = _safe_get(usage, "completion_tokens") or _safe_get(usage, "output_tokens")
|
||||
if completion_tokens:
|
||||
safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_COMPLETION, completion_tokens)
|
||||
prompt_tokens = _safe_get(usage, "prompt_tokens") or _safe_get(usage, "input_tokens")
|
||||
prompt_tokens: Final = _safe_get(usage, "prompt_tokens") or _safe_get(usage, "input_tokens")
|
||||
if prompt_tokens:
|
||||
safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_PROMPT, prompt_tokens)
|
||||
|
||||
|
|
@ -259,8 +259,8 @@ def _set_usage_outputs(span: "Span", response_obj, span_attrs):
|
|||
# API (Usage) and in `output_tokens_details` for Responses API
|
||||
# (ResponseAPIUsage). Both nested objects may be plain Pydantic models
|
||||
# without `.get`.
|
||||
token_details = _safe_get(usage, "completion_tokens_details") or _safe_get(usage, "output_tokens_details")
|
||||
reasoning_tokens = _safe_get(token_details, "reasoning_tokens")
|
||||
token_details: Final = _safe_get(usage, "completion_tokens_details") or _safe_get(usage, "output_tokens_details")
|
||||
reasoning_tokens: Final = _safe_get(token_details, "reasoning_tokens")
|
||||
if reasoning_tokens:
|
||||
safe_set_attribute(
|
||||
span,
|
||||
|
|
@ -275,8 +275,8 @@ def _set_usage_outputs(span: "Span", response_obj, span_attrs):
|
|||
# `cache_creation_input_tokens`
|
||||
# All emits are conditional, so when none of these fields exist (the
|
||||
# situation in the existing test fixtures) no extra attributes are set.
|
||||
prompt_token_details = _safe_get(usage, "prompt_tokens_details") or _safe_get(usage, "input_tokens_details")
|
||||
cache_read = _safe_get(prompt_token_details, "cached_tokens") or _safe_get(usage, "cache_read_input_tokens")
|
||||
prompt_token_details: Final = _safe_get(usage, "prompt_tokens_details") or _safe_get(usage, "input_tokens_details")
|
||||
cache_read: Final = _safe_get(prompt_token_details, "cached_tokens") or _safe_get(usage, "cache_read_input_tokens")
|
||||
if cache_read:
|
||||
safe_set_attribute(
|
||||
span,
|
||||
|
|
@ -285,7 +285,7 @@ def _set_usage_outputs(span: "Span", response_obj, span_attrs):
|
|||
)
|
||||
# Anthropic / Bedrock-Anthropic only — OpenAI's `prompt_tokens_details`
|
||||
# does not expose a cache-write count, so we read straight off `usage`.
|
||||
cache_write = _safe_get(usage, "cache_creation_input_tokens")
|
||||
cache_write: Final = _safe_get(usage, "cache_creation_input_tokens")
|
||||
if cache_write:
|
||||
safe_set_attribute(
|
||||
span,
|
||||
|
|
@ -293,7 +293,7 @@ def _set_usage_outputs(span: "Span", response_obj, span_attrs):
|
|||
cache_write,
|
||||
)
|
||||
|
||||
audio_prompt_tokens = _safe_get(prompt_token_details, "audio_tokens")
|
||||
audio_prompt_tokens: Final = _safe_get(prompt_token_details, "audio_tokens")
|
||||
if audio_prompt_tokens:
|
||||
safe_set_attribute(
|
||||
span,
|
||||
|
|
@ -310,7 +310,7 @@ def _infer_open_inference_span_kind(call_type: str | None) -> str:
|
|||
if not call_type:
|
||||
return OpenInferenceSpanKindValues.UNKNOWN.value
|
||||
|
||||
lowered = str(call_type).lower()
|
||||
lowered: Final = str(call_type).lower()
|
||||
|
||||
if "embed" in lowered:
|
||||
return OpenInferenceSpanKindValues.EMBEDDING.value
|
||||
|
|
@ -416,7 +416,7 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMO
|
|||
# routes) into a dict so downstream `.get()` calls don't crash. Existing
|
||||
# dict / `.get()`-bearing objects (incl. Pydantic OpenAI Responses API
|
||||
# models) are returned unchanged, preserving the existing test behavior.
|
||||
response_obj_for_attrs = _coerce_response_obj_for_attrs(response_obj)
|
||||
response_obj_for_attrs: Final = _coerce_response_obj_for_attrs(response_obj)
|
||||
|
||||
# Set span.kind defensively before anything else. If a downstream step
|
||||
# throws, the span still has a kind so Arize can render it correctly
|
||||
|
|
@ -425,17 +425,17 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMO
|
|||
_safe_emit("early span kind", _set_early_span_kind, span, kwargs)
|
||||
|
||||
try:
|
||||
optional_params = _sanitize_optional_params(kwargs.get("optional_params"))
|
||||
litellm_params = kwargs.get("litellm_params", {}) or {}
|
||||
standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object")
|
||||
optional_params: Final = _sanitize_optional_params(kwargs.get("optional_params"))
|
||||
litellm_params: Final = kwargs.get("litellm_params", {}) or {}
|
||||
standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object")
|
||||
if standard_logging_payload is None:
|
||||
raise ValueError("standard_logging_object not found in kwargs")
|
||||
|
||||
metadata = standard_logging_payload.get("metadata") if standard_logging_payload else None
|
||||
metadata: Final = standard_logging_payload.get("metadata") if standard_logging_payload else None
|
||||
_set_metadata_attributes(span, metadata, SpanAttributes)
|
||||
|
||||
metadata_tools = _extract_metadata_tools(metadata)
|
||||
optional_tools = _extract_optional_tools(optional_params)
|
||||
metadata_tools: Final = _extract_metadata_tools(metadata)
|
||||
optional_tools: Final = _extract_optional_tools(optional_params)
|
||||
|
||||
_set_request_attributes(
|
||||
span=span,
|
||||
|
|
@ -455,7 +455,7 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMO
|
|||
_set_tool_attributes(span, optional_tools, metadata_tools)
|
||||
attributes.set_messages(span, kwargs)
|
||||
|
||||
model_params = standard_logging_payload.get("model_parameters") if standard_logging_payload else None
|
||||
model_params: Final = standard_logging_payload.get("model_parameters") if standard_logging_payload else None
|
||||
_set_model_params(span, model_params, SpanAttributes)
|
||||
|
||||
_set_response_attributes(span=span, response_obj=response_obj_for_attrs)
|
||||
|
|
@ -468,7 +468,7 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMO
|
|||
# Additive emitters. Each is independently guarded so a failure can never
|
||||
# blank the attributes set by the main try-block above. New attributes are
|
||||
# written under new keys; existing attributes are not overwritten.
|
||||
slp = kwargs.get("standard_logging_object")
|
||||
slp: Final = kwargs.get("standard_logging_object")
|
||||
_safe_emit("session/user attrs", _set_session_and_user_attrs, span, kwargs, slp)
|
||||
_safe_emit("response cost", _set_response_cost_attr, span, slp)
|
||||
_safe_emit(
|
||||
|
|
@ -497,7 +497,7 @@ def _set_metadata_attributes(span: "Span", metadata: Any | None, span_attrs) ->
|
|||
def _extract_metadata_tools(metadata: Any | None) -> list | None:
|
||||
if not isinstance(metadata, dict):
|
||||
return None
|
||||
llm_obj = metadata.get("llm")
|
||||
llm_obj: Final = metadata.get("llm")
|
||||
if isinstance(llm_obj, dict):
|
||||
return llm_obj.get("tools")
|
||||
return None
|
||||
|
|
@ -550,7 +550,7 @@ def _set_model_params(span: "Span", model_params: dict | None, span_attrs) -> No
|
|||
|
||||
safe_set_attribute(span, span_attrs.LLM_INVOCATION_PARAMETERS, safe_dumps(model_params))
|
||||
if model_params.get("user"):
|
||||
user_id = model_params.get("user")
|
||||
user_id: Final = model_params.get("user")
|
||||
if user_id is not None:
|
||||
safe_set_attribute(span, span_attrs.USER_ID, user_id)
|
||||
|
||||
|
|
@ -573,8 +573,8 @@ def _safe_emit(label: str, fn, *args, **kwargs) -> None:
|
|||
|
||||
def _set_early_span_kind(span: "Span", kwargs: dict) -> None:
|
||||
"""Defensively set OPENINFERENCE_SPAN_KIND before any other logic runs."""
|
||||
slp = kwargs.get("standard_logging_object")
|
||||
call_type = slp.get("call_type") if isinstance(slp, dict) else None
|
||||
slp: Final = kwargs.get("standard_logging_object")
|
||||
call_type: Final = slp.get("call_type") if isinstance(slp, dict) else None
|
||||
safe_set_attribute(
|
||||
span,
|
||||
SpanAttributes.OPENINFERENCE_SPAN_KIND,
|
||||
|
|
@ -595,10 +595,10 @@ def _coerce_response_obj_for_attrs(response_obj):
|
|||
"""
|
||||
if response_obj is None or hasattr(response_obj, "get"):
|
||||
return response_obj
|
||||
text = getattr(response_obj, "text", None)
|
||||
text: Final = getattr(response_obj, "text", None)
|
||||
if isinstance(text, str) and text:
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
parsed: Final = json.loads(text)
|
||||
if isinstance(parsed, dict):
|
||||
return parsed
|
||||
except Exception:
|
||||
|
|
@ -620,7 +620,7 @@ def _coerce_text(value) -> str | None:
|
|||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, list):
|
||||
parts = []
|
||||
parts: Final = []
|
||||
for part in value:
|
||||
if isinstance(part, str):
|
||||
parts.append(part)
|
||||
|
|
@ -641,7 +641,7 @@ def _to_plain_dict(value):
|
|||
"""
|
||||
if value is None or isinstance(value, dict):
|
||||
return value
|
||||
model_dump = getattr(value, "model_dump", None)
|
||||
model_dump: Final = getattr(value, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
try:
|
||||
return model_dump()
|
||||
|
|
@ -655,7 +655,7 @@ def _get_tool_calls(message) -> list | None:
|
|||
|
||||
Works for dicts and Pydantic message objects via ``_safe_get``.
|
||||
"""
|
||||
tool_calls = _safe_get(message, "tool_calls")
|
||||
tool_calls: Final = _safe_get(message, "tool_calls")
|
||||
return tool_calls if isinstance(tool_calls, list) and tool_calls else None
|
||||
|
||||
|
||||
|
|
@ -667,11 +667,11 @@ def _normalize_tool_call(raw_tc) -> dict[str, Any] | None:
|
|||
Arguments are coerced to a JSON string per OpenInference convention.
|
||||
Returns ``None`` when ``raw_tc`` cannot be coerced to a dict.
|
||||
"""
|
||||
tc = _to_plain_dict(raw_tc)
|
||||
tc: Final = _to_plain_dict(raw_tc)
|
||||
if not isinstance(tc, dict):
|
||||
return None
|
||||
function = _to_plain_dict(tc.get("function"))
|
||||
name = function.get("name") if isinstance(function, dict) else None
|
||||
function: Final = _to_plain_dict(tc.get("function"))
|
||||
name: Final = function.get("name") if isinstance(function, dict) else None
|
||||
args = function.get("arguments") if isinstance(function, dict) else None
|
||||
if args is not None and not isinstance(args, str):
|
||||
try:
|
||||
|
|
@ -692,7 +692,7 @@ def _summarize_tool_calls_for_output(tool_calls) -> str:
|
|||
so OUTPUT_VALUE is never blanked on a malformed payload.
|
||||
"""
|
||||
try:
|
||||
normalized = [n for n in (_normalize_tool_call(tc) for tc in tool_calls) if n]
|
||||
normalized: Final = [n for n in (_normalize_tool_call(tc) for tc in tool_calls) if n]
|
||||
return json.dumps({"tool_calls": normalized})
|
||||
except Exception:
|
||||
return str(tool_calls)
|
||||
|
|
@ -705,7 +705,7 @@ def _emit_message_tool_calls(span: "Span", prefix: str, message) -> None:
|
|||
Accepts dicts or Pydantic message objects (e.g. ``litellm.Message``); the
|
||||
same applies to each tool_call entry.
|
||||
"""
|
||||
tool_calls = _get_tool_calls(message)
|
||||
tool_calls: Final = _get_tool_calls(message)
|
||||
if not tool_calls:
|
||||
return
|
||||
for tc_idx, raw_tc in enumerate(tool_calls):
|
||||
|
|
@ -744,11 +744,11 @@ def _emit_input_message_extras(span: "Span", prefix: str, message: dict) -> None
|
|||
if not isinstance(message, dict):
|
||||
return
|
||||
|
||||
name = message.get("name")
|
||||
name: Final = message.get("name")
|
||||
if name:
|
||||
safe_set_attribute(span, f"{prefix}.{MessageAttributes.MESSAGE_NAME}", name)
|
||||
|
||||
tool_call_id = message.get("tool_call_id")
|
||||
tool_call_id: Final = message.get("tool_call_id")
|
||||
if tool_call_id:
|
||||
safe_set_attribute(
|
||||
span,
|
||||
|
|
@ -758,9 +758,9 @@ def _emit_input_message_extras(span: "Span", prefix: str, message: dict) -> None
|
|||
|
||||
_emit_message_tool_calls(span, prefix, message)
|
||||
|
||||
content = message.get("content")
|
||||
content: Final = message.get("content")
|
||||
if isinstance(content, list):
|
||||
contents_prefix = f"{prefix}.{MessageAttributes.MESSAGE_CONTENTS}"
|
||||
contents_prefix: Final = f"{prefix}.{MessageAttributes.MESSAGE_CONTENTS}"
|
||||
for part_idx, part in enumerate(content):
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
|
|
@ -823,36 +823,36 @@ def _set_session_and_user_attrs(span: "Span", kwargs: dict, standard_logging_pay
|
|||
"""
|
||||
if not isinstance(standard_logging_payload, dict):
|
||||
return
|
||||
metadata = standard_logging_payload.get("metadata") or {}
|
||||
metadata: Final = standard_logging_payload.get("metadata") or {}
|
||||
if not isinstance(metadata, dict):
|
||||
return
|
||||
|
||||
session_id = metadata.get("user_api_key_end_user_id")
|
||||
session_id: Final = metadata.get("user_api_key_end_user_id")
|
||||
if session_id:
|
||||
safe_set_attribute(span, SpanAttributes.SESSION_ID, str(session_id))
|
||||
|
||||
trace_id = standard_logging_payload.get("trace_id")
|
||||
trace_id: Final = standard_logging_payload.get("trace_id")
|
||||
if trace_id:
|
||||
safe_set_attribute(span, "litellm.trace_id", str(trace_id))
|
||||
|
||||
optional_params = kwargs.get("optional_params") or {}
|
||||
model_params = standard_logging_payload.get("model_parameters") or {}
|
||||
has_user_already = bool(
|
||||
optional_params: Final = kwargs.get("optional_params") or {}
|
||||
model_params: Final = standard_logging_payload.get("model_parameters") or {}
|
||||
has_user_already: Final = bool(
|
||||
(isinstance(optional_params, dict) and optional_params.get("user"))
|
||||
or (isinstance(model_params, dict) and model_params.get("user"))
|
||||
)
|
||||
if not has_user_already:
|
||||
user_id = metadata.get("user_api_key_user_id")
|
||||
user_id: Final = metadata.get("user_api_key_user_id")
|
||||
if user_id:
|
||||
safe_set_attribute(span, SpanAttributes.USER_ID, str(user_id))
|
||||
|
||||
team_id = metadata.get("user_api_key_team_id")
|
||||
team_id: Final = metadata.get("user_api_key_team_id")
|
||||
if team_id:
|
||||
safe_set_attribute(span, "litellm.team_id", str(team_id))
|
||||
team_alias = metadata.get("user_api_key_team_alias")
|
||||
team_alias: Final = metadata.get("user_api_key_team_alias")
|
||||
if team_alias:
|
||||
safe_set_attribute(span, "litellm.team_alias", str(team_alias))
|
||||
key_alias = metadata.get("user_api_key_alias")
|
||||
key_alias: Final = metadata.get("user_api_key_alias")
|
||||
if key_alias:
|
||||
safe_set_attribute(span, "litellm.key_alias", str(key_alias))
|
||||
|
||||
|
|
@ -868,11 +868,11 @@ def _set_response_cost_attr(span: "Span", standard_logging_payload) -> None:
|
|||
"""
|
||||
if not isinstance(standard_logging_payload, dict):
|
||||
return
|
||||
cost = standard_logging_payload.get("response_cost")
|
||||
cost: Final = standard_logging_payload.get("response_cost")
|
||||
if cost is None:
|
||||
return
|
||||
try:
|
||||
cost_value = float(cost)
|
||||
cost_value: Final = float(cost)
|
||||
except (TypeError, ValueError):
|
||||
return
|
||||
safe_set_attribute(span, "llm.cost.total", cost_value)
|
||||
|
|
@ -882,7 +882,7 @@ def _set_response_cost_attr(span: "Span", standard_logging_payload) -> None:
|
|||
def _is_passthrough_call_type(call_type: str | None) -> bool:
|
||||
if not call_type:
|
||||
return False
|
||||
lowered = str(call_type).lower()
|
||||
lowered: Final = str(call_type).lower()
|
||||
return "passthrough" in lowered or "pass_through" in lowered
|
||||
|
||||
|
||||
|
|
@ -913,7 +913,7 @@ def _maybe_normalize_passthrough(
|
|||
passthrough I/O (with central redaction) for free and this helper's
|
||||
`complete_input_dict` fallback can be deleted. See follow-up issue.
|
||||
"""
|
||||
call_type = standard_logging_payload.get("call_type") if isinstance(standard_logging_payload, dict) else None
|
||||
call_type: Final = standard_logging_payload.get("call_type") if isinstance(standard_logging_payload, dict) else None
|
||||
if not _is_passthrough_call_type(call_type):
|
||||
return
|
||||
|
||||
|
|
@ -927,13 +927,13 @@ def _maybe_normalize_passthrough(
|
|||
return
|
||||
|
||||
# --- INPUT --------------------------------------------------------------
|
||||
additional_args = kwargs.get("additional_args") or {}
|
||||
additional_args: Final = kwargs.get("additional_args") or {}
|
||||
complete_input_dict = additional_args.get("complete_input_dict") if isinstance(additional_args, dict) else None
|
||||
if isinstance(complete_input_dict, dict):
|
||||
_set_passthrough_input_attributes(span, complete_input_dict.get("messages"))
|
||||
|
||||
# --- OUTPUT -------------------------------------------------------------
|
||||
parsed_response = _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs)
|
||||
parsed_response: Final = _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs)
|
||||
if not isinstance(parsed_response, dict):
|
||||
return
|
||||
|
||||
|
|
@ -977,13 +977,13 @@ def _set_passthrough_input_attributes(span: "Span", messages) -> None:
|
|||
def _set_passthrough_output_attributes(span: "Span", parsed_response: dict) -> None:
|
||||
"""Render passthrough response into OUTPUT_VALUE + LLM_OUTPUT_MESSAGES."""
|
||||
# Anthropic / Bedrock-Anthropic: `content` is a list of typed parts.
|
||||
content_list = parsed_response.get("content")
|
||||
content_list: Final = parsed_response.get("content")
|
||||
if isinstance(content_list, list) and content_list:
|
||||
texts = []
|
||||
texts: Final = []
|
||||
for part in content_list:
|
||||
if isinstance(part, dict) and isinstance(part.get("text"), str):
|
||||
texts.append(part["text"])
|
||||
joined = "\n\n".join(t for t in texts if t)
|
||||
joined: Final = "\n\n".join(t for t in texts if t)
|
||||
if joined:
|
||||
safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, joined)
|
||||
prefix = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0"
|
||||
|
|
@ -999,13 +999,13 @@ def _set_passthrough_output_attributes(span: "Span", parsed_response: dict) -> N
|
|||
)
|
||||
|
||||
# OpenAI-style passthrough: `choices[0].message.content`
|
||||
choices = parsed_response.get("choices")
|
||||
choices: Final = parsed_response.get("choices")
|
||||
if isinstance(choices, list) and choices:
|
||||
first = choices[0]
|
||||
first: Final = choices[0]
|
||||
if isinstance(first, dict):
|
||||
msg = first.get("message")
|
||||
msg: Final = first.get("message")
|
||||
if isinstance(msg, dict):
|
||||
text = _coerce_text(msg.get("content"))
|
||||
text: Final = _coerce_text(msg.get("content"))
|
||||
if text:
|
||||
safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, text)
|
||||
prefix = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0"
|
||||
|
|
@ -1024,7 +1024,7 @@ def _set_passthrough_output_attributes(span: "Span", parsed_response: dict) -> N
|
|||
def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs):
|
||||
"""Return a dict view of the provider response for passthrough routes."""
|
||||
# Prefer the coerced view (already JSON-parsed for httpx.Response).
|
||||
candidates = []
|
||||
candidates: Final = []
|
||||
if isinstance(coerced_response_obj, dict):
|
||||
candidates.append(coerced_response_obj)
|
||||
if isinstance(raw_response_obj, dict) and raw_response_obj is not coerced_response_obj:
|
||||
|
|
@ -1047,7 +1047,7 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs):
|
|||
return candidate
|
||||
|
||||
# Fallback: kwargs["original_response"] from the OTel base path.
|
||||
original = kwargs.get("original_response") if isinstance(kwargs, dict) else None
|
||||
original: Final = kwargs.get("original_response") if isinstance(kwargs, dict) else None
|
||||
if isinstance(original, dict):
|
||||
return original
|
||||
if isinstance(original, str):
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ this file has Arize ai specific helper functions
|
|||
|
||||
import os
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Union
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
|
||||
from litellm.integrations.arize import _utils
|
||||
from litellm.integrations.arize._utils import ArizeOTELAttributes
|
||||
|
|
@ -50,7 +50,7 @@ class ArizeLogger(OpenTelemetry):
|
|||
self.span_kind = SpanKind
|
||||
return
|
||||
|
||||
provider = TracerProvider(resource=self._get_litellm_resource(self.config))
|
||||
provider: Final = TracerProvider(resource=self._get_litellm_resource(self.config))
|
||||
provider.add_span_processor(self._get_span_processor())
|
||||
self.tracer = provider.get_tracer("litellm")
|
||||
self.span_kind = SpanKind
|
||||
|
|
@ -80,13 +80,13 @@ class ArizeLogger(OpenTelemetry):
|
|||
Raises:
|
||||
ValueError: If required environment variables are not set.
|
||||
"""
|
||||
space_id = os.environ.get("ARIZE_SPACE_ID")
|
||||
space_key = os.environ.get("ARIZE_SPACE_KEY")
|
||||
api_key = os.environ.get("ARIZE_API_KEY")
|
||||
project_name = os.environ.get("ARIZE_PROJECT_NAME")
|
||||
space_id: Final = os.environ.get("ARIZE_SPACE_ID")
|
||||
space_key: Final = os.environ.get("ARIZE_SPACE_KEY")
|
||||
api_key: Final = os.environ.get("ARIZE_API_KEY")
|
||||
project_name: Final = os.environ.get("ARIZE_PROJECT_NAME")
|
||||
|
||||
grpc_endpoint = os.environ.get("ARIZE_ENDPOINT")
|
||||
http_endpoint = os.environ.get("ARIZE_HTTP_ENDPOINT")
|
||||
grpc_endpoint: Final = os.environ.get("ARIZE_ENDPOINT")
|
||||
http_endpoint: Final = os.environ.get("ARIZE_HTTP_ENDPOINT")
|
||||
|
||||
endpoint = None
|
||||
protocol: Protocol = "otlp_grpc"
|
||||
|
|
@ -147,7 +147,7 @@ class ArizeLogger(OpenTelemetry):
|
|||
dict: Health check result with status and message
|
||||
"""
|
||||
try:
|
||||
config = self.get_arize_config()
|
||||
config: Final = self.get_arize_config()
|
||||
|
||||
if not config.space_id and not config.space_key:
|
||||
return {
|
||||
|
|
@ -183,7 +183,7 @@ class ArizeLogger(OpenTelemetry):
|
|||
Returns:
|
||||
dict: A dictionary of dynamic Arize headers
|
||||
"""
|
||||
dynamic_headers = {}
|
||||
dynamic_headers: Final = {}
|
||||
|
||||
#########################################################
|
||||
# `arize-space-id` handling
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import os
|
||||
import threading
|
||||
from collections import OrderedDict
|
||||
from typing import TYPE_CHECKING, Any, Union
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.arize import _utils
|
||||
|
|
@ -43,8 +43,8 @@ else:
|
|||
OpenTelemetry = None # type: ignore
|
||||
|
||||
|
||||
ARIZE_HOSTED_PHOENIX_ENDPOINT = "https://otlp.arize.com/v1/traces"
|
||||
_MAX_PROJECT_PROVIDERS = 64
|
||||
ARIZE_HOSTED_PHOENIX_ENDPOINT: Final = "https://otlp.arize.com/v1/traces"
|
||||
_MAX_PROJECT_PROVIDERS: Final = 64
|
||||
|
||||
|
||||
class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
||||
|
|
@ -80,7 +80,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
self._shared_span_processor = self._get_span_processor()
|
||||
self.span_kind = SpanKind
|
||||
|
||||
default_project = self._resolve_project_name({})
|
||||
default_project: Final = self._resolve_project_name({})
|
||||
self.tracer = self._get_tracer_for(default_project)
|
||||
verbose_logger.debug(
|
||||
"ArizePhoenixLogger: Initialized per-project TracerProvider cache "
|
||||
|
|
@ -100,7 +100,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
if getattr(self, "_use_injected_tracer_provider", False):
|
||||
return
|
||||
|
||||
shared_processor = getattr(self, "_shared_span_processor", None)
|
||||
shared_processor: Final = getattr(self, "_shared_span_processor", None)
|
||||
if shared_processor is not None:
|
||||
try:
|
||||
shared_processor.force_flush()
|
||||
|
|
@ -111,7 +111,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
)
|
||||
|
||||
with getattr(self, "_project_providers_lock", threading.Lock()):
|
||||
providers = list(getattr(self, "_project_providers", {}).values())
|
||||
providers: Final = list(getattr(self, "_project_providers", {}).values())
|
||||
|
||||
for provider in providers:
|
||||
try:
|
||||
|
|
@ -129,24 +129,24 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
"""
|
||||
from opentelemetry.sdk.resources import OTELResourceDetector, Resource
|
||||
|
||||
project_attributes: dict[str, str] = {
|
||||
project_attributes: Final[dict[str, str]] = {
|
||||
"openinference.project.name": project_name,
|
||||
"model_id": project_name,
|
||||
"service.name": project_name,
|
||||
}
|
||||
deployment_environment = getattr(self.config, "deployment_environment", None)
|
||||
deployment_environment: Final = getattr(self.config, "deployment_environment", None)
|
||||
if deployment_environment is not None:
|
||||
project_attributes["deployment.environment"] = deployment_environment
|
||||
|
||||
env_resource = OTELResourceDetector().detect()
|
||||
project_resource = Resource.create(project_attributes) # type: ignore[arg-type]
|
||||
env_resource: Final = OTELResourceDetector().detect()
|
||||
project_resource: Final = Resource.create(project_attributes) # type: ignore[arg-type]
|
||||
return env_resource.merge(project_resource)
|
||||
|
||||
def _build_tracer_provider_for_project(self, project_name: str) -> TracerProvider:
|
||||
"""Create a TracerProvider for *project_name* (caller holds no cache lock)."""
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
|
||||
provider = TracerProvider(resource=self._get_litellm_resource_for_project(project_name))
|
||||
provider: Final = TracerProvider(resource=self._get_litellm_resource_for_project(project_name))
|
||||
provider.add_span_processor(self._shared_span_processor)
|
||||
return provider
|
||||
|
||||
|
|
@ -162,7 +162,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
|
||||
# OTELResourceDetector().detect() is synchronous; build outside the lock so
|
||||
# concurrent requests for other projects are not blocked on cache misses.
|
||||
new_provider = self._build_tracer_provider_for_project(project_name)
|
||||
new_provider: Final = self._build_tracer_provider_for_project(project_name)
|
||||
|
||||
with self._project_providers_lock:
|
||||
if project_name in self._project_providers:
|
||||
|
|
@ -177,7 +177,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
|
||||
def _resolve_tracer_for_kwargs(self, kwargs: dict) -> tuple[str, Tracer]:
|
||||
"""Resolve project name once and return the matching tracer."""
|
||||
project_name = self._resolve_project_name(kwargs)
|
||||
project_name: Final = self._resolve_project_name(kwargs)
|
||||
return project_name, self._get_tracer_for(project_name)
|
||||
|
||||
def get_tracer_to_use_for_request(self, kwargs: dict) -> Tracer:
|
||||
|
|
@ -204,7 +204,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
def _normalize_project_name(name: str | None) -> str | None:
|
||||
if name is None:
|
||||
return None
|
||||
normalized = str(name).strip()
|
||||
normalized: Final = str(name).strip()
|
||||
return normalized if normalized else None
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -228,7 +228,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
user-supplied and would let an authenticated caller fake proxy-mode
|
||||
detection to route their telemetry into arbitrary Arize/Phoenix projects.
|
||||
"""
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
litellm_params: Final = kwargs.get("litellm_params")
|
||||
return isinstance(litellm_params, dict) and bool(litellm_params.get("proxy_server_request"))
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -240,9 +240,9 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
select the project. SDK callers may still set project fields directly on
|
||||
``metadata``.
|
||||
"""
|
||||
auth_metadata = metadata.get("user_api_key_auth_metadata")
|
||||
auth_metadata: Final = metadata.get("user_api_key_auth_metadata")
|
||||
if isinstance(auth_metadata, dict):
|
||||
project = ArizePhoenixLogger._normalize_project_name(auth_metadata.get(metadata_key))
|
||||
project: Final = ArizePhoenixLogger._normalize_project_name(auth_metadata.get(metadata_key))
|
||||
if project:
|
||||
return project
|
||||
|
||||
|
|
@ -252,7 +252,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
|
||||
@staticmethod
|
||||
def _metadata_project_from_kwargs(kwargs: dict, metadata_key: str) -> str | None:
|
||||
proxy_mode = ArizePhoenixLogger._is_proxy_request(kwargs)
|
||||
proxy_mode: Final = ArizePhoenixLogger._is_proxy_request(kwargs)
|
||||
for metadata in ArizePhoenixLogger._iter_metadata_dicts_from_kwargs(kwargs):
|
||||
project = ArizePhoenixLogger._project_from_metadata_dict(metadata, metadata_key, proxy_mode=proxy_mode)
|
||||
if project:
|
||||
|
|
@ -268,15 +268,15 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
``user_api_key_auth_metadata.phoenix_project_name``, env, then ``default``.
|
||||
SDK priority: request metadata fields, then env, then ``default``.
|
||||
"""
|
||||
override = ArizePhoenixLogger._metadata_project_from_kwargs(kwargs, "phoenix_project_name_override")
|
||||
override: Final = ArizePhoenixLogger._metadata_project_from_kwargs(kwargs, "phoenix_project_name_override")
|
||||
if override:
|
||||
return override
|
||||
|
||||
phoenix_name = ArizePhoenixLogger._metadata_project_from_kwargs(kwargs, "phoenix_project_name")
|
||||
phoenix_name: Final = ArizePhoenixLogger._metadata_project_from_kwargs(kwargs, "phoenix_project_name")
|
||||
if phoenix_name:
|
||||
return phoenix_name
|
||||
|
||||
env_name = ArizePhoenixLogger._normalize_project_name(
|
||||
env_name: Final = ArizePhoenixLogger._normalize_project_name(
|
||||
os.environ.get("PHOENIX_PROJECT_NAME") or os.environ.get("ARIZE_PROJECT_NAME")
|
||||
)
|
||||
if env_name:
|
||||
|
|
@ -304,23 +304,23 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
if tracer is None:
|
||||
tracer = self._resolve_tracer_for_kwargs(kwargs)[1]
|
||||
|
||||
litellm_params = kwargs.get("litellm_params", {}) or {}
|
||||
proxy_server_request = litellm_params.get("proxy_server_request", {}) or {}
|
||||
headers = proxy_server_request.get("headers", {}) or {}
|
||||
litellm_params: Final = kwargs.get("litellm_params", {}) or {}
|
||||
proxy_server_request: Final = litellm_params.get("proxy_server_request", {}) or {}
|
||||
headers: Final = proxy_server_request.get("headers", {}) or {}
|
||||
|
||||
traceparent_ctx = self.get_traceparent_from_header(headers=headers) if headers.get("traceparent") else None
|
||||
|
||||
is_proxy_mode = bool(proxy_server_request)
|
||||
is_proxy_mode: Final = bool(proxy_server_request)
|
||||
|
||||
if is_proxy_mode:
|
||||
start_time_val = kwargs.get("start_time", kwargs.get("api_call_start_time"))
|
||||
parent_span = tracer.start_span(
|
||||
start_time_val: Final = kwargs.get("start_time", kwargs.get("api_call_start_time"))
|
||||
parent_span: Final = tracer.start_span(
|
||||
name="litellm_proxy_request",
|
||||
start_time=(self._to_ns(start_time_val) if start_time_val is not None else None),
|
||||
context=traceparent_ctx,
|
||||
kind=self.span_kind.SERVER,
|
||||
)
|
||||
ctx = trace.set_span_in_context(parent_span)
|
||||
ctx: Final = trace.set_span_in_context(parent_span)
|
||||
return ctx, parent_span
|
||||
|
||||
return traceparent_ctx, None
|
||||
|
|
@ -352,9 +352,9 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
_project_name, tracer = self._resolve_tracer_for_kwargs(kwargs)
|
||||
ctx, parent_span = self._get_phoenix_context(kwargs, tracer=tracer)
|
||||
|
||||
status = Status(StatusCode.OK if success else StatusCode.ERROR)
|
||||
status: Final = Status(StatusCode.OK if success else StatusCode.ERROR)
|
||||
|
||||
span = tracer.start_span(
|
||||
span: Final = tracer.start_span(
|
||||
name=self._get_span_name(kwargs),
|
||||
start_time=self._to_ns(start_time),
|
||||
context=ctx,
|
||||
|
|
@ -389,13 +389,13 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
Retrieves the Arize Phoenix configuration based on environment variables.
|
||||
Returns:
|
||||
"""
|
||||
api_key = os.environ.get("PHOENIX_API_KEY", None)
|
||||
api_key: Final = os.environ.get("PHOENIX_API_KEY", None)
|
||||
|
||||
collector_endpoint = os.environ.get("PHOENIX_COLLECTOR_HTTP_ENDPOINT", None)
|
||||
|
||||
if not collector_endpoint:
|
||||
grpc_endpoint = os.environ.get("PHOENIX_COLLECTOR_ENDPOINT", None)
|
||||
http_endpoint = os.environ.get("PHOENIX_COLLECTOR_HTTP_ENDPOINT", None)
|
||||
grpc_endpoint: Final = os.environ.get("PHOENIX_COLLECTOR_ENDPOINT", None)
|
||||
http_endpoint: Final = os.environ.get("PHOENIX_COLLECTOR_HTTP_ENDPOINT", None)
|
||||
collector_endpoint = http_endpoint or grpc_endpoint
|
||||
|
||||
endpoint = None
|
||||
|
|
@ -434,7 +434,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
elif "app.phoenix.arize.com" in endpoint:
|
||||
raise ValueError("PHOENIX_API_KEY must be set when using Phoenix Cloud (app.phoenix.arize.com).")
|
||||
|
||||
project_name = os.environ.get("PHOENIX_PROJECT_NAME") or "default"
|
||||
project_name: Final = os.environ.get("PHOENIX_PROJECT_NAME") or "default"
|
||||
|
||||
return ArizePhoenixConfig(
|
||||
otlp_auth_headers=otlp_auth_headers,
|
||||
|
|
@ -444,7 +444,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
)
|
||||
|
||||
async def async_health_check(self):
|
||||
config = self.get_arize_phoenix_config()
|
||||
config: Final = self.get_arize_phoenix_config()
|
||||
|
||||
if not config.otlp_auth_headers:
|
||||
return {
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ Arize Phoenix API client for fetching prompt versions from Arize Phoenix.
|
|||
"""
|
||||
|
||||
import urllib.parse
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
|
|
@ -63,15 +63,15 @@ class ArizePhoenixClient:
|
|||
Returns:
|
||||
Dictionary containing prompt version data, or None if not found
|
||||
"""
|
||||
safe_id = _sanitize_id(prompt_version_id)
|
||||
url = f"{self.api_base}/v1/prompt_versions/{safe_id}"
|
||||
safe_id: Final = _sanitize_id(prompt_version_id)
|
||||
url: Final = f"{self.api_base}/v1/prompt_versions/{safe_id}"
|
||||
|
||||
try:
|
||||
# Use the underlying httpx client directly to avoid query param extraction
|
||||
response = self.http_handler.get(url, headers=self.headers)
|
||||
response.raise_for_status()
|
||||
|
||||
data = response.json()
|
||||
data: Final = response.json()
|
||||
return data.get("data")
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -100,8 +100,8 @@ class ArizePhoenixClient:
|
|||
"""
|
||||
try:
|
||||
# Try to access the prompt_versions endpoint to test connection
|
||||
url = f"{self.api_base}/prompt_versions"
|
||||
response = self.http_handler.client.get(url, headers=self.headers)
|
||||
url: Final = f"{self.api_base}/prompt_versions"
|
||||
response: Final = self.http_handler.client.get(url, headers=self.headers)
|
||||
response.raise_for_status()
|
||||
return True
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ Arize Phoenix prompt manager that integrates with LiteLLM's prompt management sy
|
|||
Fetches prompt versions from Arize Phoenix and provides workspace-based access control.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
|
||||
from jinja2 import DictLoader, select_autoescape
|
||||
from jinja2.sandbox import ImmutableSandboxedEnvironment
|
||||
|
|
@ -97,10 +97,10 @@ class ArizePhoenixTemplateManager:
|
|||
"""Load a specific prompt version from Arize Phoenix."""
|
||||
try:
|
||||
# Fetch the prompt version from Arize Phoenix
|
||||
prompt_data = self.arize_client.get_prompt_version(prompt_version_id)
|
||||
prompt_data: Final = self.arize_client.get_prompt_version(prompt_version_id)
|
||||
|
||||
if prompt_data:
|
||||
template = self._parse_prompt_data(prompt_data, prompt_version_id)
|
||||
template: Final = self._parse_prompt_data(prompt_data, prompt_version_id)
|
||||
self.prompts[prompt_version_id] = template
|
||||
else:
|
||||
raise ValueError(f"Prompt version '{prompt_version_id}' not found")
|
||||
|
|
@ -109,11 +109,11 @@ class ArizePhoenixTemplateManager:
|
|||
|
||||
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 = data.get("template", {})
|
||||
messages = template_data.get("messages", [])
|
||||
template_data: Final = data.get("template", {})
|
||||
messages: Final = template_data.get("messages", [])
|
||||
|
||||
# Extract invocation parameters
|
||||
invocation_params = data.get("invocation_parameters", {})
|
||||
invocation_params: Final = data.get("invocation_parameters", {})
|
||||
provider_params = {}
|
||||
|
||||
# Extract provider-specific parameters
|
||||
|
|
@ -129,7 +129,7 @@ class ArizePhoenixTemplateManager:
|
|||
break
|
||||
|
||||
# Build metadata dictionary
|
||||
metadata = {
|
||||
metadata: Final = {
|
||||
"model_name": data.get("model_name"),
|
||||
"model_provider": data.get("model_provider"),
|
||||
"description": data.get("description", ""),
|
||||
|
|
@ -151,8 +151,8 @@ class ArizePhoenixTemplateManager:
|
|||
if template_id not in self.prompts:
|
||||
raise ValueError(f"Template '{template_id}' not found")
|
||||
|
||||
template = self.prompts[template_id]
|
||||
rendered_messages: list[AllMessageValues] = []
|
||||
template: Final = self.prompts[template_id]
|
||||
rendered_messages: Final[list[AllMessageValues]] = []
|
||||
|
||||
for message in template.messages:
|
||||
role = message.get("role", "user")
|
||||
|
|
@ -257,22 +257,22 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
Returns:
|
||||
Tuple of (rendered_messages, metadata)
|
||||
"""
|
||||
template = self.prompt_manager.get_template(prompt_id)
|
||||
template: Final = self.prompt_manager.get_template(prompt_id)
|
||||
if not template:
|
||||
raise ValueError(f"Prompt template '{prompt_id}' not found")
|
||||
|
||||
# Render the template
|
||||
rendered_messages = self.prompt_manager.render_template(prompt_id, prompt_variables or {})
|
||||
rendered_messages: Final = self.prompt_manager.render_template(prompt_id, prompt_variables or {})
|
||||
|
||||
# Extract metadata
|
||||
metadata = {
|
||||
metadata: Final = {
|
||||
"model": template.model,
|
||||
"temperature": template.temperature,
|
||||
"max_tokens": template.max_tokens,
|
||||
}
|
||||
|
||||
# Add additional invocation parameters
|
||||
invocation_params = template.invocation_parameters
|
||||
invocation_params: Final = template.invocation_parameters
|
||||
provider_params = {}
|
||||
|
||||
if "openai" in invocation_params:
|
||||
|
|
@ -395,10 +395,10 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
rendered_messages, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables)
|
||||
|
||||
# Extract model from metadata (if specified)
|
||||
template_model = prompt_metadata.get("model")
|
||||
template_model: Final = prompt_metadata.get("model")
|
||||
|
||||
# Extract optional parameters from metadata
|
||||
optional_params = {}
|
||||
optional_params: Final = {}
|
||||
for param in [
|
||||
"temperature",
|
||||
"max_tokens",
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import datetime
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
|
||||
|
|
@ -34,11 +35,11 @@ class AthinaLogger:
|
|||
import traceback
|
||||
|
||||
try:
|
||||
is_stream = kwargs.get("stream", False)
|
||||
is_stream: Final = kwargs.get("stream", False)
|
||||
if is_stream:
|
||||
if "complete_streaming_response" in kwargs:
|
||||
# Log the completion response in streaming mode
|
||||
completion_response = kwargs["complete_streaming_response"]
|
||||
completion_response: Final = kwargs["complete_streaming_response"]
|
||||
response_json = completion_response.model_dump() if completion_response else {}
|
||||
else:
|
||||
# Skip logging if the completion response is not available
|
||||
|
|
@ -46,7 +47,7 @@ class AthinaLogger:
|
|||
else:
|
||||
# Log the completion response in non streaming mode
|
||||
response_json = response_obj.model_dump() if response_obj else {}
|
||||
data = {
|
||||
data: Final = {
|
||||
"language_model_id": kwargs.get("model"),
|
||||
"request": kwargs,
|
||||
"response": response_json,
|
||||
|
|
@ -62,16 +63,16 @@ class AthinaLogger:
|
|||
data["prompt"] = kwargs.get("messages", None)
|
||||
|
||||
# Directly add tools or functions if present
|
||||
optional_params = kwargs.get("optional_params", {})
|
||||
optional_params: Final = kwargs.get("optional_params", {})
|
||||
data.update((k, v) for k, v in optional_params.items() if k in ["tools", "functions"])
|
||||
|
||||
# Add additional metadata keys
|
||||
metadata = kwargs.get("litellm_params", {}).get("metadata", {})
|
||||
metadata: Final = kwargs.get("litellm_params", {}).get("metadata", {})
|
||||
if metadata:
|
||||
for key in self.additional_keys:
|
||||
if key in metadata:
|
||||
data[key] = metadata[key]
|
||||
response = litellm.module_level_client.post(
|
||||
response: Final = litellm.module_level_client.post(
|
||||
self.athina_logging_url,
|
||||
headers=self.headers,
|
||||
data=json.dumps(data, default=str),
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ import asyncio
|
|||
import os
|
||||
import time
|
||||
import traceback
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
|
|
@ -64,15 +65,15 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
"""
|
||||
self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
|
||||
|
||||
resolved_dcr_immutable_id = dcr_immutable_id or os.getenv("AZURE_SENTINEL_DCR_IMMUTABLE_ID")
|
||||
resolved_stream_name = stream_name or os.getenv("AZURE_SENTINEL_STREAM_NAME") or "Custom-LiteLLM"
|
||||
resolved_audit_stream_name = (
|
||||
resolved_dcr_immutable_id: Final = dcr_immutable_id or os.getenv("AZURE_SENTINEL_DCR_IMMUTABLE_ID")
|
||||
resolved_stream_name: Final = stream_name or os.getenv("AZURE_SENTINEL_STREAM_NAME") or "Custom-LiteLLM"
|
||||
resolved_audit_stream_name: Final = (
|
||||
audit_stream_name or os.getenv("AZURE_SENTINEL_AUDIT_STREAM_NAME") or resolved_stream_name
|
||||
)
|
||||
resolved_endpoint = endpoint or os.getenv("AZURE_SENTINEL_ENDPOINT")
|
||||
resolved_tenant_id = tenant_id or os.getenv("AZURE_SENTINEL_TENANT_ID") or os.getenv("AZURE_TENANT_ID")
|
||||
resolved_client_id = client_id or os.getenv("AZURE_SENTINEL_CLIENT_ID") or os.getenv("AZURE_CLIENT_ID")
|
||||
resolved_client_secret = (
|
||||
resolved_endpoint: Final = endpoint or os.getenv("AZURE_SENTINEL_ENDPOINT")
|
||||
resolved_tenant_id: Final = tenant_id or os.getenv("AZURE_SENTINEL_TENANT_ID") or os.getenv("AZURE_TENANT_ID")
|
||||
resolved_client_id: Final = client_id or os.getenv("AZURE_SENTINEL_CLIENT_ID") or os.getenv("AZURE_CLIENT_ID")
|
||||
resolved_client_secret: Final = (
|
||||
client_secret or os.getenv("AZURE_SENTINEL_CLIENT_SECRET") or os.getenv("AZURE_CLIENT_SECRET")
|
||||
)
|
||||
|
||||
|
|
@ -149,16 +150,16 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
assert self.client_id is not None, "client_id is required"
|
||||
assert self.client_secret is not None, "client_secret is required"
|
||||
|
||||
token_url = f"https://login.microsoftonline.com/{self.tenant_id}/oauth2/v2.0/token"
|
||||
token_url: Final = f"https://login.microsoftonline.com/{self.tenant_id}/oauth2/v2.0/token"
|
||||
|
||||
token_data = {
|
||||
token_data: Final = {
|
||||
"client_id": self.client_id,
|
||||
"client_secret": self.client_secret,
|
||||
"scope": self.oauth_scope,
|
||||
"grant_type": "client_credentials",
|
||||
}
|
||||
|
||||
response = await self.async_httpx_client.post(
|
||||
response: Final = await self.async_httpx_client.post(
|
||||
url=token_url,
|
||||
data=token_data,
|
||||
headers={"Content-Type": "application/x-www-form-urlencoded"},
|
||||
|
|
@ -167,9 +168,9 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
if response.status_code != 200:
|
||||
raise Exception(f"Failed to get OAuth2 token: {response.status_code} - {response.text}")
|
||||
|
||||
token_response = response.json()
|
||||
token_response: Final = response.json()
|
||||
self.oauth_token = token_response.get("access_token")
|
||||
expires_in = token_response.get("expires_in", 3600)
|
||||
expires_in: Final = token_response.get("expires_in", 3600)
|
||||
|
||||
if not self.oauth_token:
|
||||
raise Exception("OAuth2 token response did not contain access_token")
|
||||
|
|
@ -191,7 +192,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
"""
|
||||
try:
|
||||
verbose_logger.debug("Azure Sentinel: Logging - Enters logging function for model %s", kwargs)
|
||||
standard_logging_payload = kwargs.get("standard_logging_object", None)
|
||||
standard_logging_payload: Final = kwargs.get("standard_logging_object", None)
|
||||
|
||||
if standard_logging_payload is None:
|
||||
verbose_logger.warning("Azure Sentinel: standard_logging_object not found in kwargs")
|
||||
|
|
@ -221,7 +222,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
"Azure Sentinel: Logging - Enters failure logging function for model %s",
|
||||
kwargs,
|
||||
)
|
||||
standard_logging_payload = kwargs.get("standard_logging_object", None)
|
||||
standard_logging_payload: Final = kwargs.get("standard_logging_object", None)
|
||||
|
||||
if standard_logging_payload is None:
|
||||
verbose_logger.warning("Azure Sentinel: standard_logging_object not found in kwargs")
|
||||
|
|
@ -294,14 +295,14 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
verbose_logger.debug("Azure Sentinel - about to flush %s %s", len(log_queue), log_type)
|
||||
|
||||
# Get OAuth2 token
|
||||
bearer_token = await self._get_oauth_token()
|
||||
bearer_token: Final = await self._get_oauth_token()
|
||||
|
||||
# Convert log queue to JSON array format expected by Logs Ingestion API
|
||||
# Each log entry should be a JSON object in the array
|
||||
body = safe_dumps(log_queue)
|
||||
body: Final = safe_dumps(log_queue)
|
||||
|
||||
# Set headers for Logs Ingestion API
|
||||
headers = {
|
||||
headers: Final = {
|
||||
"Authorization": f"Bearer {bearer_token}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
import os
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -32,11 +33,11 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
self.azure_storage_account_key: str | None = os.getenv("AZURE_STORAGE_ACCOUNT_KEY")
|
||||
|
||||
# Required Env Variables for Azure Storage
|
||||
_azure_storage_account_name = os.getenv("AZURE_STORAGE_ACCOUNT_NAME")
|
||||
_azure_storage_account_name: Final = os.getenv("AZURE_STORAGE_ACCOUNT_NAME")
|
||||
if not _azure_storage_account_name:
|
||||
raise ValueError("Missing required environment variable: AZURE_STORAGE_ACCOUNT_NAME")
|
||||
self.azure_storage_account_name: str = _azure_storage_account_name
|
||||
_azure_storage_file_system = os.getenv("AZURE_STORAGE_FILE_SYSTEM")
|
||||
_azure_storage_file_system: Final = os.getenv("AZURE_STORAGE_FILE_SYSTEM")
|
||||
if not _azure_storage_file_system:
|
||||
raise ValueError("Missing required environment variable: AZURE_STORAGE_FILE_SYSTEM")
|
||||
self.azure_storage_file_system: str = _azure_storage_file_system
|
||||
|
|
@ -71,7 +72,7 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
"AzureBlobStorageLogger: Logging - Enters logging function for model %s",
|
||||
kwargs,
|
||||
)
|
||||
standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object")
|
||||
standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object")
|
||||
|
||||
if standard_logging_payload is None:
|
||||
raise ValueError("standard_logging_payload is not set")
|
||||
|
|
@ -94,7 +95,7 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
"AzureBlobStorageLogger: Logging - Enters logging function for model %s",
|
||||
kwargs,
|
||||
)
|
||||
standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object")
|
||||
standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object")
|
||||
|
||||
if standard_logging_payload is None:
|
||||
raise ValueError("standard_logging_payload is not set")
|
||||
|
|
@ -139,10 +140,10 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
else:
|
||||
# Get a valid token instead of always requesting a new one
|
||||
await self.set_valid_azure_ad_token()
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
|
||||
json_payload = safe_dumps(payload) + "\n" # Add newline for each log entry
|
||||
payload_bytes = json_payload.encode("utf-8")
|
||||
filename = f"{payload.get('id') or str(uuid.uuid4())}.json"
|
||||
async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
|
||||
json_payload: Final = safe_dumps(payload) + "\n" # Add newline for each log entry
|
||||
payload_bytes: Final = json_payload.encode("utf-8")
|
||||
filename: Final = f"{payload.get('id') or str(uuid.uuid4())}.json"
|
||||
base_url = f"https://{self.azure_storage_account_name}.dfs.core.windows.net/{self.azure_storage_file_system}/{filename}"
|
||||
|
||||
# Execute the 3-step upload process
|
||||
|
|
@ -160,12 +161,12 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
"""Helper method to create the file resource"""
|
||||
try:
|
||||
verbose_logger.debug("Creating file resource at: %s", base_url)
|
||||
headers = {
|
||||
headers: Final = {
|
||||
"x-ms-version": AZURE_STORAGE_MSFT_VERSION,
|
||||
"Content-Length": "0",
|
||||
"Authorization": f"Bearer {self.azure_auth_token}",
|
||||
}
|
||||
response = await client.put(f"{base_url}?resource=file", headers=headers)
|
||||
response: Final = await client.put(f"{base_url}?resource=file", headers=headers)
|
||||
response.raise_for_status()
|
||||
verbose_logger.debug("Successfully created file resource")
|
||||
except Exception as e:
|
||||
|
|
@ -176,12 +177,12 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
"""Helper method to append data to the file"""
|
||||
try:
|
||||
verbose_logger.debug("Appending data to file: %s", base_url)
|
||||
headers = {
|
||||
headers: Final = {
|
||||
"x-ms-version": AZURE_STORAGE_MSFT_VERSION,
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {self.azure_auth_token}",
|
||||
}
|
||||
response = await client.patch(
|
||||
response: Final = await client.patch(
|
||||
f"{base_url}?action=append&position=0",
|
||||
headers=headers,
|
||||
data=json_payload,
|
||||
|
|
@ -196,12 +197,12 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
"""Helper method to flush the data"""
|
||||
try:
|
||||
verbose_logger.debug("Flushing data at position %s", position)
|
||||
headers = {
|
||||
headers: Final = {
|
||||
"x-ms-version": AZURE_STORAGE_MSFT_VERSION,
|
||||
"Content-Length": "0",
|
||||
"Authorization": f"Bearer {self.azure_auth_token}",
|
||||
}
|
||||
response = await client.patch(f"{base_url}?action=flush&position={position}", headers=headers)
|
||||
response: Final = await client.patch(f"{base_url}?action=flush&position={position}", headers=headers)
|
||||
response.raise_for_status()
|
||||
verbose_logger.debug("Successfully flushed data")
|
||||
except Exception as e:
|
||||
|
|
@ -253,13 +254,13 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
if client_secret is None:
|
||||
raise ValueError("Missing required environment variable: AZURE_STORAGE_CLIENT_SECRET")
|
||||
|
||||
token_provider = get_azure_ad_token_from_entra_id(
|
||||
token_provider: Final = get_azure_ad_token_from_entra_id(
|
||||
tenant_id=tenant_id,
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
scope="https://storage.azure.com/.default",
|
||||
)
|
||||
token = token_provider()
|
||||
token: Final = token_provider()
|
||||
|
||||
verbose_logger.debug("azure auth token %s", token)
|
||||
|
||||
|
|
@ -310,16 +311,16 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
|
||||
# Create an async service client
|
||||
|
||||
service_client = await self.get_service_client()
|
||||
service_client: Final = await self.get_service_client()
|
||||
# Get file system client
|
||||
file_system_client = service_client.get_file_system_client(file_system=self.azure_storage_file_system)
|
||||
file_system_client: Final = service_client.get_file_system_client(file_system=self.azure_storage_file_system)
|
||||
|
||||
try:
|
||||
# Create directory with today's date
|
||||
from datetime import datetime
|
||||
|
||||
today = datetime.now().strftime("%Y-%m-%d")
|
||||
directory_client = file_system_client.get_directory_client(today)
|
||||
today: Final = datetime.now().strftime("%Y-%m-%d")
|
||||
directory_client: Final = file_system_client.get_directory_client(today)
|
||||
|
||||
# check if the directory exists
|
||||
if not await directory_client.exists():
|
||||
|
|
@ -327,14 +328,14 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
verbose_logger.debug("Created directory: %s", today)
|
||||
|
||||
# Create a file client
|
||||
file_name = f"{payload.get('id') or str(uuid.uuid4())}.json"
|
||||
file_client = directory_client.get_file_client(file_name)
|
||||
file_name: Final = f"{payload.get('id') or str(uuid.uuid4())}.json"
|
||||
file_client: Final = directory_client.get_file_client(file_name)
|
||||
|
||||
# Create the file
|
||||
await file_client.create_file()
|
||||
|
||||
# Content to append
|
||||
content = safe_dumps(payload).encode("utf-8")
|
||||
content: Final = safe_dumps(payload).encode("utf-8")
|
||||
|
||||
# Append content to the file
|
||||
await file_client.append_data(data=content, offset=0, length=len(content))
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.custom_prompt_management import CustomPromptManagement
|
||||
|
|
@ -11,7 +11,7 @@ from litellm.types.prompts.init_prompts import SupportedPromptIntegrations
|
|||
from .bitbucket_prompt_manager import BitBucketPromptManager
|
||||
|
||||
# Global instances
|
||||
global_bitbucket_config: dict | None = None
|
||||
global_bitbucket_config: Final[dict | None] = None
|
||||
|
||||
|
||||
def set_global_bitbucket_config(config: dict) -> None:
|
||||
|
|
@ -34,14 +34,14 @@ def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "Prom
|
|||
"""
|
||||
Initialize a prompt from a BitBucket repository.
|
||||
"""
|
||||
bitbucket_config = getattr(litellm_params, "bitbucket_config", None)
|
||||
prompt_id = getattr(litellm_params, "prompt_id", None)
|
||||
bitbucket_config: Final = getattr(litellm_params, "bitbucket_config", None)
|
||||
prompt_id: Final = getattr(litellm_params, "prompt_id", None)
|
||||
|
||||
if not bitbucket_config:
|
||||
raise ValueError("bitbucket_config is required for BitBucket prompt integration")
|
||||
|
||||
try:
|
||||
bitbucket_prompt_manager = BitBucketPromptManager(
|
||||
bitbucket_prompt_manager: Final = BitBucketPromptManager(
|
||||
bitbucket_config=bitbucket_config,
|
||||
prompt_id=prompt_id,
|
||||
)
|
||||
|
|
@ -51,7 +51,7 @@ def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "Prom
|
|||
raise e
|
||||
|
||||
|
||||
prompt_initializer_registry = {
|
||||
prompt_initializer_registry: Final = {
|
||||
SupportedPromptIntegrations.BITBUCKET.value: prompt_initializer,
|
||||
}
|
||||
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue