mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
Merge branch 'litellm_internal_staging' into litellm_tool_call_obs_error
This commit is contained in:
commit
3c149ccea0
1356 changed files with 20775 additions and 62027 deletions
|
|
@ -11,3 +11,9 @@
|
|||
|
||||
# style(ui): run prettier --write across the dashboard (#29622)
|
||||
7edf3a9cb55548b143df1692f4ed7c4681d7fcf7
|
||||
|
||||
# style: reformat litellm/ with ruff format (#31317)
|
||||
430b5b8f1b12dc261a49fda99ac5d1b22381a428
|
||||
|
||||
# style: unify ruff format width on 120 (#31518)
|
||||
3dfbeabe626d203ac9de86024519d9a96c484ce4
|
||||
|
|
|
|||
2
.github/workflows/test-linting.yml
vendored
2
.github/workflows/test-linting.yml
vendored
|
|
@ -59,7 +59,7 @@ jobs:
|
|||
echo "No changed litellm Python files to check with ruff format."
|
||||
exit 0
|
||||
fi
|
||||
xargs uv run --no-sync ruff format --check --line-length 88 --exclude '/enterprise/' < "$RUNNER_TEMP/ruff_format_files.txt"
|
||||
xargs uv run --no-sync ruff format --check --exclude '/enterprise/' < "$RUNNER_TEMP/ruff_format_files.txt"
|
||||
|
||||
- name: Debug - Check file state
|
||||
run: |
|
||||
|
|
|
|||
8
Makefile
8
Makefile
|
|
@ -82,13 +82,13 @@ install-hooks:
|
|||
./scripts/install_git_hooks.sh
|
||||
|
||||
# Formatting
|
||||
# 88-column wrap matches the Black width the whole repo is formatted to; ruff.toml's
|
||||
# global line-length is 120 (for E501/isort), so 88 is forced here.
|
||||
# Wrap width is ruff.toml's single source of truth (line-length = 120), shared by the
|
||||
# formatter, E501, and the import sorter so there's no 88-vs-120 split to reconcile.
|
||||
format: install-dev
|
||||
cd litellm && $(UV_RUN) ruff format --line-length 88 --exclude '/enterprise/' . && cd ..
|
||||
cd litellm && $(UV_RUN) ruff format --exclude '/enterprise/' . && cd ..
|
||||
|
||||
format-check: install-dev
|
||||
cd litellm && $(UV_RUN) ruff format --check --line-length 88 --exclude '/enterprise/' . && cd ..
|
||||
cd litellm && $(UV_RUN) ruff format --check --exclude '/enterprise/' . && cd ..
|
||||
|
||||
# Linting targets
|
||||
lint-ruff: install-dev
|
||||
|
|
|
|||
|
|
@ -6,9 +6,7 @@ import warnings
|
|||
warnings.filterwarnings("ignore", message=".*conflict with protected namespace.*")
|
||||
# Suppress Pydantic 2.11+ deprecation warning about accessing model_fields on instances
|
||||
# This warning can accumulate during streaming and cause memory leaks
|
||||
warnings.filterwarnings(
|
||||
"ignore", message=".*Accessing the.*attribute on the instance is deprecated.*"
|
||||
)
|
||||
warnings.filterwarnings("ignore", message=".*Accessing the.*attribute on the instance is deprecated.*")
|
||||
### INIT VARIABLES #########################
|
||||
import threading
|
||||
import os
|
||||
|
|
@ -166,13 +164,9 @@ _custom_logger_compatible_callbacks_literal = Literal[
|
|||
]
|
||||
cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None
|
||||
logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None
|
||||
_known_custom_logger_compatible_callbacks: List = list(
|
||||
get_args(_custom_logger_compatible_callbacks_literal)
|
||||
)
|
||||
_known_custom_logger_compatible_callbacks: List = list(get_args(_custom_logger_compatible_callbacks_literal))
|
||||
callbacks: List[
|
||||
Union[
|
||||
Callable, _custom_logger_compatible_callbacks_literal, "CustomLogger"
|
||||
] # CustomLogger is lazy-loaded
|
||||
Union[Callable, _custom_logger_compatible_callbacks_literal, "CustomLogger"] # CustomLogger is lazy-loaded
|
||||
] = []
|
||||
callback_settings: Dict[str, Dict[str, Any]] = {}
|
||||
initialized_langfuse_clients: int = 0
|
||||
|
|
@ -183,26 +177,16 @@ prometheus_latency_buckets: Optional[List[float]] = None
|
|||
require_auth_for_metrics_endpoint: Optional[bool] = True
|
||||
argilla_batch_size: Optional[int] = None
|
||||
datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged payload.
|
||||
gcs_pub_sub_use_v1: Optional[bool] = (
|
||||
False # if you want to use v1 gcs pubsub logged payload
|
||||
)
|
||||
generic_api_use_v1: Optional[bool] = (
|
||||
False # if you want to use v1 generic api logged payload
|
||||
)
|
||||
gcs_pub_sub_use_v1: Optional[bool] = False # if you want to use v1 gcs pubsub logged payload
|
||||
generic_api_use_v1: Optional[bool] = False # if you want to use v1 generic api logged payload
|
||||
argilla_transformation_object: Optional[Dict[str, Any]] = None
|
||||
_async_input_callback: List[
|
||||
Union[str, Callable, "CustomLogger"]
|
||||
] = ( # CustomLogger is lazy-loaded
|
||||
_async_input_callback: List[Union[str, Callable, "CustomLogger"]] = ( # CustomLogger is lazy-loaded
|
||||
[]
|
||||
) # internal variable - async custom callbacks are routed here.
|
||||
_async_success_callback: List[
|
||||
Union[str, Callable, "CustomLogger"]
|
||||
] = ( # CustomLogger is lazy-loaded
|
||||
_async_success_callback: List[Union[str, Callable, "CustomLogger"]] = ( # CustomLogger is lazy-loaded
|
||||
[]
|
||||
) # internal variable - async custom callbacks are routed here.
|
||||
_async_failure_callback: List[
|
||||
Union[str, Callable, "CustomLogger"]
|
||||
] = ( # CustomLogger is lazy-loaded
|
||||
_async_failure_callback: List[Union[str, Callable, "CustomLogger"]] = ( # CustomLogger is lazy-loaded
|
||||
[]
|
||||
) # internal variable - async custom callbacks are routed here.
|
||||
pre_call_rules: List[Callable] = []
|
||||
|
|
@ -261,9 +245,7 @@ route_all_chat_openai_to_responses: bool = (
|
|||
) # When True, routes all OpenAI /chat/completions requests through the Responses API bridge
|
||||
# When True, Gemini/Vertex Live setup is deferred until client `session.update`.
|
||||
# Default False preserves historical behavior (auto-send setup on connect).
|
||||
gemini_live_defer_setup: bool = (
|
||||
os.getenv("LITELLM_GEMINI_LIVE_DEFER_SETUP", "false").lower() == "true"
|
||||
)
|
||||
gemini_live_defer_setup: bool = os.getenv("LITELLM_GEMINI_LIVE_DEFER_SETUP", "false").lower() == "true"
|
||||
use_legacy_interactions_schema: bool = (
|
||||
os.getenv("LITELLM_USE_LEGACY_INTERACTIONS_SCHEMA", "false").lower() == "true"
|
||||
) # When True, sends Api-Revision: 2026-05-07 to Google so responses use the legacy `outputs`
|
||||
|
|
@ -317,9 +299,7 @@ common_cloud_provider_auth_params: dict = {
|
|||
"params": ["project", "region_name", "token"],
|
||||
"providers": ["vertex_ai", "bedrock", "watsonx", "azure", "vertex_ai_beta"],
|
||||
}
|
||||
use_litellm_proxy: bool = (
|
||||
False # when True, requests will be sent to the specified litellm proxy endpoint
|
||||
)
|
||||
use_litellm_proxy: bool = False # when True, requests will be sent to the specified litellm proxy endpoint
|
||||
use_client: bool = False
|
||||
ssl_verify: Union[str, bool] = True
|
||||
ssl_security_level: Optional[str] = None
|
||||
|
|
@ -327,9 +307,7 @@ ssl_certificate: Optional[str] = None
|
|||
user_url_validation: bool = True
|
||||
user_url_allowed_hosts: List[str] = []
|
||||
provider_url_destination_allowed_hosts: List[str] = []
|
||||
ssl_ecdh_curve: Optional[str] = (
|
||||
None # Set to 'X25519' to disable PQC and improve performance
|
||||
)
|
||||
ssl_ecdh_curve: Optional[str] = None # Set to 'X25519' to disable PQC and improve performance
|
||||
disable_streaming_logging: bool = False
|
||||
disable_token_counter: bool = False
|
||||
disable_add_transform_inline_image_block: bool = False
|
||||
|
|
@ -370,9 +348,7 @@ prompt_name_config_map: Dict[str, PromptSpec] = {}
|
|||
##################
|
||||
### PREVIEW FEATURES ###
|
||||
enable_preview_features: bool = False
|
||||
return_response_headers: bool = (
|
||||
False # get response headers from LLM Api providers - example x-remaining-requests,
|
||||
)
|
||||
return_response_headers: bool = False # get response headers from LLM Api providers - example x-remaining-requests,
|
||||
enable_json_schema_validation: bool = False
|
||||
enable_model_config_credential_overrides: bool = False
|
||||
enable_key_alias_format_validation: bool = (
|
||||
|
|
@ -384,17 +360,13 @@ enable_gemini_default_thinking_level_low: bool = (
|
|||
####################
|
||||
logging: bool = True
|
||||
enable_loadbalancing_on_batch_endpoints: Optional[bool] = None
|
||||
require_managed_files: bool = (
|
||||
False # proxy only - require target_model_names on POST /v1/files
|
||||
)
|
||||
require_managed_files: bool = False # proxy only - require target_model_names on POST /v1/files
|
||||
enable_caching_on_provider_specific_optional_params: bool = (
|
||||
False # feature-flag for caching on optional params - e.g. 'top_k'
|
||||
)
|
||||
caching: bool = False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
|
||||
caching_with_models: bool = False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
|
||||
cache: Optional["Cache"] = (
|
||||
None # cache object <- use this - https://docs.litellm.ai/docs/caching
|
||||
)
|
||||
cache: Optional["Cache"] = None # cache object <- use this - https://docs.litellm.ai/docs/caching
|
||||
default_in_memory_ttl: Optional[float] = None
|
||||
default_redis_ttl: Optional[float] = None
|
||||
default_redis_batch_cache_expiry: Optional[float] = None
|
||||
|
|
@ -404,15 +376,15 @@ max_budget: float = 0.0 # set the max budget across all providers
|
|||
budget_duration: Optional[str] = (
|
||||
None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
|
||||
)
|
||||
default_soft_budget: float = (
|
||||
DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0
|
||||
)
|
||||
default_soft_budget: float = DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0
|
||||
forward_traceparent_to_llm_provider: bool = False
|
||||
|
||||
|
||||
_current_cost = 0.0 # private variable, used if max budget is set
|
||||
error_logs: Dict = {}
|
||||
add_function_to_prompt: bool = False # if function calling not supported by api, append function call details to system prompt
|
||||
add_function_to_prompt: bool = (
|
||||
False # if function calling not supported by api, append function call details to system prompt
|
||||
)
|
||||
client_session: Optional[httpx.Client] = None
|
||||
aclient_session: Optional[httpx.AsyncClient] = None
|
||||
model_fallbacks: Optional[List] = None # Deprecated for 'litellm.fallbacks'
|
||||
|
|
@ -476,9 +448,7 @@ prometheus_user_budget_label_include_email_alias: bool = False
|
|||
prometheus_end_user_metrics_max_series_per_metric: Optional[int] = 10000
|
||||
prometheus_end_user_metrics_ttl_seconds: Optional[float] = 3600.0
|
||||
prometheus_end_user_metrics_cleanup_interval_seconds: Optional[float] = 60.0
|
||||
disable_add_prefix_to_prompt: bool = (
|
||||
False # used by anthropic, to disable adding prefix to prompt
|
||||
)
|
||||
disable_add_prefix_to_prompt: bool = False # used by anthropic, to disable adding prefix to prompt
|
||||
disable_copilot_system_to_assistant: bool = False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior.
|
||||
public_mcp_servers: Optional[List[str]] = None
|
||||
public_mcp_hub_strict_whitelist: bool = True
|
||||
|
|
@ -489,9 +459,7 @@ public_agent_groups: Optional[List[str]] = None
|
|||
# Old format: { "displayName": "url" } (for backward compatibility)
|
||||
public_model_groups_links: Dict[str, Union[str, Dict[str, Any]]] = {}
|
||||
#### REQUEST PRIORITIZATION #######
|
||||
priority_reservation: Optional[Dict[str, Union[float, "PriorityReservationDict"]]] = (
|
||||
None
|
||||
)
|
||||
priority_reservation: Optional[Dict[str, Union[float, "PriorityReservationDict"]]] = None
|
||||
# priority_reservation_settings is lazy-loaded via __getattr__
|
||||
# Only declare for type checking - at runtime __getattr__ handles it
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -502,9 +470,7 @@ if TYPE_CHECKING:
|
|||
use_aiohttp_transport: bool = True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead.
|
||||
aiohttp_trust_env: bool = False # set to true to use HTTP_ Proxy settings
|
||||
disable_aiohttp_transport: bool = False # Set this to true to use httpx instead
|
||||
disable_aiohttp_trust_env: bool = (
|
||||
False # When False, aiohttp will respect HTTP(S)_PROXY env vars
|
||||
)
|
||||
disable_aiohttp_trust_env: bool = False # When False, aiohttp will respect HTTP(S)_PROXY env vars
|
||||
force_ipv4: bool = False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6.
|
||||
network_mock: bool = False # When True, use mock transport — no real network calls
|
||||
|
||||
|
|
@ -520,9 +486,7 @@ context_window_fallbacks: Optional[List] = None
|
|||
content_policy_fallbacks: Optional[List] = None
|
||||
allowed_fails: int = 3
|
||||
allow_dynamic_callback_disabling: bool = True
|
||||
num_retries_per_request: Optional[int] = (
|
||||
None # for the request overall (incl. fallbacks + model retries)
|
||||
)
|
||||
num_retries_per_request: Optional[int] = None # for the request overall (incl. fallbacks + model retries)
|
||||
####### SECRET MANAGERS #####################
|
||||
secret_manager_client: Optional[Any] = (
|
||||
None # list of instantiated key management clients - e.g. azure kv, infisical, etc.
|
||||
|
|
@ -539,9 +503,7 @@ output_parse_pii: bool = False
|
|||
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
|
||||
|
||||
model_cost = get_model_cost_map(url=model_cost_map_url)
|
||||
cost_discount_config: Dict[
|
||||
str, float
|
||||
] = {} # Provider-specific cost discounts {"vertex_ai": 0.05} = 5% discount
|
||||
cost_discount_config: Dict[str, float] = {} # Provider-specific cost discounts {"vertex_ai": 0.05} = 5% discount
|
||||
cost_margin_config: Dict[
|
||||
str, Union[float, Dict[str, float]]
|
||||
] = {} # Provider-specific or global cost margins. Examples:
|
||||
|
|
@ -727,9 +689,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
|
||||
for key, value in _map.items():
|
||||
if value.get("litellm_provider") == "openai" and not is_openai_finetune_model(
|
||||
key
|
||||
):
|
||||
if value.get("litellm_provider") == "openai" and not is_openai_finetune_model(key):
|
||||
open_ai_chat_completion_models.add(key)
|
||||
elif value.get("litellm_provider") == "text-completion-openai":
|
||||
open_ai_text_completion_models.add(key)
|
||||
|
|
@ -807,9 +767,7 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
|
|||
nlp_cloud_models.add(key)
|
||||
elif value.get("litellm_provider") == "aleph_alpha":
|
||||
aleph_alpha_models.add(key)
|
||||
elif value.get(
|
||||
"litellm_provider"
|
||||
) == "bedrock" and not is_bedrock_pricing_only_model(key):
|
||||
elif value.get("litellm_provider") == "bedrock" and not is_bedrock_pricing_only_model(key):
|
||||
bedrock_models.add(key)
|
||||
elif value.get("litellm_provider") == "bedrock_converse":
|
||||
bedrock_converse_models.add(key)
|
||||
|
|
@ -1445,9 +1403,7 @@ from . import rag
|
|||
from .types.llms.custom_llm import CustomLLMItem
|
||||
|
||||
custom_provider_map: List[CustomLLMItem] = []
|
||||
_custom_providers: List[
|
||||
str
|
||||
] = [] # internal helper util, used to track names of custom providers
|
||||
_custom_providers: List[str] = [] # internal helper util, used to track names of custom providers
|
||||
disable_hf_tokenizer_download: Optional[bool] = (
|
||||
None # disable huggingface tokenizer download. Defaults to openai clk100
|
||||
)
|
||||
|
|
|
|||
|
|
@ -205,9 +205,7 @@ def _get_lazy_import_registry() -> dict[str, Callable[[str], Any]]:
|
|||
return _LAZY_IMPORT_REGISTRY
|
||||
|
||||
|
||||
def _generic_lazy_import(
|
||||
name: str, import_map: dict[str, tuple[str, str]], category: str
|
||||
) -> Any:
|
||||
def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], category: str) -> Any:
|
||||
"""
|
||||
Generic function that handles lazy importing for most attributes.
|
||||
|
||||
|
|
@ -325,9 +323,7 @@ def _lazy_import_litellm_logging(name: str) -> Any:
|
|||
|
||||
def _lazy_import_llm_provider_logic(name: str) -> Any:
|
||||
"""Handler for LLM provider logic functions (get_llm_provider, etc.)"""
|
||||
return _generic_lazy_import(
|
||||
name, _LLM_PROVIDER_LOGIC_IMPORT_MAP, "LLM provider logic"
|
||||
)
|
||||
return _generic_lazy_import(name, _LLM_PROVIDER_LOGIC_IMPORT_MAP, "LLM provider logic")
|
||||
|
||||
|
||||
def _lazy_import_utils_module(name: str) -> Any:
|
||||
|
|
|
|||
|
|
@ -17,9 +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 = os.getenv("LITELLM_DISABLE_REDACT_SECRETS", "").lower() != "true"
|
||||
|
||||
|
||||
def _redact_string(value: str) -> str:
|
||||
|
|
@ -64,9 +62,7 @@ class SecretRedactionFilter(logging.Filter):
|
|||
# Redact exception tracebacks
|
||||
if record.exc_info and record.exc_info[1] is not None:
|
||||
try:
|
||||
record.exc_text = _redact_string(
|
||||
self._formatter.formatException(record.exc_info)
|
||||
)
|
||||
record.exc_text = _redact_string(self._formatter.formatException(record.exc_info))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
|
@ -189,9 +185,7 @@ class JsonFormatter(Formatter):
|
|||
json_record["logger"] = f"{record.filename}:{record.lineno}"
|
||||
|
||||
if record.exc_info:
|
||||
json_record["stacktrace"] = record.exc_text or self.formatException(
|
||||
record.exc_info
|
||||
)
|
||||
json_record["stacktrace"] = record.exc_text or self.formatException(record.exc_info)
|
||||
|
||||
return safe_dumps(json_record)
|
||||
|
||||
|
|
|
|||
|
|
@ -187,8 +187,7 @@ def _build_azure_credential(
|
|||
)
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"azure-identity is required for Azure AD Redis authentication. "
|
||||
"Install it with: pip install azure-identity"
|
||||
"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")
|
||||
|
|
@ -292,9 +291,7 @@ def get_redis_url_from_environment():
|
|||
return os.environ["REDIS_URL"]
|
||||
|
||||
if "REDIS_HOST" not in os.environ or "REDIS_PORT" not in os.environ:
|
||||
raise ValueError(
|
||||
"Either 'REDIS_URL' or both 'REDIS_HOST' and 'REDIS_PORT' must be specified for Redis."
|
||||
)
|
||||
raise ValueError("Either 'REDIS_URL' or both 'REDIS_HOST' and 'REDIS_PORT' must be specified for Redis.")
|
||||
|
||||
if "REDIS_SSL" in os.environ and os.environ["REDIS_SSL"].lower() == "true":
|
||||
redis_protocol = "rediss"
|
||||
|
|
@ -327,9 +324,7 @@ def _get_redis_client_logic(**env_overrides):
|
|||
**env_overrides,
|
||||
}
|
||||
|
||||
_startup_nodes: Optional[Union[str, list]] = redis_kwargs.get(
|
||||
"startup_nodes", None
|
||||
) or get_secret( # type: ignore
|
||||
_startup_nodes: Optional[Union[str, list]] = redis_kwargs.get("startup_nodes", None) or get_secret( # type: ignore
|
||||
"REDIS_CLUSTER_NODES"
|
||||
)
|
||||
|
||||
|
|
@ -340,18 +335,16 @@ def _get_redis_client_logic(**env_overrides):
|
|||
elif _startup_nodes is None:
|
||||
redis_kwargs.pop("startup_nodes", None)
|
||||
|
||||
_sentinel_nodes: Optional[Union[str, list]] = redis_kwargs.get(
|
||||
"sentinel_nodes", None
|
||||
) or get_secret( # type: ignore
|
||||
_sentinel_nodes: Optional[Union[str, list]] = 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: Optional[str] = redis_kwargs.get(
|
||||
"sentinel_password", None
|
||||
) or get_secret_str("REDIS_SENTINEL_PASSWORD")
|
||||
_sentinel_password: Optional[str] = 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
|
||||
|
|
@ -364,17 +357,11 @@ 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 = 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")
|
||||
|
||||
if _gcp_service_account is not None:
|
||||
verbose_logger.debug(
|
||||
"Setting up GCP IAM authentication for Redis with service account."
|
||||
)
|
||||
verbose_logger.debug("Setting up GCP IAM authentication for Redis with service account.")
|
||||
redis_kwargs["redis_connect_func"] = create_gcp_iam_redis_connect_func(
|
||||
service_account=_gcp_service_account, ssl_ca_certs=_gcp_ssl_ca_certs
|
||||
)
|
||||
|
|
@ -390,14 +377,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 = 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 = _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(
|
||||
|
|
@ -406,15 +388,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 = 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")
|
||||
|
||||
verbose_logger.debug("Setting up Azure AD authentication for Redis.")
|
||||
redis_kwargs["redis_connect_func"] = create_azure_ad_redis_connect_func(
|
||||
|
|
@ -446,9 +422,7 @@ def _get_redis_client_logic(**env_overrides):
|
|||
redis_kwargs.pop("password", None)
|
||||
elif "startup_nodes" in redis_kwargs and redis_kwargs["startup_nodes"] is not None:
|
||||
pass
|
||||
elif (
|
||||
"sentinel_nodes" in redis_kwargs and redis_kwargs["sentinel_nodes"] is not None
|
||||
):
|
||||
elif "sentinel_nodes" in redis_kwargs and redis_kwargs["sentinel_nodes"] is not None:
|
||||
pass
|
||||
elif "host" not in redis_kwargs or redis_kwargs["host"] is None:
|
||||
raise ValueError("Either 'host' or 'url' must be specified for redis.")
|
||||
|
|
@ -505,9 +479,7 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis:
|
|||
sentinel_kwargs["password"] = sentinel_password
|
||||
|
||||
if not sentinel_nodes or not service_name:
|
||||
raise ValueError(
|
||||
"Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel."
|
||||
)
|
||||
raise ValueError("Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel.")
|
||||
|
||||
verbose_logger.debug("init_redis_sentinel: sentinel nodes are being initialized.")
|
||||
|
||||
|
|
@ -532,9 +504,7 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis:
|
|||
sentinel_kwargs["password"] = sentinel_password
|
||||
|
||||
if not sentinel_nodes or not service_name:
|
||||
raise ValueError(
|
||||
"Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel."
|
||||
)
|
||||
raise ValueError("Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel.")
|
||||
|
||||
verbose_logger.debug("init_redis_sentinel: sentinel nodes are being initialized.")
|
||||
|
||||
|
|
@ -593,9 +563,7 @@ def get_redis_async_client(
|
|||
# connection — mirrors the sync path where redis_connect_func is invoked
|
||||
# per connection. Without this, the token would expire after ~1 hour.
|
||||
if redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"):
|
||||
cluster_kwargs["credential_provider"] = GCPIAMCredentialProvider(
|
||||
redis_connect_func._gcp_service_account
|
||||
)
|
||||
cluster_kwargs["credential_provider"] = GCPIAMCredentialProvider(redis_connect_func._gcp_service_account)
|
||||
# Handle Azure AD authentication for async clusters via CredentialProvider
|
||||
# so the credential's internal cache + silent refresh runs per connection
|
||||
# (mirrors GCP IAM above; avoids static-token-baked-in-pool expiry).
|
||||
|
|
@ -629,9 +597,7 @@ def get_redis_async_client(
|
|||
url_kwargs[arg] = redis_kwargs[arg]
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
"REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format(
|
||||
arg
|
||||
)
|
||||
"REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format(arg)
|
||||
)
|
||||
return async_redis.Redis.from_url(**url_kwargs)
|
||||
|
||||
|
|
@ -650,9 +616,7 @@ def get_redis_async_client(
|
|||
username=os.environ.get("REDIS_USERNAME") or None,
|
||||
)
|
||||
elif redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"):
|
||||
redis_kwargs["credential_provider"] = GCPIAMCredentialProvider(
|
||||
redis_connect_func._gcp_service_account
|
||||
)
|
||||
redis_kwargs["credential_provider"] = GCPIAMCredentialProvider(redis_connect_func._gcp_service_account)
|
||||
|
||||
_pretty_print_redis_config(redis_kwargs=redis_kwargs)
|
||||
|
||||
|
|
@ -698,18 +662,14 @@ def get_redis_connection_pool(
|
|||
username=os.environ.get("REDIS_USERNAME") or None,
|
||||
)
|
||||
elif redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"):
|
||||
redis_kwargs["credential_provider"] = GCPIAMCredentialProvider(
|
||||
redis_connect_func._gcp_service_account
|
||||
)
|
||||
redis_kwargs["credential_provider"] = GCPIAMCredentialProvider(redis_connect_func._gcp_service_account)
|
||||
|
||||
connection_class = async_redis.Connection
|
||||
if "ssl" in redis_kwargs:
|
||||
connection_class = async_redis.SSLConnection
|
||||
redis_kwargs.pop("ssl", None)
|
||||
redis_kwargs["connection_class"] = connection_class
|
||||
return async_redis.BlockingConnectionPool(
|
||||
timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs
|
||||
)
|
||||
return async_redis.BlockingConnectionPool(timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs)
|
||||
|
||||
|
||||
def _pretty_print_redis_config(redis_kwargs: dict) -> None:
|
||||
|
|
|
|||
|
|
@ -100,9 +100,7 @@ class GCPIAMCredentialProvider(CredentialProvider):
|
|||
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 = await asyncio.to_thread(_get_cached_gcp_iam_token, self._gcp_service_account)
|
||||
return (token,)
|
||||
|
||||
|
||||
|
|
@ -128,9 +126,7 @@ class AzureADCredentialProvider(CredentialProvider):
|
|||
return (token,)
|
||||
|
||||
async def get_credentials_async(self) -> Union[Tuple[str], Tuple[str, str]]:
|
||||
token_obj = await asyncio.to_thread(
|
||||
self._credential.get_token, AZURE_REDIS_SCOPE
|
||||
)
|
||||
token_obj = await asyncio.to_thread(self._credential.get_token, AZURE_REDIS_SCOPE)
|
||||
if self._username:
|
||||
return (self._username, token_obj.token)
|
||||
return (token_obj.token,)
|
||||
|
|
|
|||
|
|
@ -79,9 +79,7 @@ class ServiceLogging(CustomLogger):
|
|||
if callback == "otel":
|
||||
from litellm.proxy.proxy_server import open_telemetry_logger
|
||||
|
||||
if open_telemetry_logger is not None and _is_otel_logger(
|
||||
open_telemetry_logger
|
||||
):
|
||||
if open_telemetry_logger is not None and _is_otel_logger(open_telemetry_logger):
|
||||
return open_telemetry_logger
|
||||
return None
|
||||
|
||||
|
|
@ -142,9 +140,7 @@ class ServiceLogging(CustomLogger):
|
|||
)
|
||||
)
|
||||
|
||||
def service_failure_hook(
|
||||
self, service: ServiceTypes, duration: float, error: Exception, call_type: str
|
||||
):
|
||||
def service_failure_hook(self, service: ServiceTypes, duration: float, error: Exception, call_type: str):
|
||||
"""
|
||||
[TODO] Not implemented for sync calls yet. V0 is focused on async monitoring (used by proxy).
|
||||
"""
|
||||
|
|
@ -186,9 +182,7 @@ class ServiceLogging(CustomLogger):
|
|||
for callback in litellm.service_callback:
|
||||
if callback == "prometheus_system":
|
||||
await self.init_prometheus_services_logger_if_none()
|
||||
await self.prometheusServicesLogger.async_service_success_hook(
|
||||
payload=payload
|
||||
)
|
||||
await self.prometheusServicesLogger.async_service_success_hook(payload=payload)
|
||||
elif callback == "datadog" or isinstance(callback, DataDogLogger):
|
||||
await self.init_datadog_logger_if_none()
|
||||
await self.dd_logger.async_service_success_hook(
|
||||
|
|
@ -205,10 +199,7 @@ class ServiceLogging(CustomLogger):
|
|||
# here is what hid those calls from traces entirely. The OTel
|
||||
# logger decides what to do with a missing parent — legacy V1
|
||||
# no-ops, V2 emits a root span (and skips metrics-only pings).
|
||||
if (
|
||||
_otel_logger_to_use is not None
|
||||
and id(_otel_logger_to_use) not in emitted_otel_logger_ids
|
||||
):
|
||||
if _otel_logger_to_use is not None and id(_otel_logger_to_use) not in emitted_otel_logger_ids:
|
||||
emitted_otel_logger_ids.add(id(_otel_logger_to_use))
|
||||
await _otel_logger_to_use.async_service_success_hook(
|
||||
payload=payload,
|
||||
|
|
@ -249,9 +240,7 @@ class ServiceLogging(CustomLogger):
|
|||
from litellm.proxy.proxy_server import open_telemetry_logger
|
||||
|
||||
if not hasattr(self, "otel_logger"):
|
||||
if open_telemetry_logger is not None and isinstance(
|
||||
open_telemetry_logger, OpenTelemetry
|
||||
):
|
||||
if open_telemetry_logger is not None and isinstance(open_telemetry_logger, OpenTelemetry):
|
||||
self.otel_logger: OpenTelemetry = open_telemetry_logger
|
||||
else:
|
||||
verbose_logger.warning(
|
||||
|
|
@ -319,10 +308,7 @@ class ServiceLogging(CustomLogger):
|
|||
|
||||
# See the success hook: no parent gate, so background failures
|
||||
# are traced too. V1 no-ops without a parent; V2 emits a root.
|
||||
if (
|
||||
_otel_logger_to_use is not None
|
||||
and id(_otel_logger_to_use) not in emitted_otel_logger_ids
|
||||
):
|
||||
if _otel_logger_to_use is not None and id(_otel_logger_to_use) not in emitted_otel_logger_ids:
|
||||
emitted_otel_logger_ids.add(id(_otel_logger_to_use))
|
||||
await _otel_logger_to_use.async_service_failure_hook(
|
||||
payload=payload,
|
||||
|
|
@ -361,9 +347,7 @@ class ServiceLogging(CustomLogger):
|
|||
pass
|
||||
else:
|
||||
raise Exception(
|
||||
"Duration={} is not a float or timedelta object. type={}".format(
|
||||
_duration, type(_duration)
|
||||
)
|
||||
"Duration={} is not a float or timedelta object. type={}".format(_duration, type(_duration))
|
||||
) # invalid _duration value
|
||||
# Batch polling callbacks (check_batch_cost) don't include call_type in kwargs.
|
||||
# Use .get() to avoid KeyError.
|
||||
|
|
|
|||
|
|
@ -119,17 +119,13 @@ class LiteLLMA2ACardResolver(_A2ACardResolver): # type: ignore[misc]
|
|||
last_error = None
|
||||
for path in paths:
|
||||
try:
|
||||
verbose_logger.debug(
|
||||
f"Attempting to fetch agent card from {self.base_url}{path}"
|
||||
)
|
||||
verbose_logger.debug(f"Attempting to fetch agent card from {self.base_url}{path}")
|
||||
return await super().get_agent_card(
|
||||
relative_card_path=path,
|
||||
http_kwargs=http_kwargs,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
f"Failed to fetch agent card from {self.base_url}{path}: {e}"
|
||||
)
|
||||
verbose_logger.debug(f"Failed to fetch agent card from {self.base_url}{path}: {e}")
|
||||
last_error = e
|
||||
continue
|
||||
|
||||
|
|
@ -138,7 +134,4 @@ class LiteLLMA2ACardResolver(_A2ACardResolver): # type: ignore[misc]
|
|||
raise last_error
|
||||
|
||||
# This shouldn't happen, but just in case
|
||||
raise Exception(
|
||||
f"Failed to fetch agent card from {self.base_url}. "
|
||||
f"Tried paths: {', '.join(paths)}"
|
||||
)
|
||||
raise Exception(f"Failed to fetch agent card from {self.base_url}. Tried paths: {', '.join(paths)}")
|
||||
|
|
|
|||
|
|
@ -87,9 +87,7 @@ class A2AClient:
|
|||
extra_headers=self.extra_headers,
|
||||
)
|
||||
|
||||
async def send_message(
|
||||
self, request: "SendMessageRequest"
|
||||
) -> LiteLLMSendMessageResponse:
|
||||
async def send_message(self, request: "SendMessageRequest") -> LiteLLMSendMessageResponse:
|
||||
"""Send a message to the A2A agent."""
|
||||
from litellm.a2a_protocol.main import asend_message
|
||||
|
||||
|
|
@ -103,7 +101,5 @@ class A2AClient:
|
|||
from litellm.a2a_protocol.main import asend_message_streaming
|
||||
|
||||
a2a_client = await self._get_client()
|
||||
async for chunk in asend_message_streaming(
|
||||
a2a_client=a2a_client, request=request
|
||||
):
|
||||
async for chunk in asend_message_streaming(a2a_client=a2a_client, request=request):
|
||||
yield chunk
|
||||
|
|
|
|||
|
|
@ -97,11 +97,7 @@ class A2ACostCalculator:
|
|||
completion_tokens = 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 = 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)
|
||||
|
||||
return input_cost + output_cost
|
||||
|
|
|
|||
|
|
@ -181,10 +181,7 @@ def handle_a2a_localhost_retry(
|
|||
ImportError: If the A2A SDK is not installed
|
||||
"""
|
||||
if not A2A_SDK_AVAILABLE or _A2AClient is None:
|
||||
raise ImportError(
|
||||
"A2A SDK is required for localhost retry handling. "
|
||||
"Install it with: pip install a2a"
|
||||
)
|
||||
raise ImportError("A2A SDK is required for localhost retry handling. Install it with: pip install a2a")
|
||||
|
||||
request_type = "streaming " if is_streaming else ""
|
||||
verbose_logger.warning(
|
||||
|
|
|
|||
|
|
@ -139,10 +139,7 @@ class A2ALocalhostURLError(A2AConnectionError):
|
|||
self.base_url = base_url
|
||||
self.original_error = original_error
|
||||
|
||||
message = (
|
||||
f"Agent card contains localhost/internal URL '{localhost_url}'. "
|
||||
f"Retrying with base URL '{base_url}'."
|
||||
)
|
||||
message = f"Agent card contains localhost/internal URL '{localhost_url}'. Retrying with base URL '{base_url}'."
|
||||
super().__init__(
|
||||
message=message,
|
||||
url=localhost_url,
|
||||
|
|
|
|||
|
|
@ -75,9 +75,7 @@ class A2ACompletionBridgeHandler:
|
|||
)
|
||||
|
||||
if a2a_provider_config is not None:
|
||||
verbose_logger.info(
|
||||
f"A2A: Using provider config for {custom_llm_provider}"
|
||||
)
|
||||
verbose_logger.info(f"A2A: Using provider config for {custom_llm_provider}")
|
||||
|
||||
return await a2a_provider_config.handle_non_streaming(
|
||||
request_id=request_id,
|
||||
|
|
@ -91,9 +89,7 @@ class A2ACompletionBridgeHandler:
|
|||
message = params.get("message", {})
|
||||
|
||||
# Transform A2A message to OpenAI format
|
||||
openai_messages = (
|
||||
A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
)
|
||||
openai_messages = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
|
||||
# Get completion params
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
|
|
@ -106,9 +102,7 @@ class A2ACompletionBridgeHandler:
|
|||
else:
|
||||
full_model = model
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A completion bridge: model={full_model}, api_base={api_base}"
|
||||
)
|
||||
verbose_logger.info(f"A2A completion bridge: model={full_model}, api_base={api_base}")
|
||||
|
||||
# Build completion params dict
|
||||
completion_params: Dict[str, Any] = {
|
||||
|
|
@ -143,11 +137,9 @@ class A2ACompletionBridgeHandler:
|
|||
response = await litellm.acompletion(**completion_params)
|
||||
|
||||
# Transform response to A2A format
|
||||
a2a_response = (
|
||||
A2ACompletionBridgeTransformation.openai_response_to_a2a_response(
|
||||
response=response,
|
||||
request_id=request_id,
|
||||
)
|
||||
a2a_response = A2ACompletionBridgeTransformation.openai_response_to_a2a_response(
|
||||
response=response,
|
||||
request_id=request_id,
|
||||
)
|
||||
|
||||
verbose_logger.info(f"A2A completion bridge completed: request_id={request_id}")
|
||||
|
|
@ -192,9 +184,7 @@ class A2ACompletionBridgeHandler:
|
|||
)
|
||||
|
||||
if a2a_provider_config is not None:
|
||||
verbose_logger.info(
|
||||
f"A2A: Using provider config for {custom_llm_provider} (streaming)"
|
||||
)
|
||||
verbose_logger.info(f"A2A: Using provider config for {custom_llm_provider} (streaming)")
|
||||
|
||||
async for chunk in a2a_provider_config.handle_streaming(
|
||||
request_id=request_id,
|
||||
|
|
@ -217,9 +207,7 @@ class A2ACompletionBridgeHandler:
|
|||
)
|
||||
|
||||
# Transform A2A message to OpenAI format
|
||||
openai_messages = (
|
||||
A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
)
|
||||
openai_messages = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
|
||||
# Get completion params
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
|
|
@ -232,9 +220,7 @@ class A2ACompletionBridgeHandler:
|
|||
else:
|
||||
full_model = model
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A completion bridge streaming: model={full_model}, api_base={api_base}"
|
||||
)
|
||||
verbose_logger.info(f"A2A completion bridge streaming: model={full_model}, api_base={api_base}")
|
||||
|
||||
# Build completion params dict
|
||||
completion_params: Dict[str, Any] = {
|
||||
|
|
@ -299,11 +285,9 @@ class A2ACompletionBridgeHandler:
|
|||
|
||||
# Emit artifact update with accumulated content
|
||||
if accumulated_text:
|
||||
artifact_event = (
|
||||
A2ACompletionBridgeTransformation.create_artifact_update_event(
|
||||
ctx=ctx,
|
||||
text=accumulated_text,
|
||||
)
|
||||
artifact_event = A2ACompletionBridgeTransformation.create_artifact_update_event(
|
||||
ctx=ctx,
|
||||
text=accumulated_text,
|
||||
)
|
||||
yield artifact_event
|
||||
|
||||
|
|
@ -315,9 +299,7 @@ class A2ACompletionBridgeHandler:
|
|||
)
|
||||
yield completed_event
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A completion bridge streaming completed: request_id={request_id}, chunks={chunk_count}"
|
||||
)
|
||||
verbose_logger.info(f"A2A completion bridge streaming completed: request_id={request_id}, chunks={chunk_count}")
|
||||
|
||||
|
||||
# Convenience functions that delegate to the class methods
|
||||
|
|
|
|||
|
|
@ -104,16 +104,12 @@ class A2ACompletionBridgeTransformation:
|
|||
# ``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 {}
|
||||
)
|
||||
existing_dict: Dict[str, Any] = existing_metadata if isinstance(existing_metadata, dict) else {}
|
||||
merged_metadata: Dict[str, Any] = {**forward_metadata, **existing_dict}
|
||||
extra_body = {**extra_body, "metadata": merged_metadata}
|
||||
completion_params["extra_body"] = extra_body
|
||||
|
||||
verbose_logger.debug(
|
||||
f"A2A -> completion forward metadata keys={list(forward_metadata.keys())}"
|
||||
)
|
||||
verbose_logger.debug(f"A2A -> completion forward metadata keys={list(forward_metadata.keys())}")
|
||||
|
||||
@staticmethod
|
||||
def a2a_message_to_openai_messages(
|
||||
|
|
@ -149,9 +145,7 @@ class A2ACompletionBridgeTransformation:
|
|||
# once at run level via extra_body.metadata (LangGraph POST /runs/wait shape).
|
||||
openai_message: Dict[str, Any] = {"role": openai_role, "content": content}
|
||||
|
||||
verbose_logger.debug(
|
||||
f"A2A -> OpenAI transform: role={role} -> {openai_role}, content_length={len(content)}"
|
||||
)
|
||||
verbose_logger.debug(f"A2A -> OpenAI transform: role={role} -> {openai_role}, content_length={len(content)}")
|
||||
|
||||
return [openai_message]
|
||||
|
||||
|
|
|
|||
|
|
@ -120,9 +120,7 @@ def _get_a2a_model_info(a2a_client: Any, kwargs: Dict[str, Any]) -> str:
|
|||
litellm_logging_obj.model = model
|
||||
litellm_logging_obj.custom_llm_provider = custom_llm_provider
|
||||
litellm_logging_obj.model_call_details["model"] = model
|
||||
litellm_logging_obj.model_call_details["custom_llm_provider"] = (
|
||||
custom_llm_provider
|
||||
)
|
||||
litellm_logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
return agent_name
|
||||
|
||||
|
|
@ -139,19 +137,13 @@ async def _send_message_via_completion_bridge(
|
|||
|
||||
Requires request; api_base is optional for providers that derive endpoint from model.
|
||||
"""
|
||||
verbose_logger.info(
|
||||
f"A2A using completion bridge: provider={custom_llm_provider}, api_base={api_base}"
|
||||
)
|
||||
verbose_logger.info(f"A2A using completion bridge: provider={custom_llm_provider}, api_base={api_base}")
|
||||
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
A2ACompletionBridgeHandler,
|
||||
)
|
||||
|
||||
params = (
|
||||
request.params.model_dump(mode="json")
|
||||
if hasattr(request.params, "model_dump")
|
||||
else dict(request.params)
|
||||
)
|
||||
params = request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params)
|
||||
|
||||
response_dict = await A2ACompletionBridgeHandler.handle_non_streaming(
|
||||
request_id=str(request.id),
|
||||
|
|
@ -161,9 +153,7 @@ async def _send_message_via_completion_bridge(
|
|||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
|
||||
return LiteLLMSendMessageResponse.from_dict(
|
||||
response_dict, request_id=str(request.id)
|
||||
)
|
||||
return LiteLLMSendMessageResponse.from_dict(response_dict, request_id=str(request.id))
|
||||
|
||||
|
||||
async def _execute_a2a_send_with_retry(
|
||||
|
|
@ -203,9 +193,7 @@ async def _execute_a2a_send_with_retry(
|
|||
except Exception:
|
||||
raise
|
||||
if a2a_response is None:
|
||||
raise RuntimeError(
|
||||
"A2A send_message failed: no response received after retry attempts."
|
||||
)
|
||||
raise RuntimeError("A2A send_message failed: no response received after retry attempts.")
|
||||
return a2a_response
|
||||
|
||||
|
||||
|
|
@ -295,9 +283,7 @@ async def asend_message(
|
|||
# Create A2A client if not provided but api_base is available
|
||||
if a2a_client is None:
|
||||
if api_base is None:
|
||||
raise ValueError(
|
||||
"Either a2a_client or api_base is required for standard A2A flow"
|
||||
)
|
||||
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}
|
||||
if agent_id:
|
||||
|
|
@ -305,9 +291,7 @@ async def asend_message(
|
|||
# Overlay agent-level headers (agent headers take precedence over LiteLLM internal ones)
|
||||
if agent_extra_headers:
|
||||
extra_headers.update(agent_extra_headers)
|
||||
a2a_client = await create_a2a_client(
|
||||
base_url=api_base, extra_headers=extra_headers
|
||||
)
|
||||
a2a_client = await create_a2a_client(base_url=api_base, extra_headers=extra_headers)
|
||||
|
||||
# Type assertion: a2a_client is guaranteed to be non-None here
|
||||
assert a2a_client is not None
|
||||
|
|
@ -317,9 +301,7 @@ async def asend_message(
|
|||
verbose_logger.info(f"A2A send_message request_id={request.id}, agent={agent_name}")
|
||||
|
||||
# Get agent card URL for localhost retry logic
|
||||
agent_card = getattr(a2a_client, "_litellm_agent_card", None) or getattr(
|
||||
a2a_client, "agent_card", None
|
||||
)
|
||||
agent_card = getattr(a2a_client, "_litellm_agent_card", None) or getattr(a2a_client, "agent_card", None)
|
||||
card_url = getattr(agent_card, "url", None) if agent_card else None
|
||||
|
||||
a2a_response = await _execute_a2a_send_with_retry(
|
||||
|
|
@ -334,9 +316,7 @@ async def asend_message(
|
|||
verbose_logger.info(f"A2A send_message completed, request_id={request.id}")
|
||||
|
||||
# Wrap in LiteLLM response type for _hidden_params support
|
||||
response = LiteLLMSendMessageResponse.from_a2a_response(
|
||||
a2a_response, request_id=str(request.id)
|
||||
)
|
||||
response = 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)
|
||||
|
|
@ -389,9 +369,7 @@ def send_message(
|
|||
if loop is not None:
|
||||
return asend_message(a2a_client=a2a_client, request=request, **kwargs)
|
||||
else:
|
||||
return asyncio.run(
|
||||
asend_message(a2a_client=a2a_client, request=request, **kwargs)
|
||||
)
|
||||
return asyncio.run(asend_message(a2a_client=a2a_client, request=request, **kwargs))
|
||||
|
||||
|
||||
def _build_streaming_logging_obj(
|
||||
|
|
@ -492,9 +470,7 @@ async def asend_message_streaming(
|
|||
raise ValueError("request is required for completion bridge")
|
||||
# api_base is optional for providers that derive endpoint from model (e.g., bedrock/agentcore)
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A streaming using completion bridge: provider={custom_llm_provider}"
|
||||
)
|
||||
verbose_logger.info(f"A2A streaming using completion bridge: provider={custom_llm_provider}")
|
||||
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
A2ACompletionBridgeHandler,
|
||||
|
|
@ -502,9 +478,7 @@ async def asend_message_streaming(
|
|||
|
||||
# Extract params from request
|
||||
params = (
|
||||
request.params.model_dump(mode="json")
|
||||
if hasattr(request.params, "model_dump")
|
||||
else dict(request.params)
|
||||
request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params)
|
||||
)
|
||||
|
||||
async for chunk in A2ACompletionBridgeHandler.handle_streaming(
|
||||
|
|
@ -524,9 +498,7 @@ async def asend_message_streaming(
|
|||
# Create A2A client if not provided but api_base is available
|
||||
if a2a_client is None:
|
||||
if api_base is None:
|
||||
raise ValueError(
|
||||
"Either a2a_client or api_base is required for standard A2A flow"
|
||||
)
|
||||
raise ValueError("Either a2a_client or api_base is required for standard A2A flow")
|
||||
# Mirror the non-streaming path: always include trace and agent-id headers
|
||||
streaming_extra_headers: Dict[str, str] = {
|
||||
"X-LiteLLM-Trace-Id": str(request.id),
|
||||
|
|
@ -535,9 +507,7 @@ async def asend_message_streaming(
|
|||
streaming_extra_headers["X-LiteLLM-Agent-Id"] = agent_id
|
||||
if agent_extra_headers:
|
||||
streaming_extra_headers.update(agent_extra_headers)
|
||||
a2a_client = await create_a2a_client(
|
||||
base_url=api_base, extra_headers=streaming_extra_headers
|
||||
)
|
||||
a2a_client = await create_a2a_client(base_url=api_base, extra_headers=streaming_extra_headers)
|
||||
|
||||
# Type assertion: a2a_client is guaranteed to be non-None here
|
||||
assert a2a_client is not None
|
||||
|
|
@ -545,9 +515,7 @@ async def asend_message_streaming(
|
|||
verbose_logger.info(f"A2A send_message_streaming request_id={request.id}")
|
||||
|
||||
# Build logging object for streaming completion callbacks
|
||||
agent_card = getattr(a2a_client, "_litellm_agent_card", None) or getattr(
|
||||
a2a_client, "agent_card", None
|
||||
)
|
||||
agent_card = getattr(a2a_client, "_litellm_agent_card", None) or getattr(a2a_client, "agent_card", None)
|
||||
card_url = getattr(agent_card, "url", None) if agent_card else None
|
||||
agent_name = getattr(agent_card, "name", "unknown") if agent_card else "unknown"
|
||||
|
||||
|
|
@ -645,8 +613,7 @@ async def create_a2a_client(
|
|||
"""
|
||||
if not A2A_SDK_AVAILABLE:
|
||||
raise ImportError(
|
||||
"The 'a2a' package is required for A2A agent invocation. "
|
||||
"Install it with: pip install a2a-sdk"
|
||||
"The 'a2a' package is required for A2A agent invocation. Install it with: pip install a2a-sdk"
|
||||
)
|
||||
|
||||
verbose_logger.info(f"Creating A2A client for {base_url}")
|
||||
|
|
@ -671,9 +638,7 @@ async def create_a2a_client(
|
|||
httpx_client = _async_handler.client
|
||||
if extra_headers:
|
||||
httpx_client.headers.update(extra_headers)
|
||||
verbose_proxy_logger.debug(
|
||||
f"A2A client created with extra_headers={list(extra_headers.keys())}"
|
||||
)
|
||||
verbose_proxy_logger.debug(f"A2A client created with extra_headers={list(extra_headers.keys())}")
|
||||
|
||||
# Resolve agent card
|
||||
resolver = A2ACardResolver(
|
||||
|
|
@ -682,9 +647,7 @@ async def create_a2a_client(
|
|||
)
|
||||
agent_card = await resolver.get_agent_card()
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Resolved agent card: {agent_card.name if hasattr(agent_card, 'name') else 'unknown'}"
|
||||
)
|
||||
verbose_logger.debug(f"Resolved agent card: {agent_card.name if hasattr(agent_card, 'name') else 'unknown'}")
|
||||
|
||||
# Create A2A client
|
||||
a2a_client = _A2AClient(
|
||||
|
|
@ -718,8 +681,7 @@ async def aget_agent_card(
|
|||
"""
|
||||
if not A2A_SDK_AVAILABLE:
|
||||
raise ImportError(
|
||||
"The 'a2a' package is required for A2A agent invocation. "
|
||||
"Install it with: pip install a2a-sdk"
|
||||
"The 'a2a' package is required for A2A agent invocation. Install it with: pip install a2a-sdk"
|
||||
)
|
||||
|
||||
verbose_logger.info(f"Fetching agent card from {base_url}")
|
||||
|
|
@ -737,7 +699,5 @@ async def aget_agent_card(
|
|||
)
|
||||
agent_card = await resolver.get_agent_card()
|
||||
|
||||
verbose_logger.info(
|
||||
f"Fetched agent card: {agent_card.name if hasattr(agent_card, 'name') else 'unknown'}"
|
||||
)
|
||||
verbose_logger.info(f"Fetched agent card: {agent_card.name if hasattr(agent_card, 'name') else 'unknown'}")
|
||||
return agent_card
|
||||
|
|
|
|||
|
|
@ -30,8 +30,7 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig):
|
|||
litellm_params = kwargs.get("litellm_params")
|
||||
if not litellm_params:
|
||||
raise ValueError(
|
||||
"litellm_params is required for BedrockAgentCoreA2AConfig "
|
||||
"(must contain model with AgentCore ARN)"
|
||||
"litellm_params is required for BedrockAgentCoreA2AConfig (must contain model with AgentCore ARN)"
|
||||
)
|
||||
return await BedrockAgentCoreA2AHandler.handle_non_streaming(
|
||||
request_id=request_id,
|
||||
|
|
@ -51,8 +50,7 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig):
|
|||
litellm_params = kwargs.get("litellm_params")
|
||||
if not litellm_params:
|
||||
raise ValueError(
|
||||
"litellm_params is required for BedrockAgentCoreA2AConfig "
|
||||
"(must contain model with AgentCore ARN)"
|
||||
"litellm_params is required for BedrockAgentCoreA2AConfig (must contain model with AgentCore ARN)"
|
||||
)
|
||||
async for chunk in BedrockAgentCoreA2AHandler.handle_streaming(
|
||||
request_id=request_id,
|
||||
|
|
|
|||
|
|
@ -44,19 +44,15 @@ class BedrockAgentCoreA2AHandler:
|
|||
Returns:
|
||||
A2A JSON-RPC response dict from the AgentCore agent
|
||||
"""
|
||||
url, headers, body = (
|
||||
BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
method="message/send",
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
url, headers, body = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
method="message/send",
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
|
||||
verbose_logger.info(
|
||||
f"BedrockAgentCore A2A: Sending non-streaming request to {url}"
|
||||
)
|
||||
verbose_logger.info(f"BedrockAgentCore A2A: Sending non-streaming request to {url}")
|
||||
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=cast(Any, httpxSpecialProvider.A2AProvider),
|
||||
|
|
@ -70,9 +66,7 @@ class BedrockAgentCoreA2AHandler:
|
|||
response_data = response.json()
|
||||
|
||||
if "error" in response_data:
|
||||
verbose_logger.warning(
|
||||
f"BedrockAgentCore A2A: Agent returned error: {response_data['error']}"
|
||||
)
|
||||
verbose_logger.warning(f"BedrockAgentCore A2A: Agent returned error: {response_data['error']}")
|
||||
|
||||
return response_data
|
||||
|
||||
|
|
@ -96,15 +90,13 @@ class BedrockAgentCoreA2AHandler:
|
|||
Yields:
|
||||
A2A streaming response events from the AgentCore agent
|
||||
"""
|
||||
url, headers, body = (
|
||||
BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
method="message/send",
|
||||
stream=True,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
url, headers, body = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
method="message/send",
|
||||
stream=True,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
|
||||
verbose_logger.info(f"BedrockAgentCore A2A: Sending streaming request to {url}")
|
||||
|
|
@ -126,15 +118,12 @@ class BedrockAgentCoreA2AHandler:
|
|||
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"
|
||||
"BedrockAgentCore A2A streaming: received JSON instead of SSE, yielding as single event"
|
||||
)
|
||||
response_body = await response.aread()
|
||||
response_data = json.loads(response_body)
|
||||
yield response_data
|
||||
else:
|
||||
# SSE stream — parse data: lines
|
||||
async for event in BedrockAgentCoreA2ATransformation.parse_sse_events(
|
||||
response
|
||||
):
|
||||
async for event in BedrockAgentCoreA2ATransformation.parse_sse_events(response):
|
||||
yield event
|
||||
|
|
|
|||
|
|
@ -50,9 +50,7 @@ def _filter_reserved_headers(
|
|||
dropped: 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
|
||||
):
|
||||
if k_lower in _RESERVED_EXACT_HEADERS or any(k_lower.startswith(prefix) for prefix in _RESERVED_PREFIX_HEADERS):
|
||||
dropped.append(k)
|
||||
continue
|
||||
filtered[k] = v
|
||||
|
|
@ -115,11 +113,7 @@ class BedrockAgentCoreA2ATransformation:
|
|||
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 = {k: v for k, v in litellm_params.items() if k not in ("model", "custom_llm_provider")}
|
||||
|
||||
agentcore_config = AmazonAgentCoreConfig()
|
||||
|
||||
|
|
@ -200,7 +194,5 @@ class BedrockAgentCoreA2ATransformation:
|
|||
event = json.loads(data_str)
|
||||
yield event
|
||||
except json.JSONDecodeError:
|
||||
verbose_logger.debug(
|
||||
f"BedrockAgentCore A2A: Skipping non-JSON SSE line: {data_str[:100]}"
|
||||
)
|
||||
verbose_logger.debug(f"BedrockAgentCore A2A: Skipping non-JSON SSE line: {data_str[:100]}")
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -22,8 +22,7 @@ class LangFlowA2AConfig(BaseA2AProviderConfig):
|
|||
litellm_params = kwargs.get("litellm_params")
|
||||
if not litellm_params:
|
||||
raise ValueError(
|
||||
"litellm_params is required for LangFlowA2AConfig "
|
||||
"(must contain custom_llm_provider and model)"
|
||||
"litellm_params is required for LangFlowA2AConfig (must contain custom_llm_provider and model)"
|
||||
)
|
||||
litellm_params = merge_a2a_session_into_litellm_params(
|
||||
litellm_params, params, litellm_params.get(A2A_USER_API_KEY_HASH_PARAM)
|
||||
|
|
@ -46,8 +45,7 @@ class LangFlowA2AConfig(BaseA2AProviderConfig):
|
|||
litellm_params = kwargs.get("litellm_params")
|
||||
if not litellm_params:
|
||||
raise ValueError(
|
||||
"litellm_params is required for LangFlowA2AConfig "
|
||||
"(must contain custom_llm_provider and model)"
|
||||
"litellm_params is required for LangFlowA2AConfig (must contain custom_llm_provider and model)"
|
||||
)
|
||||
litellm_params = merge_a2a_session_into_litellm_params(
|
||||
litellm_params, params, litellm_params.get(A2A_USER_API_KEY_HASH_PARAM)
|
||||
|
|
|
|||
|
|
@ -91,9 +91,7 @@ class PydanticAIHandler:
|
|||
"""
|
||||
if api_base is None:
|
||||
raise ValueError("api_base is required for Pydantic AI agents")
|
||||
verbose_logger.info(
|
||||
f"Pydantic AI: Faking streaming for Pydantic AI agent at {api_base}"
|
||||
)
|
||||
verbose_logger.info(f"Pydantic AI: Faking streaming for Pydantic AI agent at {api_base}")
|
||||
|
||||
# Get raw task response first (not the transformed A2A format)
|
||||
raw_response = await PydanticAITransformation.send_and_get_raw_response(
|
||||
|
|
|
|||
|
|
@ -41,17 +41,9 @@ class PydanticAITransformation:
|
|||
Cleaned object with None values removed
|
||||
"""
|
||||
if isinstance(obj, dict):
|
||||
return {
|
||||
k: PydanticAITransformation._remove_none_values(v)
|
||||
for k, v in obj.items()
|
||||
if v is not None
|
||||
}
|
||||
return {k: PydanticAITransformation._remove_none_values(v) for k, v in obj.items() if v is not None}
|
||||
elif isinstance(obj, list):
|
||||
return [
|
||||
PydanticAITransformation._remove_none_values(item)
|
||||
for item in obj
|
||||
if item is not None
|
||||
]
|
||||
return [PydanticAITransformation._remove_none_values(item) for item in obj if item is not None]
|
||||
else:
|
||||
return obj
|
||||
|
||||
|
|
@ -125,9 +117,7 @@ class PydanticAITransformation:
|
|||
status = result.get("status", {})
|
||||
state = status.get("state", "")
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Pydantic AI: Poll attempt {attempt + 1}/{max_attempts}, state={state}"
|
||||
)
|
||||
verbose_logger.debug(f"Pydantic AI: Poll attempt {attempt + 1}/{max_attempts}, state={state}")
|
||||
|
||||
if state == "completed":
|
||||
return poll_data
|
||||
|
|
@ -136,9 +126,7 @@ class PydanticAITransformation:
|
|||
|
||||
await asyncio.sleep(poll_interval)
|
||||
|
||||
raise TimeoutError(
|
||||
f"Task {task_id} did not complete within {max_attempts * poll_interval} seconds"
|
||||
)
|
||||
raise TimeoutError(f"Task {task_id} did not complete within {max_attempts * poll_interval} seconds")
|
||||
|
||||
@staticmethod
|
||||
async def _send_and_poll_raw(
|
||||
|
|
@ -211,9 +199,7 @@ class PydanticAITransformation:
|
|||
# Need to poll for completion
|
||||
task_id = result.get("id")
|
||||
if task_id:
|
||||
verbose_logger.info(
|
||||
f"Pydantic AI: Task {task_id} submitted, polling for completion..."
|
||||
)
|
||||
verbose_logger.info(f"Pydantic AI: Task {task_id} submitted, polling for completion...")
|
||||
response_data = await PydanticAITransformation._poll_for_completion(
|
||||
client=client,
|
||||
endpoint=endpoint,
|
||||
|
|
@ -222,9 +208,7 @@ class PydanticAITransformation:
|
|||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
|
||||
verbose_logger.info(
|
||||
f"Pydantic AI: Received completed response for request_id={request_id}"
|
||||
)
|
||||
verbose_logger.info(f"Pydantic AI: Received completed response for request_id={request_id}")
|
||||
|
||||
return response_data
|
||||
|
||||
|
|
@ -325,9 +309,7 @@ class PydanticAITransformation:
|
|||
Standard A2A non-streaming response format
|
||||
"""
|
||||
# Extract the agent response text
|
||||
full_text, message_id, parts = PydanticAITransformation._extract_response_text(
|
||||
response_data
|
||||
)
|
||||
full_text, message_id, parts = PydanticAITransformation._extract_response_text(response_data)
|
||||
|
||||
# Build standard A2A message
|
||||
a2a_message = {
|
||||
|
|
@ -424,9 +406,7 @@ class PydanticAITransformation:
|
|||
A2A streaming response events
|
||||
"""
|
||||
# Extract the response text from completed task
|
||||
full_text, message_id, parts = PydanticAITransformation._extract_response_text(
|
||||
response_data
|
||||
)
|
||||
full_text, message_id, parts = PydanticAITransformation._extract_response_text(response_data)
|
||||
|
||||
# Extract input message from raw response for history
|
||||
result = response_data.get("result", {})
|
||||
|
|
@ -455,9 +435,7 @@ class PydanticAITransformation:
|
|||
"contextId": context_id,
|
||||
"kind": "message",
|
||||
"messageId": input_message_id,
|
||||
"parts": input_message.get(
|
||||
"parts", [{"kind": "text", "text": ""}]
|
||||
),
|
||||
"parts": input_message.get("parts", [{"kind": "text", "text": ""}]),
|
||||
"role": "user",
|
||||
"taskId": task_id,
|
||||
}
|
||||
|
|
@ -539,6 +517,4 @@ class PydanticAITransformation:
|
|||
}
|
||||
yield completed_event
|
||||
|
||||
verbose_logger.info(
|
||||
f"Pydantic AI: Fake streaming completed for request_id={request_id}"
|
||||
)
|
||||
verbose_logger.info(f"Pydantic AI: Fake streaming completed for request_id={request_id}")
|
||||
|
|
|
|||
|
|
@ -56,9 +56,7 @@ class WatsonxOrchestrateHandler:
|
|||
return hashlib.sha256(material.encode()).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _cp4d_token_ttl_seconds(
|
||||
expiration: Any, now_wall: Optional[float] = None
|
||||
) -> int:
|
||||
def _cp4d_token_ttl_seconds(expiration: Any, now_wall: Optional[float] = 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()
|
||||
|
|
@ -72,9 +70,7 @@ class WatsonxOrchestrateHandler:
|
|||
username: Optional[str] = None,
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
) -> str:
|
||||
cache_key = WatsonxOrchestrateHandler._token_cache_key(
|
||||
auth_mode, cp4d_host, api_key, username
|
||||
)
|
||||
cache_key = WatsonxOrchestrateHandler._token_cache_key(auth_mode, cp4d_host, api_key, username)
|
||||
now = time.monotonic()
|
||||
cached = _token_cache.get(cache_key)
|
||||
if cached and cached[1] > now:
|
||||
|
|
@ -98,9 +94,7 @@ class WatsonxOrchestrateHandler:
|
|||
ttl_s = int(payload.get("expires_in", 3600))
|
||||
else:
|
||||
if not username:
|
||||
raise ValueError(
|
||||
"'username' is required in litellm_params when auth_mode='cp4d'"
|
||||
)
|
||||
raise ValueError("'username' is required in litellm_params when auth_mode='cp4d'")
|
||||
token_url = f"{cp4d_host.rstrip('/')}/icp4d-api/v1/authorize"
|
||||
response = await client.post(
|
||||
token_url,
|
||||
|
|
@ -140,15 +134,12 @@ class WatsonxOrchestrateHandler:
|
|||
response.raise_for_status()
|
||||
result: Dict[str, Any] = response.json()
|
||||
status = result.get("status", "")
|
||||
verbose_logger.debug(
|
||||
f"WXO: Poll {attempt + 1}/{max_attempts} run='{run_id}' status='{status}'"
|
||||
)
|
||||
verbose_logger.debug(f"WXO: Poll {attempt + 1}/{max_attempts} run='{run_id}' status='{status}'")
|
||||
if status in WatsonxOrchestrateTransformation.TERMINAL_STATES:
|
||||
return result
|
||||
|
||||
raise asyncio.TimeoutError(
|
||||
f"WXO run '{run_id}' did not reach a terminal state after "
|
||||
f"{max_attempts * interval_s:.0f}s"
|
||||
f"WXO run '{run_id}' did not reach a terminal state after {max_attempts * interval_s:.0f}s"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -172,9 +163,7 @@ class WatsonxOrchestrateHandler:
|
|||
status = run_data.get("status", "")
|
||||
|
||||
if status not in WatsonxOrchestrateTransformation.SUCCESS_STATES:
|
||||
raise RuntimeError(
|
||||
f"WXO run ended with non-success status '{status}': {run_data}"
|
||||
)
|
||||
raise RuntimeError(f"WXO run ended with non-success status '{status}': {run_data}")
|
||||
|
||||
return run_data
|
||||
|
||||
|
|
@ -191,9 +180,7 @@ class WatsonxOrchestrateHandler:
|
|||
event = json.loads(data_str)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
chunk_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(
|
||||
event
|
||||
)
|
||||
chunk_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(event)
|
||||
if chunk_text:
|
||||
accumulated_text += chunk_text
|
||||
return accumulated_text
|
||||
|
|
@ -208,13 +195,9 @@ class WatsonxOrchestrateHandler:
|
|||
if not cp4d_host:
|
||||
raise ValueError("'cp4d_host' is required in litellm_params for WXO agents")
|
||||
if not instance_id:
|
||||
raise ValueError(
|
||||
"'instance_id' is required in litellm_params for WXO agents"
|
||||
)
|
||||
raise ValueError("'instance_id' is required in litellm_params for WXO agents")
|
||||
if not wxo_agent_id:
|
||||
raise ValueError(
|
||||
"'wxo_agent_id' is required in litellm_params for WXO agents"
|
||||
)
|
||||
raise ValueError("'wxo_agent_id' is required in litellm_params for WXO agents")
|
||||
if not api_key:
|
||||
raise ValueError("'api_key' is required in litellm_params for WXO agents")
|
||||
|
||||
|
|
@ -244,9 +227,7 @@ class WatsonxOrchestrateHandler:
|
|||
username=wxo.username,
|
||||
client=client,
|
||||
)
|
||||
base_url = WatsonxOrchestrateTransformation.get_api_base_url(
|
||||
wxo.cp4d_host, wxo.instance_id
|
||||
)
|
||||
base_url = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id)
|
||||
auth_headers = {
|
||||
"Authorization": f"Bearer {token}",
|
||||
"Content-Type": "application/json",
|
||||
|
|
@ -273,12 +254,8 @@ class WatsonxOrchestrateHandler:
|
|||
client=client,
|
||||
)
|
||||
|
||||
response_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(
|
||||
run_data
|
||||
)
|
||||
return WatsonxOrchestrateTransformation.build_a2a_message_response(
|
||||
request_id=request_id, text=response_text
|
||||
)
|
||||
response_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(run_data)
|
||||
return WatsonxOrchestrateTransformation.build_a2a_message_response(request_id=request_id, text=response_text)
|
||||
|
||||
@staticmethod
|
||||
async def handle_streaming(
|
||||
|
|
@ -298,9 +275,7 @@ class WatsonxOrchestrateHandler:
|
|||
username=wxo.username,
|
||||
client=client,
|
||||
)
|
||||
base_url = WatsonxOrchestrateTransformation.get_api_base_url(
|
||||
wxo.cp4d_host, wxo.instance_id
|
||||
)
|
||||
base_url = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id)
|
||||
auth_headers = {
|
||||
"Authorization": f"Bearer {token}",
|
||||
"Content-Type": "application/json",
|
||||
|
|
@ -330,14 +305,8 @@ class WatsonxOrchestrateHandler:
|
|||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
response_text = (
|
||||
WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response(
|
||||
result
|
||||
)
|
||||
)
|
||||
async for (
|
||||
chunk
|
||||
) in WatsonxOrchestrateTransformation.fake_streaming_from_text(
|
||||
response_text = WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response(result)
|
||||
async for chunk in WatsonxOrchestrateTransformation.fake_streaming_from_text(
|
||||
text=response_text,
|
||||
request_id=request_id,
|
||||
chunk_size=chunk_size,
|
||||
|
|
@ -356,13 +325,9 @@ class WatsonxOrchestrateHandler:
|
|||
auth_headers=auth_headers,
|
||||
client=client,
|
||||
)
|
||||
accumulated_text = (
|
||||
WatsonxOrchestrateTransformation.extract_text_from_wxo_result(result)
|
||||
)
|
||||
accumulated_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(result)
|
||||
else:
|
||||
accumulated_text = await WatsonxOrchestrateHandler._accumulate_wxo_sse_text(
|
||||
response
|
||||
)
|
||||
accumulated_text = await WatsonxOrchestrateHandler._accumulate_wxo_sse_text(response)
|
||||
|
||||
async for chunk in WatsonxOrchestrateTransformation.fake_streaming_from_text(
|
||||
text=accumulated_text,
|
||||
|
|
|
|||
|
|
@ -19,9 +19,7 @@ class WatsonxOrchestrateTransformation:
|
|||
Handles request/response transformation between A2A and the WXO REST API.
|
||||
"""
|
||||
|
||||
TERMINAL_STATES = frozenset(
|
||||
{"completed", "succeeded", "failed", "error", "cancelled"}
|
||||
)
|
||||
TERMINAL_STATES = frozenset({"completed", "succeeded", "failed", "error", "cancelled"})
|
||||
SUCCESS_STATES = frozenset({"completed", "succeeded"})
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -114,11 +112,7 @@ class WatsonxOrchestrateTransformation:
|
|||
verbose_logger.warning("WXO: A2A result has no parts list")
|
||||
return ""
|
||||
for part in parts:
|
||||
if (
|
||||
isinstance(part, dict)
|
||||
and part.get("kind") == "text"
|
||||
and part.get("text")
|
||||
):
|
||||
if isinstance(part, dict) and part.get("kind") == "text" and part.get("text"):
|
||||
return str(part["text"])
|
||||
verbose_logger.warning("WXO: A2A result parts contained no text")
|
||||
return ""
|
||||
|
|
@ -219,6 +213,4 @@ class WatsonxOrchestrateTransformation:
|
|||
},
|
||||
}
|
||||
|
||||
verbose_logger.debug(
|
||||
f"WXO: Fake streaming completed for request_id={request_id}"
|
||||
)
|
||||
verbose_logger.debug(f"WXO: Fake streaming completed for request_id={request_id}")
|
||||
|
|
|
|||
|
|
@ -71,11 +71,7 @@ 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 {}
|
||||
)
|
||||
chunk_dict = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
|
||||
text = A2ARequestUtils.extract_text_from_response(chunk_dict)
|
||||
if text:
|
||||
self.collected_text_parts.append(text)
|
||||
|
|
@ -85,11 +81,7 @@ 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 {}
|
||||
)
|
||||
chunk_dict = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
|
||||
result = chunk_dict.get("result", {})
|
||||
if isinstance(result, dict):
|
||||
status = result.get("status", {})
|
||||
|
|
@ -110,9 +102,7 @@ class A2AStreamingIterator:
|
|||
prompt_tokens = 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 ""
|
||||
)
|
||||
output_text = self.collected_text_parts[-1] if self.collected_text_parts else ""
|
||||
completion_tokens = A2ARequestUtils.count_tokens(output_text)
|
||||
|
||||
total_tokens = prompt_tokens + completion_tokens
|
||||
|
|
@ -168,9 +158,7 @@ class A2AStreamingIterator:
|
|||
result: Dict[str, Any] = {
|
||||
"id": getattr(self.request, "id", "unknown"),
|
||||
"jsonrpc": "2.0",
|
||||
"usage": (
|
||||
usage.model_dump() if hasattr(usage, "model_dump") else dict(usage)
|
||||
),
|
||||
"usage": (usage.model_dump() if hasattr(usage, "model_dump") else dict(usage)),
|
||||
}
|
||||
|
||||
# Add final chunk result if available
|
||||
|
|
|
|||
|
|
@ -48,9 +48,7 @@ class GetAnthropicBetaHeadersConfig:
|
|||
"""Load the local backup beta headers config bundled with the package."""
|
||||
try:
|
||||
content = json.loads(
|
||||
files("litellm")
|
||||
.joinpath("anthropic_beta_headers_config.json")
|
||||
.read_text(encoding="utf-8")
|
||||
files("litellm").joinpath("anthropic_beta_headers_config.json").read_text(encoding="utf-8")
|
||||
)
|
||||
return content
|
||||
except Exception as e:
|
||||
|
|
@ -70,16 +68,14 @@ class GetAnthropicBetaHeadersConfig:
|
|||
"""Check if fetched config is a non-empty dict with expected structure."""
|
||||
if not isinstance(fetched_config, dict):
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Fetched beta headers config is not a dict (type=%s). "
|
||||
"Falling back to local backup.",
|
||||
"LiteLLM: Fetched beta headers config is not a dict (type=%s). Falling back to local backup.",
|
||||
type(fetched_config).__name__,
|
||||
)
|
||||
return False
|
||||
|
||||
if len(fetched_config) == 0:
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Fetched beta headers config is empty. "
|
||||
"Falling back to local backup.",
|
||||
"LiteLLM: Fetched beta headers config is empty. Falling back to local backup.",
|
||||
)
|
||||
return False
|
||||
|
||||
|
|
@ -95,8 +91,7 @@ class GetAnthropicBetaHeadersConfig:
|
|||
|
||||
if not has_provider:
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Fetched beta headers config missing provider keys. "
|
||||
"Falling back to local backup.",
|
||||
"LiteLLM: Fetched beta headers config missing provider keys. Falling back to local backup.",
|
||||
)
|
||||
return False
|
||||
|
||||
|
|
@ -147,20 +142,16 @@ def get_beta_headers_config(url: str) -> dict:
|
|||
content = 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.",
|
||||
"LiteLLM: Failed to fetch remote beta headers config from %s: %s. Falling back to local backup.",
|
||||
url,
|
||||
str(e),
|
||||
)
|
||||
return GetAnthropicBetaHeadersConfig.load_local_beta_headers_config()
|
||||
|
||||
# Validate the fetched config
|
||||
if not GetAnthropicBetaHeadersConfig.validate_beta_headers_config(
|
||||
fetched_config=content
|
||||
):
|
||||
if not GetAnthropicBetaHeadersConfig.validate_beta_headers_config(fetched_config=content):
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Fetched beta headers config failed integrity check. "
|
||||
"Using local backup instead. url=%s",
|
||||
"LiteLLM: Fetched beta headers config failed integrity check. Using local backup instead. url=%s",
|
||||
url,
|
||||
)
|
||||
return GetAnthropicBetaHeadersConfig.load_local_beta_headers_config()
|
||||
|
|
@ -256,9 +247,7 @@ def filter_and_transform_beta_headers(
|
|||
|
||||
# Check if header is in the mapping
|
||||
if header not in provider_mapping:
|
||||
verbose_logger.debug(
|
||||
f"Dropping unknown beta header '{header}' for provider '{provider}' (not in mapping)"
|
||||
)
|
||||
verbose_logger.debug(f"Dropping unknown beta header '{header}' for provider '{provider}' (not in mapping)")
|
||||
continue
|
||||
|
||||
# Get the mapped header value
|
||||
|
|
@ -266,9 +255,7 @@ def filter_and_transform_beta_headers(
|
|||
|
||||
# Skip if header is unsupported (null value)
|
||||
if mapped_header is None:
|
||||
verbose_logger.debug(
|
||||
f"Dropping unsupported beta header '{header}' for provider '{provider}'"
|
||||
)
|
||||
verbose_logger.debug(f"Dropping unsupported beta header '{header}' for provider '{provider}'")
|
||||
continue
|
||||
|
||||
# Add the mapped header
|
||||
|
|
|
|||
|
|
@ -148,9 +148,7 @@ class AnthropicExceptionMapping:
|
|||
parsed = None
|
||||
|
||||
# If parsed and already in Anthropic format - passthrough
|
||||
if parsed is not None and AnthropicExceptionMapping._is_anthropic_error_dict(
|
||||
parsed
|
||||
):
|
||||
if parsed is not None and AnthropicExceptionMapping._is_anthropic_error_dict(parsed):
|
||||
# Optionally add request_id if provided and not present
|
||||
if request_id and "request_id" not in parsed:
|
||||
parsed["request_id"] = request_id
|
||||
|
|
@ -158,9 +156,7 @@ class AnthropicExceptionMapping:
|
|||
|
||||
# Extract message - use parsed dict if available, otherwise raw string
|
||||
if parsed is not None:
|
||||
message = AnthropicExceptionMapping._extract_message_from_dict(
|
||||
parsed, raw_message
|
||||
)
|
||||
message = AnthropicExceptionMapping._extract_message_from_dict(parsed, raw_message)
|
||||
else:
|
||||
message = raw_message
|
||||
|
||||
|
|
|
|||
|
|
@ -102,9 +102,7 @@ def create(
|
|||
AnthropicMessagesResponse,
|
||||
Iterator[bytes],
|
||||
AsyncIterator[Any],
|
||||
Coroutine[
|
||||
Any, Any, Union[AnthropicMessagesResponse, AsyncIterator[Any], Iterator[bytes]]
|
||||
],
|
||||
Coroutine[Any, Any, Union[AnthropicMessagesResponse, AsyncIterator[Any], Iterator[bytes]]],
|
||||
]:
|
||||
"""
|
||||
Async wrapper for Anthropic's messages API
|
||||
|
|
|
|||
|
|
@ -81,12 +81,8 @@ def get_assistants(
|
|||
) -> SyncCursorPage[Assistant]:
|
||||
aget_assistants: Optional[bool] = 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
|
||||
)
|
||||
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)
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
|
|
@ -138,15 +134,9 @@ def get_assistants(
|
|||
aget_assistants=aget_assistants, # type: ignore
|
||||
) # type: ignore
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_version = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -184,9 +174,7 @@ def get_assistants(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -200,9 +188,7 @@ def get_assistants(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -266,18 +252,10 @@ def create_assistants(
|
|||
api_version: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Union[Assistant, Coroutine[Any, Any, Assistant]]:
|
||||
async_create_assistants: Optional[bool] = 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
|
||||
)
|
||||
async_create_assistants: Optional[bool] = 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)
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
|
|
@ -310,9 +288,7 @@ def create_assistants(
|
|||
}
|
||||
|
||||
# only send params that are not None
|
||||
create_assistant_data = {
|
||||
k: v for k, v in create_assistant_data.items() if v is not None
|
||||
}
|
||||
create_assistant_data = {k: v for k, v in create_assistant_data.items() if v is not None}
|
||||
|
||||
response: Optional[Union[Coroutine[Any, Any, Assistant], Assistant]] = None
|
||||
if custom_llm_provider == "openai":
|
||||
|
|
@ -348,15 +324,9 @@ def create_assistants(
|
|||
async_create_assistants=async_create_assistants, # type: ignore
|
||||
) # type: ignore
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_version = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -398,9 +368,7 @@ def create_assistants(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
if response is None:
|
||||
|
|
@ -459,21 +427,13 @@ def delete_assistant(
|
|||
api_version: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Union[AssistantDeleted, Coroutine[Any, Any, AssistantDeleted]]:
|
||||
optional_params = GenericLiteLLMParams(
|
||||
api_key=api_key, api_base=api_base, api_version=api_version, **kwargs
|
||||
)
|
||||
optional_params = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs)
|
||||
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
|
||||
async_delete_assistants: Optional[bool] = 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"
|
||||
)
|
||||
async_delete_assistants: Optional[bool] = 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")
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
|
|
@ -491,9 +451,7 @@ def delete_assistant(
|
|||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
response: Optional[
|
||||
Union[AssistantDeleted, Coroutine[Any, Any, AssistantDeleted]]
|
||||
] = None
|
||||
response: Optional[Union[AssistantDeleted, Coroutine[Any, Any, AssistantDeleted]]] = None
|
||||
if custom_llm_provider == "openai":
|
||||
api_base = (
|
||||
optional_params.api_base
|
||||
|
|
@ -503,18 +461,10 @@ def delete_assistant(
|
|||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
optional_params.organization
|
||||
or litellm.organization
|
||||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
or None
|
||||
optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) or None
|
||||
)
|
||||
# set API KEY
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
or litellm.api_key
|
||||
or litellm.openai_key
|
||||
or os.getenv("OPENAI_API_KEY")
|
||||
)
|
||||
api_key = optional_params.api_key or litellm.api_key or litellm.openai_key or os.getenv("OPENAI_API_KEY")
|
||||
|
||||
response = openai_assistants_api.delete_assistant(
|
||||
api_base=api_base,
|
||||
|
|
@ -527,15 +477,9 @@ def delete_assistant(
|
|||
async_delete_assistants=async_delete_assistants,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_version = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -577,9 +521,7 @@ def delete_assistant(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="delete_assistant", url="https://github.com/BerriAI/litellm"
|
||||
),
|
||||
request=httpx.Request(method="delete_assistant", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
if response is None:
|
||||
|
|
@ -594,9 +536,7 @@ def delete_assistant(
|
|||
### THREADS ###
|
||||
|
||||
|
||||
async def acreate_thread(
|
||||
custom_llm_provider: Literal["openai", "azure"], **kwargs
|
||||
) -> Thread:
|
||||
async def acreate_thread(custom_llm_provider: Literal["openai", "azure"], **kwargs) -> Thread:
|
||||
loop = asyncio.get_event_loop()
|
||||
### PASS ARGS TO GET ASSISTANTS ###
|
||||
kwargs["acreate_thread"] = True
|
||||
|
|
@ -716,9 +656,7 @@ def create_thread(
|
|||
acreate_thread=acreate_thread,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -729,9 +667,7 @@ def create_thread(
|
|||
) # type: ignore
|
||||
|
||||
api_version: Optional[str] = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret("AZURE_API_VERSION")
|
||||
optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
|
|
@ -767,9 +703,7 @@ def create_thread(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response # type: ignore
|
||||
|
|
@ -874,14 +808,10 @@ def get_thread(
|
|||
aget_thread=aget_thread,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_version: Optional[str] = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret("AZURE_API_VERSION")
|
||||
optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
|
||||
api_key = (
|
||||
|
|
@ -924,9 +854,7 @@ def get_thread(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response # type: ignore
|
||||
|
|
@ -1000,9 +928,7 @@ def add_message(
|
|||
) -> OpenAIMessage:
|
||||
### COMMON OBJECTS ###
|
||||
a_add_message = kwargs.pop("a_add_message", None)
|
||||
_message_data = MessageData(
|
||||
role=role, content=content, attachments=attachments, metadata=metadata
|
||||
)
|
||||
_message_data = MessageData(role=role, content=content, attachments=attachments, metadata=metadata)
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
|
|
@ -1065,14 +991,10 @@ def add_message(
|
|||
a_add_message=a_add_message,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_version: Optional[str] = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret("AZURE_API_VERSION")
|
||||
optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
|
||||
api_key = (
|
||||
|
|
@ -1113,9 +1035,7 @@ def add_message(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -1228,14 +1148,10 @@ def get_messages(
|
|||
aget_messages=aget_messages,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_version: Optional[str] = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret("AZURE_API_VERSION")
|
||||
optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
|
||||
api_key = (
|
||||
|
|
@ -1275,9 +1191,7 @@ def get_messages(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -1438,15 +1352,9 @@ def run_thread(
|
|||
event_handler=event_handler,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_version = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -1492,9 +1400,7 @@ def run_thread(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response # type: ignore
|
||||
|
|
|
|||
|
|
@ -43,11 +43,7 @@ def get_optional_params_add_message(
|
|||
"metadata": None,
|
||||
}
|
||||
|
||||
non_default_params = {
|
||||
k: v
|
||||
for k, v in passed_params.items()
|
||||
if (k in default_params and v != default_params[k])
|
||||
}
|
||||
non_default_params = {k: v for k, v in passed_params.items() if (k in default_params and v != default_params[k])}
|
||||
optional_params = {}
|
||||
|
||||
## raise exception if non-default value passed for non-openai/azure embedding calls
|
||||
|
|
@ -55,9 +51,7 @@ def get_optional_params_add_message(
|
|||
if len(non_default_params.keys()) > 0:
|
||||
keys = 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
|
||||
if litellm.drop_params is True and k not in supported_params: # drop the unsupported non-default values
|
||||
non_default_params.pop(k, None)
|
||||
elif k not in supported_params:
|
||||
raise litellm.utils.UnsupportedParamsError(
|
||||
|
|
@ -108,11 +102,7 @@ def get_optional_params_image_gen(
|
|||
"user": None,
|
||||
}
|
||||
|
||||
non_default_params = {
|
||||
k: v
|
||||
for k, v in passed_params.items()
|
||||
if (k in default_params and v != default_params[k])
|
||||
}
|
||||
non_default_params = {k: v for k, v in passed_params.items() if (k in default_params and v != default_params[k])}
|
||||
optional_params = {}
|
||||
|
||||
## raise exception if non-default value passed for non-openai/azure embedding calls
|
||||
|
|
@ -120,9 +110,7 @@ def get_optional_params_image_gen(
|
|||
if len(non_default_params.keys()) > 0:
|
||||
keys = 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
|
||||
if litellm.drop_params is True and k not in supported_params: # drop the unsupported non-default values
|
||||
non_default_params.pop(k, None)
|
||||
elif k not in supported_params:
|
||||
raise UnsupportedParamsError(
|
||||
|
|
|
|||
|
|
@ -106,9 +106,7 @@ def batch_completion(
|
|||
original_kwargs = {}
|
||||
if "kwargs" in kwargs_modified:
|
||||
original_kwargs = kwargs_modified.pop("kwargs")
|
||||
future = executor.submit(
|
||||
litellm.completion, **kwargs_modified, **original_kwargs
|
||||
)
|
||||
future = executor.submit(litellm.completion, **kwargs_modified, **original_kwargs)
|
||||
completions.append(future)
|
||||
|
||||
# Retrieve the results from the futures
|
||||
|
|
@ -153,13 +151,9 @@ def batch_completion_models(*args, **kwargs):
|
|||
futures = {}
|
||||
with ThreadPoolExecutor(max_workers=len(models)) as executor:
|
||||
for model in models:
|
||||
futures[model] = executor.submit(
|
||||
litellm.completion, *args, model=model, **kwargs
|
||||
)
|
||||
futures[model] = executor.submit(litellm.completion, *args, model=model, **kwargs)
|
||||
|
||||
for model, future in sorted(
|
||||
futures.items(), key=lambda x: models.index(x[0])
|
||||
):
|
||||
for model, future in sorted(futures.items(), key=lambda x: models.index(x[0])):
|
||||
if future.result() is not None:
|
||||
return future.result()
|
||||
elif "deployments" in kwargs:
|
||||
|
|
@ -171,14 +165,10 @@ def batch_completion_models(*args, **kwargs):
|
|||
with ThreadPoolExecutor(max_workers=len(deployments)) as executor:
|
||||
for deployment in deployments:
|
||||
for key in kwargs.keys():
|
||||
if (
|
||||
key not in deployment
|
||||
): # don't override deployment values e.g. model name, api base, etc.
|
||||
if key not in deployment: # don't override deployment values e.g. model name, api base, etc.
|
||||
deployment[key] = kwargs[key]
|
||||
kwargs = {**deployment, **nested_kwargs}
|
||||
futures[deployment["model"]] = executor.submit(
|
||||
litellm.completion, **kwargs
|
||||
)
|
||||
futures[deployment["model"]] = executor.submit(litellm.completion, **kwargs)
|
||||
|
||||
while futures:
|
||||
# wait for the first returned future
|
||||
|
|
@ -191,9 +181,7 @@ def batch_completion_models(*args, **kwargs):
|
|||
return result
|
||||
except Exception:
|
||||
# if model 1 fails, continue with response from model 2, model3
|
||||
print_verbose(
|
||||
"\n\ngot an exception, ignoring, removing from futures"
|
||||
)
|
||||
print_verbose("\n\ngot an exception, ignoring, removing from futures")
|
||||
print_verbose(futures)
|
||||
new_futures = {}
|
||||
for key, value in futures.items():
|
||||
|
|
@ -254,10 +242,7 @@ def batch_completion_models_all_responses(*args, **kwargs):
|
|||
responses = []
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=len(models)) as executor:
|
||||
futures = [
|
||||
executor.submit(litellm.completion, *args, model=model, **kwargs)
|
||||
for model in models
|
||||
]
|
||||
futures = [executor.submit(litellm.completion, *args, model=model, **kwargs) for model in models]
|
||||
|
||||
for future in futures:
|
||||
try:
|
||||
|
|
@ -265,9 +250,7 @@ def batch_completion_models_all_responses(*args, **kwargs):
|
|||
if result is not None:
|
||||
responses.append(result)
|
||||
except Exception as e:
|
||||
print_verbose(
|
||||
f"batch_completion_models_all_responses: model request failed: {str(e)}"
|
||||
)
|
||||
print_verbose(f"batch_completion_models_all_responses: model request failed: {str(e)}")
|
||||
continue
|
||||
|
||||
return responses
|
||||
|
|
|
|||
|
|
@ -10,9 +10,7 @@ from litellm.utils import token_counter
|
|||
|
||||
async def calculate_batch_cost_and_usage(
|
||||
file_content_dictionary: List[dict],
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"
|
||||
],
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],
|
||||
model_name: Optional[str] = None,
|
||||
model_info: Optional[ModelInfo] = None,
|
||||
) -> Tuple[float, Usage, List[str]]:
|
||||
|
|
@ -36,18 +34,14 @@ async def calculate_batch_cost_and_usage(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
model_name=model_name,
|
||||
)
|
||||
batch_models = _get_batch_models_from_file_content(
|
||||
file_content_dictionary, model_name
|
||||
)
|
||||
batch_models = _get_batch_models_from_file_content(file_content_dictionary, model_name)
|
||||
|
||||
return batch_cost, batch_usage, batch_models
|
||||
|
||||
|
||||
async def _handle_completed_batch(
|
||||
batch: Batch,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"
|
||||
],
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],
|
||||
model_name: Optional[str] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> Tuple[float, Usage, List[str]]:
|
||||
|
|
@ -76,9 +70,7 @@ async def _handle_completed_batch(
|
|||
model_name=model_name,
|
||||
)
|
||||
|
||||
batch_models = _get_batch_models_from_file_content(
|
||||
file_content_dictionary, model_name
|
||||
)
|
||||
batch_models = _get_batch_models_from_file_content(file_content_dictionary, model_name)
|
||||
|
||||
return batch_cost, batch_usage, batch_models
|
||||
|
||||
|
|
@ -104,9 +96,7 @@ def _get_batch_models_from_file_content(
|
|||
|
||||
def _batch_cost_calculator(
|
||||
file_content_dictionary: List[dict],
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"
|
||||
] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
|
||||
model_name: Optional[str] = None,
|
||||
model_info: Optional[ModelInfo] = None,
|
||||
) -> float:
|
||||
|
|
@ -118,9 +108,7 @@ def _batch_cost_calculator(
|
|||
and model_name
|
||||
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
|
||||
):
|
||||
batch_cost, _ = calculate_vertex_ai_batch_cost_and_usage(
|
||||
file_content_dictionary, model_name
|
||||
)
|
||||
batch_cost, _ = calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
|
||||
verbose_logger.debug("vertex_ai_total_cost=%s", batch_cost)
|
||||
return batch_cost
|
||||
|
||||
|
|
@ -181,9 +169,7 @@ def calculate_vertex_ai_batch_cost_and_usage(
|
|||
)
|
||||
total_cost += p_cost + c_cost
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"vertex_ai batch cost calculation error for line: %s", str(e)
|
||||
)
|
||||
verbose_logger.debug("vertex_ai batch cost calculation error for line: %s", str(e))
|
||||
|
||||
prompt_tokens += _prompt
|
||||
completion_tokens += _completion
|
||||
|
|
@ -206,9 +192,7 @@ def calculate_vertex_ai_batch_cost_and_usage(
|
|||
|
||||
async def _get_batch_output_file_content_as_dictionary(
|
||||
batch: Batch,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"
|
||||
] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> List[dict]:
|
||||
"""
|
||||
|
|
@ -235,12 +219,8 @@ async def _get_batch_output_file_content_as_dictionary(
|
|||
is_base64_unified_file_id = _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]
|
||||
verbose_logger.debug(
|
||||
f"Extracted LLM output file ID from unified file ID: {file_id}"
|
||||
)
|
||||
file_id = is_base64_unified_file_id.split("llm_output_file_id,")[1].split(";")[0]
|
||||
verbose_logger.debug(f"Extracted LLM output file ID from unified file ID: {file_id}")
|
||||
except (IndexError, AttributeError) as e:
|
||||
verbose_logger.error(
|
||||
f"Failed to extract LLM output file ID from unified file ID: {batch.output_file_id}, error: {e}"
|
||||
|
|
@ -380,9 +360,7 @@ def _count_entry_tokens(
|
|||
|
||||
def _get_batch_job_cost_from_file_content(
|
||||
file_content_dictionary: List[dict],
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"
|
||||
] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
|
||||
model_info: Optional[ModelInfo] = None,
|
||||
) -> float:
|
||||
"""
|
||||
|
|
@ -393,9 +371,7 @@ def _get_batch_job_cost_from_file_content(
|
|||
try:
|
||||
total_cost: float = 0.0
|
||||
# parse the file content as json
|
||||
verbose_logger.debug(
|
||||
"file_content_dictionary=%s", json.dumps(file_content_dictionary, indent=4)
|
||||
)
|
||||
verbose_logger.debug("file_content_dictionary=%s", json.dumps(file_content_dictionary, indent=4))
|
||||
for _item in file_content_dictionary:
|
||||
if _batch_response_was_successful(_item):
|
||||
_response_body = _get_response_from_batch_job_output_file(_item)
|
||||
|
|
@ -424,9 +400,7 @@ def _get_batch_job_cost_from_file_content(
|
|||
|
||||
def _get_batch_job_total_usage_from_file_content(
|
||||
file_content_dictionary: List[dict],
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"
|
||||
] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
|
||||
model_name: Optional[str] = None,
|
||||
) -> Usage:
|
||||
"""
|
||||
|
|
@ -437,9 +411,7 @@ def _get_batch_job_total_usage_from_file_content(
|
|||
and model_name
|
||||
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
|
||||
):
|
||||
_, batch_usage = calculate_vertex_ai_batch_cost_and_usage(
|
||||
file_content_dictionary, model_name
|
||||
)
|
||||
_, batch_usage = calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
|
||||
return batch_usage
|
||||
|
||||
# For other providers, use the existing logic
|
||||
|
|
@ -488,11 +460,7 @@ def _count_prompt_or_input_tokens(model: str, value: Any) -> int:
|
|||
# Nested pre-tokenized prompt: every int contributes a
|
||||
# token. Mixed string/int items still count.
|
||||
total += sum(1 if isinstance(t, int) else 0 for t in chunk)
|
||||
total += sum(
|
||||
token_counter(model=model, text=t)
|
||||
for t in chunk
|
||||
if isinstance(t, str)
|
||||
)
|
||||
total += sum(token_counter(model=model, text=t) for t in chunk if isinstance(t, str))
|
||||
return total
|
||||
return 0
|
||||
|
||||
|
|
|
|||
|
|
@ -79,11 +79,7 @@ def _resolve_timeout(
|
|||
Returns:
|
||||
Resolved timeout as float
|
||||
"""
|
||||
timeout = (
|
||||
optional_params.timeout
|
||||
or kwargs.get("request_timeout", default_timeout)
|
||||
or default_timeout
|
||||
)
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", default_timeout) or default_timeout
|
||||
|
||||
# Handle httpx.Timeout objects
|
||||
if isinstance(timeout, httpx.Timeout):
|
||||
|
|
@ -109,9 +105,7 @@ async def acreate_batch(
|
|||
completion_window: Literal["24h"],
|
||||
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions"],
|
||||
input_file_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"
|
||||
] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"] = "openai",
|
||||
metadata: Optional[Dict[str, str]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_body: Optional[Dict[str, str]] = None,
|
||||
|
|
@ -161,9 +155,7 @@ def create_batch(
|
|||
completion_window: Literal["24h"],
|
||||
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions"],
|
||||
input_file_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"
|
||||
] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"] = "openai",
|
||||
metadata: Optional[Dict[str, str]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_body: Optional[Dict[str, str]] = None,
|
||||
|
|
@ -194,9 +186,7 @@ def create_batch(
|
|||
|
||||
_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)
|
||||
)
|
||||
litellm_logging_obj: LiteLLMLoggingObj = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj", None))
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = _resolve_timeout(optional_params, kwargs, custom_llm_provider)
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
|
|
@ -224,9 +214,7 @@ def create_batch(
|
|||
extra_body=extra_body,
|
||||
)
|
||||
if output_expires_after is not None:
|
||||
_create_batch_request["output_expires_after"] = cast(
|
||||
FileExpiresAfter, output_expires_after
|
||||
)
|
||||
_create_batch_request["output_expires_after"] = cast(FileExpiresAfter, output_expires_after)
|
||||
if model is not None:
|
||||
provider_config = ProviderConfigManager.get_provider_batches_config(
|
||||
model=model,
|
||||
|
|
@ -244,12 +232,7 @@ def create_batch(
|
|||
api_key=optional_params.api_key,
|
||||
logging_obj=litellm_logging_obj,
|
||||
_is_async=_is_async,
|
||||
client=(
|
||||
client
|
||||
if client is not None
|
||||
and isinstance(client, (HTTPHandler, AsyncHTTPHandler))
|
||||
else None
|
||||
),
|
||||
client=(client if client is not None and isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None),
|
||||
timeout=timeout,
|
||||
model=model,
|
||||
)
|
||||
|
|
@ -288,16 +271,8 @@ def create_batch(
|
|||
_is_async=_is_async,
|
||||
)
|
||||
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_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_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -326,18 +301,12 @@ def create_batch(
|
|||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or ""
|
||||
vertex_ai_project = (
|
||||
optional_params.vertex_project
|
||||
or litellm.vertex_project
|
||||
or get_secret_str("VERTEXAI_PROJECT")
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
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"
|
||||
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")
|
||||
|
||||
response = vertex_ai_batches_instance.create_batch(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -351,17 +320,13 @@ def create_batch(
|
|||
)
|
||||
else:
|
||||
raise litellm.exceptions.BadRequestError(
|
||||
message="LiteLLM doesn't support custom_llm_provider={} for 'create_batch'".format(
|
||||
custom_llm_provider
|
||||
),
|
||||
message="LiteLLM doesn't support custom_llm_provider={} for 'create_batch'".format(custom_llm_provider),
|
||||
model="n/a",
|
||||
llm_provider=custom_llm_provider,
|
||||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_batch", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_batch", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -372,9 +337,7 @@ def create_batch(
|
|||
@client
|
||||
async def aretrieve_batch(
|
||||
batch_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"
|
||||
] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai",
|
||||
metadata: Optional[Dict[str, str]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_body: Optional[Dict[str, str]] = None,
|
||||
|
|
@ -420,9 +383,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
litellm_params: dict,
|
||||
_retrieve_batch_request: RetrieveBatchRequest,
|
||||
_is_async: bool,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"
|
||||
] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai",
|
||||
logging_obj: Optional[Any] = None,
|
||||
):
|
||||
api_base: Optional[str] = None
|
||||
|
|
@ -459,16 +420,8 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
max_retries=optional_params.max_retries,
|
||||
)
|
||||
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_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_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -497,18 +450,12 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or ""
|
||||
vertex_ai_project = (
|
||||
optional_params.vertex_project
|
||||
or litellm.vertex_project
|
||||
or get_secret_str("VERTEXAI_PROJECT")
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
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"
|
||||
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")
|
||||
|
||||
response = vertex_ai_batches_instance.retrieve_batch(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -528,12 +475,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
or get_secret_str("ANTHROPIC_API_BASE")
|
||||
or get_secret_str("ANTHROPIC_BASE_URL")
|
||||
)
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
or litellm.api_key
|
||||
or litellm.azure_key
|
||||
or get_secret_str("ANTHROPIC_API_KEY")
|
||||
)
|
||||
api_key = optional_params.api_key or litellm.api_key or litellm.azure_key or get_secret_str("ANTHROPIC_API_KEY")
|
||||
|
||||
response = anthropic_batches_instance.retrieve_batch(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -555,9 +497,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="retrieve_batch", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="retrieve_batch", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -566,9 +506,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
@client
|
||||
def retrieve_batch(
|
||||
batch_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"
|
||||
] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai",
|
||||
metadata: Optional[Dict[str, str]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_body: Optional[Dict[str, str]] = None,
|
||||
|
|
@ -581,9 +519,7 @@ def retrieve_batch(
|
|||
"""
|
||||
try:
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get(
|
||||
"litellm_logging_obj", None
|
||||
)
|
||||
litellm_logging_obj: Optional[LiteLLMLoggingObj] = 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(
|
||||
|
|
@ -680,12 +616,7 @@ def retrieve_batch(
|
|||
function_id="batch_retrieve",
|
||||
),
|
||||
_is_async=_is_async,
|
||||
client=(
|
||||
client
|
||||
if client is not None
|
||||
and isinstance(client, (HTTPHandler, AsyncHTTPHandler))
|
||||
else None
|
||||
),
|
||||
client=(client if client is not None and isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None),
|
||||
timeout=timeout,
|
||||
model=model,
|
||||
)
|
||||
|
|
@ -823,16 +754,8 @@ def list_batches(
|
|||
max_retries=optional_params.max_retries,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base
|
||||
or litellm.api_base
|
||||
or get_secret_str("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_version = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret_str("AZURE_API_VERSION")
|
||||
)
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -860,18 +783,12 @@ def list_batches(
|
|||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or ""
|
||||
vertex_ai_project = (
|
||||
optional_params.vertex_project
|
||||
or litellm.vertex_project
|
||||
or get_secret_str("VERTEXAI_PROJECT")
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
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"
|
||||
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")
|
||||
|
||||
response = vertex_ai_batches_instance.list_batches(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -895,9 +812,7 @@ def list_batches(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -1014,17 +929,9 @@ def cancel_batch(
|
|||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
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")
|
||||
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")
|
||||
|
||||
response = openai_batches_instance.cancel_batch(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -1036,16 +943,8 @@ def cancel_batch(
|
|||
max_retries=optional_params.max_retries,
|
||||
)
|
||||
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_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_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -1074,18 +973,12 @@ def cancel_batch(
|
|||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or None
|
||||
vertex_ai_project = (
|
||||
optional_params.vertex_project
|
||||
or litellm.vertex_project
|
||||
or get_secret_str("VERTEXAI_PROJECT")
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
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"
|
||||
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")
|
||||
|
||||
response = vertex_ai_batches_instance.cancel_batch(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -1107,9 +1000,7 @@ def cancel_batch(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="cancel_batch", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="cancel_batch", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -1117,9 +1008,7 @@ def cancel_batch(
|
|||
raise e
|
||||
|
||||
|
||||
def _handle_async_invoke_status(
|
||||
batch_id: str, aws_region_name: str, logging_obj=None, **kwargs
|
||||
) -> "LiteLLMBatch":
|
||||
def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj=None, **kwargs) -> "LiteLLMBatch":
|
||||
"""
|
||||
Handle async invoke status check for AWS Bedrock.
|
||||
|
||||
|
|
@ -1168,9 +1057,7 @@ def _handle_async_invoke_status(
|
|||
# Get output S3 URI safely
|
||||
output_s3_uri = ""
|
||||
try:
|
||||
output_s3_uri = status_response["outputDataConfig"]["s3OutputDataConfig"][
|
||||
"s3Uri"
|
||||
]
|
||||
output_s3_uri = status_response["outputDataConfig"]["s3OutputDataConfig"]["s3Uri"]
|
||||
except (KeyError, TypeError):
|
||||
pass
|
||||
|
||||
|
|
@ -1186,15 +1073,12 @@ def _handle_async_invoke_status(
|
|||
failed_at,
|
||||
_,
|
||||
_,
|
||||
) = BedrockBatchesConfig()._parse_timestamps_and_status(
|
||||
status_response, aws_status_raw
|
||||
)
|
||||
) = BedrockBatchesConfig()._parse_timestamps_and_status(status_response, aws_status_raw)
|
||||
result = LiteLLMBatch(
|
||||
id=status_response["invocationArn"],
|
||||
object="batch",
|
||||
status=normalized_status,
|
||||
created_at=created_at
|
||||
or int(time.time()), # Provide default timestamp if None
|
||||
created_at=created_at or int(time.time()), # Provide default timestamp if None
|
||||
in_progress_at=in_progress_at,
|
||||
completed_at=completed_at,
|
||||
failed_at=failed_at,
|
||||
|
|
|
|||
|
|
@ -62,9 +62,7 @@ class BudgetManager:
|
|||
# Load the user_dict from hosted db
|
||||
url = self.api_base + "/get_budget"
|
||||
data = {"project_name": self.project_name}
|
||||
response = litellm.module_level_client.post(
|
||||
url, headers=self.headers, json=data
|
||||
)
|
||||
response = litellm.module_level_client.post(url, headers=self.headers, json=data)
|
||||
response = response.json()
|
||||
if response["status"] == "error":
|
||||
self.user_dict = {} # assume this means the user dict hasn't been stored yet
|
||||
|
|
@ -91,9 +89,7 @@ class BudgetManager:
|
|||
elif duration == "yearly":
|
||||
duration_in_days = DAYS_IN_A_YEAR
|
||||
else:
|
||||
raise ValueError(
|
||||
"""duration needs to be one of ["daily", "weekly", "monthly", "yearly"]"""
|
||||
)
|
||||
raise ValueError("""duration needs to be one of ["daily", "weekly", "monthly", "yearly"]""")
|
||||
self.user_dict[user] = {
|
||||
"total_budget": total_budget,
|
||||
"duration": duration_in_days,
|
||||
|
|
@ -106,9 +102,7 @@ class BudgetManager:
|
|||
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)
|
||||
prompt_cost, _ = litellm.cost_per_token(
|
||||
model=model, prompt_tokens=prompt_tokens, completion_tokens=0
|
||||
)
|
||||
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
|
||||
return projected_cost
|
||||
|
|
@ -125,12 +119,8 @@ class BudgetManager:
|
|||
output_text: Optional[str] = None,
|
||||
):
|
||||
if model and input_text and output_text:
|
||||
prompt_tokens = litellm.token_counter(
|
||||
model=model, messages=[{"role": "user", "content": input_text}]
|
||||
)
|
||||
completion_tokens = litellm.token_counter(
|
||||
model=model, messages=[{"role": "user", "content": output_text}]
|
||||
)
|
||||
prompt_tokens = litellm.token_counter(model=model, messages=[{"role": "user", "content": input_text}])
|
||||
completion_tokens = litellm.token_counter(model=model, messages=[{"role": "user", "content": output_text}])
|
||||
(
|
||||
prompt_tokens_cost_usd_dollar,
|
||||
completion_tokens_cost_usd_dollar,
|
||||
|
|
@ -142,21 +132,15 @@ class BudgetManager:
|
|||
cost = prompt_tokens_cost_usd_dollar + completion_tokens_cost_usd_dollar
|
||||
elif completion_obj:
|
||||
cost = litellm.completion_cost(completion_response=completion_obj)
|
||||
model = completion_obj[
|
||||
"model"
|
||||
] # if this throws an error try, model = completion_obj['model']
|
||||
model = completion_obj["model"] # if this throws an error try, model = completion_obj['model']
|
||||
else:
|
||||
raise ValueError(
|
||||
"Either a chat completion object or the text response needs to be passed in. Learn more - https://docs.litellm.ai/docs/budget_manager"
|
||||
)
|
||||
|
||||
self.user_dict[user]["current_cost"] = cost + self.user_dict[user].get(
|
||||
"current_cost", 0
|
||||
)
|
||||
self.user_dict[user]["current_cost"] = cost + self.user_dict[user].get("current_cost", 0)
|
||||
if "model_cost" in self.user_dict[user]:
|
||||
self.user_dict[user]["model_cost"][model] = cost + self.user_dict[user][
|
||||
"model_cost"
|
||||
].get(model, 0)
|
||||
self.user_dict[user]["model_cost"][model] = cost + self.user_dict[user]["model_cost"].get(model, 0)
|
||||
else:
|
||||
self.user_dict[user]["model_cost"] = {model: cost}
|
||||
|
||||
|
|
@ -198,9 +182,7 @@ class BudgetManager:
|
|||
current_time = 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 = 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:
|
||||
|
|
@ -215,9 +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 = threading.Thread(target=self.save_data) # [Non-Blocking]: saves data without blocking execution
|
||||
thread.start()
|
||||
|
||||
def save_data(self):
|
||||
|
|
@ -226,15 +206,11 @@ class BudgetManager:
|
|||
|
||||
# save the user dict
|
||||
with open("user_cost.json", "w") as json_file:
|
||||
json.dump(
|
||||
self.user_dict, json_file, indent=4
|
||||
) # Indent for pretty formatting
|
||||
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}
|
||||
response = litellm.module_level_client.post(
|
||||
url, headers=self.headers, json=data
|
||||
)
|
||||
response = litellm.module_level_client.post(url, headers=self.headers, json=data)
|
||||
response = response.json()
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -27,9 +27,7 @@ def resolve_embedding_router(
|
|||
if llm_router is None:
|
||||
return None
|
||||
router_model_names: list[str] = (
|
||||
[m["model_name"] for m in llm_model_list if "model_name" in m]
|
||||
if llm_model_list is not None
|
||||
else []
|
||||
[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:
|
||||
return llm_router
|
||||
|
|
|
|||
|
|
@ -52,9 +52,7 @@ class AzureBlobCache(BaseCache):
|
|||
print_verbose(f"LiteLLM SET Cache - Azure Blob. Key={key}. Value={value}")
|
||||
serialized_value = json.dumps(value)
|
||||
try:
|
||||
await self.async_container_client.upload_blob(
|
||||
key, serialized_value, overwrite=True
|
||||
)
|
||||
await self.async_container_client.upload_blob(key, serialized_value, overwrite=True)
|
||||
except Exception as e:
|
||||
# NON blocking - notify users Azure Blob is throwing an exception
|
||||
print_verbose(f"LiteLLM set_cache() - Got exception from Azure Blob: {e}")
|
||||
|
|
|
|||
|
|
@ -171,9 +171,7 @@ class Cache:
|
|||
# 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")
|
||||
if _env_cluster_nodes is not None and isinstance(
|
||||
_env_cluster_nodes, str
|
||||
):
|
||||
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:
|
||||
|
|
@ -271,7 +269,9 @@ class Cache:
|
|||
litellm.logging_callback_manager.add_litellm_success_callback("cache")
|
||||
if "cache" not in litellm._async_success_callback:
|
||||
litellm.logging_callback_manager.add_litellm_async_success_callback("cache")
|
||||
self.supported_call_types = supported_call_types # default to ["completion", "acompletion", "embedding", "aembedding"]
|
||||
self.supported_call_types = (
|
||||
supported_call_types # default to ["completion", "acompletion", "embedding", "aembedding"]
|
||||
)
|
||||
self.type = type
|
||||
self.namespace = namespace
|
||||
self.redis_flush_size = redis_flush_size
|
||||
|
|
@ -294,9 +294,7 @@ class Cache:
|
|||
# Params whose values carry prompt content. Excluded from semantic-cache
|
||||
# scope keys so differently worded prompts share a bucket and match via
|
||||
# vector similarity rather than being split into per-wording buckets.
|
||||
_SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS: frozenset = frozenset(
|
||||
{"messages", "prompt", "input"}
|
||||
)
|
||||
_SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS: frozenset = frozenset({"messages", "prompt", "input"})
|
||||
|
||||
# Server-set identity (from proxy auth) used to isolate semantic-cache
|
||||
# buckets per tenant. Required once the prompt is out of the scope key, so a
|
||||
|
|
@ -349,11 +347,7 @@ class Cache:
|
|||
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()
|
||||
)
|
||||
scope_excluded_params = self._SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS if is_semantic_cache else frozenset()
|
||||
for param in kwargs:
|
||||
if param in scope_excluded_params:
|
||||
continue
|
||||
|
|
@ -361,12 +355,8 @@ class Cache:
|
|||
param_value: Optional[str] = self._get_param_value(param, kwargs)
|
||||
if param_value is not None:
|
||||
cache_key += f"{str(param)}: {str(param_value)}"
|
||||
elif (
|
||||
param not in litellm_param_kwargs
|
||||
): # check if user passed in optional param - e.g. top_k
|
||||
if (
|
||||
litellm.enable_caching_on_provider_specific_optional_params is True
|
||||
): # feature flagged for now
|
||||
elif param not in litellm_param_kwargs: # check if user passed in optional param - e.g. top_k
|
||||
if litellm.enable_caching_on_provider_specific_optional_params is True: # feature flagged for now
|
||||
if kwargs[param] is None:
|
||||
continue # ignore None params
|
||||
param_value = kwargs[param]
|
||||
|
|
@ -385,9 +375,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"}
|
||||
self._set_preset_cache_key_in_kwargs(
|
||||
preset_cache_key=hashed_cache_key, **kwargs_for_preset
|
||||
)
|
||||
self._set_preset_cache_key_in_kwargs(preset_cache_key=hashed_cache_key, **kwargs_for_preset)
|
||||
return hashed_cache_key
|
||||
|
||||
def _get_param_value(
|
||||
|
|
@ -415,15 +403,11 @@ class Cache:
|
|||
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: Optional[str] = metadata.get(
|
||||
"model_group"
|
||||
) or metadata_in_litellm_params.get("model_group")
|
||||
model_group: Optional[str] = metadata.get("model_group") or metadata_in_litellm_params.get("model_group")
|
||||
caching_group = 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: Optional[str]
|
||||
) -> Optional[str]:
|
||||
def _get_caching_group(self, metadata: dict, model_group: Optional[str]) -> Optional[str]:
|
||||
caching_groups: Optional[List] = metadata.get("caching_groups", [])
|
||||
if caching_groups:
|
||||
for group in caching_groups:
|
||||
|
|
@ -503,11 +487,7 @@ class Cache:
|
|||
"""
|
||||
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
|
||||
)
|
||||
namespace = 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)
|
||||
|
|
@ -537,11 +517,7 @@ class Cache:
|
|||
Common get cache logic across sync + async implementations
|
||||
"""
|
||||
# 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
|
||||
):
|
||||
if cached_result is not None and isinstance(cached_result, dict) and "timestamp" in cached_result:
|
||||
timestamp = cached_result["timestamp"]
|
||||
current_time = time.time()
|
||||
|
||||
|
|
@ -586,15 +562,11 @@ class Cache:
|
|||
) -> None:
|
||||
original_metadata = original_kwargs.get("metadata")
|
||||
cache_lookup_metadata = cache_lookup_kwargs.get("metadata")
|
||||
if not isinstance(original_metadata, dict) or not isinstance(
|
||||
cache_lookup_metadata, dict
|
||||
):
|
||||
if not isinstance(original_metadata, dict) or not isinstance(cache_lookup_metadata, dict):
|
||||
return
|
||||
|
||||
if "semantic-similarity" in cache_lookup_metadata:
|
||||
original_metadata["semantic-similarity"] = cache_lookup_metadata[
|
||||
"semantic-similarity"
|
||||
]
|
||||
original_metadata["semantic-similarity"] = cache_lookup_metadata["semantic-similarity"]
|
||||
|
||||
def get_cache(self, dynamic_cache_object: Optional[BaseCache] = None, **kwargs):
|
||||
"""
|
||||
|
|
@ -616,34 +588,22 @@ class Cache:
|
|||
cache_key = self.get_cache_key(**kwargs)
|
||||
if cache_key is not None:
|
||||
cache_control_args: DynamicCacheControl = kwargs.get("cache", {})
|
||||
max_age = (
|
||||
cache_control_args.get("s-maxage")
|
||||
or cache_control_args.get("s-max-age")
|
||||
or float("inf")
|
||||
)
|
||||
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)
|
||||
if dynamic_cache_object is not None:
|
||||
cached_result = dynamic_cache_object.get_cache(
|
||||
cache_key, **cache_lookup_kwargs
|
||||
)
|
||||
cached_result = dynamic_cache_object.get_cache(cache_key, **cache_lookup_kwargs)
|
||||
else:
|
||||
cached_result = self.cache.get_cache(
|
||||
cache_key, **cache_lookup_kwargs
|
||||
)
|
||||
cached_result = self.cache.get_cache(cache_key, **cache_lookup_kwargs)
|
||||
self._update_metadata_from_cache_lookup_kwargs(
|
||||
original_kwargs=kwargs,
|
||||
cache_lookup_kwargs=cache_lookup_kwargs,
|
||||
)
|
||||
return self._get_cache_logic(
|
||||
cached_result=cached_result, max_age=max_age
|
||||
)
|
||||
return self._get_cache_logic(cached_result=cached_result, max_age=max_age)
|
||||
except Exception:
|
||||
print_verbose(f"An exception occurred: {traceback.format_exc()}")
|
||||
return None
|
||||
|
||||
async def async_get_cache(
|
||||
self, dynamic_cache_object: Optional[BaseCache] = None, **kwargs
|
||||
):
|
||||
async def async_get_cache(self, dynamic_cache_object: Optional[BaseCache] = None, **kwargs):
|
||||
"""
|
||||
Async get cache implementation.
|
||||
|
||||
|
|
@ -660,20 +620,12 @@ class Cache:
|
|||
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"))
|
||||
)
|
||||
max_age = 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
|
||||
)
|
||||
cached_result = await dynamic_cache_object.async_get_cache(cache_key, **kwargs)
|
||||
else:
|
||||
cached_result = await self.cache.async_get_cache(
|
||||
cache_key, **kwargs
|
||||
)
|
||||
return self._get_cache_logic(
|
||||
cached_result=cached_result, max_age=max_age
|
||||
)
|
||||
cached_result = await self.cache.async_get_cache(cache_key, **kwargs)
|
||||
return self._get_cache_logic(cached_result=cached_result, max_age=max_age)
|
||||
except Exception:
|
||||
print_verbose(f"An exception occurred: {traceback.format_exc()}")
|
||||
return None
|
||||
|
|
@ -722,16 +674,12 @@ class Cache:
|
|||
try:
|
||||
if self.should_use_cache(**kwargs) is not True:
|
||||
return
|
||||
cache_key, cached_data, kwargs = self._add_cache_logic(
|
||||
result=result, **kwargs
|
||||
)
|
||||
cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs)
|
||||
self.cache.set_cache(cache_key, cached_data, **kwargs)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"LiteLLM Cache: Excepton add_cache: {str(e)}")
|
||||
|
||||
async def async_add_cache(
|
||||
self, result, dynamic_cache_object: Optional[BaseCache] = None, **kwargs
|
||||
):
|
||||
async def async_add_cache(self, result, dynamic_cache_object: Optional[BaseCache] = None, **kwargs):
|
||||
"""
|
||||
Async implementation of add_cache
|
||||
"""
|
||||
|
|
@ -742,13 +690,9 @@ class Cache:
|
|||
# high traffic - fill in results in memory and then flush
|
||||
await self.batch_cache_write(result, **kwargs)
|
||||
else:
|
||||
cache_key, cached_data, kwargs = self._add_cache_logic(
|
||||
result=result, **kwargs
|
||||
)
|
||||
cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs)
|
||||
if dynamic_cache_object is not None:
|
||||
await dynamic_cache_object.async_set_cache(
|
||||
cache_key, cached_data, **kwargs
|
||||
)
|
||||
await dynamic_cache_object.async_set_cache(cache_key, cached_data, **kwargs)
|
||||
else:
|
||||
await self.cache.async_set_cache(cache_key, cached_data, **kwargs)
|
||||
except Exception as e:
|
||||
|
|
@ -899,9 +843,7 @@ class Cache:
|
|||
)
|
||||
return cache_key, cached_data, kwargs
|
||||
|
||||
async def async_add_cache_pipeline(
|
||||
self, result, dynamic_cache_object: Optional[BaseCache] = None, **kwargs
|
||||
):
|
||||
async def async_add_cache_pipeline(self, result, dynamic_cache_object: Optional[BaseCache] = None, **kwargs):
|
||||
"""
|
||||
Async implementation of add_cache for Embedding calls
|
||||
|
||||
|
|
@ -925,19 +867,13 @@ class Cache:
|
|||
) = self.add_embedding_response_to_cache(result, i, kwargs, idx)
|
||||
cache_list.append((cache_key, cached_data))
|
||||
elif isinstance(kwargs["input"], str):
|
||||
cache_key, cached_data, kwargs = self.add_embedding_response_to_cache(
|
||||
result, kwargs["input"], kwargs
|
||||
)
|
||||
cache_key, cached_data, kwargs = self.add_embedding_response_to_cache(result, kwargs["input"], kwargs)
|
||||
cache_list.append((cache_key, cached_data))
|
||||
|
||||
if dynamic_cache_object is not None:
|
||||
await dynamic_cache_object.async_set_cache_pipeline(
|
||||
cache_list=cache_list, **kwargs
|
||||
)
|
||||
await dynamic_cache_object.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
|
||||
else:
|
||||
await self.cache.async_set_cache_pipeline(
|
||||
cache_list=cache_list, **kwargs
|
||||
)
|
||||
await self.cache.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"LiteLLM Cache: Excepton add_cache: {str(e)}")
|
||||
|
||||
|
|
|
|||
|
|
@ -163,11 +163,9 @@ class LLMCachingHandler:
|
|||
"""
|
||||
# Check if caching should be performed BEFORE doing expensive operations
|
||||
if (
|
||||
(
|
||||
(kwargs.get("caching", None) is None and litellm.cache is not None)
|
||||
or kwargs.get("caching", False) is True
|
||||
)
|
||||
and (kwargs.get("cache", {}).get("no-cache", False) is not True)
|
||||
(kwargs.get("caching", None) is None and litellm.cache is not None) or kwargs.get("caching", False) is True
|
||||
) and (
|
||||
kwargs.get("cache", {}).get("no-cache", False) is not True
|
||||
): # allow users to control returning cached responses from the completion function
|
||||
args = args or ()
|
||||
final_embedding_cached_response: Optional[EmbeddingResponse] = None
|
||||
|
|
@ -183,9 +181,7 @@ class LLMCachingHandler:
|
|||
parent_otel_span = _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
|
||||
):
|
||||
if litellm.cache is not None and self._is_call_type_supported_by_cache(original_function=original_function):
|
||||
verbose_logger.debug("Checking Async Cache")
|
||||
cached_result = await self._retrieve_from_cache(
|
||||
call_type=call_type,
|
||||
|
|
@ -204,9 +200,7 @@ class LLMCachingHandler:
|
|||
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 = (cache_check_end_time - cache_check_start_time) * 1000
|
||||
self._update_litellm_logging_obj_environment(
|
||||
logging_obj=logging_obj,
|
||||
model=model,
|
||||
|
|
@ -250,9 +244,7 @@ class LLMCachingHandler:
|
|||
and cached_result is not None
|
||||
and isinstance(cached_result, list)
|
||||
and litellm.cache is not None
|
||||
and not isinstance(
|
||||
litellm.cache.cache, S3Cache
|
||||
) # s3 doesn't support bulk writing. Exclude.
|
||||
and not isinstance(litellm.cache.cache, S3Cache) # s3 doesn't support bulk writing. Exclude.
|
||||
):
|
||||
(
|
||||
final_embedding_cached_response,
|
||||
|
|
@ -291,9 +283,7 @@ class LLMCachingHandler:
|
|||
cached_result: Optional[Any] = None
|
||||
|
||||
# Check if caching should be performed BEFORE doing expensive kwargs copy
|
||||
if litellm.cache is not None and self._is_call_type_supported_by_cache(
|
||||
original_function=original_function
|
||||
):
|
||||
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()
|
||||
|
|
@ -376,9 +366,7 @@ class LLMCachingHandler:
|
|||
else:
|
||||
raise ValueError("input must be a string or a list")
|
||||
|
||||
def _extract_model_from_cached_results(
|
||||
self, non_null_list: List[Tuple[int, CachedEmbedding]]
|
||||
) -> Optional[str]:
|
||||
def _extract_model_from_cached_results(self, non_null_list: List[Tuple[int, CachedEmbedding]]) -> Optional[str]:
|
||||
"""
|
||||
Helper method to extract the model name from cached results.
|
||||
|
||||
|
|
@ -461,9 +449,7 @@ class LLMCachingHandler:
|
|||
elif isinstance(kwargs_input_as_list[idx], str):
|
||||
from litellm.utils import token_counter
|
||||
|
||||
prompt_tokens += token_counter(
|
||||
text=kwargs_input_as_list[idx], count_response_tokens=True
|
||||
)
|
||||
prompt_tokens += token_counter(text=kwargs_input_as_list[idx], count_response_tokens=True)
|
||||
# Aggregate prompt_tokens_details from cached items
|
||||
item_details = cr.get("prompt_tokens_details")
|
||||
if item_details:
|
||||
|
|
@ -471,9 +457,7 @@ class LLMCachingHandler:
|
|||
aggregated_details = {}
|
||||
for key, value in item_details.items():
|
||||
if isinstance(value, (int, float)):
|
||||
aggregated_details[key] = (
|
||||
aggregated_details.get(key, 0) + value
|
||||
)
|
||||
aggregated_details[key] = aggregated_details.get(key, 0) + value
|
||||
else:
|
||||
aggregated_details[key] = value
|
||||
|
||||
|
|
@ -483,9 +467,7 @@ class LLMCachingHandler:
|
|||
from litellm.types.utils import PromptTokensDetailsWrapper
|
||||
|
||||
try:
|
||||
prompt_tokens_details = PromptTokensDetailsWrapper(
|
||||
**aggregated_details
|
||||
)
|
||||
prompt_tokens_details = PromptTokensDetailsWrapper(**aggregated_details)
|
||||
except Exception:
|
||||
prompt_tokens_details = None
|
||||
usage = Usage(
|
||||
|
|
@ -554,16 +536,8 @@ 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 = 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 {}
|
||||
|
||||
merged: dict = {}
|
||||
for key in set(dict1.keys()) | set(dict2.keys()):
|
||||
|
|
@ -632,9 +606,7 @@ class LLMCachingHandler:
|
|||
final_data_list.append(item)
|
||||
|
||||
_caching_handler_response.final_embedding_cached_response.data = final_data_list
|
||||
_caching_handler_response.final_embedding_cached_response._hidden_params[
|
||||
"cache_hit"
|
||||
] = True
|
||||
_caching_handler_response.final_embedding_cached_response._hidden_params["cache_hit"] = True
|
||||
_caching_handler_response.final_embedding_cached_response._response_ms = (
|
||||
end_time - start_time
|
||||
).total_seconds() * 1000
|
||||
|
|
@ -730,9 +702,7 @@ class LLMCachingHandler:
|
|||
raise ValueError("input must be a string or a list")
|
||||
tasks = []
|
||||
for idx, i in enumerate(new_kwargs["input"]):
|
||||
preset_cache_key = litellm.cache.get_cache_key(
|
||||
**{**new_kwargs, "input": i}
|
||||
)
|
||||
preset_cache_key = litellm.cache.get_cache_key(**{**new_kwargs, "input": i})
|
||||
tasks.append(
|
||||
litellm.cache.async_get_cache(
|
||||
cache_key=preset_cache_key,
|
||||
|
|
@ -750,18 +720,14 @@ class LLMCachingHandler:
|
|||
request_cache_key = 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)
|
||||
)
|
||||
self.preset_cache_key = request_cache_key or litellm.cache.get_cache_key(**request_kwargs)
|
||||
cached_result = await litellm.cache.async_get_cache(
|
||||
dynamic_cache_object=self.dual_cache,
|
||||
cache_key=self.preset_cache_key,
|
||||
**request_kwargs,
|
||||
)
|
||||
else: # fallback for caches that don't support async
|
||||
self.preset_cache_key = (
|
||||
request_cache_key or litellm.cache.get_cache_key(**request_kwargs)
|
||||
)
|
||||
self.preset_cache_key = request_cache_key or litellm.cache.get_cache_key(**request_kwargs)
|
||||
cached_result = litellm.cache.get_cache(
|
||||
dynamic_cache_object=self.dual_cache,
|
||||
cache_key=self.preset_cache_key,
|
||||
|
|
@ -808,10 +774,9 @@ class LLMCachingHandler:
|
|||
"""
|
||||
from litellm.utils import convert_to_model_response_object
|
||||
|
||||
if (
|
||||
call_type == CallTypes.acompletion.value
|
||||
or call_type == CallTypes.completion.value
|
||||
) and isinstance(cached_result, dict):
|
||||
if (call_type == CallTypes.acompletion.value or call_type == CallTypes.completion.value) and isinstance(
|
||||
cached_result, dict
|
||||
):
|
||||
if kwargs.get("stream", False) is True:
|
||||
cached_result = self._convert_cached_stream_response(
|
||||
cached_result=cached_result,
|
||||
|
|
@ -825,8 +790,7 @@ class LLMCachingHandler:
|
|||
model_response_object=ModelResponse(),
|
||||
)
|
||||
if (
|
||||
call_type == CallTypes.atext_completion.value
|
||||
or call_type == CallTypes.text_completion.value
|
||||
call_type == CallTypes.atext_completion.value or call_type == CallTypes.text_completion.value
|
||||
) and isinstance(cached_result, dict):
|
||||
if kwargs.get("stream", False) is True:
|
||||
cached_result = self._convert_cached_stream_response(
|
||||
|
|
@ -837,28 +801,26 @@ class LLMCachingHandler:
|
|||
)
|
||||
else:
|
||||
cached_result = TextCompletionResponse(**cached_result)
|
||||
elif (
|
||||
call_type == CallTypes.aembedding.value
|
||||
or call_type == CallTypes.embedding.value
|
||||
) and isinstance(cached_result, dict):
|
||||
elif (call_type == CallTypes.aembedding.value or call_type == CallTypes.embedding.value) and isinstance(
|
||||
cached_result, dict
|
||||
):
|
||||
cached_result = convert_to_model_response_object(
|
||||
response_object=cached_result,
|
||||
model_response_object=EmbeddingResponse(),
|
||||
response_type="embedding",
|
||||
)
|
||||
|
||||
elif (
|
||||
call_type == CallTypes.arerank.value or call_type == CallTypes.rerank.value
|
||||
) and isinstance(cached_result, dict):
|
||||
elif (call_type == CallTypes.arerank.value or call_type == CallTypes.rerank.value) and isinstance(
|
||||
cached_result, dict
|
||||
):
|
||||
cached_result = convert_to_model_response_object(
|
||||
response_object=cached_result,
|
||||
model_response_object=None,
|
||||
response_type="rerank",
|
||||
)
|
||||
elif (
|
||||
call_type == CallTypes.atranscription.value
|
||||
or call_type == CallTypes.transcription.value
|
||||
) and isinstance(cached_result, dict):
|
||||
elif (call_type == CallTypes.atranscription.value or call_type == CallTypes.transcription.value) and isinstance(
|
||||
cached_result, dict
|
||||
):
|
||||
hidden_params = {
|
||||
"model": "whisper-1",
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
|
|
@ -870,16 +832,12 @@ class LLMCachingHandler:
|
|||
response_type="audio_transcription",
|
||||
hidden_params=hidden_params,
|
||||
)
|
||||
elif (call_type == "aresponses" or call_type == "responses") and isinstance(
|
||||
cached_result, dict
|
||||
):
|
||||
elif (call_type == "aresponses" or call_type == "responses") and isinstance(cached_result, dict):
|
||||
use_chat_completion_cache = _is_chat_completion_cached_dict(cached_result)
|
||||
if use_chat_completion_cache:
|
||||
if kwargs.get("stream", False) is True:
|
||||
bridge_call_type = (
|
||||
CallTypes.acompletion.value
|
||||
if call_type == "aresponses"
|
||||
else CallTypes.completion.value
|
||||
CallTypes.acompletion.value if call_type == "aresponses" else CallTypes.completion.value
|
||||
)
|
||||
cached_result = self._convert_cached_stream_response(
|
||||
cached_result=cached_result,
|
||||
|
|
@ -949,10 +907,7 @@ class LLMCachingHandler:
|
|||
)
|
||||
|
||||
_stream_cached_result: Union[AsyncGenerator, Generator]
|
||||
if (
|
||||
call_type == CallTypes.acompletion.value
|
||||
or call_type == CallTypes.atext_completion.value
|
||||
):
|
||||
if call_type == CallTypes.acompletion.value or call_type == CallTypes.atext_completion.value:
|
||||
_stream_cached_result = convert_to_streaming_response_async(
|
||||
response_object=cached_result,
|
||||
)
|
||||
|
|
@ -1005,9 +960,7 @@ class LLMCachingHandler:
|
|||
parent_otel_span = _get_parent_otel_span_from_kwargs(new_kwargs)
|
||||
new_kwargs["parent_otel_span"] = parent_otel_span
|
||||
# [OPTIONAL] ADD TO CACHE
|
||||
if self._should_store_result_in_cache(
|
||||
original_function=original_function, kwargs=new_kwargs
|
||||
):
|
||||
if self._should_store_result_in_cache(original_function=original_function, kwargs=new_kwargs):
|
||||
if (
|
||||
isinstance(result, litellm.ModelResponse)
|
||||
or isinstance(result, litellm.EmbeddingResponse)
|
||||
|
|
@ -1018,9 +971,7 @@ class LLMCachingHandler:
|
|||
if (
|
||||
isinstance(result, EmbeddingResponse)
|
||||
and litellm.cache is not None
|
||||
and not isinstance(
|
||||
litellm.cache.cache, S3Cache
|
||||
) # s3 doesn't support bulk writing. Exclude.
|
||||
and not isinstance(litellm.cache.cache, S3Cache) # s3 doesn't support bulk writing. Exclude.
|
||||
):
|
||||
asyncio.create_task(
|
||||
litellm.cache.async_add_cache_pipeline(
|
||||
|
|
@ -1058,16 +1009,12 @@ class LLMCachingHandler:
|
|||
if litellm.cache is None:
|
||||
return
|
||||
|
||||
if self._should_store_result_in_cache(
|
||||
original_function=self.original_function, kwargs=new_kwargs
|
||||
):
|
||||
if self._should_store_result_in_cache(original_function=self.original_function, kwargs=new_kwargs):
|
||||
litellm.cache.add_cache(result, **new_kwargs)
|
||||
|
||||
return
|
||||
|
||||
def _should_store_result_in_cache(
|
||||
self, original_function: Callable, kwargs: Dict[str, Any]
|
||||
) -> bool:
|
||||
def _should_store_result_in_cache(self, original_function: Callable, kwargs: Dict[str, Any]) -> bool:
|
||||
"""
|
||||
Helper function to determine if the result should be stored in the cache.
|
||||
|
||||
|
|
@ -1113,15 +1060,15 @@ class LLMCachingHandler:
|
|||
|
||||
"""
|
||||
|
||||
complete_streaming_response: Optional[
|
||||
Union[ModelResponse, TextCompletionResponse]
|
||||
] = _assemble_complete_response_from_streaming_chunks(
|
||||
result=processed_chunk,
|
||||
start_time=self.start_time,
|
||||
end_time=datetime.datetime.now(),
|
||||
request_kwargs=self.request_kwargs,
|
||||
streaming_chunks=self.async_streaming_chunks,
|
||||
is_async=True,
|
||||
complete_streaming_response: Optional[Union[ModelResponse, TextCompletionResponse]] = (
|
||||
_assemble_complete_response_from_streaming_chunks(
|
||||
result=processed_chunk,
|
||||
start_time=self.start_time,
|
||||
end_time=datetime.datetime.now(),
|
||||
request_kwargs=self.request_kwargs,
|
||||
streaming_chunks=self.async_streaming_chunks,
|
||||
is_async=True,
|
||||
)
|
||||
)
|
||||
# if a complete_streaming_response is assembled, add it to the cache
|
||||
if complete_streaming_response is not None:
|
||||
|
|
@ -1135,15 +1082,15 @@ class LLMCachingHandler:
|
|||
"""
|
||||
Sync internal method to add the streaming response to the cache
|
||||
"""
|
||||
complete_streaming_response: Optional[
|
||||
Union[ModelResponse, TextCompletionResponse]
|
||||
] = _assemble_complete_response_from_streaming_chunks(
|
||||
result=processed_chunk,
|
||||
start_time=self.start_time,
|
||||
end_time=datetime.datetime.now(),
|
||||
request_kwargs=self.request_kwargs,
|
||||
streaming_chunks=self.sync_streaming_chunks,
|
||||
is_async=False,
|
||||
complete_streaming_response: Optional[Union[ModelResponse, TextCompletionResponse]] = (
|
||||
_assemble_complete_response_from_streaming_chunks(
|
||||
result=processed_chunk,
|
||||
start_time=self.start_time,
|
||||
end_time=datetime.datetime.now(),
|
||||
request_kwargs=self.request_kwargs,
|
||||
streaming_chunks=self.sync_streaming_chunks,
|
||||
is_async=False,
|
||||
)
|
||||
)
|
||||
|
||||
# if a complete_streaming_response is assembled, add it to the cache
|
||||
|
|
@ -1191,9 +1138,7 @@ class LLMCachingHandler:
|
|||
}
|
||||
|
||||
if litellm.cache is not None:
|
||||
litellm_params["preset_cache_key"] = (
|
||||
litellm.cache._get_preset_cache_key_from_kwargs(**kwargs)
|
||||
)
|
||||
litellm_params["preset_cache_key"] = litellm.cache._get_preset_cache_key_from_kwargs(**kwargs)
|
||||
else:
|
||||
litellm_params["preset_cache_key"] = None
|
||||
|
||||
|
|
@ -1202,11 +1147,7 @@ class LLMCachingHandler:
|
|||
user=kwargs.get("user", None),
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
input=(
|
||||
kwargs.get("messages", "")
|
||||
if not is_embedding
|
||||
else kwargs.get("input", "")
|
||||
),
|
||||
input=(kwargs.get("messages", "") if not is_embedding else kwargs.get("input", "")),
|
||||
api_key=kwargs.get("api_key", None),
|
||||
original_response=str(cached_result),
|
||||
additional_args=None,
|
||||
|
|
|
|||
|
|
@ -16,9 +16,7 @@ class DiskCache(BaseCache):
|
|||
try:
|
||||
import diskcache as dc
|
||||
except ModuleNotFoundError as e:
|
||||
raise ModuleNotFoundError(
|
||||
"Please install litellm with `litellm[caching]` to use disk caching."
|
||||
) from e
|
||||
raise ModuleNotFoundError("Please install litellm with `litellm[caching]` to use disk caching.") from e
|
||||
|
||||
# if users don't provider one, use the default litellm cache
|
||||
if disk_cache_dir is None:
|
||||
|
|
|
|||
|
|
@ -69,23 +69,15 @@ class DualCache(BaseCache):
|
|||
self.in_memory_cache = in_memory_cache or InMemoryCache()
|
||||
# If redis_cache is not provided, use the default RedisCache
|
||||
self.redis_cache = redis_cache
|
||||
self.last_redis_batch_access_time = LimitedSizeOrderedDict(
|
||||
max_size=default_max_redis_batch_cache_size
|
||||
)
|
||||
self.last_redis_batch_access_time = LimitedSizeOrderedDict(max_size=default_max_redis_batch_cache_size)
|
||||
self._last_redis_batch_access_time_lock = Lock()
|
||||
self.redis_batch_cache_expiry = (
|
||||
default_redis_batch_cache_expiry
|
||||
or litellm.default_redis_batch_cache_expiry
|
||||
or 10
|
||||
)
|
||||
self.default_in_memory_ttl = (
|
||||
default_in_memory_ttl or litellm.default_in_memory_ttl
|
||||
default_redis_batch_cache_expiry or litellm.default_redis_batch_cache_expiry or 10
|
||||
)
|
||||
self.default_in_memory_ttl = default_in_memory_ttl or litellm.default_in_memory_ttl
|
||||
self.default_redis_ttl = default_redis_ttl or litellm.default_redis_ttl
|
||||
|
||||
def update_cache_ttl(
|
||||
self, default_in_memory_ttl: Optional[float], default_redis_ttl: Optional[float]
|
||||
):
|
||||
def update_cache_ttl(self, default_in_memory_ttl: Optional[float], default_redis_ttl: Optional[float]):
|
||||
if default_in_memory_ttl is not None:
|
||||
self.default_in_memory_ttl = default_in_memory_ttl
|
||||
|
||||
|
|
@ -125,9 +117,7 @@ class DualCache(BaseCache):
|
|||
except Exception as e:
|
||||
print_verbose(e)
|
||||
|
||||
def increment_cache(
|
||||
self, key, value: int, local_only: bool = False, **kwargs
|
||||
) -> int:
|
||||
def increment_cache(self, key, value: int, local_only: bool = False, **kwargs) -> int:
|
||||
"""
|
||||
Key - the key in cache
|
||||
|
||||
|
|
@ -166,9 +156,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 = self.redis_cache.get_cache(
|
||||
key, parent_otel_span=parent_otel_span
|
||||
)
|
||||
redis_result = 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
|
||||
|
|
@ -196,9 +184,7 @@ class DualCache(BaseCache):
|
|||
new_loop = asyncio.new_event_loop()
|
||||
try:
|
||||
asyncio.set_event_loop(new_loop)
|
||||
return new_loop.run_until_complete(
|
||||
self.async_batch_get_cache(**received_args)
|
||||
)
|
||||
return new_loop.run_until_complete(self.async_batch_get_cache(**received_args))
|
||||
finally:
|
||||
new_loop.close()
|
||||
asyncio.set_event_loop(None)
|
||||
|
|
@ -225,14 +211,10 @@ class DualCache(BaseCache):
|
|||
):
|
||||
# Try to fetch from in-memory cache first
|
||||
try:
|
||||
print_verbose(
|
||||
f"async get cache: cache key: {key}; local_only: {local_only}"
|
||||
)
|
||||
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 = 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:
|
||||
|
|
@ -240,15 +222,11 @@ 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 = 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
|
||||
await self.in_memory_cache.async_set_cache(
|
||||
key, redis_result, **kwargs
|
||||
)
|
||||
await self.in_memory_cache.async_set_cache(key, redis_result, **kwargs)
|
||||
|
||||
result = redis_result
|
||||
|
||||
|
|
@ -277,20 +255,15 @@ class DualCache(BaseCache):
|
|||
|
||||
if (
|
||||
key not in self.last_redis_batch_access_time
|
||||
or current_time - self.last_redis_batch_access_time[key]
|
||||
>= self.redis_batch_cache_expiry
|
||||
or current_time - self.last_redis_batch_access_time[key] >= self.redis_batch_cache_expiry
|
||||
):
|
||||
sublist_keys.append(key)
|
||||
previous_access_times[key] = self.last_redis_batch_access_time.get(
|
||||
key
|
||||
)
|
||||
previous_access_times[key] = self.last_redis_batch_access_time.get(key)
|
||||
self.last_redis_batch_access_time[key] = current_time
|
||||
|
||||
return sublist_keys, previous_access_times
|
||||
|
||||
def _rollback_redis_batch_key_reservations(
|
||||
self, previous_access_times: Dict[str, Optional[float]]
|
||||
) -> None:
|
||||
def _rollback_redis_batch_key_reservations(self, previous_access_times: Dict[str, Optional[float]]) -> None:
|
||||
with self._last_redis_batch_access_time_lock:
|
||||
for key, previous_time in previous_access_times.items():
|
||||
if previous_time is None:
|
||||
|
|
@ -308,9 +281,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 = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs)
|
||||
|
||||
if in_memory_result is not None:
|
||||
result = in_memory_result
|
||||
|
|
@ -321,9 +292,7 @@ class DualCache(BaseCache):
|
|||
- check the redis cache
|
||||
"""
|
||||
current_time = time.time()
|
||||
sublist_keys, previous_access_times = self._reserve_redis_batch_keys(
|
||||
current_time, keys, result
|
||||
)
|
||||
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:
|
||||
|
|
@ -334,15 +303,11 @@ class DualCache(BaseCache):
|
|||
)
|
||||
except Exception:
|
||||
# Do not throttle subsequent callers if the Redis read fails.
|
||||
self._rollback_redis_batch_key_reservations(
|
||||
previous_access_times
|
||||
)
|
||||
self._rollback_redis_batch_key_reservations(previous_access_times)
|
||||
raise
|
||||
|
||||
# Short-circuit if redis_result is None or contains only None values
|
||||
if redis_result is None or all(
|
||||
v is None for v in redis_result.values()
|
||||
):
|
||||
if redis_result is None or all(v is None for v in redis_result.values()):
|
||||
return result
|
||||
|
||||
# Pre-compute key-to-index mapping for O(1) lookup
|
||||
|
|
@ -353,18 +318,14 @@ class DualCache(BaseCache):
|
|||
result[key_to_index[key]] = value
|
||||
|
||||
if value is not None and self.in_memory_cache is not None:
|
||||
await self.in_memory_cache.async_set_cache(
|
||||
key, value, **kwargs
|
||||
)
|
||||
await self.in_memory_cache.async_set_cache(key, value, **kwargs)
|
||||
|
||||
return result
|
||||
except Exception:
|
||||
verbose_logger.error(traceback.format_exc())
|
||||
|
||||
async def async_set_cache(self, key, value, local_only: bool = False, **kwargs):
|
||||
print_verbose(
|
||||
f"async set cache: cache key: {key}; local_only: {local_only}; value: {value}"
|
||||
)
|
||||
print_verbose(f"async set cache: cache key: {key}; local_only: {local_only}; value: {value}")
|
||||
try:
|
||||
if self.in_memory_cache is not None:
|
||||
if "ttl" not in kwargs and self.default_in_memory_ttl is not None:
|
||||
|
|
@ -374,36 +335,26 @@ class DualCache(BaseCache):
|
|||
if self.redis_cache is not None and local_only is False:
|
||||
await self.redis_cache.async_set_cache(key, value, **kwargs)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"LiteLLM Cache: Excepton async add_cache: {str(e)}"
|
||||
)
|
||||
verbose_logger.exception(f"LiteLLM Cache: Excepton async add_cache: {str(e)}")
|
||||
|
||||
# async_batch_set_cache
|
||||
async def async_set_cache_pipeline(
|
||||
self, cache_list: list, local_only: bool = False, **kwargs
|
||||
):
|
||||
async def async_set_cache_pipeline(self, cache_list: list, local_only: bool = False, **kwargs):
|
||||
"""
|
||||
Batch write values to the cache
|
||||
"""
|
||||
print_verbose(
|
||||
f"async batch set cache: cache keys: {cache_list}; local_only: {local_only}"
|
||||
)
|
||||
print_verbose(f"async batch set cache: cache keys: {cache_list}; local_only: {local_only}")
|
||||
try:
|
||||
if self.in_memory_cache is not None:
|
||||
if "ttl" not in kwargs and self.default_in_memory_ttl is not None:
|
||||
kwargs["ttl"] = self.default_in_memory_ttl
|
||||
await self.in_memory_cache.async_set_cache_pipeline(
|
||||
cache_list=cache_list, **kwargs
|
||||
)
|
||||
await self.in_memory_cache.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
|
||||
|
||||
if self.redis_cache is not None and local_only is False:
|
||||
await self.redis_cache.async_set_cache_pipeline(
|
||||
cache_list=cache_list, ttl=kwargs.pop("ttl", None), **kwargs
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"LiteLLM Cache: Excepton async add_cache: {str(e)}"
|
||||
)
|
||||
verbose_logger.exception(f"LiteLLM Cache: Excepton async add_cache: {str(e)}")
|
||||
|
||||
async def async_increment_cache(
|
||||
self,
|
||||
|
|
@ -428,9 +379,7 @@ class DualCache(BaseCache):
|
|||
result: Optional[float] = None
|
||||
try:
|
||||
if self.in_memory_cache is not None:
|
||||
result = await self.in_memory_cache.async_increment(
|
||||
key, value, **kwargs
|
||||
)
|
||||
result = await self.in_memory_cache.async_increment(key, value, **kwargs)
|
||||
|
||||
if self.redis_cache is not None and local_only is False:
|
||||
result = await self.redis_cache.async_increment(
|
||||
|
|
@ -478,9 +427,7 @@ class DualCache(BaseCache):
|
|||
)
|
||||
return result
|
||||
|
||||
async def async_set_cache_sadd(
|
||||
self, key, value: List, local_only: bool = False, **kwargs
|
||||
) -> None:
|
||||
async def async_set_cache_sadd(self, key, value: List, local_only: bool = False, **kwargs) -> None:
|
||||
"""
|
||||
Add value to a set
|
||||
|
||||
|
|
@ -492,14 +439,10 @@ class DualCache(BaseCache):
|
|||
"""
|
||||
try:
|
||||
if self.in_memory_cache is not None:
|
||||
_ = await self.in_memory_cache.async_set_cache_sadd(
|
||||
key, value, ttl=kwargs.get("ttl", None)
|
||||
)
|
||||
_ = await self.in_memory_cache.async_set_cache_sadd(key, value, ttl=kwargs.get("ttl", None))
|
||||
|
||||
if self.redis_cache is not None and local_only is False:
|
||||
_ = await self.redis_cache.async_set_cache_sadd(
|
||||
key, value, ttl=kwargs.get("ttl", None)
|
||||
)
|
||||
_ = await self.redis_cache.async_set_cache_sadd(key, value, ttl=kwargs.get("ttl", None))
|
||||
|
||||
return None
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -26,15 +26,10 @@ class GCSCache(BaseCache):
|
|||
) -> None:
|
||||
super().__init__()
|
||||
self.bucket_name = bucket_name or GCSBucketBase(bucket_name=None).BUCKET_NAME
|
||||
self.path_service_account = (
|
||||
path_service_account
|
||||
or GCSBucketBase(bucket_name=None).path_service_account_json
|
||||
)
|
||||
self.path_service_account = path_service_account or GCSBucketBase(bucket_name=None).path_service_account_json
|
||||
self.key_prefix = gcs_path.rstrip("/") + "/" if gcs_path else ""
|
||||
# create httpx clients
|
||||
self.async_client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.LoggingCallback
|
||||
)
|
||||
self.async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
|
||||
self.sync_client = _get_httpx_client()
|
||||
|
||||
def _construct_headers(self) -> dict:
|
||||
|
|
@ -64,9 +59,7 @@ class GCSCache(BaseCache):
|
|||
data = 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}"
|
||||
)
|
||||
print_verbose(f"GCS Caching: async_set_cache() - Got exception from GCS: {e}")
|
||||
|
||||
def get_cache(self, key, **kwargs):
|
||||
try:
|
||||
|
|
@ -83,9 +76,7 @@ class GCSCache(BaseCache):
|
|||
return cached_response
|
||||
return None
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"GCS Caching: get_cache() - Got exception from GCS: {e}"
|
||||
)
|
||||
verbose_logger.error(f"GCS Caching: get_cache() - Got exception from GCS: {e}")
|
||||
|
||||
async def async_get_cache(self, key, **kwargs):
|
||||
try:
|
||||
|
|
@ -98,9 +89,7 @@ class GCSCache(BaseCache):
|
|||
return json.loads(response.text)
|
||||
return None
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"GCS Caching: async_get_cache() - Got exception from GCS: {e}"
|
||||
)
|
||||
verbose_logger.error(f"GCS Caching: async_get_cache() - Got exception from GCS: {e}")
|
||||
|
||||
def flush_cache(self):
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -40,9 +40,7 @@ class InMemoryCache(BaseCache):
|
|||
max_size_in_memory if max_size_in_memory is not None else 200
|
||||
) # set an upper bound of 200 items in-memory
|
||||
self.default_ttl = default_ttl or 600
|
||||
self.max_size_per_item = (
|
||||
max_size_per_item or MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB
|
||||
) # 1MB = 1024KB
|
||||
self.max_size_per_item = max_size_per_item or MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB # 1MB = 1024KB
|
||||
|
||||
# in-memory cache
|
||||
self.cache_dict: dict = {}
|
||||
|
|
@ -58,8 +56,7 @@ class InMemoryCache(BaseCache):
|
|||
# Fast path for common primitive types that are typically small
|
||||
if (
|
||||
isinstance(value, (bool, int, float, str))
|
||||
and len(str(value))
|
||||
< self.max_size_per_item * MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB
|
||||
and len(str(value)) < self.max_size_per_item * MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB
|
||||
): # Conservative estimate
|
||||
return True
|
||||
|
||||
|
|
@ -73,9 +70,7 @@ class InMemoryCache(BaseCache):
|
|||
return size <= self.max_size_per_item
|
||||
|
||||
# Fallback for complex types
|
||||
if isinstance(value, BaseModel) and hasattr(
|
||||
value, "model_dump"
|
||||
): # Pydantic v2
|
||||
if isinstance(value, BaseModel) and hasattr(value, "model_dump"): # Pydantic v2
|
||||
value = value.model_dump()
|
||||
elif hasattr(value, "isoformat"): # datetime objects
|
||||
return True # datetime strings are always small
|
||||
|
|
@ -257,9 +252,7 @@ class InMemoryCache(BaseCache):
|
|||
) -> Optional[List[float]]:
|
||||
results = []
|
||||
for increment in increment_list:
|
||||
result = await self.async_increment(
|
||||
increment["key"], increment["increment_value"], **kwargs
|
||||
)
|
||||
result = await self.async_increment(increment["key"], increment["increment_value"], **kwargs)
|
||||
results.append(result)
|
||||
return results
|
||||
|
||||
|
|
|
|||
|
|
@ -51,34 +51,24 @@ class QdrantSemanticCache(BaseCache):
|
|||
raise Exception("collection_name must be provided, passed None")
|
||||
|
||||
self.collection_name = collection_name
|
||||
print_verbose(
|
||||
f"qdrant semantic-cache initializing COLLECTION - {self.collection_name}"
|
||||
)
|
||||
print_verbose(f"qdrant semantic-cache initializing COLLECTION - {self.collection_name}")
|
||||
|
||||
if similarity_threshold is None:
|
||||
raise Exception("similarity_threshold must be provided, passed None")
|
||||
self.similarity_threshold = similarity_threshold
|
||||
self.embedding_model = embedding_model
|
||||
self.vector_size = (
|
||||
vector_size if vector_size is not None else QDRANT_VECTOR_SIZE
|
||||
)
|
||||
self.vector_size = vector_size if vector_size is not None else QDRANT_VECTOR_SIZE
|
||||
headers = {}
|
||||
|
||||
# check if defined as os.environ/ variable
|
||||
if qdrant_api_base:
|
||||
if isinstance(qdrant_api_base, str) and qdrant_api_base.startswith(
|
||||
"os.environ/"
|
||||
):
|
||||
if isinstance(qdrant_api_base, str) and qdrant_api_base.startswith("os.environ/"):
|
||||
qdrant_api_base = get_secret_str(qdrant_api_base)
|
||||
if qdrant_api_key:
|
||||
if isinstance(qdrant_api_key, str) and qdrant_api_key.startswith(
|
||||
"os.environ/"
|
||||
):
|
||||
if isinstance(qdrant_api_key, str) and qdrant_api_key.startswith("os.environ/"):
|
||||
qdrant_api_key = get_secret_str(qdrant_api_key)
|
||||
|
||||
qdrant_api_base = (
|
||||
qdrant_api_base or os.getenv("QDRANT_URL") or os.getenv("QDRANT_API_BASE")
|
||||
)
|
||||
qdrant_api_base = qdrant_api_base or os.getenv("QDRANT_URL") or os.getenv("QDRANT_API_BASE")
|
||||
qdrant_api_key = qdrant_api_key or os.getenv("QDRANT_API_KEY")
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if qdrant_api_key:
|
||||
|
|
@ -94,22 +84,16 @@ class QdrantSemanticCache(BaseCache):
|
|||
self.headers = headers
|
||||
|
||||
self.sync_client = _get_httpx_client()
|
||||
self.async_client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.Caching
|
||||
)
|
||||
self.async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Caching)
|
||||
|
||||
if quantization_config is None:
|
||||
print_verbose(
|
||||
"Quantization config is not provided. Default binary quantization will be used."
|
||||
)
|
||||
print_verbose("Quantization config is not provided. Default binary quantization will be used.")
|
||||
collection_exists = self.sync_client.get(
|
||||
url=f"{self.qdrant_api_base}/collections/{self.collection_name}/exists",
|
||||
headers=self.headers,
|
||||
)
|
||||
if collection_exists.status_code != 200:
|
||||
raise ValueError(
|
||||
f"Error from qdrant checking if /collections exist {collection_exists.text}"
|
||||
)
|
||||
raise ValueError(f"Error from qdrant checking if /collections exist {collection_exists.text}")
|
||||
|
||||
if collection_exists.json()["result"]["exists"]:
|
||||
collection_details = self.sync_client.get(
|
||||
|
|
@ -117,9 +101,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
headers=self.headers,
|
||||
)
|
||||
self.collection_info = collection_details.json()
|
||||
print_verbose(
|
||||
f"Collection already exists.\nCollection details:{self.collection_info}"
|
||||
)
|
||||
print_verbose(f"Collection already exists.\nCollection details:{self.collection_info}")
|
||||
self._ensure_cache_key_payload_index()
|
||||
else:
|
||||
quantization_params: Dict[str, Any]
|
||||
|
|
@ -138,13 +120,9 @@ class QdrantSemanticCache(BaseCache):
|
|||
}
|
||||
}
|
||||
elif quantization_config == "product":
|
||||
quantization_params = {
|
||||
"product": {"compression": "x16", "always_ram": False}
|
||||
}
|
||||
quantization_params = {"product": {"compression": "x16", "always_ram": False}}
|
||||
else:
|
||||
raise Exception(
|
||||
"Quantization config must be one of 'scalar', 'binary' or 'product'"
|
||||
)
|
||||
raise Exception("Quantization config must be one of 'scalar', 'binary' or 'product'")
|
||||
|
||||
new_collection_status = self.sync_client.put(
|
||||
url=f"{self.qdrant_api_base}/collections/{self.collection_name}",
|
||||
|
|
@ -160,9 +138,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
headers=self.headers,
|
||||
)
|
||||
self.collection_info = collection_details.json()
|
||||
print_verbose(
|
||||
f"New collection created.\nCollection details:{self.collection_info}"
|
||||
)
|
||||
print_verbose(f"New collection created.\nCollection details:{self.collection_info}")
|
||||
self._ensure_cache_key_payload_index()
|
||||
else:
|
||||
raise Exception("Error while creating new collection")
|
||||
|
|
@ -171,9 +147,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
if cached_response is None:
|
||||
return cached_response
|
||||
try:
|
||||
cached_response = json.loads(
|
||||
cached_response
|
||||
) # Convert string to dictionary
|
||||
cached_response = json.loads(cached_response) # Convert string to dictionary
|
||||
except Exception:
|
||||
cached_response = ast.literal_eval(cached_response)
|
||||
return cached_response
|
||||
|
|
@ -202,15 +176,9 @@ class QdrantSemanticCache(BaseCache):
|
|||
},
|
||||
)
|
||||
if response.status_code not in (200, 201):
|
||||
print_verbose(
|
||||
"Qdrant semantic-cache could not create cache-key payload index: "
|
||||
f"{response.text}"
|
||||
)
|
||||
print_verbose(f"Qdrant semantic-cache could not create cache-key payload index: {response.text}")
|
||||
except Exception as exc:
|
||||
print_verbose(
|
||||
"Qdrant semantic-cache could not create cache-key payload index: "
|
||||
f"{str(exc)}"
|
||||
)
|
||||
print_verbose(f"Qdrant semantic-cache could not create cache-key payload index: {str(exc)}")
|
||||
|
||||
def _payload_matches_cache_key(self, payload: dict, key: str) -> bool:
|
||||
# Pre-isolation points stored only prompt + response with no cache-key
|
||||
|
|
@ -220,9 +188,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
cached_key = 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:
|
||||
def _get_embedding(self, prompt: str, metadata: Dict[str, Any] | None = None) -> EmbeddingResponse:
|
||||
"""Embed via the proxy Router when it serves the model, else direct."""
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_model_list, llm_router
|
||||
|
|
@ -230,9 +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 = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
if router is not None:
|
||||
return router.embedding(
|
||||
model=self.embedding_model,
|
||||
|
|
@ -246,18 +210,14 @@ class QdrantSemanticCache(BaseCache):
|
|||
cache={"no-store": True, "no-cache": True},
|
||||
)
|
||||
|
||||
async def _get_async_embedding(
|
||||
self, prompt: str, metadata: Dict[str, Any] | None = None
|
||||
) -> EmbeddingResponse:
|
||||
async def _get_async_embedding(self, prompt: str, metadata: Dict[str, Any] | None = None) -> EmbeddingResponse:
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_model_list, llm_router
|
||||
except ImportError:
|
||||
llm_model_list = None
|
||||
llm_router = None
|
||||
|
||||
router = resolve_embedding_router(
|
||||
self.embedding_model, llm_router, llm_model_list
|
||||
)
|
||||
router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
if router is not None:
|
||||
return await router.aembedding(
|
||||
model=self.embedding_model,
|
||||
|
|
@ -394,9 +354,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
# get the prompt
|
||||
messages = kwargs["messages"]
|
||||
prompt = get_str_from_messages(messages)
|
||||
embedding_response = await self._get_async_embedding(
|
||||
prompt, metadata=kwargs.get("metadata")
|
||||
)
|
||||
embedding_response = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
# get the embedding
|
||||
embedding = embedding_response["data"][0]["embedding"]
|
||||
|
|
@ -432,9 +390,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
messages = kwargs["messages"]
|
||||
prompt = get_str_from_messages(messages)
|
||||
|
||||
embedding_response = await self._get_async_embedding(
|
||||
prompt, metadata=kwargs.get("metadata")
|
||||
)
|
||||
embedding_response = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
# get the embedding
|
||||
embedding = embedding_response["data"][0]["embedding"]
|
||||
|
|
|
|||
|
|
@ -153,8 +153,7 @@ class RedisCircuitBreaker:
|
|||
if self._failure_count >= self.failure_threshold:
|
||||
if self._state != self.OPEN:
|
||||
verbose_logger.warning(
|
||||
"Redis circuit breaker OPENED after %d consecutive failures — "
|
||||
"fast-failing Redis calls for %ds",
|
||||
"Redis circuit breaker OPENED after %d consecutive failures — fast-failing Redis calls for %ds",
|
||||
self._failure_count,
|
||||
self.recovery_timeout,
|
||||
)
|
||||
|
|
@ -179,9 +178,7 @@ def _redis_circuit_breaker_guard(method): # type: ignore
|
|||
@functools.wraps(method)
|
||||
async def wrapper(self, *args, **kwargs): # type: ignore
|
||||
if self._circuit_breaker.is_open():
|
||||
raise Exception(
|
||||
f"Redis circuit breaker is open — skipping {method.__name__}"
|
||||
)
|
||||
raise Exception(f"Redis circuit breaker is open — skipping {method.__name__}")
|
||||
try:
|
||||
result = await method(self, *args, **kwargs)
|
||||
self._circuit_breaker.record_success()
|
||||
|
|
@ -233,9 +230,7 @@ class RedisCache(BaseCache):
|
|||
|
||||
redis_kwargs.update(kwargs)
|
||||
self.redis_client = get_redis_client(**redis_kwargs)
|
||||
self.redis_async_client: Optional[
|
||||
Union[async_redis_client, async_redis_cluster_client]
|
||||
] = None
|
||||
self.redis_async_client: Optional[Union[async_redis_client, async_redis_cluster_client]] = None
|
||||
self.redis_kwargs = redis_kwargs
|
||||
self.async_redis_conn_pool = get_redis_connection_pool(**redis_kwargs)
|
||||
|
||||
|
|
@ -274,9 +269,7 @@ class RedisCache(BaseCache):
|
|||
_ = asyncio.get_running_loop().create_task(self.ping())
|
||||
except Exception as e:
|
||||
if "no running event loop" in str(e):
|
||||
verbose_logger.debug(
|
||||
"Ignoring async redis ping. No running event loop."
|
||||
)
|
||||
verbose_logger.debug("Ignoring async redis ping. No running event loop.")
|
||||
else:
|
||||
verbose_logger.error(
|
||||
"Error connecting to Async Redis client - {}".format(str(e)),
|
||||
|
|
@ -289,9 +282,7 @@ class RedisCache(BaseCache):
|
|||
if hasattr(self.redis_client, "ping"):
|
||||
self.redis_client.ping() # type: ignore
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
"Error connecting to Sync Redis client", extra={"error": str(e)}
|
||||
)
|
||||
verbose_logger.error("Error connecting to Sync Redis client", extra={"error": str(e)})
|
||||
self._handle_sync_ping_error(e)
|
||||
|
||||
def _handle_async_ping_error(self, e: Exception):
|
||||
|
|
@ -350,18 +341,12 @@ class RedisCache(BaseCache):
|
|||
cache_key = self._get_async_client_cache_key()
|
||||
cached_client = in_memory_llm_clients_cache.get_cache(key=cache_key)
|
||||
if cached_client is not None:
|
||||
redis_async_client = cast(
|
||||
Union[async_redis_client, async_redis_cluster_client], cached_client
|
||||
)
|
||||
redis_async_client = cast(Union[async_redis_client, async_redis_cluster_client], cached_client)
|
||||
else:
|
||||
# Create new connection pool and client for current event loop
|
||||
self.async_redis_conn_pool = get_redis_connection_pool(**self.redis_kwargs)
|
||||
redis_async_client = get_redis_async_client(
|
||||
connection_pool=self.async_redis_conn_pool, **self.redis_kwargs
|
||||
)
|
||||
in_memory_llm_clients_cache.set_cache(
|
||||
key=cache_key, value=redis_async_client
|
||||
)
|
||||
redis_async_client = get_redis_async_client(connection_pool=self.async_redis_conn_pool, **self.redis_kwargs)
|
||||
in_memory_llm_clients_cache.set_cache(key=cache_key, value=redis_async_client)
|
||||
|
||||
self.redis_async_client = redis_async_client # type: ignore
|
||||
return redis_async_client
|
||||
|
|
@ -408,9 +393,7 @@ class RedisCache(BaseCache):
|
|||
|
||||
def set_cache(self, key, value, **kwargs):
|
||||
ttl = self.get_ttl(**kwargs)
|
||||
print_verbose(
|
||||
f"Set Redis Cache: key: {key}\nValue {value}\nttl={ttl}, redis_version={self.redis_version}"
|
||||
)
|
||||
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()
|
||||
|
|
@ -426,13 +409,9 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
except Exception as e:
|
||||
# NON blocking - notify users Redis is throwing an exception
|
||||
print_verbose(
|
||||
f"litellm.caching.caching: set() - Got exception from REDIS : {str(e)}"
|
||||
)
|
||||
print_verbose(f"litellm.caching.caching: set() - Got exception from REDIS : {str(e)}")
|
||||
|
||||
def increment_cache(
|
||||
self, key, value: int, ttl: Optional[float] = None, **kwargs
|
||||
) -> int:
|
||||
def increment_cache(self, key, value: int, ttl: Optional[float] = None, **kwargs) -> int:
|
||||
_redis_client = self.redis_client
|
||||
start_time = time.time()
|
||||
set_ttl = self.get_ttl(ttl=ttl)
|
||||
|
|
@ -571,24 +550,18 @@ class RedisCache(BaseCache):
|
|||
f"{self.namespace}-{hashlib.sha256(script.encode()).hexdigest()[:16]}"
|
||||
)
|
||||
|
||||
async def run_script(
|
||||
keys: Sequence[str], args: Sequence[Any], client: Any = None
|
||||
) -> Any:
|
||||
executor: Optional[Callable[..., Awaitable[Any]]] = (
|
||||
litellm.in_memory_llm_clients_cache.get_cache(key=script_cache_key)
|
||||
async def run_script(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any:
|
||||
executor: Optional[Callable[..., Awaitable[Any]]] = litellm.in_memory_llm_clients_cache.get_cache(
|
||||
key=script_cache_key
|
||||
)
|
||||
if executor is None:
|
||||
executor = self._register_script_for_current_loop(script)
|
||||
litellm.in_memory_llm_clients_cache.set_cache(
|
||||
key=script_cache_key, value=executor
|
||||
)
|
||||
litellm.in_memory_llm_clients_cache.set_cache(key=script_cache_key, value=executor)
|
||||
return await executor(keys=keys, args=args, client=client)
|
||||
|
||||
return run_script
|
||||
|
||||
def _register_script_for_current_loop(
|
||||
self, script: str
|
||||
) -> Callable[..., Awaitable[Any]]:
|
||||
def _register_script_for_current_loop(self, script: str) -> Callable[..., Awaitable[Any]]:
|
||||
"""
|
||||
Register the script against the current event loop's Redis client.
|
||||
|
||||
|
|
@ -599,30 +572,18 @@ class RedisCache(BaseCache):
|
|||
if hasattr(_redis_client, "register_script"):
|
||||
registered_script = _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
|
||||
)
|
||||
return await registered_script(
|
||||
keys=namespaced_keys, args=args, client=client
|
||||
)
|
||||
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)
|
||||
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)
|
||||
|
||||
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
|
||||
)
|
||||
return await _redis_client.evalsha(
|
||||
script_sha, len(namespaced_keys), *namespaced_keys, *args
|
||||
)
|
||||
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)
|
||||
return await _redis_client.evalsha(script_sha, len(namespaced_keys), *namespaced_keys, *args)
|
||||
|
||||
return cluster_executor
|
||||
|
||||
|
|
@ -678,9 +639,7 @@ class RedisCache(BaseCache):
|
|||
nx=nx,
|
||||
ex=ttl,
|
||||
)
|
||||
print_verbose(
|
||||
f"Successfully Set ASYNC Redis Cache: key: {key}\nValue {value}\nttl={ttl}"
|
||||
)
|
||||
print_verbose(f"Successfully Set ASYNC Redis Cache: key: {key}\nValue {value}\nttl={ttl}")
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
asyncio.create_task(
|
||||
|
|
@ -729,9 +688,7 @@ class RedisCache(BaseCache):
|
|||
# Iterate through each key-value pair in the cache_list and set them in the pipeline.
|
||||
for cache_key, cache_value in cache_list:
|
||||
cache_key = self.check_and_fix_namespace(key=cache_key)
|
||||
print_verbose(
|
||||
f"Set ASYNC Redis Cache PIPELINE: key: {cache_key}\nValue {cache_value}\nttl={ttl}"
|
||||
)
|
||||
print_verbose(f"Set ASYNC Redis Cache PIPELINE: key: {cache_key}\nValue {cache_value}\nttl={ttl}")
|
||||
json_cache_value = json.dumps(cache_value)
|
||||
# Set the value with a TTL if it's provided.
|
||||
_td: Optional[timedelta] = None
|
||||
|
|
@ -747,9 +704,7 @@ class RedisCache(BaseCache):
|
|||
return results
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_set_cache_pipeline(
|
||||
self, cache_list: List[Tuple[Any, Any]], ttl: Optional[float] = None, **kwargs
|
||||
):
|
||||
async def async_set_cache_pipeline(self, cache_list: List[Tuple[Any, Any]], ttl: Optional[float] = None, **kwargs):
|
||||
"""
|
||||
Use Redis Pipelines for bulk write operations
|
||||
"""
|
||||
|
|
@ -760,9 +715,7 @@ class RedisCache(BaseCache):
|
|||
_redis_client = self.init_async_client()
|
||||
start_time = time.time()
|
||||
|
||||
print_verbose(
|
||||
f"Set Async Redis Cache: key list: {cache_list}\nttl={ttl}, redis_version={self.redis_version}"
|
||||
)
|
||||
print_verbose(f"Set Async Redis Cache: key list: {cache_list}\nttl={ttl}, redis_version={self.redis_version}")
|
||||
cache_value: Any = None
|
||||
try:
|
||||
async with _redis_client.pipeline(transaction=False) as pipe:
|
||||
|
|
@ -824,9 +777,7 @@ class RedisCache(BaseCache):
|
|||
raise
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_set_cache_sadd(
|
||||
self, key, value: List, ttl: Optional[float], **kwargs
|
||||
):
|
||||
async def async_set_cache_sadd(self, key, value: List, ttl: Optional[float], **kwargs):
|
||||
from redis.asyncio import Redis
|
||||
|
||||
start_time = time.time()
|
||||
|
|
@ -857,12 +808,8 @@ class RedisCache(BaseCache):
|
|||
key = self.check_and_fix_namespace(key=key)
|
||||
print_verbose(f"Set ASYNC Redis Cache: key: {key}\nValue {value}\nttl={ttl}")
|
||||
try:
|
||||
await self._set_cache_sadd_helper(
|
||||
redis_client=_redis_client, key=key, value=value, ttl=ttl
|
||||
)
|
||||
print_verbose(
|
||||
f"Successfully Set ASYNC Redis Cache SADD: key: {key}\nValue {value}\nttl={ttl}"
|
||||
)
|
||||
await self._set_cache_sadd_helper(redis_client=_redis_client, key=key, value=value, ttl=ttl)
|
||||
print_verbose(f"Successfully Set ASYNC Redis Cache SADD: key: {key}\nValue {value}\nttl={ttl}")
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
asyncio.create_task(
|
||||
|
|
@ -1006,9 +953,7 @@ class RedisCache(BaseCache):
|
|||
return float(result)
|
||||
|
||||
async def flush_cache_buffer(self):
|
||||
print_verbose(
|
||||
f"flushing to redis....reached size of buffer {len(self.redis_batch_writing_buffer)}"
|
||||
)
|
||||
print_verbose(f"flushing to redis....reached size of buffer {len(self.redis_batch_writing_buffer)}")
|
||||
await self.async_set_cache_pipeline(self.redis_batch_writing_buffer)
|
||||
self.redis_batch_writing_buffer = []
|
||||
|
||||
|
|
@ -1021,9 +966,7 @@ class RedisCache(BaseCache):
|
|||
# cached_response is in `b{} convert it to ModelResponse
|
||||
cached_response = cached_response.decode("utf-8") # Convert bytes to string
|
||||
try:
|
||||
cached_response = json.loads(
|
||||
cached_response
|
||||
) # Convert string to dictionary
|
||||
cached_response = json.loads(cached_response) # Convert string to dictionary
|
||||
except Exception:
|
||||
cached_response = ast.literal_eval(cached_response)
|
||||
return cached_response
|
||||
|
|
@ -1044,15 +987,11 @@ class RedisCache(BaseCache):
|
|||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
print_verbose(
|
||||
f"Got Redis Cache: key: {key}, cached_response {cached_response}"
|
||||
)
|
||||
print_verbose(f"Got Redis Cache: key: {key}, cached_response {cached_response}")
|
||||
return self._get_cache_logic(cached_response=cached_response)
|
||||
except Exception as e:
|
||||
# NON blocking - notify users Redis is throwing an exception
|
||||
verbose_logger.error(
|
||||
"litellm.caching.caching: get() - Got exception from REDIS: ", e
|
||||
)
|
||||
verbose_logger.error("litellm.caching.caching: get() - Got exception from REDIS: ", e)
|
||||
|
||||
def _run_redis_mget_operation(self, keys: List[str]) -> List[Any]:
|
||||
"""
|
||||
|
|
@ -1124,9 +1063,7 @@ class RedisCache(BaseCache):
|
|||
return key_value_dict
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_get_cache(
|
||||
self, key, parent_otel_span: Optional[Span] = None, **kwargs
|
||||
):
|
||||
async def async_get_cache(self, key, parent_otel_span: Optional[Span] = None, **kwargs):
|
||||
from redis.asyncio import Redis
|
||||
|
||||
_redis_client: Redis = self.init_async_client() # type: ignore
|
||||
|
|
@ -1136,9 +1073,7 @@ class RedisCache(BaseCache):
|
|||
try:
|
||||
print_verbose(f"Get Async Redis Cache: key: {key}")
|
||||
cached_response = await _redis_client.get(key)
|
||||
print_verbose(
|
||||
f"Got Async Redis Cache: key: {key}, cached_response {cached_response}"
|
||||
)
|
||||
print_verbose(f"Got Async Redis Cache: key: {key}, cached_response {cached_response}")
|
||||
response = self._get_cache_logic(cached_response=cached_response)
|
||||
|
||||
end_time = time.time()
|
||||
|
|
@ -1170,9 +1105,7 @@ class RedisCache(BaseCache):
|
|||
event_metadata={"key": key},
|
||||
)
|
||||
)
|
||||
print_verbose(
|
||||
f"litellm.caching.caching: async get() - Got exception from REDIS: {str(e)}"
|
||||
)
|
||||
print_verbose(f"litellm.caching.caching: async get() - Got exception from REDIS: {str(e)}")
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_batch_get_cache(
|
||||
|
|
@ -1277,9 +1210,7 @@ class RedisCache(BaseCache):
|
|||
error=e,
|
||||
call_type=f"sync_ping <- {_get_call_stack_info()}",
|
||||
)
|
||||
verbose_logger.error(
|
||||
f"LiteLLM Redis Cache PING: - Got exception from REDIS : {str(e)}"
|
||||
)
|
||||
verbose_logger.error(f"LiteLLM Redis Cache PING: - Got exception from REDIS : {str(e)}")
|
||||
raise e
|
||||
|
||||
async def ping(self) -> bool:
|
||||
|
|
@ -1313,9 +1244,7 @@ class RedisCache(BaseCache):
|
|||
call_type=f"async_ping <- {_get_call_stack_info()}",
|
||||
)
|
||||
)
|
||||
verbose_logger.error(
|
||||
f"LiteLLM Redis Cache PING: - Got exception from REDIS : {str(e)}"
|
||||
)
|
||||
verbose_logger.error(f"LiteLLM Redis Cache PING: - Got exception from REDIS : {str(e)}")
|
||||
raise e
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
|
|
@ -1415,9 +1344,7 @@ class RedisCache(BaseCache):
|
|||
# Execute the pipeline and return results
|
||||
results = await pipe.execute()
|
||||
# only return float values
|
||||
verbose_logger.debug(
|
||||
f"Increment ASYNC Redis Cache PIPELINE: results: {results}"
|
||||
)
|
||||
verbose_logger.debug(f"Increment ASYNC Redis Cache PIPELINE: results: {results}")
|
||||
return [r for r in results if isinstance(r, float)]
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
|
|
@ -1441,9 +1368,7 @@ class RedisCache(BaseCache):
|
|||
_redis_client: Redis = self.init_async_client() # type: ignore
|
||||
start_time = time.time()
|
||||
|
||||
print_verbose(
|
||||
f"Increment Async Redis Cache Pipeline: increment list: {increment_list}"
|
||||
)
|
||||
print_verbose(f"Increment Async Redis Cache Pipeline: increment list: {increment_list}")
|
||||
|
||||
try:
|
||||
async with _redis_client.pipeline(transaction=False) as pipe:
|
||||
|
|
@ -1557,9 +1482,7 @@ class RedisCache(BaseCache):
|
|||
call_type=f"async_rpush <- {_get_call_stack_info()}",
|
||||
)
|
||||
)
|
||||
verbose_logger.error(
|
||||
f"LiteLLM Redis Cache RPUSH: - Got exception from REDIS : {str(e)}"
|
||||
)
|
||||
verbose_logger.error(f"LiteLLM Redis Cache RPUSH: - Got exception from REDIS : {str(e)}")
|
||||
raise e
|
||||
|
||||
async def _pipeline_rpush_helper(
|
||||
|
|
@ -1633,9 +1556,7 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
raise e
|
||||
|
||||
async def handle_lpop_count_for_older_redis_versions(
|
||||
self, pipe: pipeline, key: str, count: int
|
||||
) -> List[bytes]:
|
||||
async def handle_lpop_count_for_older_redis_versions(self, pipe: pipeline, key: str, count: int) -> List[bytes]:
|
||||
result: List[bytes] = []
|
||||
for _ in range(count):
|
||||
pipe.lpop(key)
|
||||
|
|
@ -1666,9 +1587,7 @@ class RedisCache(BaseCache):
|
|||
if count is not None and major_version < 7:
|
||||
# For Redis < 7.0, use pipeline to execute multiple LPOP commands
|
||||
async with _redis_client.pipeline(transaction=False) as pipe:
|
||||
result = await self.handle_lpop_count_for_older_redis_versions(
|
||||
pipe, key, count
|
||||
)
|
||||
result = await self.handle_lpop_count_for_older_redis_versions(pipe, key, count)
|
||||
else:
|
||||
# For Redis >= 7.0 or when count is None, use native LPOP with count
|
||||
result = await _redis_client.lpop(key, count)
|
||||
|
|
@ -1690,9 +1609,7 @@ class RedisCache(BaseCache):
|
|||
return result.decode("utf-8")
|
||||
except Exception:
|
||||
return result
|
||||
elif isinstance(result, list) and all(
|
||||
isinstance(item, bytes) for item in result
|
||||
):
|
||||
elif isinstance(result, list) and all(isinstance(item, bytes) for item in result):
|
||||
try:
|
||||
return [item.decode("utf-8") for item in result]
|
||||
except Exception:
|
||||
|
|
@ -1711,9 +1628,7 @@ class RedisCache(BaseCache):
|
|||
call_type=f"async_lpop <- {_get_call_stack_info()}",
|
||||
)
|
||||
)
|
||||
verbose_logger.error(
|
||||
f"LiteLLM Redis Cache LPOP: - Got exception from REDIS : {str(e)}"
|
||||
)
|
||||
verbose_logger.error(f"LiteLLM Redis Cache LPOP: - Got exception from REDIS : {str(e)}")
|
||||
raise e
|
||||
|
||||
async def _pipeline_lpop_helper(
|
||||
|
|
@ -1749,9 +1664,7 @@ class RedisCache(BaseCache):
|
|||
raw_results = []
|
||||
offset = 0
|
||||
for count in counts:
|
||||
key_results = [
|
||||
r for r in flat_results[offset : offset + count] if r is not None
|
||||
]
|
||||
key_results = [r for r in flat_results[offset : offset + count] if r is not None]
|
||||
raw_results.append(key_results if key_results else None)
|
||||
offset += count
|
||||
|
||||
|
|
@ -1768,11 +1681,7 @@ class RedisCache(BaseCache):
|
|||
elif isinstance(r, list):
|
||||
try:
|
||||
decoded_results.append(
|
||||
[
|
||||
item.decode("utf-8") if isinstance(item, bytes) else item
|
||||
for item in r
|
||||
if item is not None
|
||||
]
|
||||
[item.decode("utf-8") if isinstance(item, bytes) else item for item in r if item is not None]
|
||||
or None
|
||||
)
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -37,9 +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 = 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
|
||||
|
||||
|
|
|
|||
|
|
@ -97,8 +97,7 @@ class RedisSemanticCache(BaseCache):
|
|||
# Raise a more informative exception if any of the required keys are missing
|
||||
missing_var = e.args[0]
|
||||
raise ValueError(
|
||||
f"Missing required Redis configuration: {missing_var}. "
|
||||
f"Provide {missing_var} or redis_url."
|
||||
f"Missing required Redis configuration: {missing_var}. Provide {missing_var} or redis_url."
|
||||
) from e
|
||||
|
||||
redis_url = f"redis://:{password}@{host}:{port}"
|
||||
|
|
@ -158,10 +157,7 @@ class RedisSemanticCache(BaseCache):
|
|||
) -> Any:
|
||||
def _is_schema_mismatch(exc: ValueError) -> bool:
|
||||
error_message = str(exc).lower()
|
||||
return any(
|
||||
phrase in error_message
|
||||
for phrase in ("schema does not match", "index schema")
|
||||
)
|
||||
return any(phrase in error_message for phrase in ("schema does not match", "index schema"))
|
||||
|
||||
try:
|
||||
return semantic_cache_cls(
|
||||
|
|
@ -310,9 +306,7 @@ class RedisSemanticCache(BaseCache):
|
|||
return dict_method()
|
||||
return value
|
||||
|
||||
def _get_embedding(
|
||||
self, prompt: str, metadata: Dict[str, Any] | None = None
|
||||
) -> List[float]:
|
||||
def _get_embedding(self, prompt: str, metadata: Dict[str, Any] | None = None) -> List[float]:
|
||||
"""
|
||||
Routes through the proxy Router when the embedding model is a Router
|
||||
deployment so per-deployment auth (e.g. Bedrock aws_role_name) applies,
|
||||
|
|
@ -324,9 +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 = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
if router is not None:
|
||||
embedding_response = cast(
|
||||
EmbeddingResponse,
|
||||
|
|
@ -398,9 +390,7 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
value_str = str(value)
|
||||
|
||||
prompt_embedding = self._get_embedding(
|
||||
prompt, metadata=kwargs.get("metadata")
|
||||
)
|
||||
prompt_embedding = self._get_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
store_kwargs: dict[str, Any] = {
|
||||
"vector": prompt_embedding,
|
||||
|
|
@ -413,9 +403,7 @@ class RedisSemanticCache(BaseCache):
|
|||
store_kwargs["ttl"] = int(ttl)
|
||||
self.llmcache.store(prompt, value_str, **store_kwargs)
|
||||
except Exception as e:
|
||||
print_verbose(
|
||||
f"Error setting {value_str or value} in the Redis semantic cache: {str(e)}"
|
||||
)
|
||||
print_verbose(f"Error setting {value_str or value} in the Redis semantic cache: {str(e)}")
|
||||
|
||||
def get_cache(self, key: str, **kwargs) -> Any:
|
||||
"""
|
||||
|
|
@ -439,9 +427,7 @@ 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")
|
||||
)
|
||||
prompt_embedding = self._get_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
check_kwargs: dict[str, Any] = {
|
||||
"prompt": prompt,
|
||||
"vector": prompt_embedding,
|
||||
|
|
@ -485,9 +471,7 @@ class RedisSemanticCache(BaseCache):
|
|||
print_verbose(f"Error retrieving from Redis semantic cache: {str(e)}")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
|
||||
async def _get_async_embedding(
|
||||
self, prompt: str, metadata: Dict[str, Any] | None = None
|
||||
) -> List[float]:
|
||||
async def _get_async_embedding(self, prompt: str, metadata: Dict[str, Any] | None = None) -> List[float]:
|
||||
"""
|
||||
Asynchronously generate an embedding for the given prompt.
|
||||
|
||||
|
|
@ -504,9 +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 = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
try:
|
||||
if router is not None:
|
||||
embedding_response = await router.aembedding(
|
||||
|
|
@ -547,9 +529,7 @@ class RedisSemanticCache(BaseCache):
|
|||
value_str = str(value)
|
||||
|
||||
# Generate embedding for the value (response) to cache
|
||||
prompt_embedding = await self._get_async_embedding(
|
||||
prompt, metadata=kwargs.get("metadata")
|
||||
)
|
||||
prompt_embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
store_kwargs: dict[str, Any] = {
|
||||
"vector": prompt_embedding,
|
||||
|
|
@ -589,9 +569,7 @@ class RedisSemanticCache(BaseCache):
|
|||
return None
|
||||
|
||||
# Generate embedding for the prompt
|
||||
prompt_embedding = await self._get_async_embedding(
|
||||
prompt, metadata=kwargs.get("metadata")
|
||||
)
|
||||
prompt_embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
|
||||
# Check the cache for semantically similar prompts in this exact
|
||||
# LiteLLM cache-key scope.
|
||||
|
|
@ -647,9 +625,7 @@ class RedisSemanticCache(BaseCache):
|
|||
aindex = 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:
|
||||
async def async_set_cache_pipeline(self, cache_list: List[Tuple[str, Any]], **kwargs) -> None:
|
||||
"""
|
||||
Asynchronously store multiple values in the semantic cache.
|
||||
|
||||
|
|
|
|||
|
|
@ -110,9 +110,7 @@ class S3Cache(BaseCache):
|
|||
func = partial(self.set_cache, key, value, **kwargs)
|
||||
await loop.run_in_executor(None, func)
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"S3 Caching: async_set_cache() - Got exception from S3: {e}"
|
||||
)
|
||||
verbose_logger.error(f"S3 Caching: async_set_cache() - Got exception from S3: {e}")
|
||||
|
||||
def get_cache(self, key, **kwargs):
|
||||
import botocore
|
||||
|
|
@ -122,9 +120,7 @@ class S3Cache(BaseCache):
|
|||
|
||||
print_verbose(f"Get S3 Cache: key: {key}")
|
||||
# Download the data from S3
|
||||
cached_response = self.s3_client.get_object(
|
||||
Bucket=self.bucket_name, Key=key
|
||||
)
|
||||
cached_response = self.s3_client.get_object(Bucket=self.bucket_name, Key=key)
|
||||
|
||||
if cached_response is not None:
|
||||
if "Expires" in cached_response:
|
||||
|
|
@ -135,13 +131,9 @@ class S3Cache(BaseCache):
|
|||
return None
|
||||
|
||||
# cached_response is in `b{} convert it to ModelResponse
|
||||
cached_response = (
|
||||
cached_response["Body"].read().decode("utf-8")
|
||||
) # Convert bytes to string
|
||||
cached_response = cached_response["Body"].read().decode("utf-8") # Convert bytes to string
|
||||
try:
|
||||
cached_response = json.loads(
|
||||
cached_response
|
||||
) # Convert string to dictionary
|
||||
cached_response = json.loads(cached_response) # Convert string to dictionary
|
||||
except Exception:
|
||||
cached_response = ast.literal_eval(cached_response)
|
||||
if not isinstance(cached_response, dict):
|
||||
|
|
@ -153,15 +145,11 @@ class S3Cache(BaseCache):
|
|||
return cached_response
|
||||
except botocore.exceptions.ClientError as e: # type: ignore
|
||||
if e.response["Error"]["Code"] == "NoSuchKey":
|
||||
verbose_logger.debug(
|
||||
f"S3 Cache: The specified key '{key}' does not exist in the S3 bucket."
|
||||
)
|
||||
verbose_logger.debug(f"S3 Cache: The specified key '{key}' does not exist in the S3 bucket.")
|
||||
return None
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"S3 Caching: get_cache() - Got exception from S3: {e}"
|
||||
)
|
||||
verbose_logger.error(f"S3 Caching: get_cache() - Got exception from S3: {e}")
|
||||
|
||||
async def async_get_cache(self, key, **kwargs):
|
||||
"""
|
||||
|
|
@ -175,9 +163,7 @@ class S3Cache(BaseCache):
|
|||
result = await loop.run_in_executor(None, func)
|
||||
return result
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"S3 Caching: async_get_cache() - Got exception from S3: {e}"
|
||||
)
|
||||
verbose_logger.error(f"S3 Caching: async_get_cache() - Got exception from S3: {e}")
|
||||
return None
|
||||
|
||||
def flush_cache(self):
|
||||
|
|
|
|||
|
|
@ -84,31 +84,21 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
|
||||
resolved_url = None
|
||||
if sync_client is None or async_client is None:
|
||||
resolved_url = redis_url or self._build_valkey_url(
|
||||
host, port, password, ssl
|
||||
)
|
||||
resolved_url = redis_url or self._build_valkey_url(host, port, password, ssl)
|
||||
self.sync_client = (
|
||||
sync_client if sync_client is not None else Redis.from_url(resolved_url) # type: ignore[arg-type]
|
||||
)
|
||||
self.async_client = (
|
||||
async_client
|
||||
if async_client is not None
|
||||
else AsyncRedis.from_url(resolved_url) # type: ignore[arg-type]
|
||||
async_client if async_client is not None else AsyncRedis.from_url(resolved_url) # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
print_verbose(f"Valkey semantic-cache initializing index - {self.index_name}")
|
||||
|
||||
@staticmethod
|
||||
def _build_valkey_url(
|
||||
host: str | None, port: str | None, password: str | None, ssl: bool = False
|
||||
) -> str:
|
||||
def _build_valkey_url(host: str | None, port: str | None, password: str | None, ssl: bool = False) -> str:
|
||||
host = host or os.environ.get("VALKEY_HOST") or os.environ.get("REDIS_HOST")
|
||||
port = port or os.environ.get("VALKEY_PORT") or os.environ.get("REDIS_PORT")
|
||||
password = (
|
||||
password
|
||||
or os.environ.get("VALKEY_PASSWORD")
|
||||
or os.environ.get("REDIS_PASSWORD")
|
||||
)
|
||||
password = password or os.environ.get("VALKEY_PASSWORD") or os.environ.get("REDIS_PASSWORD")
|
||||
|
||||
if not host or not port:
|
||||
raise ValueError(
|
||||
|
|
@ -157,11 +147,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
for field in info.get("attributes") or []:
|
||||
if not isinstance(field, (list, tuple)):
|
||||
continue
|
||||
flat = [
|
||||
sub
|
||||
for item in field
|
||||
for sub in (item if isinstance(item, (list, tuple)) else [item])
|
||||
]
|
||||
flat = [sub for item in field for sub in (item if isinstance(item, (list, tuple)) else [item])]
|
||||
for i, marker in enumerate(flat):
|
||||
if marker in (b"dimensions", "dimensions") and i + 1 < len(flat):
|
||||
return int(flat[i + 1])
|
||||
|
|
@ -207,9 +193,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
def _doc_key(self, key: str) -> str:
|
||||
return f"{self.key_prefix}{self._scope_tag(key)}:{uuid.uuid4()}"
|
||||
|
||||
def _doc_mapping(
|
||||
self, key: str, prompt: str, value_str: str, embedding: list[float]
|
||||
) -> dict:
|
||||
def _doc_mapping(self, key: str, prompt: str, value_str: str, embedding: list[float]) -> dict:
|
||||
return {
|
||||
self.CACHE_KEY_FIELD_NAME: self._scope_tag(key),
|
||||
self.PROMPT_FIELD_NAME: prompt,
|
||||
|
|
@ -223,11 +207,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
f"(@{self.CACHE_KEY_FIELD_NAME}:{{{scope}}})"
|
||||
f"=>[KNN 1 @{self.EMBEDDING_FIELD_NAME} $vec AS {self.DISTANCE_FIELD_NAME}]"
|
||||
)
|
||||
return (
|
||||
Query(query_string)
|
||||
.return_fields(self.RESPONSE_FIELD_NAME, self.DISTANCE_FIELD_NAME)
|
||||
.dialect(2)
|
||||
)
|
||||
return Query(query_string).return_fields(self.RESPONSE_FIELD_NAME, self.DISTANCE_FIELD_NAME).dialect(2)
|
||||
|
||||
@classmethod
|
||||
def _first_hit(cls, search_result: Any) -> _ValkeyCacheHit | None:
|
||||
|
|
@ -264,9 +244,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
self._ensure_index_sync(len(embedding))
|
||||
|
||||
doc_key = self._doc_key(key)
|
||||
self.sync_client.hset(
|
||||
doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding)
|
||||
)
|
||||
self.sync_client.hset(doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding))
|
||||
ttl = self._get_ttl(**kwargs)
|
||||
if ttl is not None:
|
||||
self.sync_client.expire(doc_key, ttl)
|
||||
|
|
@ -305,9 +283,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
await self._ensure_index_async(len(embedding))
|
||||
|
||||
doc_key = self._doc_key(key)
|
||||
await self.async_client.hset(
|
||||
doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding)
|
||||
)
|
||||
await self.async_client.hset(doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding))
|
||||
ttl = self._get_ttl(**kwargs)
|
||||
if ttl is not None:
|
||||
await self.async_client.expire(doc_key, ttl)
|
||||
|
|
@ -334,20 +310,11 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
print_verbose(f"Error in async Valkey semantic-cache get_cache: {str(e)}")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
|
||||
async def async_set_cache_pipeline(
|
||||
self, cache_list: list[tuple[str, Any]], **kwargs: Any
|
||||
) -> None:
|
||||
async def async_set_cache_pipeline(self, cache_list: list[tuple[str, Any]], **kwargs: Any) -> None:
|
||||
try:
|
||||
await asyncio.gather(
|
||||
*[
|
||||
self.async_set_cache(key, value, **kwargs)
|
||||
for key, value in cache_list
|
||||
]
|
||||
)
|
||||
await asyncio.gather(*[self.async_set_cache(key, value, **kwargs) for key, value in cache_list])
|
||||
except Exception as e:
|
||||
print_verbose(
|
||||
f"Error in Valkey semantic-cache async_set_cache_pipeline: {str(e)}"
|
||||
)
|
||||
print_verbose(f"Error in Valkey semantic-cache async_set_cache_pipeline: {str(e)}")
|
||||
|
||||
async def _index_info(self) -> dict:
|
||||
return await self.async_client.ft(self.index_name).info()
|
||||
|
|
|
|||
|
|
@ -41,10 +41,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
def _is_preformatted_cached_chat_stream(result: Any) -> bool:
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
return (
|
||||
isinstance(result, CustomStreamWrapper)
|
||||
and result.custom_llm_provider == "cached_response"
|
||||
)
|
||||
return isinstance(result, CustomStreamWrapper) and result.custom_llm_provider == "cached_response"
|
||||
|
||||
@staticmethod
|
||||
def _coerce_response_object(
|
||||
|
|
@ -85,9 +82,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
raise ValueError("Stream completed response is invalid")
|
||||
return response
|
||||
|
||||
async def _collect_response_from_stream_async(
|
||||
self, stream_iter: Any
|
||||
) -> "ResponsesAPIResponse":
|
||||
async def _collect_response_from_stream_async(self, stream_iter: Any) -> "ResponsesAPIResponse":
|
||||
async for _ in stream_iter:
|
||||
pass
|
||||
|
||||
|
|
@ -102,9 +97,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
raise ValueError("Stream completed response is invalid")
|
||||
return response
|
||||
|
||||
def validate_input_kwargs(
|
||||
self, kwargs: dict
|
||||
) -> ResponsesToCompletionBridgeHandlerInputKwargs:
|
||||
def validate_input_kwargs(self, kwargs: dict) -> ResponsesToCompletionBridgeHandlerInputKwargs:
|
||||
from litellm import LiteLLMLoggingObj
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
|
@ -234,9 +227,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
)
|
||||
else:
|
||||
if self._is_preformatted_cached_chat_stream(result):
|
||||
return self._apply_post_stream_processing(
|
||||
result, model, custom_llm_provider
|
||||
)
|
||||
return self._apply_post_stream_processing(result, model, custom_llm_provider)
|
||||
completion_stream = self.transformation_handler.get_model_response_iterator(
|
||||
streaming_response=result, # type: ignore
|
||||
sync_stream=True,
|
||||
|
|
@ -248,13 +239,9 @@ class ResponsesToCompletionBridgeHandler:
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
return self._apply_post_stream_processing(
|
||||
streamwrapper, model, custom_llm_provider
|
||||
)
|
||||
return self._apply_post_stream_processing(streamwrapper, model, custom_llm_provider)
|
||||
|
||||
async def acompletion(
|
||||
self, *args, **kwargs
|
||||
) -> Union["ModelResponse", "CustomStreamWrapper"]:
|
||||
async def acompletion(self, *args, **kwargs) -> Union["ModelResponse", "CustomStreamWrapper"]:
|
||||
from litellm import aresponses
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
|
|
@ -314,9 +301,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
elif isinstance(result, ModelResponse):
|
||||
return result
|
||||
elif not stream:
|
||||
responses_api_response = await self._collect_response_from_stream_async(
|
||||
result
|
||||
)
|
||||
responses_api_response = await self._collect_response_from_stream_async(result)
|
||||
return self.transformation_handler.transform_response(
|
||||
model=model,
|
||||
raw_response=responses_api_response,
|
||||
|
|
@ -332,9 +317,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
)
|
||||
else:
|
||||
if self._is_preformatted_cached_chat_stream(result):
|
||||
return self._apply_post_stream_processing(
|
||||
result, model, custom_llm_provider
|
||||
)
|
||||
return self._apply_post_stream_processing(result, model, custom_llm_provider)
|
||||
completion_stream = self.transformation_handler.get_model_response_iterator(
|
||||
streaming_response=result, # type: ignore
|
||||
sync_stream=False,
|
||||
|
|
@ -346,9 +329,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
return self._apply_post_stream_processing(
|
||||
streamwrapper, model, custom_llm_provider
|
||||
)
|
||||
return self._apply_post_stream_processing(streamwrapper, model, custom_llm_provider)
|
||||
|
||||
@staticmethod
|
||||
def _apply_post_stream_processing(
|
||||
|
|
|
|||
|
|
@ -83,9 +83,7 @@ def _build_reasoning_item(
|
|||
summary: 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", "")}
|
||||
)
|
||||
summary.append({"type": s.get("type", "summary_text"), "text": s.get("text", "")})
|
||||
else:
|
||||
summary.append(
|
||||
{
|
||||
|
|
@ -138,9 +136,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
return {"type": "function", "name": fn_name}
|
||||
return tool_choice
|
||||
|
||||
def _handle_raw_dict_response_item(
|
||||
self, item: Dict[str, Any], index: int
|
||||
) -> Tuple[Optional[Any], int]:
|
||||
def _handle_raw_dict_response_item(self, item: Dict[str, Any], index: int) -> Tuple[Optional[Any], int]:
|
||||
"""
|
||||
Handle raw dict response items from Responses API (e.g., GPT-5 Codex format).
|
||||
|
||||
|
|
@ -183,13 +179,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if item_type == "function_call":
|
||||
# Extract provider_specific_fields if present and pass through as-is
|
||||
provider_specific_fields = item.get("provider_specific_fields")
|
||||
if provider_specific_fields and not isinstance(
|
||||
provider_specific_fields, dict
|
||||
):
|
||||
if provider_specific_fields and not isinstance(provider_specific_fields, dict):
|
||||
provider_specific_fields = (
|
||||
dict(provider_specific_fields)
|
||||
if hasattr(provider_specific_fields, "__dict__")
|
||||
else {}
|
||||
dict(provider_specific_fields) if hasattr(provider_specific_fields, "__dict__") else {}
|
||||
)
|
||||
|
||||
tool_call_dict = {
|
||||
|
|
@ -205,9 +197,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if provider_specific_fields:
|
||||
tool_call_dict["provider_specific_fields"] = provider_specific_fields
|
||||
# Also add to function's provider_specific_fields for consistency
|
||||
tool_call_dict["function"]["provider_specific_fields"] = (
|
||||
provider_specific_fields
|
||||
)
|
||||
tool_call_dict["function"]["provider_specific_fields"] = provider_specific_fields
|
||||
|
||||
msg = Message(
|
||||
content=None,
|
||||
|
|
@ -301,9 +291,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
{
|
||||
"type": "message",
|
||||
"role": role,
|
||||
"content": self._convert_content_to_responses_format(
|
||||
content, cast(str, role)
|
||||
), # type: ignore[arg-type]
|
||||
"content": self._convert_content_to_responses_format(content, cast(str, role)), # type: ignore[arg-type]
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -321,10 +309,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if key in ("max_tokens", "max_completion_tokens"):
|
||||
responses_api_request["max_output_tokens"] = value
|
||||
elif key == "tools" and value is not None:
|
||||
responses_api_request["tools"] = (
|
||||
self._convert_tools_to_responses_format(
|
||||
cast(List[Dict[str, Any]], value)
|
||||
)
|
||||
responses_api_request["tools"] = self._convert_tools_to_responses_format(
|
||||
cast(List[Dict[str, Any]], value)
|
||||
)
|
||||
elif key == "response_format":
|
||||
text_format = self._transform_response_format_to_text_format(value)
|
||||
|
|
@ -345,13 +331,9 @@ 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()
|
||||
)
|
||||
responses_optional_param_keys = set(ResponsesAPIOptionalRequestParams.__annotations__.keys())
|
||||
sanitized: Dict[str, Any] = {
|
||||
key: value
|
||||
for key, value in litellm_params.items()
|
||||
if key not in responses_optional_param_keys
|
||||
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")
|
||||
|
|
@ -427,9 +409,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if instructions:
|
||||
responses_api_request["instructions"] = instructions
|
||||
|
||||
self._map_optional_params_to_responses_api_request(
|
||||
optional_params, responses_api_request
|
||||
)
|
||||
self._map_optional_params_to_responses_api_request(optional_params, responses_api_request)
|
||||
|
||||
stream = optional_params.get("stream") or litellm_params.get("stream", False)
|
||||
verbose_logger.debug(f"Chat provider: Stream parameter: {stream}")
|
||||
|
|
@ -442,9 +422,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
previous_response_id = optional_params.get("previous_response_id")
|
||||
if previous_response_id:
|
||||
# Use the existing session handler for responses API
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: Warning ignoring previous response ID: {previous_response_id}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: Warning ignoring previous response ID: {previous_response_id}")
|
||||
|
||||
# Convert back to responses API format for the actual request
|
||||
|
||||
|
|
@ -464,13 +442,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
"client": client,
|
||||
}
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: Final request model={api_model}, input_items={len(input_items)}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: Final request model={api_model}, input_items={len(input_items)}")
|
||||
|
||||
self._merge_responses_api_request_into_request_data(
|
||||
request_data, responses_api_request, instructions
|
||||
)
|
||||
self._merge_responses_api_request_into_request_data(request_data, responses_api_request, instructions)
|
||||
|
||||
if headers:
|
||||
request_data["extra_headers"] = headers
|
||||
|
|
@ -524,11 +498,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
encrypted_content=getattr(item, "encrypted_content", None),
|
||||
summary_raw=item.summary,
|
||||
)
|
||||
reasoning_content = " ".join(
|
||||
s["text"]
|
||||
for s in pending_reasoning_item["summary"]
|
||||
if s.get("text")
|
||||
)
|
||||
reasoning_content = " ".join(s["text"] for s in pending_reasoning_item["summary"] if s.get("text"))
|
||||
|
||||
elif isinstance(item, ResponseOutputMessage):
|
||||
for content in item.content:
|
||||
|
|
@ -545,11 +515,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
annotations=annotations,
|
||||
reasoning_items=cast(
|
||||
Optional[List[ChatCompletionReasoningItem]],
|
||||
(
|
||||
[pending_reasoning_item]
|
||||
if pending_reasoning_item is not None
|
||||
else None
|
||||
),
|
||||
([pending_reasoning_item] if pending_reasoning_item is not None else None),
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -570,23 +536,25 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
|
||||
tool_call_dict = LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
|
||||
tool_call_item=item,
|
||||
index=tool_call_index,
|
||||
tool_call_dict = (
|
||||
LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
|
||||
tool_call_item=item,
|
||||
index=tool_call_index,
|
||||
)
|
||||
)
|
||||
accumulated_tool_calls.append(tool_call_dict)
|
||||
tool_call_index += 1
|
||||
|
||||
elif ResponseApplyPatchToolCall is not None and isinstance(
|
||||
item, ResponseApplyPatchToolCall
|
||||
):
|
||||
elif ResponseApplyPatchToolCall is not None and isinstance(item, ResponseApplyPatchToolCall):
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
|
||||
tool_call_dict = LiteLLMCompletionResponsesConfig.convert_apply_patch_tool_call_to_chat_completion_tool_call(
|
||||
tool_call_item=item,
|
||||
index=tool_call_index,
|
||||
tool_call_dict = (
|
||||
LiteLLMCompletionResponsesConfig.convert_apply_patch_tool_call_to_chat_completion_tool_call(
|
||||
tool_call_item=item,
|
||||
index=tool_call_index,
|
||||
)
|
||||
)
|
||||
accumulated_tool_calls.append(tool_call_dict)
|
||||
tool_call_index += 1
|
||||
|
|
@ -607,25 +575,17 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
reasoning_content=reasoning_content,
|
||||
reasoning_items=cast(
|
||||
Optional[List[ChatCompletionReasoningItem]],
|
||||
(
|
||||
[pending_reasoning_item]
|
||||
if pending_reasoning_item is not None
|
||||
else None
|
||||
),
|
||||
([pending_reasoning_item] if pending_reasoning_item is not None else None),
|
||||
),
|
||||
)
|
||||
choices.append(
|
||||
Choices(message=msg, finish_reason="tool_calls", index=index)
|
||||
)
|
||||
choices.append(Choices(message=msg, finish_reason="tool_calls", index=index))
|
||||
reasoning_content = None
|
||||
pending_reasoning_item = None
|
||||
|
||||
return choices
|
||||
|
||||
@classmethod
|
||||
def _extract_output_from_completed_event(
|
||||
cls, parsed_chunk: Dict[str, Any]
|
||||
) -> Optional[List[Dict[str, Any]]]:
|
||||
def _extract_output_from_completed_event(cls, parsed_chunk: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]:
|
||||
response_payload = parsed_chunk.get("response")
|
||||
if not isinstance(response_payload, dict):
|
||||
return None
|
||||
|
|
@ -635,9 +595,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
return cast(List[Dict[str, Any]], response_output)
|
||||
|
||||
@classmethod
|
||||
def _recover_output_items_from_raw_sse(
|
||||
cls, raw_sse: Optional[str]
|
||||
) -> List[Dict[str, Any]]:
|
||||
def _recover_output_items_from_raw_sse(cls, raw_sse: Optional[str]) -> List[Dict[str, Any]]:
|
||||
if not raw_sse or not isinstance(raw_sse, str):
|
||||
return []
|
||||
|
||||
|
|
@ -652,9 +610,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
event_type = parsed_chunk.get("type")
|
||||
|
||||
if event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
recovered_output = cls._extract_output_from_completed_event(
|
||||
parsed_chunk
|
||||
)
|
||||
recovered_output = cls._extract_output_from_completed_event(parsed_chunk)
|
||||
if recovered_output is not None:
|
||||
return recovered_output
|
||||
continue
|
||||
|
|
@ -688,9 +644,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
return []
|
||||
|
||||
@classmethod
|
||||
def _recover_output_items_from_logging(
|
||||
cls, logging_obj: "LiteLLMLoggingObj"
|
||||
) -> List[Dict[str, Any]]:
|
||||
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")
|
||||
return cls._recover_output_items_from_raw_sse(original_response)
|
||||
|
|
@ -721,9 +675,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 = 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)
|
||||
|
|
@ -739,17 +691,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
)
|
||||
|
||||
if len(choices) == 0:
|
||||
if (
|
||||
raw_response.incomplete_details is not None
|
||||
and raw_response.incomplete_details.reason is not None
|
||||
):
|
||||
raise ValueError(
|
||||
f"{model} unable to complete request: {raw_response.incomplete_details.reason}"
|
||||
)
|
||||
if raw_response.incomplete_details is not None and raw_response.incomplete_details.reason is not None:
|
||||
raise ValueError(f"{model} unable to complete request: {raw_response.incomplete_details.reason}")
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown items in responses API response: {output_items}"
|
||||
)
|
||||
raise ValueError(f"Unknown items in responses API response: {output_items}")
|
||||
|
||||
setattr(model_response, "choices", choices)
|
||||
|
||||
|
|
@ -758,28 +703,21 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
setattr(
|
||||
model_response,
|
||||
"usage",
|
||||
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
raw_response.usage
|
||||
),
|
||||
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_response.usage),
|
||||
)
|
||||
|
||||
# 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", {})
|
||||
if raw_response_hidden_params:
|
||||
if (
|
||||
not hasattr(model_response, "_hidden_params")
|
||||
or model_response._hidden_params is None
|
||||
):
|
||||
if not hasattr(model_response, "_hidden_params") or model_response._hidden_params is None:
|
||||
model_response._hidden_params = {}
|
||||
# Merge the raw_response hidden params with model_response hidden params
|
||||
# Preserve existing keys in model_response but add/override with raw_response params
|
||||
for key, value in raw_response_hidden_params.items():
|
||||
if key == "additional_headers" and key in model_response._hidden_params:
|
||||
# Merge additional_headers to preserve both sets
|
||||
existing_additional_headers = model_response._hidden_params.get(
|
||||
"additional_headers", {}
|
||||
)
|
||||
existing_additional_headers = model_response._hidden_params.get("additional_headers", {})
|
||||
merged_headers = {**value, **existing_additional_headers}
|
||||
model_response._hidden_params[key] = merged_headers
|
||||
else:
|
||||
|
|
@ -789,19 +727,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
def get_model_response_iterator(
|
||||
self,
|
||||
streaming_response: Union[
|
||||
Iterator[str], AsyncIterator[str], "ModelResponse", "BaseModel"
|
||||
],
|
||||
streaming_response: Union[Iterator[str], AsyncIterator[str], "ModelResponse", "BaseModel"],
|
||||
sync_stream: bool,
|
||||
json_mode: Optional[bool] = False,
|
||||
) -> BaseModelResponseIterator:
|
||||
return OpenAiResponsesToChatCompletionStreamIterator(
|
||||
streaming_response, sync_stream, json_mode
|
||||
)
|
||||
return OpenAiResponsesToChatCompletionStreamIterator(streaming_response, sync_stream, json_mode)
|
||||
|
||||
def _convert_content_str_to_input_text(
|
||||
self, content: str, role: str
|
||||
) -> Dict[str, Any]:
|
||||
def _convert_content_str_to_input_text(self, content: str, role: str) -> Dict[str, Any]:
|
||||
if role == "user" or role == "system" or role == "tool":
|
||||
return {"type": "input_text", "text": content}
|
||||
else:
|
||||
|
|
@ -828,9 +760,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 = ResponseInputImageParam(image_url=actual_image_url, detail="auto", type="input_image")
|
||||
|
||||
if detail:
|
||||
image_param["detail"] = detail
|
||||
|
|
@ -857,9 +787,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
"""Convert chat completion content to responses API format"""
|
||||
from litellm.types.llms.openai import ChatCompletionImageObject
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: Converting content to responses format - input type: {type(content)}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: Converting content to responses format - input type: {type(content)}")
|
||||
|
||||
if content is None:
|
||||
return [self._convert_content_str_to_input_text("", role)]
|
||||
|
|
@ -870,9 +798,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
elif isinstance(content, list):
|
||||
result = []
|
||||
for i, item in enumerate(content):
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: Processing content item {i}: {type(item)} = {item}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: Processing content item {i}: {type(item)} = {item}")
|
||||
if isinstance(item, str):
|
||||
converted = self._convert_content_str_to_input_text(item, role)
|
||||
result.append(converted)
|
||||
|
|
@ -881,9 +807,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
# Handle multimodal content
|
||||
original_type = item.get("type")
|
||||
if original_type == "text":
|
||||
converted = self._convert_content_str_to_input_text(
|
||||
item.get("text", ""), role
|
||||
)
|
||||
converted = self._convert_content_str_to_input_text(item.get("text", ""), role)
|
||||
result.append(converted)
|
||||
verbose_logger.debug(f"Chat provider: text -> {converted}")
|
||||
elif original_type == "image_url":
|
||||
|
|
@ -895,18 +819,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
),
|
||||
)
|
||||
result.append(converted)
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: image_url -> {converted}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: image_url -> {converted}")
|
||||
else:
|
||||
# Try to map other types to responses API format
|
||||
item_type = original_type or "input_text"
|
||||
if item_type == "image":
|
||||
converted = {"type": "input_image", **item}
|
||||
result.append(converted)
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: image -> {converted}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: image -> {converted}")
|
||||
elif item_type == "file":
|
||||
# Map Chat Completion file to Responses API input_file
|
||||
# {"type": "file", "file": {"file_data": "...", "filename": "..."}}
|
||||
|
|
@ -918,9 +838,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if key in file_data:
|
||||
converted[key] = file_data[key]
|
||||
result.append(converted)
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: file -> {converted}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: file -> {converted}")
|
||||
elif item_type in [
|
||||
"input_text",
|
||||
"input_image",
|
||||
|
|
@ -932,18 +850,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
]:
|
||||
# Already in responses API format
|
||||
result.append(item)
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: passthrough -> {item}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: passthrough -> {item}")
|
||||
else:
|
||||
# Default to input_text for unknown types
|
||||
converted = self._convert_content_str_to_input_text(
|
||||
str(item.get("text", item)), role
|
||||
)
|
||||
converted = self._convert_content_str_to_input_text(str(item.get("text", item)), role)
|
||||
result.append(converted)
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: unknown({original_type}) -> {converted}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: unknown({original_type}) -> {converted}")
|
||||
verbose_logger.debug(f"Chat provider: Final converted content: {result}")
|
||||
return result
|
||||
else:
|
||||
|
|
@ -951,17 +863,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
verbose_logger.debug(f"Chat provider: Other content type -> {result}")
|
||||
return result
|
||||
|
||||
def _convert_tools_to_responses_format(
|
||||
self, tools: List[Dict[str, Any]]
|
||||
) -> List["ALL_RESPONSES_API_TOOL_PARAMS"]:
|
||||
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"] = []
|
||||
for tool in tools:
|
||||
# convert function tool from chat completion to responses API format
|
||||
if tool.get("type") == "function":
|
||||
function_tool = cast(
|
||||
ChatCompletionToolParamFunctionChunk, tool.get("function")
|
||||
)
|
||||
function_tool = cast(ChatCompletionToolParamFunctionChunk, tool.get("function"))
|
||||
responses_tools.append(
|
||||
FunctionToolParam(
|
||||
name=function_tool["name"],
|
||||
|
|
@ -987,9 +895,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if not extra_body:
|
||||
return optional_params
|
||||
|
||||
supported_responses_api_params = set(
|
||||
ResponsesAPIOptionalRequestParams.__annotations__.keys()
|
||||
)
|
||||
supported_responses_api_params = set(ResponsesAPIOptionalRequestParams.__annotations__.keys())
|
||||
# Also include params we handle specially
|
||||
supported_responses_api_params.update(
|
||||
{
|
||||
|
|
@ -1007,9 +913,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
return optional_params
|
||||
|
||||
def _map_reasoning_effort(
|
||||
self, reasoning_effort: Union[str, Dict[str, Any]]
|
||||
) -> Optional[Reasoning]:
|
||||
def _map_reasoning_effort(self, reasoning_effort: Union[str, Dict[str, Any]]) -> Optional[Reasoning]:
|
||||
# If dict is passed, convert it directly to Reasoning object
|
||||
if isinstance(reasoning_effort, dict):
|
||||
return Reasoning(**reasoning_effort) # type: ignore[typeddict-item]
|
||||
|
|
@ -1017,46 +921,25 @@ 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 = (
|
||||
litellm.reasoning_auto_summary
|
||||
or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true"
|
||||
litellm.reasoning_auto_summary or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true"
|
||||
)
|
||||
|
||||
# If string is passed, map with optional summary based on flag/env var
|
||||
if reasoning_effort == "none":
|
||||
return (
|
||||
Reasoning(effort="none", summary="detailed")
|
||||
if auto_summary_enabled
|
||||
else Reasoning(effort="none")
|
||||
) # type: ignore
|
||||
return Reasoning(effort="none", summary="detailed") if auto_summary_enabled else Reasoning(effort="none") # type: ignore
|
||||
elif reasoning_effort == "high":
|
||||
return (
|
||||
Reasoning(effort="high", summary="detailed")
|
||||
if auto_summary_enabled
|
||||
else Reasoning(effort="high")
|
||||
)
|
||||
return Reasoning(effort="high", summary="detailed") if auto_summary_enabled else Reasoning(effort="high")
|
||||
elif reasoning_effort == "xhigh":
|
||||
return (
|
||||
Reasoning(effort="xhigh", summary="detailed")
|
||||
if auto_summary_enabled
|
||||
else Reasoning(effort="xhigh")
|
||||
) # type: ignore[typeddict-item]
|
||||
return Reasoning(effort="xhigh", summary="detailed") if auto_summary_enabled else Reasoning(effort="xhigh") # type: ignore[typeddict-item]
|
||||
elif reasoning_effort == "medium":
|
||||
return (
|
||||
Reasoning(effort="medium", summary="detailed")
|
||||
if auto_summary_enabled
|
||||
else Reasoning(effort="medium")
|
||||
Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium")
|
||||
)
|
||||
elif reasoning_effort == "low":
|
||||
return (
|
||||
Reasoning(effort="low", summary="detailed")
|
||||
if auto_summary_enabled
|
||||
else Reasoning(effort="low")
|
||||
)
|
||||
return Reasoning(effort="low", summary="detailed") if auto_summary_enabled else Reasoning(effort="low")
|
||||
elif reasoning_effort == "minimal":
|
||||
return (
|
||||
Reasoning(effort="minimal", summary="detailed")
|
||||
if auto_summary_enabled
|
||||
else Reasoning(effort="minimal")
|
||||
Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal")
|
||||
)
|
||||
return None
|
||||
|
||||
|
|
@ -1072,10 +955,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
responses_api_request: The responses API request dict to modify
|
||||
web_search_options: Web search configuration (dict or other value)
|
||||
"""
|
||||
if (
|
||||
"tools" not in responses_api_request
|
||||
or responses_api_request["tools"] is None
|
||||
):
|
||||
if "tools" not in responses_api_request or responses_api_request["tools"] is None:
|
||||
responses_api_request["tools"] = []
|
||||
|
||||
# Get the tools list with proper type narrowing
|
||||
|
|
@ -1165,17 +1045,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
annotation_dict = annotation
|
||||
else:
|
||||
# Skip unsupported annotation types
|
||||
verbose_logger.debug(
|
||||
f"Skipping unsupported annotation type: {type(annotation)}"
|
||||
)
|
||||
verbose_logger.debug(f"Skipping unsupported annotation type: {type(annotation)}")
|
||||
continue
|
||||
|
||||
result.append(annotation_dict) # type: ignore
|
||||
except Exception as e:
|
||||
# Skip malformed annotations
|
||||
verbose_logger.debug(
|
||||
f"Skipping malformed annotation: {annotation}, error: {e}"
|
||||
)
|
||||
verbose_logger.debug(f"Skipping malformed annotation: {annotation}, error: {e}")
|
||||
continue
|
||||
|
||||
return result if result else None
|
||||
|
|
@ -1196,9 +1072,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
|
||||
class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
||||
def __init__(
|
||||
self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False
|
||||
):
|
||||
def __init__(self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False):
|
||||
super().__init__(streaming_response, sync_stream, json_mode)
|
||||
|
||||
def _handle_string_chunk(
|
||||
|
|
@ -1211,9 +1085,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
|
||||
)
|
||||
return GenericStreamingChunk(text="", tool_use=None, is_finished=False, finish_reason="", usage=None)
|
||||
index = str_line.find("data:")
|
||||
if index != -1:
|
||||
str_line = str_line[index + 5 :]
|
||||
|
|
@ -1258,9 +1130,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
event_type = event_type.value
|
||||
|
||||
if parsed_chunk.get("object") == "chat.completion.chunk" or (
|
||||
event_type is None
|
||||
and isinstance(parsed_chunk.get("choices"), list)
|
||||
and parsed_chunk.get("choices")
|
||||
event_type is None and isinstance(parsed_chunk.get("choices"), list) and parsed_chunk.get("choices")
|
||||
):
|
||||
return ModelResponseStream(**parsed_chunk)
|
||||
|
||||
|
|
@ -1284,13 +1154,9 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
if output_item.get("type") == "function_call":
|
||||
# Extract provider_specific_fields if present
|
||||
provider_specific_fields = output_item.get("provider_specific_fields")
|
||||
if provider_specific_fields and not isinstance(
|
||||
provider_specific_fields, dict
|
||||
):
|
||||
if provider_specific_fields and not isinstance(provider_specific_fields, dict):
|
||||
provider_specific_fields = (
|
||||
dict(provider_specific_fields)
|
||||
if hasattr(provider_specific_fields, "__dict__")
|
||||
else {}
|
||||
dict(provider_specific_fields) if hasattr(provider_specific_fields, "__dict__") else {}
|
||||
)
|
||||
|
||||
function_chunk = ChatCompletionToolCallFunctionChunk(
|
||||
|
|
@ -1299,9 +1165,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
)
|
||||
|
||||
if provider_specific_fields:
|
||||
function_chunk["provider_specific_fields"] = (
|
||||
provider_specific_fields
|
||||
)
|
||||
function_chunk["provider_specific_fields"] = provider_specific_fields
|
||||
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
|
|
@ -1344,9 +1208,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
id=None,
|
||||
index=tool_call_index,
|
||||
type="function",
|
||||
function=ChatCompletionToolCallFunctionChunk(
|
||||
name=None, arguments=content_part
|
||||
),
|
||||
function=ChatCompletionToolCallFunctionChunk(name=None, arguments=content_part),
|
||||
)
|
||||
]
|
||||
),
|
||||
|
|
@ -1355,22 +1217,16 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
]
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Chat provider: Invalid function argument delta {parsed_chunk}"
|
||||
)
|
||||
raise ValueError(f"Chat provider: Invalid function argument delta {parsed_chunk}")
|
||||
elif event_type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE:
|
||||
# New output item added
|
||||
output_item = parsed_chunk.get("item", {})
|
||||
if output_item.get("type") == "function_call":
|
||||
# Extract provider_specific_fields if present
|
||||
provider_specific_fields = output_item.get("provider_specific_fields")
|
||||
if provider_specific_fields and not isinstance(
|
||||
provider_specific_fields, dict
|
||||
):
|
||||
if provider_specific_fields and not isinstance(provider_specific_fields, dict):
|
||||
provider_specific_fields = (
|
||||
dict(provider_specific_fields)
|
||||
if hasattr(provider_specific_fields, "__dict__")
|
||||
else {}
|
||||
dict(provider_specific_fields) if hasattr(provider_specific_fields, "__dict__") else {}
|
||||
)
|
||||
|
||||
function_chunk = ChatCompletionToolCallFunctionChunk(
|
||||
|
|
@ -1380,9 +1236,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
|
||||
# Add provider_specific_fields to function if present
|
||||
if provider_specific_fields:
|
||||
function_chunk["provider_specific_fields"] = (
|
||||
provider_specific_fields
|
||||
)
|
||||
function_chunk["provider_specific_fields"] = provider_specific_fields
|
||||
|
||||
tool_call_index = parsed_chunk.get("output_index", 0)
|
||||
tool_call_chunk = ChatCompletionToolCallChunk(
|
||||
|
|
@ -1458,9 +1312,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
output_items = response_data.get("output", []) if response_data else []
|
||||
|
||||
has_function_calls = any(
|
||||
item.get("type") == "function_call"
|
||||
for item in output_items
|
||||
if isinstance(item, dict)
|
||||
item.get("type") == "function_call" for item in output_items if isinstance(item, dict)
|
||||
)
|
||||
|
||||
finish_reason = "tool_calls" if has_function_calls else "stop"
|
||||
|
|
@ -1488,11 +1340,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
if response_data.get("usage"):
|
||||
from litellm.responses.utils import ResponseAPILoggingUtils
|
||||
|
||||
usage = (
|
||||
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
response_data.get("usage")
|
||||
)
|
||||
)
|
||||
usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(response_data.get("usage"))
|
||||
return ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
|
|
@ -1509,9 +1357,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
else:
|
||||
pass
|
||||
# For any unhandled event types, create a minimal valid chunk or skip
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: Unhandled event type '{event_type}', creating empty chunk"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: Unhandled event type '{event_type}', creating empty chunk")
|
||||
|
||||
# Return a minimal valid chunk for unknown events
|
||||
return ModelResponseStream(
|
||||
|
|
@ -1534,9 +1380,5 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
Returns:
|
||||
ModelResponseStream: OpenAI-formatted streaming chunk
|
||||
"""
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: transform_streaming_response called with chunk: {chunk}"
|
||||
)
|
||||
return OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(
|
||||
chunk
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: transform_streaming_response called with chunk: {chunk}")
|
||||
return OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(chunk)
|
||||
|
|
|
|||
|
|
@ -107,8 +107,7 @@ def _normalize_messages_for_compression(
|
|||
"""
|
||||
if call_type not in _SUPPORTED_CALL_TYPES:
|
||||
raise ValueError(
|
||||
f"Unsupported call_type={call_type!r} for compression. "
|
||||
f"Expected one of: {sorted(_SUPPORTED_CALL_TYPES)}."
|
||||
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]
|
||||
|
|
@ -334,9 +333,7 @@ def _select_kept_indices_for_budget(
|
|||
return kept_indices, truncated_overrides
|
||||
|
||||
|
||||
def _get_dropped_tool_span_indices(
|
||||
kept_indices: Set[int], tool_exchange_spans: List[Set[int]]
|
||||
) -> Set[int]:
|
||||
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()
|
||||
for span in tool_exchange_spans:
|
||||
if not any(idx in kept_indices for idx in span):
|
||||
|
|
@ -440,9 +437,7 @@ def compress(
|
|||
|
||||
tool_exchange_spans: List[Set[int]] = []
|
||||
if _is_anthropic_call_type(call_type_str):
|
||||
tool_exchange_spans, tool_sequence_error = (
|
||||
_extract_anthropic_tool_exchange_spans(original_messages)
|
||||
)
|
||||
tool_exchange_spans, tool_sequence_error = _extract_anthropic_tool_exchange_spans(original_messages)
|
||||
if tool_sequence_error is not None:
|
||||
return CompressedResult(
|
||||
messages=original_messages,
|
||||
|
|
@ -484,9 +479,7 @@ def compress(
|
|||
# Use the truncated version if we made one, otherwise the original
|
||||
compressed_messages.append(truncated_overrides.get(i, msg))
|
||||
else:
|
||||
key = extract_key(
|
||||
normalized_messages[i], fallback_index=i, used_keys=used_keys
|
||||
)
|
||||
key = extract_key(normalized_messages[i], fallback_index=i, used_keys=used_keys)
|
||||
content = _content_to_text(msg.get("content", ""))
|
||||
cache[key] = content
|
||||
compressed_messages.append(stub_message(msg, key))
|
||||
|
|
@ -503,11 +496,7 @@ def compress(
|
|||
messages=compressed_messages,
|
||||
original_tokens=original_tokens,
|
||||
compressed_tokens=compressed_tokens,
|
||||
compression_ratio=(
|
||||
round(1 - (compressed_tokens / original_tokens), 4)
|
||||
if original_tokens > 0
|
||||
else 0.0
|
||||
),
|
||||
compression_ratio=(round(1 - (compressed_tokens / original_tokens), 4) if original_tokens > 0 else 0.0),
|
||||
cache=cache,
|
||||
tools=tools,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -33,9 +33,7 @@ def detect_content_type(content: str) -> str:
|
|||
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()
|
||||
)
|
||||
indented_lines = 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):
|
||||
|
|
|
|||
|
|
@ -26,9 +26,7 @@ def extract_key(message: dict, fallback_index: int, used_keys: Set[str]) -> str:
|
|||
"""
|
||||
content = message.get("content", "")
|
||||
if isinstance(content, list):
|
||||
content = " ".join(
|
||||
p.get("text", "") if isinstance(p, dict) else str(p) for p in content
|
||||
)
|
||||
content = " ".join(p.get("text", "") if isinstance(p, dict) else str(p) for p in content)
|
||||
|
||||
key = None
|
||||
for pattern in _FILE_PATH_PATTERNS:
|
||||
|
|
@ -62,9 +60,7 @@ def stub_message(message: dict, key: str) -> dict:
|
|||
"""
|
||||
content = message.get("content", "")
|
||||
if isinstance(content, list):
|
||||
content = " ".join(
|
||||
p.get("text", "") if isinstance(p, dict) else str(p) for p in content
|
||||
)
|
||||
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)
|
||||
|
|
@ -91,9 +87,7 @@ def truncate_message(message: dict, max_tokens: int) -> dict:
|
|||
"""
|
||||
content = message.get("content", "")
|
||||
if isinstance(content, list):
|
||||
content = " ".join(
|
||||
p.get("text", "") if isinstance(p, dict) else str(p) for p in content
|
||||
)
|
||||
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)
|
||||
|
|
@ -113,8 +107,6 @@ def truncate_message(message: dict, max_tokens: int) -> dict:
|
|||
first_count = (target_lines * 7) // 10
|
||||
last_count = target_lines - first_count
|
||||
truncated = (
|
||||
"\n".join(lines[:first_count])
|
||||
+ "\n...[truncated for context window]...\n"
|
||||
+ "\n".join(lines[-last_count:])
|
||||
"\n".join(lines[:first_count]) + "\n...[truncated for context window]...\n" + "\n".join(lines[-last_count:])
|
||||
)
|
||||
return {**message, "content": truncated}
|
||||
|
|
|
|||
|
|
@ -17,8 +17,7 @@ def build_retrieval_tool(available_keys: List[str]) -> dict:
|
|||
"description": (
|
||||
"Retrieve the full content of a file or message that was "
|
||||
"compressed to save tokens. Use this when you need the complete "
|
||||
"content to answer accurately. Available keys: "
|
||||
+ ", ".join(available_keys)
|
||||
"content to answer accurately. Available keys: " + ", ".join(available_keys)
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
|
|
|
|||
|
|
@ -91,11 +91,7 @@ def bm25_score_messages(
|
|||
return exact
|
||||
if len(query_term) < 4:
|
||||
return 0
|
||||
return sum(
|
||||
count
|
||||
for token, count in tf_counts.items()
|
||||
if token != query_term and token.startswith(query_term)
|
||||
)
|
||||
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] = []
|
||||
|
|
|
|||
|
|
@ -4,28 +4,16 @@ from typing import List, Literal, Optional
|
|||
|
||||
from litellm.litellm_core_utils.env_utils import get_env_int
|
||||
|
||||
DEFAULT_HEALTH_CHECK_PROMPT = str(
|
||||
os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm")
|
||||
)
|
||||
AZURE_DEFAULT_RESPONSES_API_VERSION = str(
|
||||
os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview")
|
||||
)
|
||||
DEFAULT_HEALTH_CHECK_PROMPT = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"))
|
||||
AZURE_DEFAULT_RESPONSES_API_VERSION = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview"))
|
||||
ROUTER_MAX_FALLBACKS = int(os.getenv("ROUTER_MAX_FALLBACKS", 5))
|
||||
DEFAULT_BATCH_SIZE = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
|
||||
DEFAULT_FLUSH_INTERVAL_SECONDS = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5))
|
||||
DEFAULT_S3_FLUSH_INTERVAL_SECONDS = int(
|
||||
os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10)
|
||||
)
|
||||
DEFAULT_S3_FLUSH_INTERVAL_SECONDS = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10))
|
||||
DEFAULT_S3_BATCH_SIZE = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512))
|
||||
DEFAULT_SQS_FLUSH_INTERVAL_SECONDS = int(
|
||||
os.getenv("DEFAULT_SQS_FLUSH_INTERVAL_SECONDS", 10)
|
||||
)
|
||||
DEFAULT_NUM_WORKERS_LITELLM_PROXY = int(
|
||||
os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", 1)
|
||||
)
|
||||
DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE = int(
|
||||
os.getenv("DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE", 1)
|
||||
)
|
||||
DEFAULT_SQS_FLUSH_INTERVAL_SECONDS = int(os.getenv("DEFAULT_SQS_FLUSH_INTERVAL_SECONDS", 10))
|
||||
DEFAULT_NUM_WORKERS_LITELLM_PROXY = int(os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", 1))
|
||||
DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE = int(os.getenv("DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE", 1))
|
||||
DEFAULT_SQS_BATCH_SIZE = int(os.getenv("DEFAULT_SQS_BATCH_SIZE", 512))
|
||||
SQS_SEND_MESSAGE_ACTION = "SendMessage"
|
||||
SQS_API_VERSION = "2012-11-05"
|
||||
|
|
@ -34,9 +22,7 @@ DEFAULT_MAX_RETRIES = int(os.getenv("DEFAULT_MAX_RETRIES", 2))
|
|||
# radius: each record fans out to spend logs + every callback integration.
|
||||
MAX_CALLBACK_LOG_RECORDS = 1000
|
||||
DEFAULT_MAX_RECURSE_DEPTH = int(os.getenv("DEFAULT_MAX_RECURSE_DEPTH", 100))
|
||||
DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER = int(
|
||||
os.getenv("DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER", 10)
|
||||
)
|
||||
DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER = int(os.getenv("DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER", 10))
|
||||
DEFAULT_FAILURE_THRESHOLD_PERCENT = float(
|
||||
os.getenv("DEFAULT_FAILURE_THRESHOLD_PERCENT", 0.5)
|
||||
) # default cooldown a deployment if 50% of requests fail in a given minute
|
||||
|
|
@ -44,12 +30,8 @@ DEFAULT_MAX_TOKENS = int(os.getenv("DEFAULT_MAX_TOKENS", 4096))
|
|||
DEFAULT_ALLOWED_FAILS = int(os.getenv("DEFAULT_ALLOWED_FAILS", 3))
|
||||
DEFAULT_REDIS_SYNC_INTERVAL = int(os.getenv("DEFAULT_REDIS_SYNC_INTERVAL", 1))
|
||||
DEFAULT_COOLDOWN_TIME_SECONDS = int(os.getenv("DEFAULT_COOLDOWN_TIME_SECONDS", 5))
|
||||
DEFAULT_REPLICATE_POLLING_RETRIES = int(
|
||||
os.getenv("DEFAULT_REPLICATE_POLLING_RETRIES", 5)
|
||||
)
|
||||
DEFAULT_REPLICATE_POLLING_DELAY_SECONDS = int(
|
||||
os.getenv("DEFAULT_REPLICATE_POLLING_DELAY_SECONDS", 1)
|
||||
)
|
||||
DEFAULT_REPLICATE_POLLING_RETRIES = int(os.getenv("DEFAULT_REPLICATE_POLLING_RETRIES", 5))
|
||||
DEFAULT_REPLICATE_POLLING_DELAY_SECONDS = int(os.getenv("DEFAULT_REPLICATE_POLLING_DELAY_SECONDS", 1))
|
||||
DEFAULT_IMAGE_TOKEN_COUNT = int(os.getenv("DEFAULT_IMAGE_TOKEN_COUNT", 250))
|
||||
|
||||
# Maximum wall-clock seconds a streaming response is allowed to run.
|
||||
|
|
@ -67,9 +49,7 @@ MAX_BASE64_LENGTH_FOR_LOGGING = int(os.getenv("MAX_BASE64_LENGTH_FOR_LOGGING", 6
|
|||
|
||||
# When true, adds detailed per-phase timing breakdown headers to responses.
|
||||
# Headers: x-litellm-timing-{pre-processing,llm-api,post-processing,message-copy}-ms
|
||||
LITELLM_DETAILED_TIMING = (
|
||||
os.getenv("LITELLM_DETAILED_TIMING", "false").lower() == "true"
|
||||
)
|
||||
LITELLM_DETAILED_TIMING = os.getenv("LITELLM_DETAILED_TIMING", "false").lower() == "true"
|
||||
|
||||
# Model cost map validation constants
|
||||
MODEL_COST_MAP_MIN_MODEL_COUNT = int(
|
||||
|
|
@ -91,9 +71,7 @@ MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB = int(
|
|||
# Surrogate-repair fallback in _read_request_body runs two full-body re.sub passes
|
||||
# that block the event loop on multi-MB malformed bodies. Skip the repair above this
|
||||
# size and raise the existing 400 immediately. Set to 0 to disable the cap.
|
||||
MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB = get_env_int(
|
||||
"MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB", 1
|
||||
)
|
||||
MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB = get_env_int("MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB", 1)
|
||||
SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD = int(
|
||||
os.getenv("SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD", 1000)
|
||||
) # Minimum number of requests to consider "reasonable traffic". Used for single-deployment cooldown logic.
|
||||
|
|
@ -101,42 +79,28 @@ DEFAULT_FAILURE_THRESHOLD_MINIMUM_REQUESTS = int(
|
|||
os.getenv("DEFAULT_FAILURE_THRESHOLD_MINIMUM_REQUESTS", 5)
|
||||
) # Minimum number of requests before applying error rate cooldown. Prevents cooldown from triggering on first failure.
|
||||
|
||||
DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET = int(
|
||||
os.getenv("DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET", 0)
|
||||
)
|
||||
DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET = int(os.getenv("DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET", 0))
|
||||
|
||||
# MCP Semantic Tool Filter Defaults
|
||||
DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL = str(
|
||||
os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL", "text-embedding-3-small")
|
||||
)
|
||||
DEFAULT_MCP_SEMANTIC_FILTER_TOP_K = int(
|
||||
os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_TOP_K", 10)
|
||||
)
|
||||
DEFAULT_MCP_SEMANTIC_FILTER_TOP_K = int(os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_TOP_K", 10))
|
||||
DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD = float(
|
||||
os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD", 0.3)
|
||||
)
|
||||
MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH = int(
|
||||
os.getenv("MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH", 150)
|
||||
)
|
||||
MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH = int(os.getenv("MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH", 150))
|
||||
|
||||
# Semantic Guard Defaults
|
||||
DEFAULT_SEMANTIC_GUARD_EMBEDDING_MODEL = str(
|
||||
os.getenv("DEFAULT_SEMANTIC_GUARD_EMBEDDING_MODEL", "text-embedding-3-small")
|
||||
)
|
||||
DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD = float(
|
||||
os.getenv("DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD", 0.75)
|
||||
)
|
||||
DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD = float(os.getenv("DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD", 0.75))
|
||||
|
||||
# MCP OAuth2 Client Credentials Defaults
|
||||
MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS = int(
|
||||
os.getenv("MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS", "60")
|
||||
)
|
||||
MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE = int(
|
||||
os.getenv("MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE", "200")
|
||||
)
|
||||
MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL = int(
|
||||
os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600")
|
||||
)
|
||||
MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS = int(os.getenv("MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS", "60"))
|
||||
MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE", "200"))
|
||||
MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600"))
|
||||
|
||||
# Default npm cache directory for STDIO MCP servers.
|
||||
# npm/npx needs a writable cache dir; in containers the default (~/.npm)
|
||||
|
|
@ -149,9 +113,7 @@ MCP_PER_USER_TOKEN_REDIS_KEY_PREFIX = "mcp:per_user_token"
|
|||
MCP_PER_USER_TOKEN_DEFAULT_TTL = int(
|
||||
os.getenv("MCP_PER_USER_TOKEN_DEFAULT_TTL", "43200") # 12 hours
|
||||
)
|
||||
MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS = int(
|
||||
os.getenv("MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS", "60")
|
||||
)
|
||||
MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS = int(os.getenv("MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS", "60"))
|
||||
|
||||
# MCP timeout defaults (seconds). Override via env vars for slow/custom MCP servers.
|
||||
MCP_CLIENT_TIMEOUT = float(os.getenv("LITELLM_MCP_CLIENT_TIMEOUT", "60.0"))
|
||||
|
|
@ -166,14 +128,11 @@ MCP_HEALTH_CHECK_TIMEOUT = float(os.getenv("LITELLM_MCP_HEALTH_CHECK_TIMEOUT", "
|
|||
# Extend via LITELLM_MCP_STDIO_EXTRA_COMMANDS env var (comma-separated).
|
||||
_MCP_STDIO_EXTRA_COMMANDS = os.getenv("LITELLM_MCP_STDIO_EXTRA_COMMANDS", "")
|
||||
MCP_STDIO_ALLOWED_COMMANDS: frozenset = frozenset(
|
||||
{"npx", "uvx", "python", "python3", "node", "docker", "deno"}
|
||||
| (set(_MCP_STDIO_EXTRA_COMMANDS.split(",")) - {""})
|
||||
{"npx", "uvx", "python", "python3", "node", "docker", "deno"} | (set(_MCP_STDIO_EXTRA_COMMANDS.split(",")) - {""})
|
||||
)
|
||||
|
||||
# MCP OAuth2 Token Exchange (OBO) Defaults
|
||||
MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE = int(
|
||||
os.getenv("MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE", "500")
|
||||
)
|
||||
MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE = int(os.getenv("MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE", "500"))
|
||||
|
||||
LITELLM_UI_ALLOW_HEADERS = [
|
||||
"x-litellm-semantic-filter",
|
||||
|
|
@ -189,9 +148,7 @@ DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO = int(
|
|||
os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO", 128)
|
||||
)
|
||||
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE = int(
|
||||
os.getenv(
|
||||
"DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE", 512
|
||||
)
|
||||
os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE", 512)
|
||||
)
|
||||
|
||||
# Maximum number of callbacks that can be registered
|
||||
|
|
@ -223,31 +180,19 @@ OPEN_SANDBOX_DEFAULT_TIMEOUT = 300
|
|||
OPEN_SANDBOX_READY_TIMEOUT = 30.0
|
||||
OPEN_SANDBOX_POLL_INTERVAL = 0.2
|
||||
|
||||
DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET = int(
|
||||
os.getenv("DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET", 1024)
|
||||
)
|
||||
DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET = int(os.getenv("DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET", 1024))
|
||||
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET = int(
|
||||
os.getenv("DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET", 2048)
|
||||
)
|
||||
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET = int(
|
||||
os.getenv("DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET", 4096)
|
||||
)
|
||||
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET = int(
|
||||
os.getenv("DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET", 8192)
|
||||
)
|
||||
DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET = int(
|
||||
os.getenv("DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET", 16384)
|
||||
)
|
||||
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET = int(os.getenv("DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET", 4096))
|
||||
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET = int(os.getenv("DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET", 8192))
|
||||
DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET = int(os.getenv("DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET", 16384))
|
||||
MAX_TOKEN_TRIMMING_ATTEMPTS = int(
|
||||
os.getenv("MAX_TOKEN_TRIMMING_ATTEMPTS", 10)
|
||||
) # Maximum number of attempts to trim the message
|
||||
|
||||
RUNWAYML_DEFAULT_API_VERSION = str(
|
||||
os.getenv("RUNWAYML_DEFAULT_API_VERSION", "2024-11-06")
|
||||
)
|
||||
RUNWAYML_POLLING_TIMEOUT = int(
|
||||
os.getenv("RUNWAYML_POLLING_TIMEOUT", 600)
|
||||
) # 10 minutes default for image generation
|
||||
RUNWAYML_DEFAULT_API_VERSION = str(os.getenv("RUNWAYML_DEFAULT_API_VERSION", "2024-11-06"))
|
||||
RUNWAYML_POLLING_TIMEOUT = int(os.getenv("RUNWAYML_POLLING_TIMEOUT", 600)) # 10 minutes default for image generation
|
||||
|
||||
########## Networking constants ##############################################################
|
||||
_DEFAULT_TTL_FOR_HTTPX_CLIENTS = 3600 # 1 hour, re-use the same httpx client for 1 hour
|
||||
|
|
@ -255,9 +200,7 @@ _DEFAULT_TTL_FOR_HTTPX_CLIENTS = 3600 # 1 hour, re-use the same httpx client fo
|
|||
# Aiohttp connection pooling - prevents memory leaks from unbounded connection growth
|
||||
# Set to 0 for unlimited (not recommended for production)
|
||||
AIOHTTP_CONNECTOR_LIMIT = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT", 1000))
|
||||
AIOHTTP_CONNECTOR_LIMIT_PER_HOST = int(
|
||||
os.getenv("AIOHTTP_CONNECTOR_LIMIT_PER_HOST", 500)
|
||||
)
|
||||
AIOHTTP_CONNECTOR_LIMIT_PER_HOST = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT_PER_HOST", 500))
|
||||
AIOHTTP_KEEPALIVE_TIMEOUT = int(os.getenv("AIOHTTP_KEEPALIVE_TIMEOUT", 120))
|
||||
AIOHTTP_TTL_DNS_CACHE = int(os.getenv("AIOHTTP_TTL_DNS_CACHE", 300))
|
||||
# TCP keep-alive (SO_KEEPALIVE) — opt-in. Required when running behind NAT/LBs
|
||||
|
|
@ -283,9 +226,7 @@ AIOHTTP_NEEDS_CLEANUP_CLOSED = (3, 13, 0) <= sys.version_info < (
|
|||
# Default to None (unlimited) to match OpenAI's official agents SDK behavior
|
||||
# https://github.com/openai/openai-agents-python/blob/cf1b933660e44fd37b4350c41febab8221801409/src/agents/realtime/openai_realtime.py#L235
|
||||
_max_size_env = os.getenv("REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES")
|
||||
REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES = (
|
||||
int(_max_size_env) if _max_size_env is not None else None
|
||||
)
|
||||
REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES = int(_max_size_env) if _max_size_env is not None else None
|
||||
|
||||
# SSL/TLS cipher configuration for faster handshakes
|
||||
# Strategy: Strongly prefer fast modern ciphers, but allow fallback to commonly supported ones
|
||||
|
|
@ -316,9 +257,7 @@ REDIS_UPDATE_BUFFER_KEY = "litellm_spend_update_buffer"
|
|||
REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_spend_update_buffer"
|
||||
REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_team_spend_update_buffer"
|
||||
REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_org_spend_update_buffer"
|
||||
REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY = (
|
||||
"litellm_daily_end_user_spend_update_buffer"
|
||||
)
|
||||
REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_end_user_spend_update_buffer"
|
||||
REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_agent_spend_update_buffer"
|
||||
REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_tag_spend_update_buffer"
|
||||
MAX_REDIS_BUFFER_DEQUEUE_COUNT = int(os.getenv("MAX_REDIS_BUFFER_DEQUEUE_COUNT", 100))
|
||||
|
|
@ -327,12 +266,8 @@ LITELLM_ASYNCIO_QUEUE_MAXSIZE = int(os.getenv("LITELLM_ASYNCIO_QUEUE_MAXSIZE", 1
|
|||
TOOL_POLICY_CACHE_TTL_SECONDS = int(os.getenv("TOOL_POLICY_CACHE_TTL_SECONDS", 60))
|
||||
# Aggregation threshold: default to 80% of the asyncio queue maxsize so the check can always trigger.
|
||||
# Must be < LITELLM_ASYNCIO_QUEUE_MAXSIZE; if set higher the aggregation logic will never fire.
|
||||
MAX_SIZE_IN_MEMORY_QUEUE = int(
|
||||
os.getenv("MAX_SIZE_IN_MEMORY_QUEUE", int(LITELLM_ASYNCIO_QUEUE_MAXSIZE * 0.8))
|
||||
)
|
||||
MAX_IN_MEMORY_QUEUE_FLUSH_COUNT = int(
|
||||
os.getenv("MAX_IN_MEMORY_QUEUE_FLUSH_COUNT", 1000)
|
||||
)
|
||||
MAX_SIZE_IN_MEMORY_QUEUE = int(os.getenv("MAX_SIZE_IN_MEMORY_QUEUE", int(LITELLM_ASYNCIO_QUEUE_MAXSIZE * 0.8)))
|
||||
MAX_IN_MEMORY_QUEUE_FLUSH_COUNT = int(os.getenv("MAX_IN_MEMORY_QUEUE_FLUSH_COUNT", 1000))
|
||||
###############################################################################################
|
||||
MINIMUM_PROMPT_CACHE_TOKEN_COUNT = int(
|
||||
os.getenv("MINIMUM_PROMPT_CACHE_TOKEN_COUNT", 1024)
|
||||
|
|
@ -344,49 +279,31 @@ HOURS_IN_A_DAY = int(os.getenv("HOURS_IN_A_DAY", 24))
|
|||
DAYS_IN_A_WEEK = int(os.getenv("DAYS_IN_A_WEEK", 7))
|
||||
DAYS_IN_A_MONTH = int(os.getenv("DAYS_IN_A_MONTH", 28))
|
||||
DAYS_IN_A_YEAR = int(os.getenv("DAYS_IN_A_YEAR", 365))
|
||||
REPLICATE_MODEL_NAME_WITH_ID_LENGTH = int(
|
||||
os.getenv("REPLICATE_MODEL_NAME_WITH_ID_LENGTH", 64)
|
||||
)
|
||||
REPLICATE_MODEL_NAME_WITH_ID_LENGTH = int(os.getenv("REPLICATE_MODEL_NAME_WITH_ID_LENGTH", 64))
|
||||
#### TOKEN COUNTING ####
|
||||
FUNCTION_DEFINITION_TOKEN_COUNT = int(os.getenv("FUNCTION_DEFINITION_TOKEN_COUNT", 9))
|
||||
SYSTEM_MESSAGE_TOKEN_COUNT = int(os.getenv("SYSTEM_MESSAGE_TOKEN_COUNT", 4))
|
||||
TOOL_CHOICE_OBJECT_TOKEN_COUNT = int(os.getenv("TOOL_CHOICE_OBJECT_TOKEN_COUNT", 4))
|
||||
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT = int(
|
||||
os.getenv("DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT", 10)
|
||||
)
|
||||
DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT = int(
|
||||
os.getenv("DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT", 20)
|
||||
)
|
||||
MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES = int(
|
||||
os.getenv("MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES", 768)
|
||||
)
|
||||
MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES = int(
|
||||
os.getenv("MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES", 2000)
|
||||
)
|
||||
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT = int(os.getenv("DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT", 10))
|
||||
DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT = int(os.getenv("DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT", 20))
|
||||
MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES = int(os.getenv("MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES", 768))
|
||||
MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES = int(os.getenv("MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES", 2000))
|
||||
MAX_TILE_WIDTH = int(os.getenv("MAX_TILE_WIDTH", 512))
|
||||
MAX_TILE_HEIGHT = int(os.getenv("MAX_TILE_HEIGHT", 512))
|
||||
OPENAI_FILE_SEARCH_COST_PER_1K_CALLS = float(
|
||||
os.getenv("OPENAI_FILE_SEARCH_COST_PER_1K_CALLS", 2.5 / 1000)
|
||||
)
|
||||
OPENAI_FILE_SEARCH_COST_PER_1K_CALLS = float(os.getenv("OPENAI_FILE_SEARCH_COST_PER_1K_CALLS", 2.5 / 1000))
|
||||
# Azure OpenAI Assistants feature costs
|
||||
# Source: https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/
|
||||
AZURE_FILE_SEARCH_COST_PER_GB_PER_DAY = float(
|
||||
os.getenv("AZURE_FILE_SEARCH_COST_PER_GB_PER_DAY", 0.1) # $0.1 USD per 1 GB/Day
|
||||
)
|
||||
AZURE_COMPUTER_USE_INPUT_COST_PER_1K_TOKENS = float(
|
||||
os.getenv(
|
||||
"AZURE_COMPUTER_USE_INPUT_COST_PER_1K_TOKENS", 3.0
|
||||
) # $0.003 USD per 1K Tokens
|
||||
os.getenv("AZURE_COMPUTER_USE_INPUT_COST_PER_1K_TOKENS", 3.0) # $0.003 USD per 1K Tokens
|
||||
)
|
||||
AZURE_COMPUTER_USE_OUTPUT_COST_PER_1K_TOKENS = float(
|
||||
os.getenv(
|
||||
"AZURE_COMPUTER_USE_OUTPUT_COST_PER_1K_TOKENS", 12.0
|
||||
) # $0.012 USD per 1K Tokens
|
||||
os.getenv("AZURE_COMPUTER_USE_OUTPUT_COST_PER_1K_TOKENS", 12.0) # $0.012 USD per 1K Tokens
|
||||
)
|
||||
AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY = float(
|
||||
os.getenv(
|
||||
"AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY", 0.1
|
||||
) # $0.1 USD per 1 GB/Day (same as file search)
|
||||
os.getenv("AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY", 0.1) # $0.1 USD per 1 GB/Day (same as file search)
|
||||
)
|
||||
MIN_NON_ZERO_TEMPERATURE = float(os.getenv("MIN_NON_ZERO_TEMPERATURE", 0.0001))
|
||||
#### RELIABILITY ####
|
||||
|
|
@ -400,9 +317,7 @@ _REALTIME_BODY_CACHE_SIZE = 1000 # Keep realtime helper caches bounded; workloa
|
|||
INITIAL_RETRY_DELAY = float(os.getenv("INITIAL_RETRY_DELAY", 0.5))
|
||||
MAX_RETRY_DELAY = float(os.getenv("MAX_RETRY_DELAY", 8.0))
|
||||
JITTER = float(os.getenv("JITTER", 0.75))
|
||||
DEFAULT_IN_MEMORY_TTL = int(
|
||||
os.getenv("DEFAULT_IN_MEMORY_TTL", 5)
|
||||
) # default time to live for the in-memory cache
|
||||
DEFAULT_IN_MEMORY_TTL = int(os.getenv("DEFAULT_IN_MEMORY_TTL", 5)) # default time to live for the in-memory cache
|
||||
DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE = int(
|
||||
os.getenv("DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE", 1000)
|
||||
) # default max size for redis batch cache
|
||||
|
|
@ -410,23 +325,13 @@ DEFAULT_POLLING_INTERVAL = float(
|
|||
os.getenv("DEFAULT_POLLING_INTERVAL", 0.03)
|
||||
) # default polling interval for the scheduler
|
||||
AZURE_OPERATION_POLLING_TIMEOUT = int(os.getenv("AZURE_OPERATION_POLLING_TIMEOUT", 120))
|
||||
AZURE_DOCUMENT_INTELLIGENCE_API_VERSION = str(
|
||||
os.getenv("AZURE_DOCUMENT_INTELLIGENCE_API_VERSION", "2024-11-30")
|
||||
)
|
||||
AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI = int(
|
||||
os.getenv("AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI", 96)
|
||||
)
|
||||
AZURE_DOCUMENT_INTELLIGENCE_API_VERSION = str(os.getenv("AZURE_DOCUMENT_INTELLIGENCE_API_VERSION", "2024-11-30"))
|
||||
AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI = int(os.getenv("AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI", 96))
|
||||
REDIS_SOCKET_TIMEOUT = float(os.getenv("REDIS_SOCKET_TIMEOUT", 0.1))
|
||||
REDIS_CONNECTION_POOL_TIMEOUT = int(os.getenv("REDIS_CONNECTION_POOL_TIMEOUT", 5))
|
||||
REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD = int(
|
||||
os.getenv("REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD", 5)
|
||||
)
|
||||
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT = int(
|
||||
os.getenv("REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT", 60)
|
||||
)
|
||||
REDIS_CIRCUIT_BREAKER_ENABLED = (
|
||||
os.getenv("REDIS_CIRCUIT_BREAKER_ENABLED", "true").lower() == "true"
|
||||
)
|
||||
REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD = int(os.getenv("REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD", 5))
|
||||
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT = int(os.getenv("REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT", 60))
|
||||
REDIS_CIRCUIT_BREAKER_ENABLED = os.getenv("REDIS_CIRCUIT_BREAKER_ENABLED", "true").lower() == "true"
|
||||
# Default Redis major version to assume when version cannot be determined
|
||||
# Using 7 as it's the modern version that supports LPOP with count parameter
|
||||
DEFAULT_REDIS_MAJOR_VERSION = int(os.getenv("DEFAULT_REDIS_MAJOR_VERSION", 7))
|
||||
|
|
@ -436,17 +341,11 @@ NON_LLM_CONNECTION_TIMEOUT = int(
|
|||
MAX_EXCEPTION_MESSAGE_LENGTH = int(os.getenv("MAX_EXCEPTION_MESSAGE_LENGTH", 2000))
|
||||
MAX_STRING_LENGTH_PROMPT_IN_DB = int(os.getenv("MAX_STRING_LENGTH_PROMPT_IN_DB", 2048))
|
||||
BEDROCK_MAX_POLICY_SIZE = int(os.getenv("BEDROCK_MAX_POLICY_SIZE", 75))
|
||||
BEDROCK_MIN_THINKING_BUDGET_TOKENS = int(
|
||||
os.getenv("BEDROCK_MIN_THINKING_BUDGET_TOKENS", 1024)
|
||||
)
|
||||
BEDROCK_MIN_THINKING_BUDGET_TOKENS = int(os.getenv("BEDROCK_MIN_THINKING_BUDGET_TOKENS", 1024))
|
||||
# Anthropic's Messages API rejects thinking.budget_tokens < 1024.
|
||||
ANTHROPIC_MIN_THINKING_BUDGET_TOKENS = 1024
|
||||
REPLICATE_POLLING_DELAY_SECONDS = float(
|
||||
os.getenv("REPLICATE_POLLING_DELAY_SECONDS", 0.5)
|
||||
)
|
||||
DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS = int(
|
||||
os.getenv("DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS", 4096)
|
||||
)
|
||||
REPLICATE_POLLING_DELAY_SECONDS = float(os.getenv("REPLICATE_POLLING_DELAY_SECONDS", 0.5))
|
||||
DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS = int(os.getenv("DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS", 4096))
|
||||
DEFAULT_OCI_CHAT_MAX_TOKENS = 4096
|
||||
TOGETHER_AI_4_B = int(os.getenv("TOGETHER_AI_4_B", 4))
|
||||
TOGETHER_AI_8_B = int(os.getenv("TOGETHER_AI_8_B", 8))
|
||||
|
|
@ -475,13 +374,9 @@ DEFAULT_REQUEST_TIMEOUT_SECONDS: float = 6000.0
|
|||
# deadline and connect handshake (see ``http_handler`` cached handler paths).
|
||||
COMPLETION_HTTP_FALLBACK_SECONDS: float = 600.0
|
||||
HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS: float = 5.0
|
||||
request_timeout: float = float(
|
||||
os.getenv("REQUEST_TIMEOUT", str(int(DEFAULT_REQUEST_TIMEOUT_SECONDS)))
|
||||
)
|
||||
request_timeout: float = float(os.getenv("REQUEST_TIMEOUT", str(int(DEFAULT_REQUEST_TIMEOUT_SECONDS))))
|
||||
request_timeout_explicitly_set: bool = "REQUEST_TIMEOUT" in os.environ
|
||||
DEFAULT_A2A_AGENT_TIMEOUT: float = float(
|
||||
os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)
|
||||
) # 10 minutes
|
||||
DEFAULT_A2A_AGENT_TIMEOUT: float = float(os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)) # 10 minutes
|
||||
# Patterns that indicate a localhost/internal URL in A2A agent cards that should be
|
||||
# replaced with the original base_url. This is a common misconfiguration where
|
||||
# developers deploy agents with development URLs in their agent cards.
|
||||
|
|
@ -511,16 +406,10 @@ FIREWORKS_AI_16_B = int(os.getenv("FIREWORKS_AI_16_B", 16))
|
|||
FIREWORKS_AI_80_B = int(os.getenv("FIREWORKS_AI_80_B", 80))
|
||||
#### Logging callback constants ####
|
||||
REDACTED_BY_LITELM_STRING = "REDACTED_BY_LITELM"
|
||||
MAX_LANGFUSE_INITIALIZED_CLIENTS = int(
|
||||
os.getenv("MAX_LANGFUSE_INITIALIZED_CLIENTS", 50)
|
||||
)
|
||||
LOGGING_WORKER_CONCURRENCY = int(
|
||||
os.getenv("LOGGING_WORKER_CONCURRENCY", 100)
|
||||
) # Must be above 0
|
||||
MAX_LANGFUSE_INITIALIZED_CLIENTS = int(os.getenv("MAX_LANGFUSE_INITIALIZED_CLIENTS", 50))
|
||||
LOGGING_WORKER_CONCURRENCY = int(os.getenv("LOGGING_WORKER_CONCURRENCY", 100)) # Must be above 0
|
||||
LOGGING_WORKER_MAX_QUEUE_SIZE = int(os.getenv("LOGGING_WORKER_MAX_QUEUE_SIZE", 50_000))
|
||||
LOGGING_WORKER_MAX_TIME_PER_COROUTINE = float(
|
||||
os.getenv("LOGGING_WORKER_MAX_TIME_PER_COROUTINE", 20.0)
|
||||
)
|
||||
LOGGING_WORKER_MAX_TIME_PER_COROUTINE = float(os.getenv("LOGGING_WORKER_MAX_TIME_PER_COROUTINE", 20.0))
|
||||
LOGGING_WORKER_CLEAR_PERCENTAGE = int(
|
||||
os.getenv("LOGGING_WORKER_CLEAR_PERCENTAGE", 50)
|
||||
) # Percentage of queue to clear (default: 50%)
|
||||
|
|
@ -535,17 +424,13 @@ DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE = os.getenv(
|
|||
|
||||
LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED = 499
|
||||
|
||||
EMAIL_BUDGET_ALERT_TTL = int(
|
||||
os.getenv("EMAIL_BUDGET_ALERT_TTL", 24 * 60 * 60)
|
||||
) # 24 hours in seconds
|
||||
EMAIL_BUDGET_ALERT_TTL = int(os.getenv("EMAIL_BUDGET_ALERT_TTL", 24 * 60 * 60)) # 24 hours in seconds
|
||||
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE = float(
|
||||
os.getenv("EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE", 0.8)
|
||||
) # 80% of max budget
|
||||
############### LLM Provider Constants ###############
|
||||
### ANTHROPIC CONSTANTS ###
|
||||
ANTHROPIC_TOKEN_COUNTING_BETA_VERSION = os.getenv(
|
||||
"ANTHROPIC_TOKEN_COUNTING_BETA_VERSION", "token-counting-2024-11-01"
|
||||
)
|
||||
ANTHROPIC_TOKEN_COUNTING_BETA_VERSION = os.getenv("ANTHROPIC_TOKEN_COUNTING_BETA_VERSION", "token-counting-2024-11-01")
|
||||
ANTHROPIC_SKILLS_API_BETA_VERSION = "skills-2025-10-02"
|
||||
ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES = {
|
||||
"low": 1,
|
||||
|
|
@ -560,9 +445,7 @@ LITELLM_WEB_SEARCH_TOOL_NAME = "litellm_web_search"
|
|||
DEFAULT_IMAGE_ENDPOINT_MODEL = "dall-e-2"
|
||||
DEFAULT_VIDEO_ENDPOINT_MODEL = "sora-2"
|
||||
|
||||
DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS = int(
|
||||
os.getenv("DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS", 8)
|
||||
)
|
||||
DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS = int(os.getenv("DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS", 8))
|
||||
|
||||
### DATAFORSEO CONSTANTS ###
|
||||
DEFAULT_DATAFORSEO_LOCATION_CODE = int(
|
||||
|
|
@ -945,8 +828,7 @@ clarifai_models: set = set(
|
|||
"clarifai/qwen.qwenLM.Qwen3-30B-A3B-Instruct-2507",
|
||||
"clarifai/qwen.qwen3.qwen3-next-80B-A3B-Thinking",
|
||||
"clarifai/openai.chat-completion.gpt-oss-120b",
|
||||
"clarifai/qwen.qwenLM.Qwen3-30B-A3B-Thinking-2507"
|
||||
"clarifai/openai.chat-completion.gpt-5-nano",
|
||||
"clarifai/qwen.qwenLM.Qwen3-30B-A3B-Thinking-2507clarifai/openai.chat-completion.gpt-5-nano",
|
||||
"clarifai/openai.chat-completion.gpt-4o",
|
||||
"clarifai/gcp.generate.gemini-2_5-pro",
|
||||
"clarifai/anthropic.completion.claude-sonnet-4",
|
||||
|
|
@ -1375,9 +1257,7 @@ OPENAI_FINISH_REASONS = [
|
|||
"tool_calls",
|
||||
"content_filter",
|
||||
]
|
||||
HUMANLOOP_PROMPT_CACHE_TTL_SECONDS = int(
|
||||
os.getenv("HUMANLOOP_PROMPT_CACHE_TTL_SECONDS", 60)
|
||||
) # 1 minute
|
||||
HUMANLOOP_PROMPT_CACHE_TTL_SECONDS = int(os.getenv("HUMANLOOP_PROMPT_CACHE_TTL_SECONDS", 60)) # 1 minute
|
||||
RESPONSE_FORMAT_TOOL_NAME = "json_tool_call" # default tool name used when converting response format to tool call
|
||||
|
||||
########################### Logging Callback Constants ###########################
|
||||
|
|
@ -1385,9 +1265,7 @@ AZURE_STORAGE_MSFT_VERSION = "2019-07-07"
|
|||
PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES = int(
|
||||
os.getenv("PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES", 5)
|
||||
)
|
||||
CLOUDZERO_EXPORT_INTERVAL_MINUTES = int(
|
||||
os.getenv("CLOUDZERO_EXPORT_INTERVAL_MINUTES", 60)
|
||||
)
|
||||
CLOUDZERO_EXPORT_INTERVAL_MINUTES = int(os.getenv("CLOUDZERO_EXPORT_INTERVAL_MINUTES", 60))
|
||||
MCP_TOOL_NAME_PREFIX = "mcp_tool"
|
||||
MAXIMUM_TRACEBACK_LINES_TO_LOG = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100))
|
||||
|
||||
|
|
@ -1450,37 +1328,23 @@ PASS_THROUGH_HEADER_PREFIX = "x-pass-"
|
|||
|
||||
BASE_MCP_ROUTE = "/mcp"
|
||||
|
||||
BATCH_STATUS_POLL_INTERVAL_SECONDS = int(
|
||||
os.getenv("BATCH_STATUS_POLL_INTERVAL_SECONDS", 3600)
|
||||
) # 1 hour
|
||||
BATCH_STATUS_POLL_MAX_ATTEMPTS = int(
|
||||
os.getenv("BATCH_STATUS_POLL_MAX_ATTEMPTS", 24)
|
||||
) # for 24 hours
|
||||
BATCH_STATUS_POLL_INTERVAL_SECONDS = int(os.getenv("BATCH_STATUS_POLL_INTERVAL_SECONDS", 3600)) # 1 hour
|
||||
BATCH_STATUS_POLL_MAX_ATTEMPTS = int(os.getenv("BATCH_STATUS_POLL_MAX_ATTEMPTS", 24)) # for 24 hours
|
||||
|
||||
HEALTH_CHECK_TIMEOUT_SECONDS = int(
|
||||
os.getenv("HEALTH_CHECK_TIMEOUT_SECONDS", 60)
|
||||
) # 60 seconds
|
||||
_background_health_check_max_tokens_env = os.getenv(
|
||||
"BACKGROUND_HEALTH_CHECK_MAX_TOKENS"
|
||||
)
|
||||
HEALTH_CHECK_TIMEOUT_SECONDS = int(os.getenv("HEALTH_CHECK_TIMEOUT_SECONDS", 60)) # 60 seconds
|
||||
_background_health_check_max_tokens_env = os.getenv("BACKGROUND_HEALTH_CHECK_MAX_TOKENS")
|
||||
try:
|
||||
_raw_background_health_check_max_tokens = (
|
||||
_background_health_check_max_tokens_env.strip()
|
||||
if _background_health_check_max_tokens_env is not None
|
||||
else ""
|
||||
_background_health_check_max_tokens_env.strip() if _background_health_check_max_tokens_env is not None else ""
|
||||
)
|
||||
BACKGROUND_HEALTH_CHECK_MAX_TOKENS: Optional[int] = (
|
||||
int(_raw_background_health_check_max_tokens)
|
||||
if _raw_background_health_check_max_tokens
|
||||
else None
|
||||
int(_raw_background_health_check_max_tokens) if _raw_background_health_check_max_tokens else None
|
||||
)
|
||||
except (ValueError, TypeError):
|
||||
BACKGROUND_HEALTH_CHECK_MAX_TOKENS = None
|
||||
|
||||
|
||||
_background_health_check_max_tokens_reasoning_env = os.getenv(
|
||||
"BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING"
|
||||
)
|
||||
_background_health_check_max_tokens_reasoning_env = os.getenv("BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING")
|
||||
try:
|
||||
_raw_background_health_check_max_tokens_reasoning = (
|
||||
_background_health_check_max_tokens_reasoning_env.strip()
|
||||
|
|
@ -1522,9 +1386,7 @@ LITELLM_KEY_ROTATION_LOCK_TTL_SECONDS = int(
|
|||
os.getenv("LITELLM_KEY_ROTATION_LOCK_TTL_SECONDS", 600)
|
||||
) # 10 minutes default — caps the deadlock window if a pod crashes mid-rotation
|
||||
UI_SESSION_TOKEN_TEAM_ID = "litellm-dashboard"
|
||||
LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_ENABLED = os.getenv(
|
||||
"LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_ENABLED", "false"
|
||||
)
|
||||
LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_ENABLED = os.getenv("LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_ENABLED", "false")
|
||||
LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_INTERVAL_SECONDS = int(
|
||||
os.getenv("LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_INTERVAL_SECONDS", 86400)
|
||||
) # 24 hours default
|
||||
|
|
@ -1541,15 +1403,11 @@ CLI_SSO_SESSION_TTL_SECONDS = 600
|
|||
CLI_SESSION_KEY_PREFIX = "cli-session"
|
||||
# Support both CLI_JWT_EXPIRATION_HOURS and LITELLM_CLI_JWT_EXPIRATION_HOURS for backwards compatibility
|
||||
CLI_JWT_EXPIRATION_HOURS = int(
|
||||
os.getenv("CLI_JWT_EXPIRATION_HOURS")
|
||||
or os.getenv("LITELLM_CLI_JWT_EXPIRATION_HOURS")
|
||||
or 24
|
||||
os.getenv("CLI_JWT_EXPIRATION_HOURS") or os.getenv("LITELLM_CLI_JWT_EXPIRATION_HOURS") or 24
|
||||
)
|
||||
# Comma-separated allowlisted OIDC claim map for CLI SSO polling, e.g.
|
||||
# "employment_type->acme_employment_type,org_info.department->department"
|
||||
CLI_SSO_CLAIM_MAP = (
|
||||
os.getenv("CLI_SSO_CLAIM_MAP") or os.getenv("LITELLM_CLI_SSO_CLAIM_MAP") or ""
|
||||
)
|
||||
CLI_SSO_CLAIM_MAP = os.getenv("CLI_SSO_CLAIM_MAP") or os.getenv("LITELLM_CLI_SSO_CLAIM_MAP") or ""
|
||||
CLI_SSO_CLAIM_MAX_SCALAR_LENGTH = 1024
|
||||
|
||||
########################### UI SESSION DURATION ###########################
|
||||
|
|
@ -1563,54 +1421,34 @@ DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME = "db_daily_tag_spend_update_job"
|
|||
PROMETHEUS_EMIT_BUDGET_METRICS_JOB_NAME = "prometheus_emit_budget_metrics"
|
||||
CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME = "cloudzero_export_usage_data"
|
||||
MAVVRIK_FOCUS_EXPORT_JOB_NAME = "mavvrik_focus_export_usage_data"
|
||||
CLOUDZERO_MAX_FETCHED_DATA_RECORDS = int(
|
||||
os.getenv("CLOUDZERO_MAX_FETCHED_DATA_RECORDS", 50000)
|
||||
)
|
||||
CLOUDZERO_MAX_FETCHED_DATA_RECORDS = int(os.getenv("CLOUDZERO_MAX_FETCHED_DATA_RECORDS", 50000))
|
||||
SPEND_LOG_CLEANUP_JOB_NAME = "spend_log_cleanup"
|
||||
KEY_ROTATION_JOB_NAME = "litellm_key_rotation_job"
|
||||
EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME = "litellm_expired_ui_session_key_cleanup_job"
|
||||
SPEND_LOG_RUN_LOOPS = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500))
|
||||
SPEND_LOG_CLEANUP_BATCH_SIZE = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000))
|
||||
SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(
|
||||
os.getenv("SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3)
|
||||
)
|
||||
SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(os.getenv("SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3))
|
||||
SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS = float(
|
||||
os.getenv("SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.5)
|
||||
)
|
||||
SPEND_LOG_PARTITION_INTERVAL = os.getenv("SPEND_LOG_PARTITION_INTERVAL", "day")
|
||||
SPEND_LOG_PARTITION_PRECREATE_AHEAD = int(
|
||||
os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7)
|
||||
)
|
||||
SPEND_LOG_PARTITION_PRECREATE_AHEAD = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7))
|
||||
SPEND_LOG_QUEUE_SIZE_THRESHOLD = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100))
|
||||
SPEND_LOG_QUEUE_POLL_INTERVAL = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0))
|
||||
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = int(
|
||||
os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000)
|
||||
)
|
||||
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS = int(
|
||||
os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)
|
||||
) # 1 minute
|
||||
PROXY_BUDGET_RESCHEDULER_MIN_TIME = int(
|
||||
os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597)
|
||||
)
|
||||
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000))
|
||||
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute
|
||||
PROXY_BUDGET_RESCHEDULER_MIN_TIME = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597))
|
||||
PROXY_BATCH_POLLING_INTERVAL = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600))
|
||||
MAX_OBJECTS_PER_POLL_CYCLE = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50)))
|
||||
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS = max(
|
||||
1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7))
|
||||
)
|
||||
STALE_OBJECT_CLEANUP_BATCH_SIZE = max(
|
||||
1, int(os.getenv("STALE_OBJECT_CLEANUP_BATCH_SIZE", 1000))
|
||||
)
|
||||
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS = max(1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7)))
|
||||
STALE_OBJECT_CLEANUP_BATCH_SIZE = max(1, int(os.getenv("STALE_OBJECT_CLEANUP_BATCH_SIZE", 1000)))
|
||||
# Set PROXY_BATCH_POLLING_ENABLED=false to disable the CheckBatchCost and
|
||||
# CheckResponsesCost background polling jobs entirely (e.g. to avoid DB load on
|
||||
# installations with large numbers of stale managed objects).
|
||||
_batch_polling_env = os.getenv("PROXY_BATCH_POLLING_ENABLED", "true").lower()
|
||||
PROXY_BATCH_POLLING_ENABLED = _batch_polling_env == "true"
|
||||
PROXY_BUDGET_RESCHEDULER_MAX_TIME = int(
|
||||
os.getenv("PROXY_BUDGET_RESCHEDULER_MAX_TIME", 605)
|
||||
)
|
||||
PROXY_BATCH_WRITE_AT = int(
|
||||
os.getenv("PROXY_BATCH_WRITE_AT", 10)
|
||||
) # in seconds, increased from 10
|
||||
PROXY_BUDGET_RESCHEDULER_MAX_TIME = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MAX_TIME", 605))
|
||||
PROXY_BATCH_WRITE_AT = int(os.getenv("PROXY_BATCH_WRITE_AT", 10)) # in seconds, increased from 10
|
||||
|
||||
# APScheduler Configuration - MEMORY LEAK FIX
|
||||
# These settings prevent memory leaks in APScheduler's normalize() and _apply_jitter() functions
|
||||
|
|
@ -1621,12 +1459,8 @@ APSCHEDULER_COALESCE = os.getenv("APSCHEDULER_COALESCE", "True").lower() in [
|
|||
APSCHEDULER_MISFIRE_GRACE_TIME = int(
|
||||
os.getenv("APSCHEDULER_MISFIRE_GRACE_TIME", 3600)
|
||||
) # ignore runs older than 1 hour (was 120)
|
||||
APSCHEDULER_MAX_INSTANCES = int(
|
||||
os.getenv("APSCHEDULER_MAX_INSTANCES", 1)
|
||||
) # prevent concurrent job instances
|
||||
APSCHEDULER_REPLACE_EXISTING = os.getenv(
|
||||
"APSCHEDULER_REPLACE_EXISTING", "True"
|
||||
).lower() in [
|
||||
APSCHEDULER_MAX_INSTANCES = int(os.getenv("APSCHEDULER_MAX_INSTANCES", 1)) # prevent concurrent job instances
|
||||
APSCHEDULER_REPLACE_EXISTING = os.getenv("APSCHEDULER_REPLACE_EXISTING", "True").lower() in [
|
||||
"true",
|
||||
"1",
|
||||
] # always replace existing jobs
|
||||
|
|
@ -1635,38 +1469,24 @@ APSCHEDULER_REPLACE_EXISTING = os.getenv(
|
|||
# This will run tag spcific tasks at a later time to smooth QPS
|
||||
DAILY_TAG_SPEND_BATCH_MULTIPLIER = 2.3
|
||||
|
||||
DEFAULT_HEALTH_CHECK_INTERVAL = int(
|
||||
os.getenv("DEFAULT_HEALTH_CHECK_INTERVAL", 300)
|
||||
) # 5 minutes
|
||||
DEFAULT_HEALTH_CHECK_INTERVAL = int(os.getenv("DEFAULT_HEALTH_CHECK_INTERVAL", 300)) # 5 minutes
|
||||
DEFAULT_SHARED_HEALTH_CHECK_TTL = int(
|
||||
os.getenv("DEFAULT_SHARED_HEALTH_CHECK_TTL", 300)
|
||||
) # 5 minutes - TTL for cached health check results
|
||||
DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL = int(
|
||||
os.getenv("DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL", 60)
|
||||
) # 1 minute - TTL for health check lock
|
||||
DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER = (
|
||||
2 # health state is stale after interval * this
|
||||
)
|
||||
PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS = int(
|
||||
os.getenv("PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS", 9)
|
||||
)
|
||||
DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER = 2 # health state is stale after interval * this
|
||||
PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS = int(os.getenv("PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS", 9))
|
||||
DEFAULT_MODEL_CREATED_AT_TIME = int(
|
||||
os.getenv("DEFAULT_MODEL_CREATED_AT_TIME", 1677610602)
|
||||
) # returns on `/models` endpoint
|
||||
DEFAULT_SLACK_ALERTING_THRESHOLD = int(
|
||||
os.getenv("DEFAULT_SLACK_ALERTING_THRESHOLD", 300)
|
||||
)
|
||||
DEFAULT_SLACK_ALERTING_THRESHOLD = int(os.getenv("DEFAULT_SLACK_ALERTING_THRESHOLD", 300))
|
||||
MAX_TEAM_LIST_LIMIT = int(os.getenv("MAX_TEAM_LIST_LIMIT", 20))
|
||||
MAX_POLICY_ESTIMATE_IMPACT_ROWS = int(
|
||||
os.getenv("MAX_POLICY_ESTIMATE_IMPACT_ROWS", 1000)
|
||||
)
|
||||
DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD = float(
|
||||
os.getenv("DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD", 0.7)
|
||||
)
|
||||
MAX_POLICY_ESTIMATE_IMPACT_ROWS = int(os.getenv("MAX_POLICY_ESTIMATE_IMPACT_ROWS", 1000))
|
||||
DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD = float(os.getenv("DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD", 0.7))
|
||||
LENGTH_OF_LITELLM_GENERATED_KEY = int(os.getenv("LENGTH_OF_LITELLM_GENERATED_KEY", 16))
|
||||
SECRET_MANAGER_REFRESH_INTERVAL = int(
|
||||
os.getenv("SECRET_MANAGER_REFRESH_INTERVAL", 86400)
|
||||
)
|
||||
SECRET_MANAGER_REFRESH_INTERVAL = int(os.getenv("SECRET_MANAGER_REFRESH_INTERVAL", 86400))
|
||||
LITELLM_SETTINGS_SAFE_DB_OVERRIDES = [
|
||||
"default_internal_user_params",
|
||||
"default_team_params",
|
||||
|
|
@ -1678,9 +1498,7 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES = [
|
|||
"cost_margin_config",
|
||||
]
|
||||
SPECIAL_LITELLM_AUTH_TOKEN = ["ui-token"]
|
||||
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(
|
||||
os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60)
|
||||
)
|
||||
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60))
|
||||
DEFAULT_ACCESS_GROUP_CACHE_TTL = int(os.getenv("DEFAULT_ACCESS_GROUP_CACHE_TTL", 600))
|
||||
# Short TTL for negative MCP access-group existence lookups. Keeps unauthenticated
|
||||
# callers from forcing a DB query per request for unknown names, while bounding
|
||||
|
|
@ -1760,9 +1578,7 @@ SENTRY_PII_DENYLIST = [
|
|||
]
|
||||
|
||||
# CoroutineChecker cache configuration
|
||||
COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY = int(
|
||||
os.getenv("COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY", 1000)
|
||||
)
|
||||
COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY = int(os.getenv("COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY", 1000))
|
||||
|
||||
########################### RAG Text Splitter Constants ###########################
|
||||
DEFAULT_CHUNK_SIZE = int(os.getenv("DEFAULT_CHUNK_SIZE", 1000))
|
||||
|
|
@ -1770,31 +1586,19 @@ DEFAULT_CHUNK_OVERLAP = int(os.getenv("DEFAULT_CHUNK_OVERLAP", 200))
|
|||
|
||||
########################### S3 Vectors RAG Constants ###########################
|
||||
S3_VECTORS_DEFAULT_DIMENSION = int(os.getenv("S3_VECTORS_DEFAULT_DIMENSION", 1024))
|
||||
S3_VECTORS_DEFAULT_DISTANCE_METRIC = str(
|
||||
os.getenv("S3_VECTORS_DEFAULT_DISTANCE_METRIC", "cosine")
|
||||
)
|
||||
S3_VECTORS_DEFAULT_DISTANCE_METRIC = str(os.getenv("S3_VECTORS_DEFAULT_DISTANCE_METRIC", "cosine"))
|
||||
S3_VECTORS_DEFAULT_NON_FILTERABLE_METADATA_KEYS = ["source_text"]
|
||||
|
||||
########################### Microsoft SSO Constants ###########################
|
||||
MICROSOFT_USER_EMAIL_ATTRIBUTE = str(
|
||||
os.getenv("MICROSOFT_USER_EMAIL_ATTRIBUTE", "userPrincipalName")
|
||||
)
|
||||
MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE = str(
|
||||
os.getenv("MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE", "displayName")
|
||||
)
|
||||
MICROSOFT_USER_EMAIL_ATTRIBUTE = str(os.getenv("MICROSOFT_USER_EMAIL_ATTRIBUTE", "userPrincipalName"))
|
||||
MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE = str(os.getenv("MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE", "displayName"))
|
||||
MICROSOFT_USER_ID_ATTRIBUTE = str(os.getenv("MICROSOFT_USER_ID_ATTRIBUTE", "id"))
|
||||
MICROSOFT_USER_FIRST_NAME_ATTRIBUTE = str(
|
||||
os.getenv("MICROSOFT_USER_FIRST_NAME_ATTRIBUTE", "givenName")
|
||||
)
|
||||
MICROSOFT_USER_LAST_NAME_ATTRIBUTE = str(
|
||||
os.getenv("MICROSOFT_USER_LAST_NAME_ATTRIBUTE", "surname")
|
||||
)
|
||||
MICROSOFT_USER_FIRST_NAME_ATTRIBUTE = str(os.getenv("MICROSOFT_USER_FIRST_NAME_ATTRIBUTE", "givenName"))
|
||||
MICROSOFT_USER_LAST_NAME_ATTRIBUTE = str(os.getenv("MICROSOFT_USER_LAST_NAME_ATTRIBUTE", "surname"))
|
||||
|
||||
# Maximum payload size (in bytes) to fully serialize for DEBUG logging.
|
||||
# Payloads larger than this are truncated to avoid multi-second json.dumps blocking the response.
|
||||
MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG = int(
|
||||
os.getenv("MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG", 102400)
|
||||
) # 100 KB
|
||||
MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG = int(os.getenv("MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG", 102400)) # 100 KB
|
||||
|
||||
# Policy template enrichment
|
||||
MAX_COMPETITOR_NAMES = int(os.getenv("MAX_COMPETITOR_NAMES", 100))
|
||||
|
|
|
|||
|
|
@ -97,9 +97,7 @@ def create_sync_endpoint_function(endpoint_config: Dict) -> Callable:
|
|||
)
|
||||
|
||||
if container_provider_config is None:
|
||||
raise ValueError(
|
||||
f"Container provider config not found for: {resolved_custom_llm_provider}"
|
||||
)
|
||||
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}
|
||||
|
|
@ -239,9 +237,5 @@ 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"
|
||||
)
|
||||
retrieve_container_file_content = _generated_endpoints.get("retrieve_container_file_content")
|
||||
aretrieve_container_file_content = _generated_endpoints.get("aretrieve_container_file_content")
|
||||
|
|
|
|||
|
|
@ -211,31 +211,23 @@ def create_container(
|
|||
**kwargs,
|
||||
)
|
||||
# get provider config
|
||||
container_provider_config: Optional[BaseContainerConfig] = (
|
||||
ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if container_provider_config is None:
|
||||
raise ValueError(
|
||||
f"container operations are not supported for {custom_llm_provider}"
|
||||
)
|
||||
raise ValueError(f"container operations are not supported for {custom_llm_provider}")
|
||||
|
||||
local_vars.update(kwargs)
|
||||
# Get ContainerCreateOptionalRequestParams with only valid parameters
|
||||
container_create_optional_params: ContainerCreateOptionalRequestParams = (
|
||||
ContainerRequestUtils.get_requested_container_create_optional_param(
|
||||
local_vars
|
||||
)
|
||||
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_provider_config=container_provider_config,
|
||||
container_create_optional_params=container_create_optional_params,
|
||||
)
|
||||
container_create_request_params: Dict = ContainerRequestUtils.get_optional_params_container_create(
|
||||
container_provider_config=container_provider_config,
|
||||
container_create_optional_params=container_create_optional_params,
|
||||
)
|
||||
|
||||
# Pre Call logging
|
||||
|
|
@ -440,22 +432,16 @@ def list_containers(
|
|||
**kwargs,
|
||||
)
|
||||
# get provider config
|
||||
container_provider_config: Optional[BaseContainerConfig] = (
|
||||
ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if container_provider_config is None:
|
||||
raise ValueError(
|
||||
f"Container provider config not found for provider: {custom_llm_provider}"
|
||||
)
|
||||
raise ValueError(f"Container provider config not found for provider: {custom_llm_provider}")
|
||||
|
||||
# Get container list request parameters
|
||||
container_list_optional_params: ContainerListOptionalRequestParams = (
|
||||
ContainerRequestUtils.get_requested_container_list_optional_param(
|
||||
local_vars
|
||||
)
|
||||
ContainerRequestUtils.get_requested_container_list_optional_param(local_vars)
|
||||
)
|
||||
|
||||
# Pre Call logging
|
||||
|
|
@ -641,27 +627,21 @@ def retrieve_container(
|
|||
)
|
||||
|
||||
# Decode container ID and extract provider info
|
||||
original_container_id, resolved_custom_llm_provider, litellm_params = (
|
||||
decode_managed_container_id_for_request(
|
||||
container_id=container_id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
original_container_id, resolved_custom_llm_provider, litellm_params = decode_managed_container_id_for_request(
|
||||
container_id=container_id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
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
|
||||
|
||||
# get provider config
|
||||
container_provider_config: Optional[BaseContainerConfig] = (
|
||||
ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(resolved_custom_llm_provider),
|
||||
)
|
||||
container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(resolved_custom_llm_provider),
|
||||
)
|
||||
|
||||
if container_provider_config is None:
|
||||
raise ValueError(
|
||||
f"Container provider config not found for provider: {resolved_custom_llm_provider}"
|
||||
)
|
||||
raise ValueError(f"Container provider config not found for provider: {resolved_custom_llm_provider}")
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
|
|
@ -865,27 +845,21 @@ def delete_container(
|
|||
)
|
||||
|
||||
# Decode container ID and extract provider info
|
||||
original_container_id, resolved_custom_llm_provider, litellm_params = (
|
||||
decode_managed_container_id_for_request(
|
||||
container_id=container_id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
original_container_id, resolved_custom_llm_provider, litellm_params = decode_managed_container_id_for_request(
|
||||
container_id=container_id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
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
|
||||
|
||||
# get provider config
|
||||
container_provider_config: Optional[BaseContainerConfig] = (
|
||||
ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(resolved_custom_llm_provider),
|
||||
)
|
||||
container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(resolved_custom_llm_provider),
|
||||
)
|
||||
|
||||
if container_provider_config is None:
|
||||
raise ValueError(
|
||||
f"Container provider config not found for provider: {resolved_custom_llm_provider}"
|
||||
)
|
||||
raise ValueError(f"Container provider config not found for provider: {resolved_custom_llm_provider}")
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
|
|
@ -1103,25 +1077,19 @@ def list_container_files(
|
|||
)
|
||||
|
||||
# Decode container ID and extract provider info
|
||||
original_container_id, resolved_custom_llm_provider, litellm_params = (
|
||||
decode_managed_container_id_for_request(
|
||||
container_id=container_id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
original_container_id, resolved_custom_llm_provider, litellm_params = decode_managed_container_id_for_request(
|
||||
container_id=container_id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
# get provider config
|
||||
container_provider_config: Optional[BaseContainerConfig] = (
|
||||
ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(resolved_custom_llm_provider),
|
||||
)
|
||||
container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(resolved_custom_llm_provider),
|
||||
)
|
||||
|
||||
if container_provider_config is None:
|
||||
raise ValueError(
|
||||
f"Container provider config not found for provider: {resolved_custom_llm_provider}"
|
||||
)
|
||||
raise ValueError(f"Container provider config not found for provider: {resolved_custom_llm_provider}")
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
|
|
@ -1363,25 +1331,19 @@ def upload_container_file(
|
|||
)
|
||||
|
||||
# Decode container ID and extract provider info
|
||||
original_container_id, resolved_custom_llm_provider, litellm_params = (
|
||||
decode_managed_container_id_for_request(
|
||||
container_id=container_id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
original_container_id, resolved_custom_llm_provider, litellm_params = decode_managed_container_id_for_request(
|
||||
container_id=container_id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
# get provider config
|
||||
container_provider_config: Optional[BaseContainerConfig] = (
|
||||
ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(resolved_custom_llm_provider),
|
||||
)
|
||||
container_provider_config: Optional[BaseContainerConfig] = ProviderConfigManager.get_provider_container_config(
|
||||
provider=litellm.LlmProviders(resolved_custom_llm_provider),
|
||||
)
|
||||
|
||||
if container_provider_config is None:
|
||||
raise ValueError(
|
||||
f"Container provider config not found for provider: {resolved_custom_llm_provider}"
|
||||
)
|
||||
raise ValueError(f"Container provider config not found for provider: {resolved_custom_llm_provider}")
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
|
|
|
|||
|
|
@ -66,11 +66,7 @@ class ContainerRequestUtils:
|
|||
supported_params = 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 = {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
|
||||
|
|
|
|||
|
|
@ -317,9 +317,7 @@ def cost_per_token(
|
|||
### SERVICE TIER ###
|
||||
service_tier: Optional[str] = None, # for OpenAI service tier pricing
|
||||
### DATA RESIDENCY ###
|
||||
data_residency: Optional[
|
||||
str
|
||||
] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
|
||||
data_residency: Optional[str] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
|
||||
response: Optional[Any] = None,
|
||||
### REQUEST MODEL ###
|
||||
request_model: Optional[str] = None, # original request model for router detection
|
||||
|
|
@ -376,9 +374,7 @@ def cost_per_token(
|
|||
# either `cache_write_tokens` (kimi-k2) or `cache_creation_tokens`.
|
||||
# Mirror db_spend_update_writer to stay symmetric.
|
||||
_cache_creation_tokens = float(
|
||||
getattr(_pt_details, "cache_write_tokens", 0)
|
||||
or getattr(_pt_details, "cache_creation_tokens", 0)
|
||||
or 0
|
||||
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)
|
||||
|
|
@ -451,12 +447,8 @@ def cost_per_token(
|
|||
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}"
|
||||
)
|
||||
if (
|
||||
model_with_provider_and_region in model_cost_ref
|
||||
): # use region based pricing, if it's available
|
||||
model_with_provider_and_region = 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:
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
|
||||
|
|
@ -475,9 +467,7 @@ def cost_per_token(
|
|||
Option2. model = "openai/gpt-4" - model = provider/model
|
||||
Option3. model = "anthropic.claude-3" - model = model
|
||||
"""
|
||||
if (
|
||||
model_with_provider in model_cost_ref
|
||||
): # Option 2. use model with provider, model = "openai/gpt-4"
|
||||
if model_with_provider in model_cost_ref: # Option 2. use model with provider, model = "openai/gpt-4"
|
||||
model = model_with_provider
|
||||
elif model in model_cost_ref: # Option 1. use model passed, model="gpt-4"
|
||||
model = model
|
||||
|
|
@ -488,9 +478,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
|
||||
)
|
||||
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)
|
||||
prompt_cost: float = 0.0
|
||||
completion_cost: float = 0.0
|
||||
|
|
@ -587,11 +575,7 @@ def cost_per_token(
|
|||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
number_of_queries=number_of_queries or 1,
|
||||
optional_params=(
|
||||
response._hidden_params
|
||||
if response and hasattr(response, "_hidden_params")
|
||||
else None
|
||||
),
|
||||
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(
|
||||
|
|
@ -615,13 +599,9 @@ def cost_per_token(
|
|||
service_tier=service_tier,
|
||||
)
|
||||
elif custom_llm_provider == "anthropic":
|
||||
return anthropic_cost_per_token(
|
||||
model=model, usage=usage_block, service_tier=service_tier
|
||||
)
|
||||
return anthropic_cost_per_token(model=model, usage=usage_block, service_tier=service_tier)
|
||||
elif custom_llm_provider == "bedrock":
|
||||
return bedrock_cost_per_token(
|
||||
model=model, usage=usage_block, service_tier=service_tier
|
||||
)
|
||||
return bedrock_cost_per_token(model=model, usage=usage_block, service_tier=service_tier)
|
||||
elif custom_llm_provider == "openai":
|
||||
return openai_cost_per_token(
|
||||
model=model,
|
||||
|
|
@ -641,9 +621,7 @@ def cost_per_token(
|
|||
service_tier=service_tier,
|
||||
)
|
||||
elif custom_llm_provider == "gemini":
|
||||
return gemini_cost_per_token(
|
||||
model=model, usage=usage_block, service_tier=service_tier
|
||||
)
|
||||
return gemini_cost_per_token(model=model, usage=usage_block, service_tier=service_tier)
|
||||
elif custom_llm_provider == "deepseek":
|
||||
return deepseek_cost_per_token(model=model, usage=usage_block)
|
||||
elif custom_llm_provider == "perplexity":
|
||||
|
|
@ -667,13 +645,9 @@ 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 = _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:
|
||||
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(
|
||||
model=model,
|
||||
usage=usage_block,
|
||||
|
|
@ -682,10 +656,7 @@ def cost_per_token(
|
|||
data_residency=data_residency,
|
||||
)
|
||||
|
||||
if (
|
||||
model_info.get("input_cost_per_second", None) is not None
|
||||
and response_time_ms is not None
|
||||
):
|
||||
if model_info.get("input_cost_per_second", None) is not None and response_time_ms is not None:
|
||||
verbose_logger.debug(
|
||||
"For model=%s - input_cost_per_second: %s; response time: %s",
|
||||
model,
|
||||
|
|
@ -697,10 +668,7 @@ def cost_per_token(
|
|||
model_info["input_cost_per_second"] * response_time_ms / 1000 # type: ignore
|
||||
)
|
||||
|
||||
if (
|
||||
model_info.get("output_cost_per_second", None) is not None
|
||||
and response_time_ms is not None
|
||||
):
|
||||
if model_info.get("output_cost_per_second", None) is not None and response_time_ms is not None:
|
||||
verbose_logger.debug(
|
||||
"For model=%s - output_cost_per_second: %s; response time: %s",
|
||||
model,
|
||||
|
|
@ -724,7 +692,9 @@ 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 = DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND # assume all calls sent to A100 80GB for now
|
||||
a100_80gb_price_per_second_public = (
|
||||
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())
|
||||
|
|
@ -773,9 +743,7 @@ def _select_model_name_for_cost_calc(
|
|||
|
||||
return_model: Optional[str] = None
|
||||
region_name: Optional[str] = None
|
||||
custom_llm_provider = _get_provider_for_cost_calc(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
custom_llm_provider = _get_provider_for_cost_calc(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
completion_response_model: Optional[str] = None
|
||||
if completion_response is not None:
|
||||
|
|
@ -788,10 +756,7 @@ def _select_model_name_for_cost_calc(
|
|||
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]
|
||||
if (
|
||||
entry.get("input_cost_per_token") is not None
|
||||
or entry.get("input_cost_per_second") is not None
|
||||
):
|
||||
if entry.get("input_cost_per_token") is not None or entry.get("input_cost_per_second") is not None:
|
||||
return_model = router_model_id
|
||||
else:
|
||||
return_model = model
|
||||
|
|
@ -802,14 +767,9 @@ def _select_model_name_for_cost_calc(
|
|||
return_model = base_model
|
||||
|
||||
elif completion_response_model is None and hidden_params is not None:
|
||||
if (
|
||||
hidden_params.get("model", None) is not None
|
||||
and len(hidden_params["model"]) > 0
|
||||
):
|
||||
if hidden_params.get("model", None) is not None and len(hidden_params["model"]) > 0:
|
||||
return_model = hidden_params.get("model", model)
|
||||
elif (
|
||||
hidden_params is not None and hidden_params.get("region_name", None) is not None
|
||||
):
|
||||
elif hidden_params is not None and hidden_params.get("region_name", None) is not None:
|
||||
region_name = hidden_params.get("region_name", None)
|
||||
|
||||
if return_model is None and completion_response_model is not None:
|
||||
|
|
@ -897,10 +857,7 @@ def _normalize_service_tier(service_tier: object) -> str | None:
|
|||
on the response usage) instead of crashing the downstream cost-key lookup,
|
||||
which calls service_tier.lower()
|
||||
"""
|
||||
if (
|
||||
not isinstance(service_tier, str)
|
||||
or service_tier.lower() == ServiceTier.AUTO.value
|
||||
):
|
||||
if not isinstance(service_tier, str) or service_tier.lower() == ServiceTier.AUTO.value:
|
||||
return None
|
||||
return service_tier
|
||||
|
||||
|
|
@ -926,20 +883,12 @@ def _get_usage_object(
|
|||
and (isinstance(usage_obj, dict) or isinstance(usage_obj, ResponseAPIUsage))
|
||||
and ResponseAPILoggingUtils._is_response_api_usage(usage_obj)
|
||||
):
|
||||
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
usage_obj
|
||||
)
|
||||
elif TranscriptionUsageObjectTransformation.is_transcription_usage_object(
|
||||
usage_obj
|
||||
):
|
||||
return (
|
||||
TranscriptionUsageObjectTransformation.transform_transcription_usage_object(
|
||||
cast(
|
||||
Union[
|
||||
TranscriptionUsageDurationObject, TranscriptionUsageTokensObject
|
||||
],
|
||||
usage_obj,
|
||||
)
|
||||
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage_obj)
|
||||
elif TranscriptionUsageObjectTransformation.is_transcription_usage_object(usage_obj):
|
||||
return TranscriptionUsageObjectTransformation.transform_transcription_usage_object(
|
||||
cast(
|
||||
Union[TranscriptionUsageDurationObject, TranscriptionUsageTokensObject],
|
||||
usage_obj,
|
||||
)
|
||||
)
|
||||
elif isinstance(usage_obj, dict):
|
||||
|
|
@ -947,9 +896,7 @@ def _get_usage_object(
|
|||
elif isinstance(usage_obj, BaseModel):
|
||||
return Usage(**usage_obj.model_dump())
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
f"Unknown usage object type: {type(usage_obj)}, usage_obj: {usage_obj}"
|
||||
)
|
||||
verbose_logger.debug(f"Unknown usage object type: {type(usage_obj)}, usage_obj: {usage_obj}")
|
||||
return None
|
||||
|
||||
|
||||
|
|
@ -958,24 +905,18 @@ def _is_known_usage_objects(usage_obj):
|
|||
return (
|
||||
isinstance(usage_obj, litellm.Usage)
|
||||
or isinstance(usage_obj, ResponseAPIUsage)
|
||||
or TranscriptionUsageObjectTransformation.is_transcription_usage_object(
|
||||
usage_obj
|
||||
)
|
||||
or TranscriptionUsageObjectTransformation.is_transcription_usage_object(usage_obj)
|
||||
)
|
||||
|
||||
|
||||
def _infer_call_type(
|
||||
call_type: Optional[CallTypesLiteral], completion_response: Any
|
||||
) -> Optional[CallTypesLiteral]:
|
||||
def _infer_call_type(call_type: Optional[CallTypesLiteral], completion_response: Any) -> Optional[CallTypesLiteral]:
|
||||
if call_type is not None:
|
||||
return call_type
|
||||
|
||||
if completion_response is None:
|
||||
return None
|
||||
|
||||
if isinstance(completion_response, ModelResponse) or isinstance(
|
||||
completion_response, ModelResponseStream
|
||||
):
|
||||
if isinstance(completion_response, ModelResponse) or isinstance(completion_response, ModelResponseStream):
|
||||
return "completion"
|
||||
elif isinstance(completion_response, EmbeddingResponse):
|
||||
return "embedding"
|
||||
|
|
@ -1053,9 +994,7 @@ def _apply_cost_margin(
|
|||
if custom_llm_provider and custom_llm_provider in litellm.cost_margin_config:
|
||||
margin_config = litellm.cost_margin_config[custom_llm_provider]
|
||||
if verbose_logger.isEnabledFor(logging.DEBUG):
|
||||
verbose_logger.debug(
|
||||
f"Found provider-specific margin config for {custom_llm_provider}: {margin_config}"
|
||||
)
|
||||
verbose_logger.debug(f"Found provider-specific margin config for {custom_llm_provider}: {margin_config}")
|
||||
elif "global" in litellm.cost_margin_config:
|
||||
margin_config = litellm.cost_margin_config["global"]
|
||||
if verbose_logger.isEnabledFor(logging.DEBUG):
|
||||
|
|
@ -1184,9 +1123,7 @@ def completion_cost(
|
|||
### SERVICE TIER ###
|
||||
service_tier: Optional[str] = None, # for OpenAI service tier pricing
|
||||
### DATA RESIDENCY ###
|
||||
data_residency: Optional[
|
||||
str
|
||||
] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
|
||||
data_residency: Optional[str] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
|
||||
) -> float:
|
||||
"""
|
||||
Calculate the cost of a given completion call fot GPT-3.5-turbo, llama2, any litellm supported llm.
|
||||
|
|
@ -1235,9 +1172,7 @@ def completion_cost(
|
|||
cache_creation_input_tokens: Optional[int] = None
|
||||
cache_read_input_tokens: Optional[int] = None
|
||||
audio_transcription_file_duration: float = 0.0
|
||||
cost_per_token_usage_object: Optional[Usage] = _get_usage_object(
|
||||
completion_response=completion_response
|
||||
)
|
||||
cost_per_token_usage_object: Optional[Usage] = _get_usage_object(completion_response=completion_response)
|
||||
rerank_billed_units: Optional[RerankBilledUnits] = None
|
||||
|
||||
# Extract service_tier from optional_params if not provided directly
|
||||
|
|
@ -1258,9 +1193,7 @@ def completion_cost(
|
|||
# Extract service_tier from usage object if not provided
|
||||
if service_tier is None and cost_per_token_usage_object is not None:
|
||||
if isinstance(cost_per_token_usage_object, BaseModel):
|
||||
service_tier = getattr(
|
||||
cost_per_token_usage_object, "service_tier", None
|
||||
)
|
||||
service_tier = getattr(cost_per_token_usage_object, "service_tier", None)
|
||||
elif isinstance(cost_per_token_usage_object, dict):
|
||||
service_tier = cost_per_token_usage_object.get("service_tier")
|
||||
|
||||
|
|
@ -1285,23 +1218,16 @@ def completion_cost(
|
|||
for idx, model in enumerate(potential_model_names):
|
||||
try:
|
||||
if verbose_logger.isEnabledFor(logging.DEBUG):
|
||||
verbose_logger.debug(
|
||||
f"selected model name for cost calculation: {model}"
|
||||
)
|
||||
verbose_logger.debug(f"selected model name for cost calculation: {model}")
|
||||
|
||||
if completion_response is not None and (
|
||||
isinstance(completion_response, BaseModel)
|
||||
or isinstance(completion_response, dict)
|
||||
isinstance(completion_response, BaseModel) or isinstance(completion_response, dict)
|
||||
): # tts returns a custom class
|
||||
if isinstance(completion_response, dict):
|
||||
usage_obj: Optional[Union[dict, Usage]] = (
|
||||
completion_response.get("usage", {})
|
||||
)
|
||||
usage_obj: Optional[Union[dict, Usage]] = completion_response.get("usage", {})
|
||||
else:
|
||||
usage_obj = getattr(completion_response, "usage", {})
|
||||
if isinstance(usage_obj, BaseModel) and not _is_known_usage_objects(
|
||||
usage_obj=usage_obj
|
||||
):
|
||||
if isinstance(usage_obj, BaseModel) and not _is_known_usage_objects(usage_obj=usage_obj):
|
||||
_usage_for_dump = cast(BaseModel, usage_obj)
|
||||
setattr(
|
||||
completion_response,
|
||||
|
|
@ -1319,9 +1245,7 @@ def completion_cost(
|
|||
_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
_usage
|
||||
).model_dump()
|
||||
elif TranscriptionUsageObjectTransformation.is_transcription_usage_object(
|
||||
_usage
|
||||
):
|
||||
elif TranscriptionUsageObjectTransformation.is_transcription_usage_object(_usage):
|
||||
tr_usage = TranscriptionUsageObjectTransformation.transform_transcription_usage_object(
|
||||
cast(
|
||||
Union[
|
||||
|
|
@ -1339,29 +1263,21 @@ def completion_cost(
|
|||
# get input/output tokens from completion_response
|
||||
prompt_tokens = _usage.get("prompt_tokens", 0)
|
||||
completion_tokens = _usage.get("completion_tokens", 0)
|
||||
cache_creation_input_tokens = _usage.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
cache_creation_input_tokens = _usage.get("cache_creation_input_tokens", 0)
|
||||
cache_read_input_tokens = _usage.get("cache_read_input_tokens", 0)
|
||||
if (
|
||||
"prompt_tokens_details" in _usage
|
||||
and _usage["prompt_tokens_details"] != {}
|
||||
and _usage["prompt_tokens_details"]
|
||||
):
|
||||
prompt_tokens_details = (
|
||||
_usage.get("prompt_tokens_details") or {}
|
||||
)
|
||||
cache_read_input_tokens = prompt_tokens_details.get(
|
||||
"cached_tokens", 0
|
||||
)
|
||||
prompt_tokens_details = _usage.get("prompt_tokens_details") or {}
|
||||
cache_read_input_tokens = prompt_tokens_details.get("cached_tokens", 0)
|
||||
|
||||
total_time = getattr(completion_response, "_response_ms", 0)
|
||||
|
||||
hidden_params = getattr(completion_response, "_hidden_params", None)
|
||||
if hidden_params is not None:
|
||||
custom_llm_provider = hidden_params.get(
|
||||
"custom_llm_provider", custom_llm_provider or None
|
||||
)
|
||||
custom_llm_provider = hidden_params.get("custom_llm_provider", custom_llm_provider or None)
|
||||
region_name = hidden_params.get("region_name", region_name)
|
||||
|
||||
# For Gemini/Vertex AI responses, trafficType is stored in
|
||||
|
|
@ -1369,14 +1285,10 @@ def completion_cost(
|
|||
# by the cost key lookup (_priority / _flex suffixes) so that
|
||||
# ON_DEMAND_PRIORITY requests are billed at priority prices.
|
||||
if service_tier is None:
|
||||
provider_specific = (
|
||||
hidden_params.get("provider_specific_fields") or {}
|
||||
)
|
||||
provider_specific = hidden_params.get("provider_specific_fields") or {}
|
||||
raw_traffic_type = provider_specific.get("traffic_type")
|
||||
if raw_traffic_type:
|
||||
service_tier = _map_traffic_type_to_service_tier(
|
||||
raw_traffic_type
|
||||
)
|
||||
service_tier = _map_traffic_type_to_service_tier(raw_traffic_type)
|
||||
else:
|
||||
if model is None:
|
||||
raise ValueError(
|
||||
|
|
@ -1392,9 +1304,7 @@ def completion_cost(
|
|||
if call_type in _A2A_CALL_TYPES:
|
||||
from litellm.a2a_protocol.cost_calculator import A2ACostCalculator
|
||||
|
||||
return A2ACostCalculator.calculate_a2a_cost(
|
||||
litellm_logging_obj=litellm_logging_obj
|
||||
)
|
||||
return A2ACostCalculator.calculate_a2a_cost(litellm_logging_obj=litellm_logging_obj)
|
||||
|
||||
if model is None:
|
||||
raise ValueError(
|
||||
|
|
@ -1411,9 +1321,9 @@ def completion_cost(
|
|||
str(e)
|
||||
)
|
||||
)
|
||||
if CostCalculatorUtils._call_type_has_image_response(
|
||||
call_type
|
||||
) and isinstance(completion_response, ImageResponse):
|
||||
if CostCalculatorUtils._call_type_has_image_response(call_type) and isinstance(
|
||||
completion_response, ImageResponse
|
||||
):
|
||||
### IMAGE GENERATION COST CALCULATION ###
|
||||
return CostCalculatorUtils.route_image_generation_cost_calculator(
|
||||
model=model,
|
||||
|
|
@ -1430,9 +1340,7 @@ def completion_cost(
|
|||
# Extract custom model_info for deployment-specific pricing
|
||||
_video_model_info: Optional[ModelInfo] = None
|
||||
if custom_pricing and litellm_logging_obj is not None:
|
||||
_litellm_params = getattr(
|
||||
litellm_logging_obj, "litellm_params", None
|
||||
)
|
||||
_litellm_params = getattr(litellm_logging_obj, "litellm_params", None)
|
||||
if _litellm_params is not None:
|
||||
_metadata = _litellm_params.get("metadata", {}) or {}
|
||||
_video_model_info = _metadata.get("model_info", None)
|
||||
|
|
@ -1446,9 +1354,7 @@ def completion_cost(
|
|||
duration_seconds = usage_obj.get("duration_seconds", None)
|
||||
_vr = usage_obj.get("video_resolution", None)
|
||||
else:
|
||||
duration_seconds = getattr(
|
||||
usage_obj, "duration_seconds", None
|
||||
)
|
||||
duration_seconds = getattr(usage_obj, "duration_seconds", None)
|
||||
_vr = getattr(usage_obj, "video_resolution", None)
|
||||
if _vr is not None:
|
||||
video_resolution = str(_vr).strip().lower()
|
||||
|
|
@ -1487,9 +1393,7 @@ def completion_cost(
|
|||
getattr(completion_response, "duration", 0.0),
|
||||
)
|
||||
elif call_type in _RERANK_CALL_TYPES:
|
||||
if completion_response is not None and isinstance(
|
||||
completion_response, RerankResponse
|
||||
):
|
||||
if completion_response is not None and isinstance(completion_response, RerankResponse):
|
||||
meta_obj = completion_response.meta
|
||||
if meta_obj is not None:
|
||||
billed_units = meta_obj.get("billed_units", {}) or {}
|
||||
|
|
@ -1501,9 +1405,7 @@ def completion_cost(
|
|||
total_tokens=billed_units.get("total_tokens"),
|
||||
)
|
||||
|
||||
search_units = (
|
||||
billed_units.get("search_units") or 1
|
||||
) # cohere charges per request by default.
|
||||
search_units = billed_units.get("search_units") or 1 # cohere charges per request by default.
|
||||
completion_tokens = search_units
|
||||
elif call_type in _SEARCH_CALL_TYPES:
|
||||
from litellm.search import search_provider_cost_per_query
|
||||
|
|
@ -1577,10 +1479,7 @@ def completion_cost(
|
|||
elif call_type == _AREALTIME_CALL_TYPE and isinstance(
|
||||
completion_response, LiteLLMRealtimeStreamLoggingObject
|
||||
):
|
||||
if (
|
||||
cost_per_token_usage_object is None
|
||||
or custom_llm_provider is None
|
||||
):
|
||||
if cost_per_token_usage_object is None or custom_llm_provider is None:
|
||||
raise ValueError(
|
||||
"usage object and custom_llm_provider must be provided for realtime stream cost calculation. Got cost_per_token_usage_object={}, custom_llm_provider={}".format(
|
||||
cost_per_token_usage_object,
|
||||
|
|
@ -1599,59 +1498,36 @@ def completion_cost(
|
|||
MCPCostCalculator,
|
||||
)
|
||||
|
||||
return MCPCostCalculator.calculate_mcp_tool_call_cost(
|
||||
litellm_logging_obj=litellm_logging_obj
|
||||
)
|
||||
return MCPCostCalculator.calculate_mcp_tool_call_cost(litellm_logging_obj=litellm_logging_obj)
|
||||
# Calculate cost based on prompt_tokens, completion_tokens
|
||||
if (
|
||||
"togethercomputer" in model
|
||||
or "together_ai" in model
|
||||
or custom_llm_provider == "together_ai"
|
||||
):
|
||||
if "togethercomputer" in model or "together_ai" in model or custom_llm_provider == "together_ai":
|
||||
# together ai prices based on size of llm
|
||||
# get_model_params_and_category takes a model name and returns the category of LLM size it is in model_prices_and_context_window.json
|
||||
|
||||
model = get_model_params_and_category(
|
||||
model, call_type=CallTypes(call_type)
|
||||
)
|
||||
model = get_model_params_and_category(model, call_type=CallTypes(call_type))
|
||||
|
||||
# replicate llms are calculate based on time for request running
|
||||
# see https://replicate.com/pricing
|
||||
elif (
|
||||
model in litellm.replicate_models or "replicate" in model
|
||||
) and model not in litellm.model_cost:
|
||||
elif (model in litellm.replicate_models or "replicate" in model) and model not in litellm.model_cost:
|
||||
# for unmapped replicate model, default to replicate's time tracking logic
|
||||
return get_replicate_completion_pricing(
|
||||
completion_response, total_time
|
||||
) # type: ignore
|
||||
return get_replicate_completion_pricing(completion_response, total_time) # type: ignore
|
||||
|
||||
if model is None:
|
||||
raise ValueError(
|
||||
f"Model is None and does not exist in passed completion_response. Passed completion_response={completion_response}, model={model}"
|
||||
)
|
||||
|
||||
if (
|
||||
custom_llm_provider is not None
|
||||
and custom_llm_provider == "vertex_ai"
|
||||
):
|
||||
if custom_llm_provider is not None and custom_llm_provider == "vertex_ai":
|
||||
# Calculate the prompt characters + response characters
|
||||
if len(messages) > 0:
|
||||
prompt_string = litellm.utils.get_formatted_prompt(
|
||||
data={"messages": messages}, call_type="completion"
|
||||
)
|
||||
|
||||
prompt_characters = litellm.utils._count_characters(
|
||||
text=prompt_string
|
||||
)
|
||||
if completion_response is not None and isinstance(
|
||||
completion_response, ModelResponse
|
||||
):
|
||||
completion_string = litellm.utils.get_response_string(
|
||||
response_obj=completion_response
|
||||
)
|
||||
completion_characters = litellm.utils._count_characters(
|
||||
text=completion_string
|
||||
)
|
||||
prompt_characters = litellm.utils._count_characters(text=prompt_string)
|
||||
if completion_response is not None and isinstance(completion_response, ModelResponse):
|
||||
completion_string = litellm.utils.get_response_string(response_obj=completion_response)
|
||||
completion_characters = litellm.utils._count_characters(text=completion_string)
|
||||
|
||||
# Get the original request model for router detection
|
||||
request_model_for_cost = None
|
||||
|
|
@ -1688,12 +1564,8 @@ def completion_cost(
|
|||
if custom_llm_provider == "azure_ai":
|
||||
model_for_additional_costs = request_model_for_cost
|
||||
if completion_response is not None:
|
||||
hidden_params = (
|
||||
getattr(completion_response, "_hidden_params", None) or {}
|
||||
)
|
||||
hidden_model = hidden_params.get("model") or hidden_params.get(
|
||||
"litellm_model_name"
|
||||
)
|
||||
hidden_params = getattr(completion_response, "_hidden_params", None) or {}
|
||||
hidden_model = hidden_params.get("model") or hidden_params.get("litellm_model_name")
|
||||
if hidden_model and (
|
||||
"model_router" in (hidden_model or "").lower()
|
||||
or "model-router" in (hidden_model or "").lower()
|
||||
|
|
@ -1712,17 +1584,13 @@ def completion_cost(
|
|||
else:
|
||||
additional_costs = None
|
||||
|
||||
_final_cost = (
|
||||
prompt_tokens_cost_usd_dollar + completion_tokens_cost_usd_dollar
|
||||
)
|
||||
cost_for_built_in_tools = (
|
||||
StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
|
||||
model=model,
|
||||
response_object=completion_response,
|
||||
usage=cost_per_token_usage_object,
|
||||
standard_built_in_tools_params=standard_built_in_tools_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
_final_cost = prompt_tokens_cost_usd_dollar + completion_tokens_cost_usd_dollar
|
||||
cost_for_built_in_tools = StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
|
||||
model=model,
|
||||
response_object=completion_response,
|
||||
usage=cost_per_token_usage_object,
|
||||
standard_built_in_tools_params=standard_built_in_tools_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
_final_cost += cost_for_built_in_tools
|
||||
if additional_costs:
|
||||
|
|
@ -1763,23 +1631,17 @@ def completion_cost(
|
|||
_cache_read_cost: Optional[float] = None
|
||||
_cache_creation_cost: Optional[float] = None
|
||||
if cost_per_token_usage_object is not None:
|
||||
_cr = getattr(
|
||||
cost_per_token_usage_object, "cache_read_input_tokens", None
|
||||
) or (cost_per_token_usage_object.model_extra or {}).get(
|
||||
"cache_read_input_tokens"
|
||||
)
|
||||
_cr = getattr(cost_per_token_usage_object, "cache_read_input_tokens", None) or (
|
||||
cost_per_token_usage_object.model_extra or {}
|
||||
).get("cache_read_input_tokens")
|
||||
_cc = getattr(
|
||||
cost_per_token_usage_object,
|
||||
"cache_creation_input_tokens",
|
||||
None,
|
||||
) or (cost_per_token_usage_object.model_extra or {}).get(
|
||||
"cache_creation_input_tokens"
|
||||
)
|
||||
) or (cost_per_token_usage_object.model_extra or {}).get("cache_creation_input_tokens")
|
||||
if (_cr or _cc) and model:
|
||||
try:
|
||||
_mi = litellm.get_model_info(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
_mi = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
_cr_rate = _mi.get("cache_read_input_token_cost")
|
||||
if _cr and _cr_rate is not None:
|
||||
_cache_read_cost = float(_cr) * float(_cr_rate)
|
||||
|
|
@ -1814,11 +1676,7 @@ def completion_cost(
|
|||
)
|
||||
if idx == len(potential_model_names) - 1:
|
||||
raise e
|
||||
raise Exception(
|
||||
"Unable to calculat cost for received potential model names - {}".format(
|
||||
potential_model_names
|
||||
)
|
||||
)
|
||||
raise Exception("Unable to calculat cost for received potential model names - {}".format(potential_model_names))
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
|
@ -1832,10 +1690,7 @@ def get_response_cost_from_hidden_params(
|
|||
_hidden_params_dict = hidden_params
|
||||
|
||||
additional_headers = _hidden_params_dict.get("additional_headers", {})
|
||||
if (
|
||||
additional_headers
|
||||
and "llm_provider-x-litellm-response-cost" in 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"]
|
||||
if response_cost is None:
|
||||
return None
|
||||
|
|
@ -1892,9 +1747,7 @@ def response_cost_calculator(
|
|||
### SERVICE TIER ###
|
||||
service_tier: Optional[str] = None, # for OpenAI service tier pricing
|
||||
### DATA RESIDENCY ###
|
||||
data_residency: Optional[
|
||||
str
|
||||
] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
|
||||
data_residency: Optional[str] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
|
||||
) -> float:
|
||||
"""
|
||||
Returns
|
||||
|
|
@ -1908,9 +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 = get_response_cost_from_hidden_params(response_object._hidden_params)
|
||||
if provider_response_cost is not None:
|
||||
return provider_response_cost
|
||||
|
||||
|
|
@ -1957,17 +1808,13 @@ def ocr_cost(
|
|||
# validate it's an OCR response
|
||||
#########################################################
|
||||
if response is None or not isinstance(response, OCRResponse):
|
||||
raise ValueError(
|
||||
f"response must be of type OCRResponse got type={type(response)}"
|
||||
)
|
||||
raise ValueError(f"response must be of type OCRResponse got type={type(response)}")
|
||||
|
||||
if response.usage_info is None:
|
||||
raise ValueError("OCR response usage_info is None")
|
||||
|
||||
try:
|
||||
model_info: Optional[ModelInfo] = litellm.get_model_info(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
model_info: Optional[ModelInfo] = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
except Exception:
|
||||
model_info = None
|
||||
|
||||
|
|
@ -2043,9 +1890,7 @@ def vector_store_search_cost(
|
|||
)
|
||||
|
||||
if config is None:
|
||||
verbose_logger.debug(
|
||||
f"Vector store search is not supported for {custom_llm_provider}"
|
||||
)
|
||||
verbose_logger.debug(f"Vector store search is not supported for {custom_llm_provider}")
|
||||
return 0.0, 0.0
|
||||
|
||||
return config.calculate_vector_store_cost(
|
||||
|
|
@ -2062,9 +1907,7 @@ def rerank_cost(
|
|||
Returns
|
||||
- float or None: cost of response OR none if error.
|
||||
"""
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
try:
|
||||
config = ProviderConfigManager.get_provider_rerank_config(
|
||||
|
|
@ -2091,12 +1934,8 @@ def rerank_cost(
|
|||
raise e
|
||||
|
||||
|
||||
def transcription_cost(
|
||||
model: str, custom_llm_provider: Optional[str], duration: float
|
||||
) -> Tuple[float, float]:
|
||||
return openai_cost_per_second(
|
||||
model=model, custom_llm_provider=custom_llm_provider, duration=duration
|
||||
)
|
||||
def transcription_cost(model: str, custom_llm_provider: Optional[str], duration: float) -> Tuple[float, float]:
|
||||
return openai_cost_per_second(model=model, custom_llm_provider=custom_llm_provider, duration=duration)
|
||||
|
||||
|
||||
def default_image_cost_calculator(
|
||||
|
|
@ -2125,11 +1964,7 @@ def default_image_cost_calculator(
|
|||
"""
|
||||
# Standardize size format to use "-x-"
|
||||
size_str: str = size or "1024-x-1024"
|
||||
size_str = (
|
||||
size_str.replace("x", "-x-")
|
||||
if "x" in size_str and "-x-" not in size_str
|
||||
else size_str
|
||||
)
|
||||
size_str = size_str.replace("x", "-x-") if "x" in size_str and "-x-" not in size_str else size_str
|
||||
|
||||
# Parse dimensions
|
||||
height, width = map(int, size_str.split("-x-"))
|
||||
|
|
@ -2138,29 +1973,17 @@ def default_image_cost_calculator(
|
|||
base_model_name = f"{size_str}/{model}"
|
||||
model_name_without_custom_llm_provider: Optional[str] = None
|
||||
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_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
|
||||
|
||||
# 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 = f"{ImageGenerationRequestQuality.HIGH.value}/{base_model_name}"
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Looking up cost for models: {model_name_with_quality}, {base_model_name}"
|
||||
)
|
||||
verbose_logger.debug(f"Looking up cost for models: {model_name_with_quality}, {base_model_name}")
|
||||
|
||||
model_without_provider = f"{size_str}/{model.split('/')[-1]}"
|
||||
model_with_quality_without_provider = (
|
||||
f"{quality}/{model_without_provider}" if quality else model_without_provider
|
||||
)
|
||||
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: Optional[dict] = None
|
||||
|
|
@ -2178,26 +2001,16 @@ def default_image_cost_calculator(
|
|||
cost_info = litellm.model_cost[_model]
|
||||
break
|
||||
if cost_info is None:
|
||||
raise Exception(
|
||||
f"Model not found in cost map. Tried checking {models_to_check}"
|
||||
)
|
||||
raise Exception(f"Model not found in cost map. Tried checking {models_to_check}")
|
||||
|
||||
# Priority 1: Use per-image pricing if available (for gpt-image-1 and similar models)
|
||||
if (
|
||||
"input_cost_per_image" in cost_info
|
||||
and cost_info["input_cost_per_image"] is not None
|
||||
):
|
||||
if "input_cost_per_image" in cost_info and cost_info["input_cost_per_image"] is not None:
|
||||
return cost_info["input_cost_per_image"] * n
|
||||
# Priority 2: Fall back to per-pixel pricing for backward compatibility
|
||||
elif (
|
||||
"input_cost_per_pixel" in cost_info
|
||||
and cost_info["input_cost_per_pixel"] is not None
|
||||
):
|
||||
elif "input_cost_per_pixel" in cost_info and cost_info["input_cost_per_pixel"] is not None:
|
||||
return cost_info["input_cost_per_pixel"] * height * width * n
|
||||
else:
|
||||
raise Exception(
|
||||
f"No pricing information found for model {model}. Tried checking {models_to_check}"
|
||||
)
|
||||
raise Exception(f"No pricing information found for model {model}. Tried checking {models_to_check}")
|
||||
|
||||
|
||||
def default_video_cost_calculator(
|
||||
|
|
@ -2234,12 +2047,8 @@ def default_video_cost_calculator(
|
|||
base_model_name = model
|
||||
model_name_without_custom_llm_provider: Optional[str] = None
|
||||
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}/{model_name_without_custom_llm_provider}"
|
||||
)
|
||||
model_name_without_custom_llm_provider = model.replace(f"{custom_llm_provider}/", "")
|
||||
base_model_name = f"{custom_llm_provider}/{model_name_without_custom_llm_provider}"
|
||||
|
||||
verbose_logger.debug(f"Looking up cost for video model: {base_model_name}")
|
||||
|
||||
|
|
@ -2299,9 +2108,7 @@ def batch_cost_calculator(
|
|||
deployment-specific pricing is used.
|
||||
"""
|
||||
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
verbose_logger.debug(
|
||||
"Calculating batch cost per token. model=%s, custom_llm_provider=%s",
|
||||
|
|
@ -2311,9 +2118,7 @@ def batch_cost_calculator(
|
|||
|
||||
if model_info is None:
|
||||
try:
|
||||
model_info = litellm.get_model_info(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
except Exception:
|
||||
model_info = None
|
||||
elif not any(
|
||||
|
|
@ -2329,9 +2134,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 = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
if global_info:
|
||||
model_info = global_info
|
||||
except Exception:
|
||||
|
|
@ -2358,13 +2161,8 @@ def batch_cost_calculator(
|
|||
# Add cache read cost if applicable
|
||||
details = _parse_prompt_tokens_details(usage)
|
||||
cache_read_tokens = details["cache_hit_tokens"]
|
||||
cache_read_cost_key = _get_service_tier_cost_key(
|
||||
"cache_read_input_token_cost", None
|
||||
)
|
||||
total_prompt_cost += (
|
||||
calculate_cost_component(model_info, cache_read_cost_key, cache_read_tokens)
|
||||
/ 2
|
||||
)
|
||||
cache_read_cost_key = _get_service_tier_cost_key("cache_read_input_token_cost", None)
|
||||
total_prompt_cost += calculate_cost_component(model_info, cache_read_cost_key, cache_read_tokens) / 2
|
||||
if output_cost_per_token_batches:
|
||||
total_completion_cost = usage.completion_tokens * output_cost_per_token_batches
|
||||
elif output_cost_per_token:
|
||||
|
|
@ -2409,10 +2207,7 @@ class BaseTokenUsageProcessor:
|
|||
setattr(combined, attr, current_val + new_val)
|
||||
# Handle nested prompt_tokens_details
|
||||
if hasattr(usage, "prompt_tokens_details") and usage.prompt_tokens_details:
|
||||
if (
|
||||
not hasattr(combined, "prompt_tokens_details")
|
||||
or not combined.prompt_tokens_details
|
||||
):
|
||||
if not hasattr(combined, "prompt_tokens_details") or not combined.prompt_tokens_details:
|
||||
combined.prompt_tokens_details = PromptTokensDetailsWrapper()
|
||||
|
||||
# Check what keys exist in the model's prompt_tokens_details
|
||||
|
|
@ -2423,9 +2218,7 @@ class BaseTokenUsageProcessor:
|
|||
and not attr.startswith("_")
|
||||
and not callable(getattr(usage.prompt_tokens_details, attr))
|
||||
):
|
||||
current_val = (
|
||||
getattr(combined.prompt_tokens_details, attr, 0) or 0
|
||||
)
|
||||
current_val = getattr(combined.prompt_tokens_details, attr, 0) or 0
|
||||
new_val = getattr(usage.prompt_tokens_details, attr, 0) or 0
|
||||
if new_val is not None and isinstance(new_val, (int, float)):
|
||||
setattr(
|
||||
|
|
@ -2435,27 +2228,15 @@ class BaseTokenUsageProcessor:
|
|||
)
|
||||
|
||||
# Handle nested completion_tokens_details
|
||||
if (
|
||||
hasattr(usage, "completion_tokens_details")
|
||||
and usage.completion_tokens_details
|
||||
):
|
||||
if (
|
||||
not hasattr(combined, "completion_tokens_details")
|
||||
or not combined.completion_tokens_details
|
||||
):
|
||||
combined.completion_tokens_details = (
|
||||
CompletionTokensDetailsWrapper()
|
||||
)
|
||||
if hasattr(usage, "completion_tokens_details") and usage.completion_tokens_details:
|
||||
if not hasattr(combined, "completion_tokens_details") or not combined.completion_tokens_details:
|
||||
combined.completion_tokens_details = CompletionTokensDetailsWrapper()
|
||||
|
||||
# Check what keys exist in the model's completion_tokens_details
|
||||
# Access model_fields on the class, not the instance, to avoid Pydantic 2.11+ deprecation warnings
|
||||
for attr in type(usage.completion_tokens_details).model_fields:
|
||||
if not attr.startswith("_") and not callable(
|
||||
getattr(usage.completion_tokens_details, attr)
|
||||
):
|
||||
current_val = (
|
||||
getattr(combined.completion_tokens_details, attr, 0) or 0
|
||||
)
|
||||
if not attr.startswith("_") and not callable(getattr(usage.completion_tokens_details, attr)):
|
||||
current_val = getattr(combined.completion_tokens_details, attr, 0) or 0
|
||||
new_val = getattr(usage.completion_tokens_details, attr, 0) or 0
|
||||
if isinstance(new_val, (int, float)):
|
||||
setattr(
|
||||
|
|
@ -2481,10 +2262,8 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor):
|
|||
)
|
||||
usage_objects: List[Usage] = []
|
||||
for result in response_done_events:
|
||||
usage_object = (
|
||||
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
result["response"].get("usage", {})
|
||||
)
|
||||
usage_object = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
result["response"].get("usage", {})
|
||||
)
|
||||
usage_objects.append(usage_object)
|
||||
return usage_objects
|
||||
|
|
@ -2496,14 +2275,8 @@ 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
|
||||
)
|
||||
collected_usage_objects = RealtimeAPITokenUsageProcessor.collect_usage_from_realtime_stream_results(results)
|
||||
combined_usage_object = RealtimeAPITokenUsageProcessor.combine_usage_objects(collected_usage_objects)
|
||||
return combined_usage_object
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -2516,9 +2289,7 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor):
|
|||
)
|
||||
|
||||
|
||||
_TRANSCRIPTION_COMPLETED_EVENT_TYPE = (
|
||||
"conversation.item.input_audio_transcription.completed"
|
||||
)
|
||||
_TRANSCRIPTION_COMPLETED_EVENT_TYPE = "conversation.item.input_audio_transcription.completed"
|
||||
|
||||
|
||||
def handle_realtime_stream_cost_calculation(
|
||||
|
|
@ -2540,9 +2311,7 @@ def handle_realtime_stream_cost_calculation(
|
|||
potential_model_names = []
|
||||
for result in results:
|
||||
if result["type"] == "session.created":
|
||||
received_model = cast(OpenAIRealtimeStreamSessionEvents, result)[
|
||||
"session"
|
||||
].get("model", None)
|
||||
received_model = cast(OpenAIRealtimeStreamSessionEvents, result)["session"].get("model", None)
|
||||
potential_model_names.append(received_model)
|
||||
|
||||
potential_model_names.append(litellm_model_name)
|
||||
|
|
@ -2591,20 +2360,14 @@ def handle_realtime_transcription_cost_calculation(
|
|||
- {"type": "tokens", "input_tokens": ...} → priced via input/audio token cost
|
||||
"""
|
||||
completed_events = [
|
||||
cast(dict, result)
|
||||
for result in results
|
||||
if result.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE
|
||||
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 = _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
|
||||
)
|
||||
model_info = litellm.get_model_info(model=model_name, custom_llm_provider=custom_llm_provider)
|
||||
except Exception:
|
||||
model_info = None
|
||||
|
||||
|
|
@ -2627,9 +2390,9 @@ def _get_transcription_model_name_from_results(
|
|||
"session.updated",
|
||||
):
|
||||
session = cast(dict, result).get("session", {}) or {}
|
||||
transcription = (
|
||||
(session.get("audio", {}) or {}).get("input", {}) or {}
|
||||
).get("transcription", {}) or session.get("input_audio_transcription", {})
|
||||
transcription = ((session.get("audio", {}) or {}).get("input", {}) or {}).get(
|
||||
"transcription", {}
|
||||
) or session.get("input_audio_transcription", {})
|
||||
model = (transcription or {}).get("model") or session.get("model")
|
||||
if model:
|
||||
return model
|
||||
|
|
@ -2650,15 +2413,9 @@ def _transcription_usage_cost(usage: dict, model_info: Optional[ModelInfo]) -> f
|
|||
text_tokens = input_token_details.get("text_tokens") or 0
|
||||
output_tokens = usage.get("output_tokens") or 0
|
||||
audio_cost = 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
|
||||
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)
|
||||
return audio_cost + text_cost + output_cost
|
||||
return 0.0
|
||||
|
|
|
|||
|
|
@ -29,9 +29,7 @@ class SpeechToCompletionBridgeHandler:
|
|||
super().__init__()
|
||||
self.transformation_handler = SpeechToCompletionBridgeTransformationHandler()
|
||||
|
||||
def validate_input_kwargs(
|
||||
self, kwargs: dict
|
||||
) -> SpeechToCompletionBridgeHandlerInputKwargs:
|
||||
def validate_input_kwargs(self, kwargs: dict) -> SpeechToCompletionBridgeHandlerInputKwargs:
|
||||
from litellm import LiteLLMLoggingObj
|
||||
|
||||
model = kwargs.get("model")
|
||||
|
|
|
|||
|
|
@ -29,9 +29,7 @@ class SpeechToCompletionBridgeTransformationHandler:
|
|||
if isinstance(voice, str):
|
||||
passed_optional_params["audio"] = {"voice": voice}
|
||||
if "response_format" in optional_params:
|
||||
passed_optional_params["audio"]["format"] = optional_params[
|
||||
"response_format"
|
||||
]
|
||||
passed_optional_params["audio"]["format"] = optional_params["response_format"]
|
||||
|
||||
return_kwargs = {
|
||||
"model": model,
|
||||
|
|
@ -53,9 +51,7 @@ class SpeechToCompletionBridgeTransformationHandler:
|
|||
return_kwargs = {k: v for k, v in return_kwargs.items() if v is not None}
|
||||
return return_kwargs
|
||||
|
||||
def _convert_pcm16_to_wav(
|
||||
self, pcm_data: bytes, sample_rate: int = 24000, channels: int = 1
|
||||
) -> bytes:
|
||||
def _convert_pcm16_to_wav(self, pcm_data: bytes, sample_rate: int = 24000, channels: int = 1) -> bytes:
|
||||
"""
|
||||
Convert raw PCM16 data to WAV format.
|
||||
|
||||
|
|
@ -97,13 +93,9 @@ class SpeechToCompletionBridgeTransformationHandler:
|
|||
|
||||
def _is_gemini_tts_model(self, model: str) -> bool:
|
||||
"""Check if the model is a Gemini TTS model that returns PCM16 data."""
|
||||
return "gemini" in model.lower() and (
|
||||
"tts" in model.lower() or "preview-tts" in model.lower()
|
||||
)
|
||||
return "gemini" in model.lower() and ("tts" in model.lower() or "preview-tts" in model.lower())
|
||||
|
||||
def transform_response(
|
||||
self, model_response: "ModelResponse"
|
||||
) -> "HttpxBinaryResponseContent":
|
||||
def transform_response(self, model_response: "ModelResponse") -> "HttpxBinaryResponseContent":
|
||||
import base64
|
||||
|
||||
import httpx
|
||||
|
|
|
|||
|
|
@ -152,10 +152,8 @@ def create_eval(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if evals_api_provider_config is None:
|
||||
|
|
@ -175,9 +173,7 @@ def create_eval(
|
|||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
headers = evals_api_provider_config.validate_environment(
|
||||
headers=headers, litellm_params=litellm_params
|
||||
)
|
||||
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(
|
||||
|
|
@ -188,9 +184,7 @@ def create_eval(
|
|||
|
||||
# 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"
|
||||
)
|
||||
url = evals_api_provider_config.get_complete_url(api_base=api_base, endpoint="evals")
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
|
|
@ -343,10 +337,8 @@ def list_evals(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if evals_api_provider_config is None:
|
||||
|
|
@ -371,9 +363,7 @@ def list_evals(
|
|||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
headers = evals_api_provider_config.validate_environment(
|
||||
headers=headers, litellm_params=litellm_params
|
||||
)
|
||||
headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
# Transform request
|
||||
url, query_params = evals_api_provider_config.transform_list_evals_request(
|
||||
|
|
@ -513,10 +503,8 @@ def get_eval(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if evals_api_provider_config is None:
|
||||
|
|
@ -524,9 +512,7 @@ def get_eval(
|
|||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
headers = evals_api_provider_config.validate_environment(
|
||||
headers=headers, litellm_params=litellm_params
|
||||
)
|
||||
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
|
||||
|
|
@ -682,10 +668,8 @@ def update_eval(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if evals_api_provider_config is None:
|
||||
|
|
@ -732,9 +716,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 = {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
|
||||
|
||||
|
|
@ -744,9 +726,7 @@ def update_eval(
|
|||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
headers = evals_api_provider_config.validate_environment(
|
||||
headers=headers, litellm_params=litellm_params
|
||||
)
|
||||
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
|
||||
|
|
@ -893,10 +873,8 @@ def delete_eval(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if evals_api_provider_config is None:
|
||||
|
|
@ -904,9 +882,7 @@ def delete_eval(
|
|||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
headers = evals_api_provider_config.validate_environment(
|
||||
headers=headers, litellm_params=litellm_params
|
||||
)
|
||||
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
|
||||
|
|
@ -1047,10 +1023,8 @@ def cancel_eval(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if evals_api_provider_config is None:
|
||||
|
|
@ -1058,9 +1032,7 @@ def cancel_eval(
|
|||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
headers = evals_api_provider_config.validate_environment(
|
||||
headers=headers, litellm_params=litellm_params
|
||||
)
|
||||
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
|
||||
|
|
@ -1230,10 +1202,8 @@ def create_run(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if evals_api_provider_config is None:
|
||||
|
|
@ -1254,9 +1224,7 @@ def create_run(
|
|||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
headers = evals_api_provider_config.validate_environment(
|
||||
headers=headers, litellm_params=litellm_params
|
||||
)
|
||||
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
|
||||
|
|
@ -1418,10 +1386,8 @@ def list_runs(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if evals_api_provider_config is None:
|
||||
|
|
@ -1444,9 +1410,7 @@ def list_runs(
|
|||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
headers = evals_api_provider_config.validate_environment(
|
||||
headers=headers, litellm_params=litellm_params
|
||||
)
|
||||
headers = evals_api_provider_config.validate_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
# Transform request
|
||||
url, query_params = evals_api_provider_config.transform_list_runs_request(
|
||||
|
|
@ -1592,10 +1556,8 @@ def get_run(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if evals_api_provider_config is None:
|
||||
|
|
@ -1603,9 +1565,7 @@ def get_run(
|
|||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
headers = evals_api_provider_config.validate_environment(
|
||||
headers=headers, litellm_params=litellm_params
|
||||
)
|
||||
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
|
||||
|
|
@ -1752,10 +1712,8 @@ def cancel_run(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if evals_api_provider_config is None:
|
||||
|
|
@ -1763,9 +1721,7 @@ def cancel_run(
|
|||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
headers = evals_api_provider_config.validate_environment(
|
||||
headers=headers, litellm_params=litellm_params
|
||||
)
|
||||
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
|
||||
|
|
@ -1921,10 +1877,8 @@ def delete_run(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
evals_api_provider_config: Optional[BaseEvalsAPIConfig] = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if evals_api_provider_config is None:
|
||||
|
|
@ -1932,9 +1886,7 @@ def delete_run(
|
|||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
headers = evals_api_provider_config.validate_environment(
|
||||
headers=headers, litellm_params=litellm_params
|
||||
)
|
||||
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
|
||||
|
|
|
|||
|
|
@ -146,9 +146,7 @@ class AuthenticationError(openai.AuthenticationError): # type: ignore
|
|||
self.num_retries = num_retries
|
||||
self.response = response or httpx.Response(
|
||||
status_code=self.status_code,
|
||||
request=httpx.Request(
|
||||
method="GET", url="https://litellm.ai"
|
||||
), # mock request object
|
||||
request=httpx.Request(method="GET", url="https://litellm.ai"), # mock request object
|
||||
)
|
||||
super().__init__(
|
||||
self.message, response=self.response, body=None
|
||||
|
|
@ -192,9 +190,7 @@ class NotFoundError(openai.NotFoundError): # type: ignore
|
|||
self.num_retries = num_retries
|
||||
self.response = response or httpx.Response(
|
||||
status_code=self.status_code,
|
||||
request=httpx.Request(
|
||||
method="GET", url="https://litellm.ai"
|
||||
), # mock request object
|
||||
request=httpx.Request(method="GET", url="https://litellm.ai"), # mock request object
|
||||
)
|
||||
super().__init__(
|
||||
self.message, response=self.response, body=None
|
||||
|
|
@ -347,9 +343,7 @@ class Timeout(openai.APITimeoutError): # type: ignore
|
|||
method="POST",
|
||||
url="https://api.openai.com/v1",
|
||||
)
|
||||
super().__init__(
|
||||
request=request
|
||||
) # Call the base class constructor with the parameters it needs
|
||||
super().__init__(request=request) # Call the base class constructor with the parameters it needs
|
||||
self.status_code = exception_status_code or 408
|
||||
self.message = "litellm.Timeout: {}".format(message)
|
||||
self.model = model
|
||||
|
|
@ -438,9 +432,7 @@ class RateLimitError(openai.RateLimitError): # type: ignore
|
|||
litellm_debug_info: Optional[str] = None,
|
||||
max_retries: Optional[int] = None,
|
||||
num_retries: Optional[int] = None,
|
||||
category: Union[str, RateLimitErrorCategory] = (
|
||||
RateLimitErrorCategory.VENDOR_RATE_LIMIT
|
||||
),
|
||||
category: Union[str, RateLimitErrorCategory] = (RateLimitErrorCategory.VENDOR_RATE_LIMIT),
|
||||
rate_limit_type: Optional[Union[str, RateLimitType]] = None,
|
||||
headers: Optional[Dict[str, str]] = None,
|
||||
detail: Any = None,
|
||||
|
|
@ -452,16 +444,12 @@ class RateLimitError(openai.RateLimitError): # type: ignore
|
|||
self.litellm_debug_info = litellm_debug_info
|
||||
self.max_retries = max_retries
|
||||
self.num_retries = num_retries
|
||||
self.category = (
|
||||
category.value if isinstance(category, RateLimitErrorCategory) else category
|
||||
)
|
||||
self.category = category.value if isinstance(category, RateLimitErrorCategory) else category
|
||||
# Which dimension was exceeded — request count, token count, parallel
|
||||
# requests, budget, max iterations. None when the source didn't
|
||||
# classify the failure (e.g. legacy vendor 429 with no header hints).
|
||||
self.rate_limit_type: Optional[str] = (
|
||||
rate_limit_type.value
|
||||
if isinstance(rate_limit_type, RateLimitType)
|
||||
else rate_limit_type
|
||||
rate_limit_type.value if isinstance(rate_limit_type, RateLimitType) else rate_limit_type
|
||||
)
|
||||
# Headers explicitly attached to the error (e.g. retry-after,
|
||||
# rate_limit_type, reset_at). Preserved across the proxy boundary so
|
||||
|
|
@ -476,12 +464,8 @@ 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
|
||||
)
|
||||
self.headers: Optional[Dict[str, str]] = (
|
||||
{k: str(v) for k, v in headers.items()} if headers else None
|
||||
)
|
||||
_response_headers = getattr(response, "headers", None) if response is not None else None
|
||||
self.headers: Optional[Dict[str, str]] = {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.
|
||||
self.detail = detail if detail is not None else self.message
|
||||
|
|
@ -664,9 +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 = getattr(response, "headers", None) if response is not None else None
|
||||
self.response = httpx.Response(
|
||||
status_code=self.status_code,
|
||||
headers=_response_headers,
|
||||
|
|
@ -714,9 +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 = getattr(response, "headers", None) if response is not None else None
|
||||
self.response = httpx.Response(
|
||||
status_code=self.status_code,
|
||||
headers=_response_headers,
|
||||
|
|
@ -764,9 +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 = getattr(response, "headers", None) if response is not None else None
|
||||
self.response = httpx.Response(
|
||||
status_code=self.status_code,
|
||||
headers=_response_headers,
|
||||
|
|
@ -915,9 +893,7 @@ class APIResponseValidationError(openai.APIResponseValidationError): # type: ig
|
|||
|
||||
|
||||
class JSONSchemaValidationError(APIResponseValidationError):
|
||||
def __init__(
|
||||
self, model: str, llm_provider: str, raw_response: str, schema: str
|
||||
) -> None:
|
||||
def __init__(self, model: str, llm_provider: str, raw_response: str, schema: str) -> None:
|
||||
self.raw_response = raw_response
|
||||
self.schema = schema
|
||||
self.model = model
|
||||
|
|
@ -953,9 +929,7 @@ class UnsupportedParamsError(BadRequestError):
|
|||
self.litellm_debug_info = litellm_debug_info
|
||||
response = response or httpx.Response(
|
||||
status_code=self.status_code,
|
||||
request=httpx.Request(
|
||||
method="GET", url="https://litellm.ai"
|
||||
), # mock request object
|
||||
request=httpx.Request(method="GET", url="https://litellm.ai"), # mock request object
|
||||
)
|
||||
self.max_retries = max_retries
|
||||
self.num_retries = num_retries
|
||||
|
|
@ -1005,10 +979,7 @@ class BudgetExceededError(Exception):
|
|||
# to match the normalization RateLimitError.__init__ performs.
|
||||
self.category: str = RateLimitErrorCategory.LITELLM_RATE_LIMIT.value
|
||||
self.rate_limit_type: str = RateLimitType.BUDGET.value
|
||||
message = (
|
||||
message
|
||||
or f"Budget has been exceeded! Current cost: {current_cost}, Max budget: {max_budget}"
|
||||
)
|
||||
message = message or f"Budget has been exceeded! Current cost: {current_cost}, Max budget: {max_budget}"
|
||||
self.message = message
|
||||
super().__init__(message)
|
||||
|
||||
|
|
@ -1022,9 +993,7 @@ class InvalidRequestError(openai.BadRequestError): # type: ignore
|
|||
self.llm_provider = llm_provider
|
||||
self.response = httpx.Response(
|
||||
status_code=400,
|
||||
request=httpx.Request(
|
||||
method="GET", url="https://litellm.ai"
|
||||
), # mock request object
|
||||
request=httpx.Request(method="GET", url="https://litellm.ai"), # mock request object
|
||||
)
|
||||
super().__init__(
|
||||
message=self.message, response=self.response, body=None
|
||||
|
|
@ -1061,9 +1030,7 @@ class LiteLLMUnknownProvider(BadRequestError):
|
|||
self.message = LiteLLMCommonStrings.llm_provider_not_provided.value.format(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
super().__init__(
|
||||
self.message, model=model, llm_provider=custom_llm_provider, response=None
|
||||
)
|
||||
super().__init__(self.message, model=model, llm_provider=custom_llm_provider, response=None)
|
||||
|
||||
def __str__(self):
|
||||
return self.message
|
||||
|
|
@ -1248,8 +1215,5 @@ class SensitiveDataRouteException(Exception):
|
|||
self.guardrail_name = guardrail_name
|
||||
self.detection_info = detection_info or {}
|
||||
self.sticky_session_routing = sticky_session_routing
|
||||
self.message = (
|
||||
message
|
||||
or f"Sensitive data detected by {guardrail_name}. Routing to model: {route_to_model}"
|
||||
)
|
||||
self.message = message or f"Sensitive data detected by {guardrail_name}. Routing to model: {route_to_model}"
|
||||
super().__init__(self.message)
|
||||
|
|
|
|||
|
|
@ -26,9 +26,7 @@ streamable_http_client: Optional[Any] = None
|
|||
try:
|
||||
import mcp.client.streamable_http as streamable_http_module # type: ignore
|
||||
|
||||
streamable_http_client = getattr(
|
||||
streamable_http_module, "streamable_http_client", None
|
||||
)
|
||||
streamable_http_client = getattr(streamable_http_module, "streamable_http_client", None)
|
||||
except ImportError:
|
||||
pass
|
||||
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
|
||||
|
|
@ -62,9 +60,7 @@ def to_basic_auth(auth_value: str) -> str:
|
|||
|
||||
def _strip_header_whitespace(headers: Dict[str, str]) -> Dict[str, str]:
|
||||
return {
|
||||
(key.strip() if isinstance(key, str) else key): (
|
||||
value.strip() if isinstance(value, str) else value
|
||||
)
|
||||
(key.strip() if isinstance(key, str) else key): (value.strip() if isinstance(value, str) else value)
|
||||
for key, value in headers.items()
|
||||
}
|
||||
|
||||
|
|
@ -107,10 +103,7 @@ class MCPSigV4Auth(httpx.Auth):
|
|||
try:
|
||||
from botocore.credentials import Credentials
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"Missing botocore to use AWS SigV4 authentication. "
|
||||
"Run 'pip install boto3'."
|
||||
)
|
||||
raise ImportError("Missing botocore to use AWS SigV4 authentication. Run 'pip install boto3'.")
|
||||
self.service_name = aws_service_name or "bedrock-agentcore"
|
||||
self.region_name = aws_region_name or "us-east-1"
|
||||
# Note: os.environ/ prefixed values are already resolved by
|
||||
|
|
@ -157,9 +150,7 @@ class MCPSigV4Auth(httpx.Auth):
|
|||
import boto3
|
||||
from botocore.credentials import Credentials
|
||||
|
||||
session_name = (
|
||||
aws_session_name or f"litellm-mcp-{int(__import__('time').time())}"
|
||||
)
|
||||
session_name = aws_session_name or f"litellm-mcp-{int(__import__('time').time())}"
|
||||
sts_kwargs: 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
|
||||
|
|
@ -178,9 +169,7 @@ class MCPSigV4Auth(httpx.Auth):
|
|||
token=sts_creds["SessionToken"],
|
||||
)
|
||||
|
||||
def auth_flow(
|
||||
self, request: httpx.Request
|
||||
) -> Generator[httpx.Request, httpx.Response, None]:
|
||||
def auth_flow(self, request: httpx.Request) -> Generator[httpx.Request, httpx.Response, None]:
|
||||
from botocore.auth import SigV4Auth
|
||||
from botocore.awsrequest import AWSRequest
|
||||
|
||||
|
|
@ -282,10 +271,7 @@ class MCPClient:
|
|||
)
|
||||
# HTTP transport (default)
|
||||
if streamable_http_client is None:
|
||||
raise ImportError(
|
||||
"streamable_http_client is not available. "
|
||||
"Please install mcp with HTTP support."
|
||||
)
|
||||
raise ImportError("streamable_http_client is not available. Please install mcp with HTTP support.")
|
||||
headers = self._get_auth_headers()
|
||||
httpx_client_factory = self._create_httpx_client_factory()
|
||||
verbose_logger.debug("litellm headers for streamable_http_client: %s", headers)
|
||||
|
|
@ -299,9 +285,7 @@ class MCPClient:
|
|||
)
|
||||
return transport_ctx, http_client
|
||||
|
||||
def _get_safe_stdio_env(
|
||||
self, provided_env: Optional[Dict[str, str]]
|
||||
) -> Optional[Dict[str, str]]:
|
||||
def _get_safe_stdio_env(self, provided_env: Optional[Dict[str, str]]) -> Optional[Dict[str, str]]:
|
||||
"""
|
||||
Return a safe environment for the stdio subprocess.
|
||||
|
||||
|
|
@ -393,18 +377,12 @@ class MCPClient:
|
|||
try:
|
||||
await transport_ctx.__aexit__(None, None, None)
|
||||
except BaseException as exit_error:
|
||||
verbose_logger.debug(
|
||||
f"Error during transport context exit: {exit_error}"
|
||||
)
|
||||
verbose_logger.debug(f"Error during transport context exit: {exit_error}")
|
||||
root_cause = _first_non_cancelled_cause(exit_error)
|
||||
if root_cause is not None and isinstance(
|
||||
in_flight_error, asyncio.CancelledError
|
||||
):
|
||||
if root_cause is not None and isinstance(in_flight_error, asyncio.CancelledError):
|
||||
raise root_cause from in_flight_error
|
||||
|
||||
async def run_with_session(
|
||||
self, operation: Callable[[ClientSession], Awaitable[TSessionResult]]
|
||||
) -> TSessionResult:
|
||||
async def run_with_session(self, operation: Callable[[ClientSession], Awaitable[TSessionResult]]) -> TSessionResult:
|
||||
"""Open a session, run the provided coroutine, and clean up."""
|
||||
http_client: Optional[httpx.AsyncClient] = None
|
||||
try:
|
||||
|
|
@ -412,9 +390,7 @@ class MCPClient:
|
|||
transport_ctx, http_client = self._create_transport_context()
|
||||
return await self._execute_session_operation(transport_ctx, operation)
|
||||
except Exception:
|
||||
verbose_logger.warning(
|
||||
"MCP client run_with_session failed for %s", self.server_url or "stdio"
|
||||
)
|
||||
verbose_logger.warning("MCP client run_with_session failed for %s", self.server_url or "stdio")
|
||||
raise
|
||||
finally:
|
||||
if http_client is not None:
|
||||
|
|
@ -483,17 +459,11 @@ class MCPClient:
|
|||
"""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)
|
||||
verbose_logger.debug(
|
||||
f"MCP client using SSL configuration: {type(ssl_config).__name__}"
|
||||
)
|
||||
verbose_logger.debug(f"MCP client using SSL configuration: {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
|
||||
)
|
||||
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
|
||||
return httpx.AsyncClient(
|
||||
headers=headers,
|
||||
|
|
@ -515,9 +485,7 @@ class MCPClient:
|
|||
MCP client (triggering the upstream OAuth flow) rather than
|
||||
masking them as "connected, no tools".
|
||||
"""
|
||||
verbose_logger.debug(
|
||||
f"MCP client listing tools from {self.server_url or 'stdio'}"
|
||||
)
|
||||
verbose_logger.debug(f"MCP client listing tools from {self.server_url or 'stdio'}")
|
||||
|
||||
async def _list_tools_operation(session: ClientSession):
|
||||
return await session.list_tools()
|
||||
|
|
@ -526,9 +494,7 @@ class MCPClient:
|
|||
result = await self.run_with_session(_list_tools_operation)
|
||||
tool_count = len(result.tools)
|
||||
tool_names = [tool.name for tool in result.tools]
|
||||
verbose_logger.info(
|
||||
f"MCP client listed {tool_count} tools from {self.server_url or 'stdio'}: {tool_names}"
|
||||
)
|
||||
verbose_logger.info(f"MCP client listed {tool_count} tools from {self.server_url or 'stdio'}: {tool_names}")
|
||||
return result.tools
|
||||
except asyncio.CancelledError:
|
||||
verbose_logger.warning("MCP client list_tools was cancelled")
|
||||
|
|
@ -562,13 +528,9 @@ class MCPClient:
|
|||
"""
|
||||
Call an MCP Tool.
|
||||
"""
|
||||
verbose_logger.info(
|
||||
f"MCP client calling tool '{call_tool_request_params.name}'"
|
||||
)
|
||||
verbose_logger.info(f"MCP client calling tool '{call_tool_request_params.name}'")
|
||||
|
||||
async def on_progress(
|
||||
progress: float, total: float | None, message: str | None
|
||||
):
|
||||
async def on_progress(progress: float, total: float | None, message: str | None):
|
||||
percentage = (progress / total * 100) if total else 0
|
||||
verbose_logger.info(
|
||||
f"MCP Tool '{call_tool_request_params.name}' progress: "
|
||||
|
|
@ -591,14 +553,10 @@ class MCPClient:
|
|||
|
||||
try:
|
||||
tool_result = await self.run_with_session(_call_tool_operation)
|
||||
verbose_logger.info(
|
||||
f"MCP client tool call '{call_tool_request_params.name}' completed successfully"
|
||||
)
|
||||
verbose_logger.info(f"MCP client tool call '{call_tool_request_params.name}' completed successfully")
|
||||
return tool_result
|
||||
except asyncio.CancelledError:
|
||||
verbose_logger.warning(
|
||||
f"MCP client tool call timed out after {self.timeout}s for {self.server_url}"
|
||||
)
|
||||
verbose_logger.warning(f"MCP client tool call timed out after {self.timeout}s for {self.server_url}")
|
||||
raise
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
|
@ -623,17 +581,13 @@ class MCPClient:
|
|||
)
|
||||
# Return a default error result instead of raising
|
||||
return MCPCallToolResult(
|
||||
content=[
|
||||
TextContent(type="text", text=f"{error_type}: {str(e)}")
|
||||
], # Empty content for error case
|
||||
content=[TextContent(type="text", text=f"{error_type}: {str(e)}")], # Empty content for error case
|
||||
isError=True,
|
||||
)
|
||||
|
||||
async def list_prompts(self) -> List[Prompt]:
|
||||
"""List available prompts from the server."""
|
||||
verbose_logger.debug(
|
||||
f"MCP client listing tools from {self.server_url or 'stdio'}"
|
||||
)
|
||||
verbose_logger.debug(f"MCP client listing tools from {self.server_url or 'stdio'}")
|
||||
|
||||
async def _list_prompts_operation(session: ClientSession):
|
||||
return await session.list_prompts()
|
||||
|
|
@ -667,13 +621,9 @@ class MCPClient:
|
|||
# Return empty list instead of raising to allow graceful degradation
|
||||
return []
|
||||
|
||||
async def get_prompt(
|
||||
self, get_prompt_request_params: GetPromptRequestParams
|
||||
) -> GetPromptResult:
|
||||
async def get_prompt(self, get_prompt_request_params: GetPromptRequestParams) -> GetPromptResult:
|
||||
"""Fetch a prompt definition from the MCP server."""
|
||||
verbose_logger.info(
|
||||
f"MCP client fetching prompt '{get_prompt_request_params.name}'"
|
||||
)
|
||||
verbose_logger.info(f"MCP client fetching prompt '{get_prompt_request_params.name}'")
|
||||
|
||||
async def _get_prompt_operation(session: ClientSession):
|
||||
verbose_logger.debug("MCP client sending get_prompt request to session")
|
||||
|
|
@ -684,9 +634,7 @@ class MCPClient:
|
|||
|
||||
try:
|
||||
get_prompt_result = await self.run_with_session(_get_prompt_operation)
|
||||
verbose_logger.info(
|
||||
f"MCP client get_prompt '{get_prompt_request_params.name}' completed successfully"
|
||||
)
|
||||
verbose_logger.info(f"MCP client get_prompt '{get_prompt_request_params.name}' completed successfully")
|
||||
return get_prompt_result
|
||||
except asyncio.CancelledError:
|
||||
verbose_logger.warning("MCP client get_prompt was cancelled")
|
||||
|
|
@ -716,9 +664,7 @@ class MCPClient:
|
|||
|
||||
async def list_resources(self) -> list[Resource]:
|
||||
"""List available resources from the server."""
|
||||
verbose_logger.debug(
|
||||
f"MCP client listing resources from {self.server_url or 'stdio'}"
|
||||
)
|
||||
verbose_logger.debug(f"MCP client listing resources from {self.server_url or 'stdio'}")
|
||||
|
||||
async def _list_resources_operation(session: ClientSession):
|
||||
return await session.list_resources()
|
||||
|
|
@ -754,9 +700,7 @@ class MCPClient:
|
|||
|
||||
async def list_resource_templates(self) -> list[ResourceTemplate]:
|
||||
"""List available resource templates from the server."""
|
||||
verbose_logger.debug(
|
||||
f"MCP client listing resource templates from {self.server_url or 'stdio'}"
|
||||
)
|
||||
verbose_logger.debug(f"MCP client listing resource templates from {self.server_url or 'stdio'}")
|
||||
|
||||
async def _list_resource_templates_operation(session: ClientSession):
|
||||
return await session.list_resource_templates()
|
||||
|
|
@ -764,9 +708,7 @@ class MCPClient:
|
|||
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
|
||||
]
|
||||
resource_template_names = [resourceTemplate.name for resourceTemplate in result.resourceTemplates]
|
||||
verbose_logger.info(
|
||||
f"MCP client listed {resource_template_count} resource templates from {self.server_url or 'stdio'}: {resource_template_names}"
|
||||
)
|
||||
|
|
@ -802,9 +744,7 @@ class MCPClient:
|
|||
|
||||
try:
|
||||
read_resource_result = await self.run_with_session(_read_resource_operation)
|
||||
verbose_logger.info(
|
||||
f"MCP client read_resource '{url}' completed successfully"
|
||||
)
|
||||
verbose_logger.info(f"MCP client read_resource '{url}' completed successfully")
|
||||
return read_resource_result
|
||||
except asyncio.CancelledError:
|
||||
verbose_logger.warning("MCP client read_resource was cancelled")
|
||||
|
|
|
|||
|
|
@ -90,9 +90,7 @@ async def load_mcp_tools(
|
|||
"""
|
||||
tools = await session.list_tools()
|
||||
if format == "openai":
|
||||
return [
|
||||
transform_mcp_tool_to_openai_tool(mcp_tool=tool) for tool in tools.tools
|
||||
]
|
||||
return [transform_mcp_tool_to_openai_tool(mcp_tool=tool) for tool in tools.tools]
|
||||
return tools.tools
|
||||
|
||||
|
||||
|
|
@ -148,10 +146,8 @@ 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(
|
||||
openai_tool=openai_tool,
|
||||
)
|
||||
mcp_tool_call_request_params = transform_openai_tool_call_request_to_mcp_tool_call_request(
|
||||
openai_tool=openai_tool,
|
||||
)
|
||||
return await call_mcp_tool(
|
||||
session=session,
|
||||
|
|
|
|||
|
|
@ -26,9 +26,7 @@ FileCreateProvider = Literal[
|
|||
"manus",
|
||||
"anthropic",
|
||||
]
|
||||
FileRetrieveProvider = Literal[
|
||||
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "manus", "anthropic"
|
||||
]
|
||||
FileRetrieveProvider = Literal["openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "manus", "anthropic"]
|
||||
FileDeleteProvider = Literal["openai", "azure", "gemini", "manus", "anthropic"]
|
||||
FileListProvider = Literal["openai", "azure", "manus", "anthropic"]
|
||||
import litellm
|
||||
|
|
@ -91,9 +89,7 @@ def _add_trusted_model_credentials_to_litellm_params(
|
|||
) -> None:
|
||||
trusted_model_credentials = kwargs.get("_litellm_internal_model_credentials")
|
||||
if isinstance(trusted_model_credentials, type(MappingProxyType({}))):
|
||||
litellm_params_dict["_litellm_internal_model_credentials"] = (
|
||||
trusted_model_credentials
|
||||
)
|
||||
litellm_params_dict["_litellm_internal_model_credentials"] = trusted_model_credentials
|
||||
|
||||
|
||||
@client
|
||||
|
|
@ -162,9 +158,7 @@ def create_file(
|
|||
_is_async = kwargs.pop("acreate_file", False) is True
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict = dict(**kwargs)
|
||||
logging_obj = cast(
|
||||
Optional[LiteLLMLoggingObj], kwargs.get("litellm_logging_obj")
|
||||
)
|
||||
logging_obj = cast(Optional[LiteLLMLoggingObj], kwargs.get("litellm_logging_obj"))
|
||||
if logging_obj is None:
|
||||
raise ValueError("logging_obj is required")
|
||||
client = kwargs.get("client")
|
||||
|
|
@ -215,12 +209,7 @@ def create_file(
|
|||
api_key=optional_params.api_key,
|
||||
logging_obj=logging_obj,
|
||||
_is_async=_is_async,
|
||||
client=(
|
||||
client
|
||||
if client is not None
|
||||
and isinstance(client, (HTTPHandler, AsyncHTTPHandler))
|
||||
else None
|
||||
),
|
||||
client=(client if client is not None and isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None),
|
||||
timeout=timeout,
|
||||
)
|
||||
elif custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
|
|
@ -264,9 +253,7 @@ def create_file(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_file", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_file", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -405,9 +392,7 @@ def file_retrieve(
|
|||
stream=False,
|
||||
call_type="afile_retrieve" if _is_async else "file_retrieve",
|
||||
start_time=time.time(),
|
||||
litellm_call_id=kwargs.get(
|
||||
"litellm_call_id", str(uuid_module.uuid4())
|
||||
),
|
||||
litellm_call_id=kwargs.get("litellm_call_id", str(uuid_module.uuid4())),
|
||||
function_id=str(kwargs.get("id") or ""),
|
||||
)
|
||||
|
||||
|
|
@ -420,10 +405,7 @@ def file_retrieve(
|
|||
logging_obj=logging_obj,
|
||||
_is_async=_is_async,
|
||||
client=(
|
||||
client
|
||||
if client is not None
|
||||
and isinstance(client, (HTTPHandler, AsyncHTTPHandler))
|
||||
else None
|
||||
client if client is not None and isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None
|
||||
),
|
||||
timeout=timeout,
|
||||
)
|
||||
|
|
@ -510,9 +492,7 @@ def file_delete(
|
|||
try:
|
||||
try:
|
||||
if model is not None:
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(
|
||||
model, custom_llm_provider
|
||||
)
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model, custom_llm_provider)
|
||||
except Exception:
|
||||
pass
|
||||
optional_params = GenericLiteLLMParams(**kwargs)
|
||||
|
|
@ -592,9 +572,7 @@ def file_delete(
|
|||
stream=False,
|
||||
call_type="afile_delete" if _is_async else "file_delete",
|
||||
start_time=time.time(),
|
||||
litellm_call_id=kwargs.get(
|
||||
"litellm_call_id", str(uuid_module.uuid4())
|
||||
),
|
||||
litellm_call_id=kwargs.get("litellm_call_id", str(uuid_module.uuid4())),
|
||||
function_id=str(kwargs.get("id") or ""),
|
||||
)
|
||||
|
||||
|
|
@ -606,10 +584,7 @@ def file_delete(
|
|||
logging_obj=logging_obj,
|
||||
_is_async=_is_async,
|
||||
client=(
|
||||
client
|
||||
if client is not None
|
||||
and isinstance(client, (HTTPHandler, AsyncHTTPHandler))
|
||||
else None
|
||||
client if client is not None and isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None
|
||||
),
|
||||
timeout=timeout,
|
||||
)
|
||||
|
|
@ -731,9 +706,7 @@ def file_list(
|
|||
stream=False,
|
||||
call_type="afile_list" if _is_async else "file_list",
|
||||
start_time=time.time(),
|
||||
litellm_call_id=kwargs.get(
|
||||
"litellm_call_id", str(uuid_module.uuid4())
|
||||
),
|
||||
litellm_call_id=kwargs.get("litellm_call_id", str(uuid_module.uuid4())),
|
||||
function_id=str(kwargs.get("id", "")),
|
||||
)
|
||||
|
||||
|
|
@ -745,12 +718,7 @@ def file_list(
|
|||
headers=extra_headers or {},
|
||||
logging_obj=logging_obj,
|
||||
_is_async=_is_async,
|
||||
client=(
|
||||
client
|
||||
if client is not None
|
||||
and isinstance(client, (HTTPHandler, AsyncHTTPHandler))
|
||||
else None
|
||||
),
|
||||
client=(client if client is not None and isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None),
|
||||
timeout=timeout,
|
||||
)
|
||||
return response
|
||||
|
|
@ -794,9 +762,7 @@ def file_list(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="file_list", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="file_list", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -886,9 +852,7 @@ def file_content(
|
|||
|
||||
try:
|
||||
if model is not None:
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(
|
||||
model, custom_llm_provider
|
||||
)
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model, custom_llm_provider)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
|
@ -922,9 +886,7 @@ def file_content(
|
|||
chunk_size=chunk_size,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
logging_obj=cast(
|
||||
Optional[LiteLLMLoggingObj], kwargs.get("litellm_logging_obj")
|
||||
),
|
||||
logging_obj=cast(Optional[LiteLLMLoggingObj], kwargs.get("litellm_logging_obj")),
|
||||
_is_async=_is_async,
|
||||
client=client,
|
||||
)
|
||||
|
|
@ -946,9 +908,7 @@ def file_content(
|
|||
stream=False,
|
||||
call_type="afile_content" if _is_async else "file_content",
|
||||
start_time=time.time(),
|
||||
litellm_call_id=kwargs.get(
|
||||
"litellm_call_id", str(uuid_module.uuid4())
|
||||
),
|
||||
litellm_call_id=kwargs.get("litellm_call_id", str(uuid_module.uuid4())),
|
||||
function_id=str(kwargs.get("id") or ""),
|
||||
)
|
||||
|
||||
|
|
@ -959,12 +919,7 @@ def file_content(
|
|||
headers=extra_headers or {},
|
||||
logging_obj=logging_obj,
|
||||
_is_async=_is_async,
|
||||
client=(
|
||||
client
|
||||
if client is not None
|
||||
and isinstance(client, (HTTPHandler, AsyncHTTPHandler))
|
||||
else None
|
||||
),
|
||||
client=(client if client is not None and isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None),
|
||||
timeout=timeout,
|
||||
)
|
||||
return response
|
||||
|
|
@ -1004,18 +959,12 @@ def file_content(
|
|||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or ""
|
||||
vertex_ai_project = (
|
||||
optional_params.vertex_project
|
||||
or litellm.vertex_project
|
||||
or get_secret_str("VERTEXAI_PROJECT")
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
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"
|
||||
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")
|
||||
|
||||
response = vertex_ai_files_instance.file_content(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -1047,9 +996,7 @@ def file_content(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -1095,9 +1042,9 @@ def file_content_streaming(
|
|||
headers=response.headers,
|
||||
)
|
||||
|
||||
response: Union[
|
||||
FileContentStreamingResult, Coroutine[Any, Any, FileContentStreamingResult]
|
||||
] = FileContentStreamingResult(stream_iterator=iter(()), headers={})
|
||||
response: Union[FileContentStreamingResult, Coroutine[Any, Any, FileContentStreamingResult]] = (
|
||||
FileContentStreamingResult(stream_iterator=iter(()), headers={})
|
||||
)
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
openai_creds = get_openai_credentials(
|
||||
api_base=optional_params.api_base,
|
||||
|
|
@ -1130,9 +1077,7 @@ def file_content_streaming(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -94,9 +94,7 @@ class FileContentStreamingResponse:
|
|||
self._close_completed = True
|
||||
self._logging_completed = True
|
||||
stream_to_close = self.stream_iterator
|
||||
self.stream_iterator = cast(
|
||||
Union[Iterator[bytes], AsyncIterator[bytes]], iter(())
|
||||
)
|
||||
self.stream_iterator = cast(Union[Iterator[bytes], AsyncIterator[bytes]], iter(()))
|
||||
|
||||
# Shield cleanup from request cancellation so upstream HTTP connections
|
||||
# are released promptly on client disconnects.
|
||||
|
|
@ -115,9 +113,7 @@ class FileContentStreamingResponse:
|
|||
self._close_completed = True
|
||||
self._logging_completed = True
|
||||
stream_to_close = self.stream_iterator
|
||||
self.stream_iterator = cast(
|
||||
Union[Iterator[bytes], AsyncIterator[bytes]], iter(())
|
||||
)
|
||||
self.stream_iterator = cast(Union[Iterator[bytes], AsyncIterator[bytes]], iter(()))
|
||||
|
||||
if hasattr(stream_to_close, "close"):
|
||||
cast(Iterator[bytes], stream_to_close).close() # type: ignore[attr-defined]
|
||||
|
|
@ -134,9 +130,7 @@ class FileContentStreamingResponse:
|
|||
def _sync_hidden_params(self) -> None:
|
||||
litellm_params: dict[str, Any] = {}
|
||||
if self.logging_obj is not None:
|
||||
litellm_params = (
|
||||
self.logging_obj.model_call_details.get("litellm_params", {}) or {}
|
||||
)
|
||||
litellm_params = self.logging_obj.model_call_details.get("litellm_params", {}) or {}
|
||||
|
||||
if "api_base" not in self._hidden_params and litellm_params.get("api_base"):
|
||||
self._hidden_params["api_base"] = litellm_params["api_base"]
|
||||
|
|
@ -232,12 +226,8 @@ class FileContentStreamingResponse:
|
|||
self._logging_completed = True
|
||||
end_time = datetime.datetime.now()
|
||||
traceback_str = 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
|
||||
)
|
||||
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)
|
||||
|
||||
def _log_failure_sync(self, error: Exception) -> None:
|
||||
if self._logging_completed or self.logging_obj is None:
|
||||
|
|
@ -245,6 +235,4 @@ class FileContentStreamingResponse:
|
|||
|
||||
self._logging_completed = True
|
||||
end_time = datetime.datetime.now()
|
||||
self.logging_obj.failure_handler(
|
||||
error, traceback.format_exc(), self._start_time, end_time
|
||||
)
|
||||
self.logging_obj.failure_handler(error, traceback.format_exc(), self._start_time, end_time)
|
||||
|
|
|
|||
|
|
@ -1,8 +1,6 @@
|
|||
from typing import AsyncIterator, Dict, Iterator, Literal, NamedTuple, Union
|
||||
|
||||
FileContentProvider = Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic", "manus"
|
||||
]
|
||||
FileContentProvider = Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic", "manus"]
|
||||
|
||||
|
||||
class FileContentStreamingResult(NamedTuple):
|
||||
|
|
|
|||
|
|
@ -26,24 +26,18 @@ class FilesAPIUtils:
|
|||
"""
|
||||
|
||||
@staticmethod
|
||||
def is_batch_jsonl_file(
|
||||
create_file_data: CreateFileRequest, extracted_file_data: ExtractedFileData
|
||||
) -> bool:
|
||||
def is_batch_jsonl_file(create_file_data: CreateFileRequest, extracted_file_data: ExtractedFileData) -> bool:
|
||||
"""
|
||||
Check if the file is a batch jsonl file
|
||||
"""
|
||||
return (
|
||||
create_file_data.get("purpose") == "batch"
|
||||
and FilesAPIUtils.valid_content_type(
|
||||
extracted_file_data.get("content_type")
|
||||
)
|
||||
and FilesAPIUtils.valid_content_type(extracted_file_data.get("content_type"))
|
||||
and extracted_file_data.get("content") is not None
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def is_batch_jsonl_request(
|
||||
create_file_data: CreateFileRequest, content_type: Optional[str]
|
||||
) -> bool:
|
||||
def is_batch_jsonl_request(create_file_data: CreateFileRequest, content_type: Optional[str]) -> bool:
|
||||
"""
|
||||
Batch-jsonl check from metadata only, so the body can stay a streamable
|
||||
Path/handle instead of being read into memory.
|
||||
|
|
|
|||
|
|
@ -90,9 +90,7 @@ async def acreate_fine_tuning_job(
|
|||
Async: Creates and executes a batch from an uploaded file of request
|
||||
|
||||
"""
|
||||
verbose_logger.debug(
|
||||
"inside acreate_fine_tuning_job model=%s and kwargs=%s", model, kwargs
|
||||
)
|
||||
verbose_logger.debug("inside acreate_fine_tuning_job model=%s and kwargs=%s", model, kwargs)
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["acreate_fine_tuning_job"] = True
|
||||
|
|
@ -126,9 +124,7 @@ async def acreate_fine_tuning_job(
|
|||
raise e
|
||||
|
||||
|
||||
def _build_fine_tuning_job_data(
|
||||
model, training_file, hyperparameters, suffix, validation_file, integrations, seed
|
||||
):
|
||||
def _build_fine_tuning_job_data(model, training_file, hyperparameters, suffix, validation_file, integrations, seed):
|
||||
return FineTuningJobCreate(
|
||||
model=model,
|
||||
training_file=training_file,
|
||||
|
|
@ -245,17 +241,9 @@ def create_fine_tuning_job(
|
|||
)
|
||||
# Azure OpenAI
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base
|
||||
or litellm.api_base
|
||||
or get_secret_str("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_version = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret_str("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") # type: ignore
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -272,9 +260,7 @@ def create_fine_tuning_job(
|
|||
get_secret_str("AZURE_AD_TOKEN") # type: ignore
|
||||
|
||||
# Prepare Azure-specific parameters for extra_body
|
||||
extra_body = _prepare_azure_extra_body(
|
||||
extra_body, kwargs, azure_specific_hyperparams
|
||||
)
|
||||
extra_body = _prepare_azure_extra_body(extra_body, kwargs, azure_specific_hyperparams)
|
||||
|
||||
create_fine_tuning_job_data_dict = _build_fine_tuning_job_data(
|
||||
model,
|
||||
|
|
@ -303,18 +289,12 @@ def create_fine_tuning_job(
|
|||
elif custom_llm_provider == "vertex_ai":
|
||||
api_base = optional_params.api_base or ""
|
||||
vertex_ai_project = (
|
||||
optional_params.vertex_project
|
||||
or litellm.vertex_project
|
||||
or get_secret_str("VERTEXAI_PROJECT")
|
||||
optional_params.vertex_project or litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_ai_location = (
|
||||
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"
|
||||
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")
|
||||
response = vertex_fine_tuning_apis_instance.create_fine_tuning_job(
|
||||
_is_async=_is_async,
|
||||
create_fine_tuning_job_data=_build_fine_tuning_job_data(
|
||||
|
|
@ -344,9 +324,7 @@ def create_fine_tuning_job(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -464,17 +442,9 @@ def cancel_fine_tuning_job(
|
|||
)
|
||||
# Azure OpenAI
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base
|
||||
or litellm.api_base
|
||||
or get_secret("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_version = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret_str("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") # type: ignore
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -510,9 +480,7 @@ def cancel_fine_tuning_job(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -633,17 +601,9 @@ def list_fine_tuning_jobs(
|
|||
)
|
||||
# Azure OpenAI
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base
|
||||
or litellm.api_base
|
||||
or get_secret_str("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_version = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret_str("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") # type: ignore
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -680,9 +640,7 @@ def list_fine_tuning_jobs(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(
|
||||
method="create_thread", url="https://github.com/BerriAI/litellm"
|
||||
), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -769,17 +727,9 @@ def retrieve_fine_tuning_job(
|
|||
or "https://api.openai.com/v1"
|
||||
)
|
||||
organization = (
|
||||
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")
|
||||
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")
|
||||
|
||||
response = openai_fine_tuning_apis_instance.retrieve_fine_tuning_job(
|
||||
api_base=api_base,
|
||||
|
|
@ -794,17 +744,9 @@ def retrieve_fine_tuning_job(
|
|||
)
|
||||
# Azure OpenAI
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = (
|
||||
optional_params.api_base
|
||||
or litellm.api_base
|
||||
or get_secret_str("AZURE_API_BASE")
|
||||
) # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") # type: ignore
|
||||
|
||||
api_version = (
|
||||
optional_params.api_version
|
||||
or litellm.api_version
|
||||
or get_secret_str("AZURE_API_VERSION")
|
||||
) # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") # type: ignore
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
|
|||
|
|
@ -25,14 +25,12 @@ 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(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
litellm_params=litellm_params,
|
||||
**(extra_kwargs or {}),
|
||||
)
|
||||
completion_request = GOOGLE_GENAI_ADAPTER.translate_generate_content_to_completion(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
litellm_params=litellm_params,
|
||||
**(extra_kwargs or {}),
|
||||
)
|
||||
|
||||
completion_kwargs: Dict[str, Any] = dict(completion_request)
|
||||
|
|
@ -62,15 +60,13 @@ class GenerateContentToCompletionHandler:
|
|||
) -> Union[Dict[str, Any], AsyncIterator[bytes]]:
|
||||
"""Handle generate_content call asynchronously using completion adapter"""
|
||||
|
||||
completion_kwargs = (
|
||||
GenerateContentToCompletionHandler._prepare_completion_kwargs(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
completion_kwargs = GenerateContentToCompletionHandler._prepare_completion_kwargs(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
try:
|
||||
|
|
@ -81,10 +77,8 @@ class GenerateContentToCompletionHandler:
|
|||
# This can happen in error cases or when stream is not properly supported
|
||||
if not hasattr(completion_response, "__aiter__"):
|
||||
# If it's not a stream, treat it as a regular response
|
||||
generate_content_response = (
|
||||
GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
|
||||
cast(ModelResponse, completion_response)
|
||||
)
|
||||
generate_content_response = GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
|
||||
cast(ModelResponse, completion_response)
|
||||
)
|
||||
return generate_content_response
|
||||
else:
|
||||
|
|
@ -97,17 +91,13 @@ class GenerateContentToCompletionHandler:
|
|||
raise ValueError("Failed to transform streaming response")
|
||||
else:
|
||||
# Transform completion response back to generate_content format
|
||||
generate_content_response = (
|
||||
GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
|
||||
cast(ModelResponse, completion_response)
|
||||
)
|
||||
generate_content_response = GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
|
||||
cast(ModelResponse, completion_response)
|
||||
)
|
||||
return generate_content_response
|
||||
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"Error calling litellm.acompletion for generate_content: {str(e)}"
|
||||
)
|
||||
raise ValueError(f"Error calling litellm.acompletion for generate_content: {str(e)}")
|
||||
|
||||
@staticmethod
|
||||
def generate_content_handler(
|
||||
|
|
@ -135,15 +125,13 @@ class GenerateContentToCompletionHandler:
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
completion_kwargs = (
|
||||
GenerateContentToCompletionHandler._prepare_completion_kwargs(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
completion_kwargs = GenerateContentToCompletionHandler._prepare_completion_kwargs(
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
try:
|
||||
|
|
@ -154,10 +142,8 @@ class GenerateContentToCompletionHandler:
|
|||
# This can happen in error cases or when stream is not properly supported
|
||||
if not hasattr(completion_response, "__iter__"):
|
||||
# If it's not a stream, treat it as a regular response
|
||||
generate_content_response = (
|
||||
GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
|
||||
cast(ModelResponse, completion_response)
|
||||
)
|
||||
generate_content_response = GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
|
||||
cast(ModelResponse, completion_response)
|
||||
)
|
||||
return generate_content_response
|
||||
else:
|
||||
|
|
@ -170,14 +156,10 @@ class GenerateContentToCompletionHandler:
|
|||
raise ValueError("Failed to transform streaming response")
|
||||
else:
|
||||
# Transform completion response back to generate_content format
|
||||
generate_content_response = (
|
||||
GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
|
||||
cast(ModelResponse, completion_response)
|
||||
)
|
||||
generate_content_response = GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
|
||||
cast(ModelResponse, completion_response)
|
||||
)
|
||||
return generate_content_response
|
||||
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"Error calling litellm.completion for generate_content: {str(e)}"
|
||||
)
|
||||
raise ValueError(f"Error calling litellm.completion for generate_content: {str(e)}")
|
||||
|
|
|
|||
|
|
@ -49,17 +49,13 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
if self._returned_response:
|
||||
raise StopIteration
|
||||
self._returned_response = True
|
||||
return GoogleGenAIAdapter().translate_completion_to_generate_content(
|
||||
self.completion_stream
|
||||
)
|
||||
return GoogleGenAIAdapter().translate_completion_to_generate_content(self.completion_stream)
|
||||
|
||||
for chunk in self.completion_stream:
|
||||
if chunk == "None" or chunk is None:
|
||||
continue
|
||||
|
||||
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(
|
||||
chunk, self
|
||||
)
|
||||
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(chunk, self)
|
||||
if transformed_chunk:
|
||||
return transformed_chunk
|
||||
|
||||
|
|
@ -75,17 +71,13 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
if self._returned_response:
|
||||
raise StopAsyncIteration
|
||||
self._returned_response = True
|
||||
return GoogleGenAIAdapter().translate_completion_to_generate_content(
|
||||
self.completion_stream
|
||||
)
|
||||
return GoogleGenAIAdapter().translate_completion_to_generate_content(self.completion_stream)
|
||||
|
||||
async for chunk in self.completion_stream:
|
||||
if chunk == "None" or chunk is None:
|
||||
continue
|
||||
|
||||
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(
|
||||
chunk, self
|
||||
)
|
||||
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(chunk, self)
|
||||
if transformed_chunk:
|
||||
return transformed_chunk
|
||||
|
||||
|
|
@ -100,13 +92,10 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
try:
|
||||
# For tool calls with no arguments, accumulated_args will be "", which is not valid JSON.
|
||||
# We default to an empty JSON object in this case.
|
||||
parsed_args = json.loads(
|
||||
tool_call_data["arguments"] or "{}"
|
||||
)
|
||||
parsed_args = json.loads(tool_call_data["arguments"] or "{}")
|
||||
function_call_part = {
|
||||
"functionCall": {
|
||||
"name": tool_call_data["name"]
|
||||
or "undefined_tool_name",
|
||||
"name": tool_call_data["name"] or "undefined_tool_name",
|
||||
"args": parsed_args,
|
||||
}
|
||||
}
|
||||
|
|
@ -163,9 +152,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
yield payload.encode()
|
||||
elif isinstance(chunk, ModelResponseStream):
|
||||
# Transform OpenAI streaming chunk to Google GenAI format
|
||||
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(
|
||||
chunk, self
|
||||
)
|
||||
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(chunk, self)
|
||||
|
||||
if isinstance(transformed_chunk, dict): # Only return non-empty chunks
|
||||
payload = f"data: {json.dumps(transformed_chunk)}\n\n"
|
||||
|
|
@ -209,9 +196,7 @@ class GoogleGenAIAdapter:
|
|||
"""
|
||||
|
||||
# Extract top-level fields from kwargs
|
||||
system_instruction = kwargs.get("systemInstruction") or kwargs.get(
|
||||
"system_instruction"
|
||||
)
|
||||
system_instruction = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
|
||||
tools = kwargs.get("tools")
|
||||
tool_config = kwargs.get("toolConfig") or kwargs.get("tool_config")
|
||||
|
||||
|
|
@ -222,9 +207,7 @@ class GoogleGenAIAdapter:
|
|||
contents_list = contents
|
||||
|
||||
# Transform contents to OpenAI messages format
|
||||
messages = self._transform_contents_to_messages(
|
||||
contents_list, system_instruction=system_instruction
|
||||
)
|
||||
messages = self._transform_contents_to_messages(contents_list, system_instruction=system_instruction)
|
||||
|
||||
# Create base request as dict (which is compatible with ChatCompletionRequest)
|
||||
completion_request: ChatCompletionRequest = {
|
||||
|
|
@ -271,9 +254,7 @@ class GoogleGenAIAdapter:
|
|||
|
||||
# Handle tool_config (tool choice)
|
||||
if tool_config:
|
||||
tool_choice = self._transform_google_genai_tool_config_to_openai(
|
||||
tool_config
|
||||
)
|
||||
tool_choice = self._transform_google_genai_tool_config_to_openai(tool_config)
|
||||
if tool_choice:
|
||||
completion_request["tool_choice"] = tool_choice
|
||||
|
||||
|
|
@ -316,9 +297,7 @@ class GoogleGenAIAdapter:
|
|||
completion_stream: Any,
|
||||
) -> Union[AsyncIterator[bytes], None]:
|
||||
"""Transform streaming completion output to Google GenAI format"""
|
||||
google_genai_wrapper = GoogleGenAIStreamWrapper(
|
||||
completion_stream=completion_stream
|
||||
)
|
||||
google_genai_wrapper = GoogleGenAIStreamWrapper(completion_stream=completion_stream)
|
||||
# Return the SSE-wrapped version for proper event formatting
|
||||
return google_genai_wrapper.async_google_genai_sse_wrapper()
|
||||
|
||||
|
|
@ -374,11 +353,7 @@ class GoogleGenAIAdapter:
|
|||
if system_instruction:
|
||||
system_parts = system_instruction.get("parts", [])
|
||||
if system_parts and "text" in system_parts[0]:
|
||||
messages.append(
|
||||
ChatCompletionSystemMessage(
|
||||
role="system", content=system_parts[0]["text"]
|
||||
)
|
||||
)
|
||||
messages.append(ChatCompletionSystemMessage(role="system", content=system_parts[0]["text"]))
|
||||
|
||||
for content in contents:
|
||||
role = content.get("role", "user")
|
||||
|
|
@ -386,9 +361,7 @@ class GoogleGenAIAdapter:
|
|||
|
||||
if role == "user":
|
||||
# Handle user messages with potential function responses
|
||||
content_parts: List[
|
||||
Union[ChatCompletionTextObject, ChatCompletionImageObject]
|
||||
] = []
|
||||
content_parts: List[Union[ChatCompletionTextObject, ChatCompletionImageObject]] = []
|
||||
tool_messages: List[ChatCompletionToolMessage] = []
|
||||
|
||||
for part in parts:
|
||||
|
|
@ -410,9 +383,7 @@ class GoogleGenAIAdapter:
|
|||
ChatCompletionImageObject,
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": f"data:{mime_type};base64,{data}"
|
||||
},
|
||||
"image_url": {"url": f"data:{mime_type};base64,{data}"},
|
||||
},
|
||||
)
|
||||
)
|
||||
|
|
@ -426,11 +397,7 @@ class GoogleGenAIAdapter:
|
|||
)
|
||||
tool_messages.append(tool_message)
|
||||
elif isinstance(part, str):
|
||||
content_parts.append(
|
||||
cast(
|
||||
ChatCompletionTextObject, {"type": "text", "text": part}
|
||||
)
|
||||
)
|
||||
content_parts.append(cast(ChatCompletionTextObject, {"type": "text", "text": part}))
|
||||
|
||||
# Add user message if there's content
|
||||
if content_parts:
|
||||
|
|
@ -441,18 +408,10 @@ class GoogleGenAIAdapter:
|
|||
and content_parts[0].get("type") == "text"
|
||||
):
|
||||
text_part = cast(ChatCompletionTextObject, content_parts[0])
|
||||
messages.append(
|
||||
ChatCompletionUserMessage(
|
||||
role="user", content=text_part["text"]
|
||||
)
|
||||
)
|
||||
messages.append(ChatCompletionUserMessage(role="user", content=text_part["text"]))
|
||||
else:
|
||||
# Use multimodal format (array of content parts)
|
||||
messages.append(
|
||||
ChatCompletionUserMessage(
|
||||
role="user", content=content_parts
|
||||
)
|
||||
)
|
||||
messages.append(ChatCompletionUserMessage(role="user", content=content_parts))
|
||||
|
||||
# Add tool messages
|
||||
messages.extend(tool_messages)
|
||||
|
|
@ -520,15 +479,13 @@ class GoogleGenAIAdapter:
|
|||
# Handle different choice types (Choices vs StreamingChoices)
|
||||
if isinstance(choice, Choices):
|
||||
if not choice.message:
|
||||
raise ValueError(
|
||||
"Invalid completion response: no message found in choice"
|
||||
)
|
||||
raise ValueError("Invalid completion response: no message found in choice")
|
||||
parts = self._transform_openai_message_to_google_genai_parts(choice.message)
|
||||
else:
|
||||
# Fallback for generic choice objects
|
||||
message_content = getattr(choice, "message", {}).get(
|
||||
message_content = getattr(choice, "message", {}).get("content", "") or getattr(choice, "delta", {}).get(
|
||||
"content", ""
|
||||
) or getattr(choice, "delta", {}).get("content", "")
|
||||
)
|
||||
parts = [{"text": message_content}] if message_content else []
|
||||
|
||||
# Create Google GenAI format response
|
||||
|
|
@ -536,9 +493,7 @@ class GoogleGenAIAdapter:
|
|||
"candidates": [
|
||||
{
|
||||
"content": {"parts": parts, "role": "model"},
|
||||
"finishReason": self._map_finish_reason(
|
||||
getattr(choice, "finish_reason", None)
|
||||
),
|
||||
"finishReason": self._map_finish_reason(getattr(choice, "finish_reason", None)),
|
||||
"index": 0,
|
||||
"safetyRatings": [],
|
||||
}
|
||||
|
|
@ -589,9 +544,7 @@ class GoogleGenAIAdapter:
|
|||
# Handle streaming choice
|
||||
if isinstance(choice, StreamingChoices):
|
||||
if choice.delta:
|
||||
parts = self._transform_openai_delta_to_google_genai_parts_with_accumulation(
|
||||
choice.delta, wrapper
|
||||
)
|
||||
parts = self._transform_openai_delta_to_google_genai_parts_with_accumulation(choice.delta, wrapper)
|
||||
else:
|
||||
parts = []
|
||||
finish_reason = getattr(choice, "finish_reason", None)
|
||||
|
|
@ -610,11 +563,7 @@ class GoogleGenAIAdapter:
|
|||
"candidates": [
|
||||
{
|
||||
"content": {"parts": parts, "role": "model"},
|
||||
"finishReason": (
|
||||
self._map_finish_reason(finish_reason)
|
||||
if finish_reason
|
||||
else None
|
||||
),
|
||||
"finishReason": (self._map_finish_reason(finish_reason) if finish_reason else None),
|
||||
"index": 0,
|
||||
"safetyRatings": [],
|
||||
}
|
||||
|
|
@ -660,11 +609,7 @@ class GoogleGenAIAdapter:
|
|||
for tool_call in message.tool_calls:
|
||||
if hasattr(tool_call, "function") and tool_call.function:
|
||||
try:
|
||||
args = (
|
||||
json.loads(tool_call.function.arguments)
|
||||
if tool_call.function.arguments
|
||||
else {}
|
||||
)
|
||||
args = json.loads(tool_call.function.arguments) if tool_call.function.arguments else {}
|
||||
except json.JSONDecodeError:
|
||||
args = {}
|
||||
|
||||
|
|
@ -717,18 +662,14 @@ class GoogleGenAIAdapter:
|
|||
|
||||
# Optimization: Skip chunks that have no new data
|
||||
if not function_name and not args_chunk:
|
||||
verbose_logger.debug(
|
||||
f"Skipping empty tool call chunk for index: {tool_call_index}"
|
||||
)
|
||||
verbose_logger.debug(f"Skipping empty tool call chunk for index: {tool_call_index}")
|
||||
continue
|
||||
|
||||
if function_name:
|
||||
wrapper.accumulated_tool_calls[tool_call_index]["name"] = function_name
|
||||
|
||||
if args_chunk:
|
||||
wrapper.accumulated_tool_calls[tool_call_index]["arguments"] += (
|
||||
args_chunk
|
||||
)
|
||||
wrapper.accumulated_tool_calls[tool_call_index]["arguments"] += args_chunk
|
||||
|
||||
# Attempt to parse and emit a complete tool call
|
||||
accumulated_data = wrapper.accumulated_tool_calls[tool_call_index]
|
||||
|
|
@ -744,9 +685,7 @@ class GoogleGenAIAdapter:
|
|||
# The part will be created by a later chunk that brings the name.
|
||||
if accumulated_name:
|
||||
# If successful, create the part and clean up
|
||||
function_call_part = {
|
||||
"functionCall": {"name": accumulated_name, "args": parsed_args}
|
||||
}
|
||||
function_call_part = {"functionCall": {"name": accumulated_name, "args": parsed_args}}
|
||||
parts.append(function_call_part)
|
||||
|
||||
# Remove the completed tool call from the accumulator
|
||||
|
|
|
|||
|
|
@ -103,9 +103,7 @@ class GenerateContentHelper:
|
|||
Returns:
|
||||
GenerateContentSetupResult containing all setup information
|
||||
"""
|
||||
litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get(
|
||||
"litellm_logging_obj"
|
||||
)
|
||||
litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
|
||||
|
||||
# get llm provider logic
|
||||
|
|
@ -135,11 +133,11 @@ class GenerateContentHelper:
|
|||
litellm_params.custom_llm_provider = custom_llm_provider
|
||||
|
||||
# get provider config
|
||||
generate_content_provider_config: Optional[
|
||||
BaseGoogleGenAIGenerateContentConfig
|
||||
] = ProviderConfigManager.get_provider_google_genai_generate_content_config(
|
||||
model=model,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
generate_content_provider_config: Optional[BaseGoogleGenAIGenerateContentConfig] = (
|
||||
ProviderConfigManager.get_provider_google_genai_generate_content_config(
|
||||
model=model,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
)
|
||||
|
||||
if generate_content_provider_config is None:
|
||||
|
|
@ -163,30 +161,24 @@ 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=config or {},
|
||||
model=model,
|
||||
)
|
||||
generate_content_config_dict = 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 = 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] = {
|
||||
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(
|
||||
model=model,
|
||||
contents=contents,
|
||||
tools=tools,
|
||||
generate_content_config_dict=generate_content_config_dict,
|
||||
system_instruction=system_instruction,
|
||||
)
|
||||
request_body = generate_content_provider_config.transform_generate_content_request(
|
||||
model=model,
|
||||
contents=contents,
|
||||
tools=tools,
|
||||
generate_content_config_dict=generate_content_config_dict,
|
||||
system_instruction=system_instruction,
|
||||
)
|
||||
|
||||
# Pre Call logging
|
||||
|
|
@ -328,12 +320,8 @@ def generate_content(
|
|||
config = kwargs.pop("generationConfig")
|
||||
# Check for mock response first
|
||||
litellm_params = 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
|
||||
)
|
||||
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(
|
||||
|
|
@ -346,9 +334,7 @@ def generate_content(
|
|||
)
|
||||
|
||||
# Extract systemInstruction from kwargs to pass to handler
|
||||
system_instruction = kwargs.get("systemInstruction") or kwargs.get(
|
||||
"system_instruction"
|
||||
)
|
||||
system_instruction = 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:
|
||||
|
|
@ -375,9 +361,7 @@ def generate_content(
|
|||
litellm_params=setup_result.litellm_params,
|
||||
logging_obj=setup_result.litellm_logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=_merge_native_request_fields(
|
||||
setup_result.native_request_fields, extra_body
|
||||
),
|
||||
extra_body=_merge_native_request_fields(setup_result.native_request_fields, extra_body),
|
||||
timeout=timeout or request_timeout,
|
||||
_is_async=_is_async,
|
||||
client=kwargs.get("client"),
|
||||
|
|
@ -439,9 +423,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 = 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:
|
||||
|
|
@ -449,17 +431,15 @@ async def agenerate_content_stream(
|
|||
kwargs.pop("stream", None)
|
||||
|
||||
# Use the adapter to convert to completion format
|
||||
return (
|
||||
await GenerateContentToCompletionHandler.async_generate_content_handler(
|
||||
model=model,
|
||||
contents=contents, # type: ignore
|
||||
config=setup_result.generate_content_config_dict,
|
||||
litellm_params=setup_result.litellm_params,
|
||||
tools=tools,
|
||||
stream=True,
|
||||
extra_headers=extra_headers,
|
||||
**kwargs,
|
||||
)
|
||||
return await GenerateContentToCompletionHandler.async_generate_content_handler(
|
||||
model=model,
|
||||
contents=contents, # type: ignore
|
||||
config=setup_result.generate_content_config_dict,
|
||||
litellm_params=setup_result.litellm_params,
|
||||
tools=tools,
|
||||
stream=True,
|
||||
extra_headers=extra_headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Call the handler with async enabled and streaming
|
||||
|
|
@ -474,9 +454,7 @@ async def agenerate_content_stream(
|
|||
litellm_params=setup_result.litellm_params,
|
||||
logging_obj=setup_result.litellm_logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=_merge_native_request_fields(
|
||||
setup_result.native_request_fields, extra_body
|
||||
),
|
||||
extra_body=_merge_native_request_fields(setup_result.native_request_fields, extra_body),
|
||||
timeout=timeout or request_timeout,
|
||||
_is_async=True,
|
||||
client=kwargs.get("client"),
|
||||
|
|
@ -533,9 +511,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 = 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:
|
||||
|
|
@ -565,9 +541,7 @@ def generate_content_stream(
|
|||
litellm_params=setup_result.litellm_params,
|
||||
logging_obj=setup_result.litellm_logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=_merge_native_request_fields(
|
||||
setup_result.native_request_fields, extra_body
|
||||
),
|
||||
extra_body=_merge_native_request_fields(setup_result.native_request_fields, extra_body),
|
||||
timeout=timeout or request_timeout,
|
||||
_is_async=_is_async,
|
||||
client=kwargs.get("client"),
|
||||
|
|
|
|||
|
|
@ -98,9 +98,7 @@ class BaseGoogleGenAIGenerateContentStreamingIterator:
|
|||
)
|
||||
|
||||
|
||||
class GoogleGenAIGenerateContentStreamingIterator(
|
||||
BaseGoogleGenAIGenerateContentStreamingIterator
|
||||
):
|
||||
class GoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateContentStreamingIterator):
|
||||
"""
|
||||
Streaming iterator specifically for Google GenAI generate content API.
|
||||
"""
|
||||
|
|
@ -148,14 +146,10 @@ class GoogleGenAIGenerateContentStreamingIterator(
|
|||
async def __anext__(self):
|
||||
# This should not be used for sync responses
|
||||
# If you need async iteration, use AsyncGoogleGenAIGenerateContentStreamingIterator
|
||||
raise NotImplementedError(
|
||||
"Use AsyncGoogleGenAIGenerateContentStreamingIterator for async iteration"
|
||||
)
|
||||
raise NotImplementedError("Use AsyncGoogleGenAIGenerateContentStreamingIterator for async iteration")
|
||||
|
||||
|
||||
class AsyncGoogleGenAIGenerateContentStreamingIterator(
|
||||
BaseGoogleGenAIGenerateContentStreamingIterator
|
||||
):
|
||||
class AsyncGoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateContentStreamingIterator):
|
||||
"""
|
||||
Async streaming iterator specifically for Google GenAI generate content API.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -111,9 +111,7 @@ async def aimage_generation(*args, **kwargs) -> ImageResponse:
|
|||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(
|
||||
model=model, api_base=kwargs.get("api_base", None)
|
||||
)
|
||||
_, 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)
|
||||
|
|
@ -127,9 +125,7 @@ async def aimage_generation(*args, **kwargs) -> ImageResponse:
|
|||
response = await init_response # type: ignore
|
||||
|
||||
if response is None:
|
||||
raise ValueError(
|
||||
"Unable to get Image Response. Please pass a valid llm_provider."
|
||||
)
|
||||
raise ValueError("Unable to get Image Response. Please pass a valid llm_provider.")
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
|
|
@ -272,15 +268,10 @@ def image_generation(
|
|||
} # model-specific params - pass them straight to the model/provider
|
||||
|
||||
image_generation_config: Optional[BaseImageGenerationConfig] = None
|
||||
if (
|
||||
custom_llm_provider is not None
|
||||
and custom_llm_provider in LlmProviders._member_map_.values()
|
||||
):
|
||||
image_generation_config = (
|
||||
ProviderConfigManager.get_provider_image_generation_config(
|
||||
model=base_model or model,
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
if custom_llm_provider is not None and custom_llm_provider in LlmProviders._member_map_.values():
|
||||
image_generation_config = ProviderConfigManager.get_provider_image_generation_config(
|
||||
model=base_model or model,
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
optional_params = get_optional_params_image_gen(
|
||||
|
|
@ -327,11 +318,7 @@ def image_generation(
|
|||
|
||||
api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
|
||||
api_version = (
|
||||
api_version
|
||||
or litellm.api_version
|
||||
or get_secret_str("AZURE_API_VERSION")
|
||||
)
|
||||
api_version = api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
api_key
|
||||
|
|
@ -341,9 +328,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 = 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:
|
||||
|
|
@ -355,10 +340,7 @@ def image_generation(
|
|||
tenant_id = litellm_params_dict.get("tenant_id")
|
||||
client_id = litellm_params_dict.get("client_id")
|
||||
client_secret = litellm_params_dict.get("client_secret")
|
||||
azure_scope = (
|
||||
litellm_params_dict.get("azure_scope")
|
||||
or "https://cognitiveservices.azure.com/.default"
|
||||
)
|
||||
azure_scope = litellm_params_dict.get("azure_scope") or "https://cognitiveservices.azure.com/.default"
|
||||
|
||||
# Create token provider if credentials are available
|
||||
if tenant_id and client_id and client_secret:
|
||||
|
|
@ -413,9 +395,7 @@ def image_generation(
|
|||
litellm.LlmProviders.DASHSCOPE,
|
||||
):
|
||||
if image_generation_config is None:
|
||||
raise ValueError(
|
||||
f"image generation config is not supported for {custom_llm_provider}"
|
||||
)
|
||||
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
|
||||
|
|
@ -524,9 +504,7 @@ def image_generation(
|
|||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
)
|
||||
elif (
|
||||
custom_llm_provider in litellm._custom_providers
|
||||
): # Assume custom LLM provider
|
||||
elif custom_llm_provider in litellm._custom_providers: # Assume custom LLM provider
|
||||
# Get the Custom Handler
|
||||
custom_handler: Optional[CustomLLM] = None
|
||||
for item in litellm.custom_provider_map:
|
||||
|
|
@ -534,9 +512,7 @@ def image_generation(
|
|||
custom_handler = item["custom_handler"]
|
||||
|
||||
if custom_handler is None:
|
||||
raise LiteLLMUnknownProvider(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
raise LiteLLMUnknownProvider(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
## ROUTE LLM CALL ##
|
||||
if aimg_generation is True:
|
||||
|
|
@ -612,15 +588,11 @@ async def aimage_variation(*args, **kwargs) -> ImageResponse:
|
|||
func_with_context = 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)
|
||||
)
|
||||
_, 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)
|
||||
if isinstance(init_response, dict) or isinstance(
|
||||
init_response, ImageResponse
|
||||
): ## CACHING SCENARIO
|
||||
if isinstance(init_response, dict) or isinstance(init_response, ImageResponse): ## CACHING SCENARIO
|
||||
if isinstance(init_response, dict):
|
||||
init_response = ImageResponse(**init_response)
|
||||
response = init_response
|
||||
|
|
@ -793,9 +765,7 @@ def image_edit(
|
|||
_is_async = 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 = 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] = {}
|
||||
|
|
@ -822,17 +792,13 @@ def image_edit(
|
|||
custom_handler = item["custom_handler"]
|
||||
|
||||
if custom_handler is None:
|
||||
raise LiteLLMUnknownProvider(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
raise LiteLLMUnknownProvider(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
model_response = ImageResponse()
|
||||
|
||||
if _is_async:
|
||||
async_custom_client: Optional[AsyncHTTPHandler] = None
|
||||
if kwargs.get("client") is not None and isinstance(
|
||||
kwargs.get("client"), AsyncHTTPHandler
|
||||
):
|
||||
if kwargs.get("client") is not None and isinstance(kwargs.get("client"), AsyncHTTPHandler):
|
||||
async_custom_client = kwargs.get("client")
|
||||
|
||||
return custom_handler.aimage_edit(
|
||||
|
|
@ -849,9 +815,7 @@ def image_edit(
|
|||
)
|
||||
else:
|
||||
custom_client: Optional[HTTPHandler] = None
|
||||
if kwargs.get("client") is not None and isinstance(
|
||||
kwargs.get("client"), HTTPHandler
|
||||
):
|
||||
if kwargs.get("client") is not None and isinstance(kwargs.get("client"), HTTPHandler):
|
||||
custom_client = kwargs.get("client")
|
||||
|
||||
return custom_handler.image_edit(
|
||||
|
|
@ -881,19 +845,15 @@ def image_edit(
|
|||
local_vars.update(kwargs)
|
||||
# Get ImageEditOptionalRequestParams with only valid parameters
|
||||
image_edit_optional_params: ImageEditOptionalRequestParams = (
|
||||
_get_ImageEditRequestUtils().get_requested_image_edit_optional_param(
|
||||
local_vars
|
||||
)
|
||||
_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(
|
||||
model=model,
|
||||
image_edit_provider_config=image_edit_provider_config,
|
||||
image_edit_optional_params=image_edit_optional_params,
|
||||
drop_params=kwargs.get("drop_params"),
|
||||
additional_drop_params=kwargs.get("additional_drop_params"),
|
||||
)
|
||||
image_edit_request_params: 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,
|
||||
drop_params=kwargs.get("drop_params"),
|
||||
additional_drop_params=kwargs.get("additional_drop_params"),
|
||||
)
|
||||
|
||||
# Pre Call logging
|
||||
|
|
|
|||
|
|
@ -39,9 +39,7 @@ class ImageEditRequestUtils:
|
|||
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 = [param for param in filtered_optional_params if param not in supported_params]
|
||||
|
||||
if unsupported_params:
|
||||
if should_drop:
|
||||
|
|
@ -54,9 +52,7 @@ class ImageEditRequestUtils:
|
|||
)
|
||||
|
||||
mapped_params = image_edit_provider_config.map_openai_params(
|
||||
image_edit_optional_params=cast(
|
||||
ImageEditOptionalRequestParams, filtered_optional_params
|
||||
),
|
||||
image_edit_optional_params=cast(ImageEditOptionalRequestParams, filtered_optional_params),
|
||||
model=model,
|
||||
drop_params=should_drop,
|
||||
)
|
||||
|
|
@ -77,9 +73,7 @@ class ImageEditRequestUtils:
|
|||
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
|
||||
}
|
||||
filtered_params = {k: v for k, v in params.items() if k in valid_keys and v is not None}
|
||||
return cast(ImageEditOptionalRequestParams, filtered_params)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -99,9 +93,7 @@ class ImageEditRequestUtils:
|
|||
# Save current position
|
||||
current_pos = image_data.tell()
|
||||
image_data.seek(0)
|
||||
bytes_data = image_data.read(
|
||||
100
|
||||
) # First 100 bytes are enough for detection
|
||||
bytes_data = image_data.read(100) # First 100 bytes are enough for detection
|
||||
# Restore position
|
||||
image_data.seek(current_pos)
|
||||
elif isinstance(image_data, BufferedReader):
|
||||
|
|
|
|||
|
|
@ -40,9 +40,7 @@ def squash_payloads(queue):
|
|||
return squashed
|
||||
|
||||
|
||||
def _print_alerting_payload_warning(
|
||||
payload: dict, slackAlertingInstance: SlackAlertingType
|
||||
):
|
||||
def _print_alerting_payload_warning(payload: dict, slackAlertingInstance: SlackAlertingType):
|
||||
"""
|
||||
Print the payload to the console when
|
||||
slackAlertingInstance.alerting_args.log_to_console is True
|
||||
|
|
@ -70,12 +68,8 @@ async def send_to_webhook(slackAlertingInstance: SlackAlertingType, item, count)
|
|||
data=json.dumps(payload),
|
||||
)
|
||||
if response.status_code != 200:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Error sending slack alert to url={item['url']}. Error={response.text}"
|
||||
)
|
||||
verbose_proxy_logger.debug(f"Error sending slack alert to url={item['url']}. Error={response.text}")
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(f"Error sending slack alert: {str(e)}")
|
||||
finally:
|
||||
_print_alerting_payload_warning(
|
||||
payload, slackAlertingInstance=slackAlertingInstance
|
||||
)
|
||||
_print_alerting_payload_warning(payload, slackAlertingInstance=slackAlertingInstance)
|
||||
|
|
|
|||
|
|
@ -41,8 +41,7 @@ class AlertingHangingRequestCheck:
|
|||
# stay cached for at least 1.5x the threshold to guarantee a check
|
||||
# happens after they cross it
|
||||
self.hanging_request_cache_ttl = int(
|
||||
self.slack_alerting_object.alerting_threshold * 1.5
|
||||
+ HANGING_ALERT_BUFFER_TIME_SECONDS
|
||||
self.slack_alerting_object.alerting_threshold * 1.5 + HANGING_ALERT_BUFFER_TIME_SECONDS
|
||||
)
|
||||
self.hanging_request_cache = InMemoryCache(
|
||||
default_ttl=self.hanging_request_cache_ttl,
|
||||
|
|
@ -62,9 +61,7 @@ class AlertingHangingRequestCheck:
|
|||
model = request_data.get("model", "")
|
||||
api_base: Optional[str] = None
|
||||
|
||||
if request_data.get("deployment", None) is not None and isinstance(
|
||||
request_data["deployment"], dict
|
||||
):
|
||||
if request_data.get("deployment", None) is not None and isinstance(request_data["deployment"], dict):
|
||||
api_base = litellm.get_api_base(
|
||||
model=model,
|
||||
optional_params=request_data["deployment"].get("litellm_params", {}),
|
||||
|
|
@ -104,9 +101,7 @@ class AlertingHangingRequestCheck:
|
|||
)
|
||||
|
||||
for request_id in hanging_requests:
|
||||
hanging_request_data: Optional[
|
||||
HangingRequestData
|
||||
] = await self.hanging_request_cache.async_get_cache(
|
||||
hanging_request_data: Optional[HangingRequestData] = await self.hanging_request_cache.async_get_cache(
|
||||
key=request_id,
|
||||
)
|
||||
|
||||
|
|
@ -116,12 +111,10 @@ class AlertingHangingRequestCheck:
|
|||
if hanging_request_data.alerted:
|
||||
continue
|
||||
|
||||
request_status = (
|
||||
await proxy_logging_obj.internal_usage_cache.async_get_cache(
|
||||
key="request_status:{}".format(hanging_request_data.request_id),
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
request_status = await proxy_logging_obj.internal_usage_cache.async_get_cache(
|
||||
key="request_status:{}".format(hanging_request_data.request_id),
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
# this means the request status was either success or fail
|
||||
# and is not hanging
|
||||
|
|
@ -141,9 +134,7 @@ class AlertingHangingRequestCheck:
|
|||
################
|
||||
# Send the Alert on Slack
|
||||
################
|
||||
await self.send_hanging_request_alert(
|
||||
hanging_request_data=hanging_request_data
|
||||
)
|
||||
await self.send_hanging_request_alert(hanging_request_data=hanging_request_data)
|
||||
# flag so the entry is skipped on later ticks; one alert per hang,
|
||||
# with the existing TTL still handling cleanup
|
||||
hanging_request_data.alerted = True
|
||||
|
|
|
|||
|
|
@ -62,9 +62,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
def __init__(
|
||||
self,
|
||||
internal_usage_cache: Optional[DualCache] = None,
|
||||
alerting_threshold: Optional[
|
||||
float
|
||||
] = None, # threshold for slow / hanging llm responses (in seconds)
|
||||
alerting_threshold: Optional[float] = None, # threshold for slow / hanging llm responses (in seconds)
|
||||
alerting: Optional[List] = [],
|
||||
alert_types: List[AlertType] = DEFAULT_ALERT_TYPES,
|
||||
alert_to_webhook_url: Optional[
|
||||
|
|
@ -81,12 +79,8 @@ class SlackAlerting(CustomBatchLogger):
|
|||
self.alerting = alerting
|
||||
self.alert_types = alert_types
|
||||
self.internal_usage_cache = internal_usage_cache or DualCache()
|
||||
self.async_http_handler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.LoggingCallback
|
||||
)
|
||||
self.alert_to_webhook_url = process_slack_alerting_variables(
|
||||
alert_to_webhook_url=alert_to_webhook_url
|
||||
)
|
||||
self.async_http_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
|
||||
self.alert_to_webhook_url = process_slack_alerting_variables(alert_to_webhook_url=alert_to_webhook_url)
|
||||
self.is_running = False
|
||||
self.alerting_args = SlackAlertingArgs(**alerting_args)
|
||||
self.default_webhook_url = default_webhook_url
|
||||
|
|
@ -98,9 +92,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
self.alert_type_config: Dict[str, AlertTypeConfig] = {}
|
||||
if alert_type_config:
|
||||
for key, val in alert_type_config.items():
|
||||
self.alert_type_config[key] = (
|
||||
AlertTypeConfig(**val) if isinstance(val, dict) else val
|
||||
)
|
||||
self.alert_type_config[key] = AlertTypeConfig(**val) if isinstance(val, dict) else val
|
||||
self.digest_buckets: Dict[str, DigestEntry] = {}
|
||||
self.digest_lock = asyncio.Lock()
|
||||
super().__init__(**kwargs, flush_lock=self.flush_lock)
|
||||
|
|
@ -130,23 +122,14 @@ class SlackAlerting(CustomBatchLogger):
|
|||
self.periodic_started = True
|
||||
if alert_type_config is not None:
|
||||
for key, val in alert_type_config.items():
|
||||
self.alert_type_config[key] = (
|
||||
AlertTypeConfig(**val) if isinstance(val, dict) else val
|
||||
)
|
||||
self.alert_type_config[key] = AlertTypeConfig(**val) if isinstance(val, dict) else val
|
||||
|
||||
if alert_to_webhook_url is not None:
|
||||
# update the dict
|
||||
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
|
||||
)
|
||||
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 = 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
|
||||
|
|
@ -161,15 +144,11 @@ class SlackAlerting(CustomBatchLogger):
|
|||
# Convert to dict for processing
|
||||
cache_value = dict(outage_value)
|
||||
|
||||
if "deployment_ids" in cache_value and isinstance(
|
||||
cache_value["deployment_ids"], set
|
||||
):
|
||||
if "deployment_ids" in cache_value and isinstance(cache_value["deployment_ids"], set):
|
||||
cache_value["deployment_ids"] = list(cache_value["deployment_ids"])
|
||||
return cache_value
|
||||
|
||||
def _restore_outage_value_from_cache(
|
||||
self, outage_value: Optional[dict]
|
||||
) -> Optional[dict]:
|
||||
def _restore_outage_value_from_cache(self, outage_value: Optional[dict]) -> Optional[dict]:
|
||||
"""
|
||||
Helper method to restore outage value after retrieving from cache.
|
||||
Converts list objects back to sets for proper handling.
|
||||
|
|
@ -234,9 +213,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
_deployment_latency_map: Optional[dict] = None
|
||||
try:
|
||||
# try sorting deployments by latency
|
||||
_deployment_latencies = sorted(
|
||||
_deployment_latencies.items(), key=lambda x: x[1]
|
||||
)
|
||||
_deployment_latencies = sorted(_deployment_latencies.items(), key=lambda x: x[1])
|
||||
_deployment_latency_map = dict(_deployment_latencies)
|
||||
except Exception:
|
||||
pass
|
||||
|
|
@ -276,23 +253,13 @@ class SlackAlerting(CustomBatchLogger):
|
|||
alerting_metadata: dict = {}
|
||||
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"]
|
||||
):
|
||||
if kwargs is not None and "litellm_params" in kwargs and "metadata" in kwargs["litellm_params"]:
|
||||
_metadata: dict = kwargs["litellm_params"]["metadata"]
|
||||
request_info = _add_key_name_and_team_to_alert(
|
||||
request_info=request_info, metadata=_metadata
|
||||
)
|
||||
request_info = _add_key_name_and_team_to_alert(request_info=request_info, metadata=_metadata)
|
||||
|
||||
_deployment_latency_map = self._get_deployment_latencies_to_alert(
|
||||
metadata=_metadata
|
||||
)
|
||||
_deployment_latency_map = 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}"
|
||||
)
|
||||
request_info += f"\nAvailable Deployment Latencies\n{_deployment_latency_map}"
|
||||
|
||||
if "alerting_metadata" in _metadata:
|
||||
alerting_metadata = _metadata["alerting_metadata"]
|
||||
|
|
@ -305,9 +272,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
api_base=api_base,
|
||||
)
|
||||
|
||||
async def async_update_daily_reports(
|
||||
self, deployment_metrics: DeploymentMetrics
|
||||
) -> int:
|
||||
async def async_update_daily_reports(self, deployment_metrics: DeploymentMetrics) -> int:
|
||||
"""
|
||||
Store the perf by deployment in cache
|
||||
- Number of failed requests per deployment
|
||||
|
|
@ -338,9 +303,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
## LATENCY ##
|
||||
if deployment_metrics.latency_per_output_token is not None:
|
||||
await self.internal_usage_cache.async_increment_cache(
|
||||
key="{}:{}".format(
|
||||
deployment_metrics.id, SlackAlertingCacheKeys.latency_key.value
|
||||
),
|
||||
key="{}:{}".format(deployment_metrics.id, SlackAlertingCacheKeys.latency_key.value),
|
||||
value=deployment_metrics.latency_per_output_token,
|
||||
parent_otel_span=None, # no attached request, this is a background operation
|
||||
)
|
||||
|
|
@ -370,13 +333,8 @@ class SlackAlerting(CustomBatchLogger):
|
|||
ids = router.get_model_ids()
|
||||
|
||||
# get keys
|
||||
failed_request_keys = [
|
||||
"{}:{}".format(id, SlackAlertingCacheKeys.failed_requests_key.value)
|
||||
for id in ids
|
||||
]
|
||||
latency_keys = [
|
||||
"{}:{}".format(id, SlackAlertingCacheKeys.latency_key.value) for id in ids
|
||||
]
|
||||
failed_request_keys = ["{}:{}".format(id, SlackAlertingCacheKeys.failed_requests_key.value) for id in ids]
|
||||
latency_keys = ["{}:{}".format(id, SlackAlertingCacheKeys.latency_key.value) for id in ids]
|
||||
|
||||
combined_metrics_keys = failed_request_keys + latency_keys # reduce cache calls
|
||||
|
||||
|
|
@ -396,18 +354,13 @@ class SlackAlerting(CustomBatchLogger):
|
|||
if all_none:
|
||||
return False
|
||||
|
||||
failed_request_values = combined_metrics_values[
|
||||
: len(failed_request_keys)
|
||||
] # # [1, 2, None, ..]
|
||||
failed_request_values = combined_metrics_values[: len(failed_request_keys)] # # [1, 2, None, ..]
|
||||
latency_values = combined_metrics_values[len(failed_request_keys) :]
|
||||
|
||||
# find top 5 failed
|
||||
## Replace None values with a placeholder value (-1 in this case)
|
||||
placeholder_value = 0
|
||||
replaced_failed_values = [
|
||||
value if value is not None else placeholder_value
|
||||
for value in failed_request_values
|
||||
]
|
||||
replaced_failed_values = [value if value is not None else placeholder_value for value in failed_request_values]
|
||||
|
||||
## Get the indices of top 5 keys with the highest numerical values (ignoring None and 0 values)
|
||||
top_5_failed = sorted(
|
||||
|
|
@ -415,17 +368,12 @@ class SlackAlerting(CustomBatchLogger):
|
|||
key=lambda i: replaced_failed_values[i],
|
||||
reverse=True,
|
||||
)[:5]
|
||||
top_5_failed = [
|
||||
index for index in top_5_failed if replaced_failed_values[index] > 0
|
||||
]
|
||||
top_5_failed = [index for index in top_5_failed if replaced_failed_values[index] > 0]
|
||||
|
||||
# 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 = [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(
|
||||
|
|
@ -433,9 +381,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
key=lambda i: replaced_slowest_values[i],
|
||||
reverse=True,
|
||||
)[:5]
|
||||
top_5_slowest = [
|
||||
index for index in top_5_slowest if replaced_slowest_values[index] > 0
|
||||
]
|
||||
top_5_slowest = [index for index in top_5_slowest if replaced_slowest_values[index] > 0]
|
||||
|
||||
# format alert -> return the litellm model name + api base
|
||||
message = f"\n\nTime: `{time.time()}`s\nHere are today's key metrics 📈: \n\n"
|
||||
|
|
@ -453,14 +399,14 @@ class SlackAlerting(CustomBatchLogger):
|
|||
|
||||
api_base = litellm.get_api_base(
|
||||
model=deployment_name,
|
||||
optional_params=(
|
||||
_deployment["litellm_params"] if _deployment is not None else {}
|
||||
),
|
||||
optional_params=(_deployment["litellm_params"] if _deployment is not None else {}),
|
||||
)
|
||||
if api_base is None:
|
||||
api_base = ""
|
||||
value = replaced_failed_values[top_5_failed[i]]
|
||||
message += f"\t{i + 1}. Deployment: `{deployment_name}`, Failed Requests: `{value}`, API Base: `{api_base}`\n"
|
||||
message += (
|
||||
f"\t{i + 1}. Deployment: `{deployment_name}`, Failed Requests: `{value}`, API Base: `{api_base}`\n"
|
||||
)
|
||||
|
||||
message += "\n\n*😅 Top Slowest Deployments:*\n\n"
|
||||
if not top_5_slowest:
|
||||
|
|
@ -474,9 +420,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
deployment_name = ""
|
||||
api_base = litellm.get_api_base(
|
||||
model=deployment_name,
|
||||
optional_params=(
|
||||
_deployment["litellm_params"] if _deployment is not None else {}
|
||||
),
|
||||
optional_params=(_deployment["litellm_params"] if _deployment is not None else {}),
|
||||
)
|
||||
value = round(replaced_slowest_values[top_5_slowest[i]], 3)
|
||||
message += f"\t{i + 1}. Deployment: `{deployment_name}`, Latency per output token: `{value}s/token`, API Base: `{api_base}`\n\n"
|
||||
|
|
@ -485,9 +429,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
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
|
||||
await self.internal_usage_cache.async_set_cache_pipeline(
|
||||
cache_list=combined_metrics_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"
|
||||
|
||||
|
|
@ -511,9 +453,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
if AlertType.llm_requests_hanging not in self.alert_types:
|
||||
return
|
||||
|
||||
await self.hanging_request_check.add_request_to_hanging_request_check(
|
||||
request_data=request_data
|
||||
)
|
||||
await self.hanging_request_check.add_request_to_hanging_request_check(request_data=request_data)
|
||||
|
||||
async def failed_tracking_alert(self, error_message: str, failing_model: str):
|
||||
"""
|
||||
|
|
@ -686,9 +626,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
if user_info.max_budget is not None:
|
||||
if user_info.spend >= user_info.max_budget:
|
||||
event = "budget_crossed"
|
||||
event_message += (
|
||||
f"Budget Crossed\n Total Budget:`{user_info.max_budget}`"
|
||||
)
|
||||
event_message += f"Budget Crossed\n Total Budget:`{user_info.max_budget}`"
|
||||
elif percent_left <= SLACK_ALERTING_THRESHOLD_5_PERCENT:
|
||||
event = "threshold_crossed"
|
||||
event_message += "5% Threshold Crossed "
|
||||
|
|
@ -755,9 +693,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
projected_spend=None,
|
||||
event="spend_tracked",
|
||||
event_group=Litellm_EntityType.END_USER,
|
||||
event_message="Customer spend tracked. Customer={}, spend={}".format(
|
||||
end_user_id, response_cost
|
||||
),
|
||||
event_message="Customer spend tracked. Customer={}, spend={}".format(end_user_id, response_cost),
|
||||
)
|
||||
|
||||
await self.send_webhook_alert(webhook_event=event)
|
||||
|
|
@ -852,9 +788,9 @@ class SlackAlerting(CustomBatchLogger):
|
|||
### UNIQUE CACHE KEY ###
|
||||
cache_key = provider + region_name
|
||||
|
||||
outage_value: Optional[
|
||||
ProviderRegionOutageModel
|
||||
] = await self.internal_usage_cache.async_get_cache(key=cache_key)
|
||||
outage_value: Optional[ProviderRegionOutageModel] = await self.internal_usage_cache.async_get_cache(
|
||||
key=cache_key
|
||||
)
|
||||
|
||||
# Convert deployment_ids back to set if it was stored as a list
|
||||
if outage_value is not None:
|
||||
|
|
@ -904,8 +840,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
## MINOR OUTAGE ALERT SENT ##
|
||||
if (
|
||||
outage_value["minor_alert_sent"] is False
|
||||
and len(outage_value["alerts"])
|
||||
>= self.alerting_args.minor_outage_alert_threshold
|
||||
and len(outage_value["alerts"]) >= self.alerting_args.minor_outage_alert_threshold
|
||||
and len(_deployment_set) > 1 # make sure it's not just 1 bad deployment
|
||||
):
|
||||
msg = self._outage_alert_msg_factory(
|
||||
|
|
@ -929,8 +864,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
## MAJOR OUTAGE ALERT SENT ##
|
||||
elif (
|
||||
outage_value["major_alert_sent"] is False
|
||||
and len(outage_value["alerts"])
|
||||
>= self.alerting_args.major_outage_alert_threshold
|
||||
and len(outage_value["alerts"]) >= self.alerting_args.major_outage_alert_threshold
|
||||
and len(_deployment_set) > 1 # make sure it's not just 1 bad deployment
|
||||
):
|
||||
msg = self._outage_alert_msg_factory(
|
||||
|
|
@ -955,9 +889,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
## update cache ##
|
||||
# Convert set to list for JSON serialization
|
||||
cache_value = self._prepare_outage_value_for_cache(outage_value)
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=cache_key, value=cache_value
|
||||
)
|
||||
await self.internal_usage_cache.async_set_cache(key=cache_key, value=cache_value)
|
||||
|
||||
async def outage_alerts(
|
||||
self,
|
||||
|
|
@ -979,9 +911,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
max_alerts_size = 10
|
||||
"""
|
||||
try:
|
||||
outage_value: Optional[
|
||||
OutageModel
|
||||
] = await self.internal_usage_cache.async_get_cache(key=deployment_id) # type: ignore
|
||||
outage_value: Optional[OutageModel] = await self.internal_usage_cache.async_get_cache(key=deployment_id) # type: ignore
|
||||
if (
|
||||
getattr(exception, "status_code", None) is None
|
||||
or (
|
||||
|
|
@ -1004,9 +934,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 = litellm.get_api_base(model=model, optional_params=deployment.litellm_params)
|
||||
|
||||
if outage_value is None:
|
||||
outage_value = OutageModel(
|
||||
|
|
@ -1025,10 +953,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
)
|
||||
return
|
||||
|
||||
if (
|
||||
len(outage_value["alerts"])
|
||||
< self.alerting_args.max_outage_alert_list_size
|
||||
):
|
||||
if len(outage_value["alerts"]) < self.alerting_args.max_outage_alert_list_size:
|
||||
outage_value["alerts"].append(exception.status_code) # type: ignore
|
||||
else: # prevent memory leaks
|
||||
pass
|
||||
|
|
@ -1038,8 +963,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
## MINOR OUTAGE ALERT SENT ##
|
||||
if (
|
||||
outage_value["minor_alert_sent"] is False
|
||||
and len(outage_value["alerts"])
|
||||
>= self.alerting_args.minor_outage_alert_threshold
|
||||
and len(outage_value["alerts"]) >= self.alerting_args.minor_outage_alert_threshold
|
||||
):
|
||||
msg = self._outage_alert_msg_factory(
|
||||
alert_type="Minor",
|
||||
|
|
@ -1060,8 +984,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
outage_value["minor_alert_sent"] = True
|
||||
elif (
|
||||
outage_value["major_alert_sent"] is False
|
||||
and len(outage_value["alerts"])
|
||||
>= self.alerting_args.major_outage_alert_threshold
|
||||
and len(outage_value["alerts"]) >= self.alerting_args.major_outage_alert_threshold
|
||||
):
|
||||
msg = self._outage_alert_msg_factory(
|
||||
alert_type="Major",
|
||||
|
|
@ -1084,15 +1007,11 @@ class SlackAlerting(CustomBatchLogger):
|
|||
## update cache ##
|
||||
# Convert set to list for JSON serialization
|
||||
cache_value = self._prepare_outage_value_for_cache(outage_value)
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=deployment_id, value=cache_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
|
||||
):
|
||||
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)
|
||||
model_info = {}
|
||||
base_model = ""
|
||||
|
|
@ -1193,14 +1112,10 @@ Model Info:
|
|||
|
||||
if premium_user is not True:
|
||||
if email_logo_url is not None or email_support_contact is not None:
|
||||
raise ValueError(
|
||||
f"Trying to Customize Email Alerting\n {CommonProxyErrors.not_premium_user.value}"
|
||||
)
|
||||
raise ValueError(f"Trying to Customize Email Alerting\n {CommonProxyErrors.not_premium_user.value}")
|
||||
return
|
||||
|
||||
async def send_key_created_or_user_invited_email(
|
||||
self, webhook_event: WebhookEvent
|
||||
) -> bool:
|
||||
async def send_key_created_or_user_invited_email(self, webhook_event: WebhookEvent) -> bool:
|
||||
try:
|
||||
from litellm.proxy.utils import send_email
|
||||
|
||||
|
|
@ -1213,13 +1128,9 @@ Model Info:
|
|||
return False
|
||||
from litellm.proxy.proxy_server import premium_user, prisma_client
|
||||
|
||||
email_logo_url = os.getenv(
|
||||
"SMTP_SENDER_LOGO", os.getenv("EMAIL_LOGO_URL", None)
|
||||
)
|
||||
email_logo_url = os.getenv("SMTP_SENDER_LOGO", os.getenv("EMAIL_LOGO_URL", None))
|
||||
email_support_contact = os.getenv("EMAIL_SUPPORT_CONTACT", None)
|
||||
await self._check_if_using_premium_email_feature(
|
||||
premium_user, email_logo_url, email_support_contact
|
||||
)
|
||||
await self._check_if_using_premium_email_feature(premium_user, email_logo_url, email_support_contact)
|
||||
if email_logo_url is None:
|
||||
email_logo_url = LITELLM_LOGO_URL
|
||||
if email_support_contact is None:
|
||||
|
|
@ -1228,14 +1139,8 @@ Model Info:
|
|||
event_name = webhook_event.event_message
|
||||
recipient_email = webhook_event.user_email
|
||||
recipient_user_id = 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 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
|
||||
|
|
@ -1265,9 +1170,7 @@ Model Info:
|
|||
team_id = 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 = 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(
|
||||
|
|
@ -1302,9 +1205,7 @@ Model Info:
|
|||
verbose_proxy_logger.error("Error sending email alert %s", str(e))
|
||||
return False
|
||||
|
||||
async def send_email_alert_using_smtp(
|
||||
self, webhook_event: WebhookEvent, alert_type: str
|
||||
) -> bool:
|
||||
async def send_email_alert_using_smtp(self, webhook_event: WebhookEvent, alert_type: str) -> bool:
|
||||
"""
|
||||
Sends structured Email alert to an SMTP server
|
||||
|
||||
|
|
@ -1315,13 +1216,9 @@ Model Info:
|
|||
from litellm.proxy.proxy_server import premium_user
|
||||
from litellm.proxy.utils import send_email
|
||||
|
||||
email_logo_url = os.getenv(
|
||||
"SMTP_SENDER_LOGO", os.getenv("EMAIL_LOGO_URL", None)
|
||||
)
|
||||
email_logo_url = os.getenv("SMTP_SENDER_LOGO", os.getenv("EMAIL_LOGO_URL", None))
|
||||
email_support_contact = os.getenv("EMAIL_SUPPORT_CONTACT", None)
|
||||
await self._check_if_using_premium_email_feature(
|
||||
premium_user, email_logo_url, email_support_contact
|
||||
)
|
||||
await self._check_if_using_premium_email_feature(premium_user, email_logo_url, email_support_contact)
|
||||
|
||||
if email_logo_url is None:
|
||||
email_logo_url = LITELLM_LOGO_URL
|
||||
|
|
@ -1334,9 +1231,7 @@ Model Info:
|
|||
max_budget = 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()
|
||||
)
|
||||
verbose_proxy_logger.error("Trying to send email alert to no recipient", extra=webhook_event.dict())
|
||||
|
||||
if webhook_event.event == "budget_crossed":
|
||||
email_html_content = f"""
|
||||
|
|
@ -1404,30 +1299,16 @@ Model Info:
|
|||
return
|
||||
|
||||
# Start periodic flush if not already started
|
||||
if (
|
||||
not self.periodic_started
|
||||
and self.alerting is not None
|
||||
and len(self.alerting) > 0
|
||||
):
|
||||
if not self.periodic_started and self.alerting is not None and len(self.alerting) > 0:
|
||||
asyncio.create_task(self.periodic_flush())
|
||||
self.periodic_started = True
|
||||
|
||||
if (
|
||||
"webhook" in self.alerting
|
||||
and alert_type == "budget_alerts"
|
||||
and user_info is not None
|
||||
):
|
||||
if "webhook" in self.alerting and alert_type == "budget_alerts" and user_info is not None:
|
||||
await self.send_webhook_alert(webhook_event=user_info)
|
||||
|
||||
if (
|
||||
"email" in self.alerting
|
||||
and alert_type == "budget_alerts"
|
||||
and user_info is not None
|
||||
):
|
||||
if "email" in self.alerting and alert_type == "budget_alerts" and user_info is not None:
|
||||
# only send budget alerts over Email
|
||||
await self.send_email_alert_using_smtp(
|
||||
webhook_event=user_info, alert_type=alert_type
|
||||
)
|
||||
await self.send_email_alert_using_smtp(webhook_event=user_info, alert_type=alert_type)
|
||||
|
||||
if "slack" not in self.alerting:
|
||||
return
|
||||
|
|
@ -1441,13 +1322,8 @@ Model Info:
|
|||
_atc = 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
|
||||
):
|
||||
_digest_webhook: Optional[Union[str, List[str]]] = (
|
||||
self.alert_to_webhook_url[alert_type]
|
||||
)
|
||||
if self.alert_to_webhook_url is not None and alert_type in self.alert_to_webhook_url:
|
||||
_digest_webhook: Optional[Union[str, List[str]]] = self.alert_to_webhook_url[alert_type]
|
||||
elif self.default_webhook_url is not None:
|
||||
_digest_webhook = self.default_webhook_url
|
||||
else:
|
||||
|
|
@ -1485,7 +1361,9 @@ Model Info:
|
|||
if alert_type == "daily_reports" or alert_type == "new_model_added":
|
||||
formatted_message = alert_type_formatted + message
|
||||
else:
|
||||
formatted_message = f"{alert_type_formatted}\nLevel: `{level}`\nTimestamp: `{current_time}`\n\nMessage: {message}"
|
||||
formatted_message = (
|
||||
f"{alert_type_formatted}\nLevel: `{level}`\nTimestamp: `{current_time}`\n\nMessage: {message}"
|
||||
)
|
||||
|
||||
if kwargs:
|
||||
for key, value in kwargs.items():
|
||||
|
|
@ -1497,13 +1375,8 @@ Model Info:
|
|||
formatted_message += f"\n\nProxy URL: `{_proxy_base_url}`"
|
||||
|
||||
# check if we find the slack webhook url in self.alert_to_webhook_url
|
||||
if (
|
||||
self.alert_to_webhook_url is not None
|
||||
and alert_type in self.alert_to_webhook_url
|
||||
):
|
||||
slack_webhook_url: Optional[Union[str, List[str]]] = (
|
||||
self.alert_to_webhook_url[alert_type]
|
||||
)
|
||||
if self.alert_to_webhook_url is not None and alert_type in self.alert_to_webhook_url:
|
||||
slack_webhook_url: Optional[Union[str, List[str]]] = self.alert_to_webhook_url[alert_type]
|
||||
elif self.default_webhook_url is not None:
|
||||
slack_webhook_url = self.default_webhook_url
|
||||
else:
|
||||
|
|
@ -1543,9 +1416,7 @@ Model Info:
|
|||
|
||||
squashed_queue = squash_payloads(self.log_queue)
|
||||
tasks = [
|
||||
send_to_webhook(
|
||||
slackAlertingInstance=self, item=item["item"], count=item["count"]
|
||||
)
|
||||
send_to_webhook(slackAlertingInstance=self, item=item["item"], count=item["count"])
|
||||
for item in squashed_queue.values()
|
||||
]
|
||||
await asyncio.gather(*tasks)
|
||||
|
|
@ -1645,9 +1516,7 @@ Model Info:
|
|||
):
|
||||
completion_tokens = 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
|
||||
)
|
||||
final_value = float(response_s.total_seconds() / completion_tokens)
|
||||
if isinstance(final_value, timedelta):
|
||||
final_value = final_value.total_seconds()
|
||||
|
||||
|
|
@ -1692,9 +1561,7 @@ Model Info:
|
|||
)
|
||||
|
||||
if "region_outage_alerts" in self.alert_types:
|
||||
await self.region_outage_alerts(
|
||||
exception=kwargs["exception"], deployment_id=model_id
|
||||
)
|
||||
await self.region_outage_alerts(exception=kwargs["exception"], deployment_id=model_id)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
|
@ -1781,7 +1648,9 @@ Model Info:
|
|||
todays_date = datetime.datetime.now().date()
|
||||
start_date = todays_date - datetime.timedelta(days=days)
|
||||
|
||||
_event_cache_key = f"weekly_spend_report_sent_{start_date.strftime('%Y-%m-%d')}_{todays_date.strftime('%Y-%m-%d')}"
|
||||
_event_cache_key = (
|
||||
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
|
||||
|
||||
|
|
@ -1800,9 +1669,7 @@ Model Info:
|
|||
_spend_message += "\n*Team Spend Report:*\n"
|
||||
for spend in spend_per_team:
|
||||
_team_spend = round(float(spend["total_spend"]), 4)
|
||||
_spend_message += (
|
||||
f"Team: `{spend['team_alias']}` | Spend: `${_team_spend}`\n"
|
||||
)
|
||||
_spend_message += f"Team: `{spend['team_alias']}` | Spend: `${_team_spend}`\n"
|
||||
|
||||
if spend_per_tag is not None:
|
||||
_spend_message += "\n*Tag Spend Report:*\n"
|
||||
|
|
@ -1840,9 +1707,7 @@ Model Info:
|
|||
todays_date = datetime.datetime.now().date()
|
||||
first_day_of_month = 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
|
||||
)
|
||||
last_day_of_month = first_day_of_month + datetime.timedelta(days=last_day_of_month - 1)
|
||||
|
||||
_event_cache_key = f"monthly_spend_report_sent_{first_day_of_month.strftime('%Y-%m-%d')}_{last_day_of_month.strftime('%Y-%m-%d')}"
|
||||
if await self.internal_usage_cache.async_get_cache(key=_event_cache_key):
|
||||
|
|
@ -1867,9 +1732,7 @@ Model Info:
|
|||
_team_spend = float(_team_spend)
|
||||
# round to 4 decimal places
|
||||
_team_spend = round(_team_spend, 4)
|
||||
_spend_message += (
|
||||
f"Team: `{spend['team_alias']}` | Spend: `${_team_spend}`\n"
|
||||
)
|
||||
_spend_message += f"Team: `{spend['team_alias']}` | Spend: `${_team_spend}`\n"
|
||||
|
||||
if monthly_spend_per_tag is not None:
|
||||
_spend_message += "\n*Tag Spend Report:*\n"
|
||||
|
|
@ -1908,13 +1771,9 @@ Model Info:
|
|||
)
|
||||
|
||||
# call prometheuslogger.
|
||||
falllback_success_info_prometheus = (
|
||||
await get_fallback_metric_from_prometheus()
|
||||
)
|
||||
falllback_success_info_prometheus = await get_fallback_metric_from_prometheus()
|
||||
|
||||
fallback_message = (
|
||||
f"*Fallback Statistics:*\n{falllback_success_info_prometheus}"
|
||||
)
|
||||
fallback_message = f"*Fallback Statistics:*\n{falllback_success_info_prometheus}"
|
||||
|
||||
await self.send_alert(
|
||||
message=fallback_message,
|
||||
|
|
@ -1969,9 +1828,7 @@ Model Info:
|
|||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"Error sending send_virtual_key_event_slack %s", e
|
||||
)
|
||||
verbose_proxy_logger.error("Error sending send_virtual_key_event_slack %s", e)
|
||||
|
||||
return
|
||||
|
||||
|
|
@ -1982,10 +1839,7 @@ Model Info:
|
|||
if request_data is None:
|
||||
return False
|
||||
|
||||
if (
|
||||
request_data.get("litellm_status", "") != "success"
|
||||
and request_data.get("litellm_status", "") != "fail"
|
||||
):
|
||||
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: Optional[str] = await self.internal_usage_cache.async_get_cache(
|
||||
|
|
|
|||
|
|
@ -34,9 +34,7 @@ def process_slack_alerting_variables(
|
|||
if "os.environ/" in webhook_url:
|
||||
_env_value = get_secret(secret_name=webhook_url)
|
||||
if not isinstance(_env_value, str):
|
||||
raise ValueError(
|
||||
f"Invalid webhook url value for: {webhook_url}. Got type={type(_env_value)}"
|
||||
)
|
||||
raise ValueError(f"Invalid webhook url value for: {webhook_url}. Got type={type(_env_value)}")
|
||||
_webhook_values.append(_env_value)
|
||||
else:
|
||||
_webhook_values.append(webhook_url)
|
||||
|
|
@ -47,9 +45,7 @@ def process_slack_alerting_variables(
|
|||
if "os.environ/" in webhook_urls:
|
||||
_env_value = get_secret(secret_name=webhook_urls)
|
||||
if not isinstance(_env_value, str):
|
||||
raise ValueError(
|
||||
f"Invalid webhook url value for: {webhook_urls}. Got type={type(_env_value)}"
|
||||
)
|
||||
raise ValueError(f"Invalid webhook url value for: {webhook_urls}. Got type={type(_env_value)}")
|
||||
_webhook_value_str = _env_value
|
||||
else:
|
||||
_webhook_value_str = webhook_urls
|
||||
|
|
@ -76,10 +72,7 @@ async def _add_langfuse_trace_id_to_alert(
|
|||
# Only run if langfuse is added as a callback
|
||||
#########################################################
|
||||
|
||||
if (
|
||||
request_data is not None
|
||||
and request_data.get("litellm_logging_obj", None) is not None
|
||||
):
|
||||
if request_data is not None and request_data.get("litellm_logging_obj", None) is not None:
|
||||
trace_id: Optional[str] = None
|
||||
litellm_logging_obj: Logging = request_data["litellm_logging_obj"]
|
||||
|
||||
|
|
@ -89,9 +82,7 @@ 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 = litellm_logging_obj._get_callback_object(service_name="langfuse")
|
||||
if langfuse_object is not None:
|
||||
base_url = langfuse_object.Langfuse.base_url
|
||||
return f"{base_url}/trace/{trace_id}"
|
||||
|
|
|
|||
|
|
@ -73,15 +73,11 @@ class SpanAttributes:
|
|||
"""
|
||||
Number of tokens in the prompt.
|
||||
"""
|
||||
LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_WRITE = (
|
||||
"llm.token_count.prompt_details.cache_write"
|
||||
)
|
||||
LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_WRITE = "llm.token_count.prompt_details.cache_write"
|
||||
"""
|
||||
Number of tokens in the prompt that were written to cache.
|
||||
"""
|
||||
LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_READ = (
|
||||
"llm.token_count.prompt_details.cache_read"
|
||||
)
|
||||
LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_READ = "llm.token_count.prompt_details.cache_read"
|
||||
"""
|
||||
Number of tokens in the prompt that were read from cache.
|
||||
"""
|
||||
|
|
@ -93,15 +89,11 @@ class SpanAttributes:
|
|||
"""
|
||||
Number of tokens in the completion.
|
||||
"""
|
||||
LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING = (
|
||||
"llm.token_count.completion_details.reasoning"
|
||||
)
|
||||
LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING = "llm.token_count.completion_details.reasoning"
|
||||
"""
|
||||
Number of tokens used for reasoning steps in the completion.
|
||||
"""
|
||||
LLM_TOKEN_COUNT_COMPLETION_DETAILS_AUDIO = (
|
||||
"llm.token_count.completion_details.audio"
|
||||
)
|
||||
LLM_TOKEN_COUNT_COMPLETION_DETAILS_AUDIO = "llm.token_count.completion_details.audio"
|
||||
"""
|
||||
The number of audio input tokens generated by the model
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -65,9 +65,7 @@ class AgentOps(OpenTelemetry):
|
|||
|
||||
headers = f"Authorization=Bearer {jwt_token}" if jwt_token else None
|
||||
|
||||
otel_config = OpenTelemetryConfig(
|
||||
exporter="otlp_http", endpoint=config.endpoint, headers=headers
|
||||
)
|
||||
otel_config = OpenTelemetryConfig(exporter="otlp_http", endpoint=config.endpoint, headers=headers)
|
||||
|
||||
# Initialize OpenTelemetry with our config
|
||||
super().__init__(config=otel_config, callback_name="agentops")
|
||||
|
|
|
|||
|
|
@ -78,11 +78,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 = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0
|
||||
|
||||
processed_messages = self._apply_message_injections(
|
||||
points=message_points,
|
||||
|
|
@ -111,10 +107,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
``max_blocks`` is reached. Injection points are honored in config order,
|
||||
so earlier points win when slots are scarce.
|
||||
"""
|
||||
used_blocks = sum(
|
||||
AnthropicCacheControlHook._count_cache_control_blocks(msg)
|
||||
for msg in messages
|
||||
)
|
||||
used_blocks = sum(AnthropicCacheControlHook._count_cache_control_blocks(msg) for msg in messages)
|
||||
|
||||
limit_reached = False
|
||||
for point in points:
|
||||
|
|
@ -122,27 +115,21 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
limit_reached = True
|
||||
break
|
||||
|
||||
control: ChatCompletionCachedContent = point.get(
|
||||
"control", None
|
||||
) or ChatCompletionCachedContent(type="ephemeral")
|
||||
control: ChatCompletionCachedContent = point.get("control", None) or ChatCompletionCachedContent(
|
||||
type="ephemeral"
|
||||
)
|
||||
|
||||
for target_index in AnthropicCacheControlHook._resolve_target_indices(
|
||||
point=point, messages=messages
|
||||
):
|
||||
for target_index in AnthropicCacheControlHook._resolve_target_indices(point=point, messages=messages):
|
||||
if used_blocks >= max_blocks:
|
||||
limit_reached = True
|
||||
break
|
||||
|
||||
if AnthropicCacheControlHook._message_has_cache_control(
|
||||
messages[target_index]
|
||||
):
|
||||
if AnthropicCacheControlHook._message_has_cache_control(messages[target_index]):
|
||||
# Client already marked this message; don't overwrite it.
|
||||
continue
|
||||
|
||||
messages[target_index] = (
|
||||
AnthropicCacheControlHook._safe_insert_cache_control_in_message(
|
||||
messages[target_index], control
|
||||
)
|
||||
messages[target_index] = AnthropicCacheControlHook._safe_insert_cache_control_in_message(
|
||||
messages[target_index], control
|
||||
)
|
||||
used_blocks += 1
|
||||
|
||||
|
|
@ -190,11 +177,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
# Case 2: Target by role
|
||||
targetted_role = point.get("role", None)
|
||||
if targetted_role is not None:
|
||||
return [
|
||||
idx
|
||||
for idx, msg in enumerate(messages)
|
||||
if msg.get("role") == targetted_role
|
||||
]
|
||||
return [idx for idx, msg in enumerate(messages) if msg.get("role") == targetted_role]
|
||||
|
||||
return []
|
||||
|
||||
|
|
@ -338,9 +321,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
_init_custom_logger_compatible_class,
|
||||
)
|
||||
|
||||
if AnthropicCacheControlHook.should_use_anthropic_cache_control_hook(
|
||||
non_default_params
|
||||
):
|
||||
if AnthropicCacheControlHook.should_use_anthropic_cache_control_hook(non_default_params):
|
||||
return _init_custom_logger_compatible_class(
|
||||
logging_integration="anthropic_cache_control_hook",
|
||||
internal_usage_cache=None,
|
||||
|
|
|
|||
|
|
@ -47,12 +47,8 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
**kwargs,
|
||||
):
|
||||
if litellm.argilla_transformation_object is None:
|
||||
raise Exception(
|
||||
"'litellm.argilla_transformation_object' is required, to log your payload to Argilla."
|
||||
)
|
||||
self.validate_argilla_transformation_object(
|
||||
litellm.argilla_transformation_object
|
||||
)
|
||||
raise Exception("'litellm.argilla_transformation_object' is required, to log your payload to Argilla.")
|
||||
self.validate_argilla_transformation_object(litellm.argilla_transformation_object)
|
||||
self.argilla_transformation_object = litellm.argilla_transformation_object
|
||||
self.default_credentials = self.get_credentials_from_env(
|
||||
argilla_api_key=argilla_api_key,
|
||||
|
|
@ -61,30 +57,21 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
)
|
||||
self.sampling_rate: float = (
|
||||
float(os.getenv("ARGILLA_SAMPLING_RATE")) # type: ignore
|
||||
if os.getenv("ARGILLA_SAMPLING_RATE") is not None
|
||||
and os.getenv("ARGILLA_SAMPLING_RATE").strip().isdigit() # type: ignore
|
||||
if os.getenv("ARGILLA_SAMPLING_RATE") is not None and os.getenv("ARGILLA_SAMPLING_RATE").strip().isdigit() # type: ignore
|
||||
else 1.0
|
||||
)
|
||||
|
||||
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
|
||||
)
|
||||
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
|
||||
if _batch_size:
|
||||
self.batch_size = int(_batch_size)
|
||||
asyncio.create_task(self.periodic_flush())
|
||||
self.flush_lock = asyncio.Lock()
|
||||
super().__init__(**kwargs, flush_lock=self.flush_lock)
|
||||
|
||||
def validate_argilla_transformation_object(
|
||||
self, argilla_transformation_object: Dict[str, Any]
|
||||
):
|
||||
def validate_argilla_transformation_object(self, argilla_transformation_object: Dict[str, Any]):
|
||||
if not isinstance(argilla_transformation_object, dict):
|
||||
raise Exception(
|
||||
"'argilla_transformation_object' must be a dictionary, to log your payload to Argilla."
|
||||
)
|
||||
raise Exception("'argilla_transformation_object' must be a dictionary, to log your payload to Argilla.")
|
||||
|
||||
for v in argilla_transformation_object.values():
|
||||
if v not in SUPPORTED_PAYLOAD_FIELDS:
|
||||
|
|
@ -102,21 +89,11 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
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 = 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."
|
||||
)
|
||||
raise Exception("Invalid Argilla Base URL given. _credentials_base_url=None.")
|
||||
|
||||
_credentials_dataset_name = (
|
||||
argilla_dataset_name
|
||||
or os.getenv("ARGILLA_DATASET_NAME")
|
||||
or "litellm-completion"
|
||||
)
|
||||
_credentials_dataset_name = argilla_dataset_name or os.getenv("ARGILLA_DATASET_NAME") or "litellm-completion"
|
||||
if _credentials_dataset_name is None:
|
||||
raise Exception("Invalid Argilla Dataset give. Value=None.")
|
||||
else:
|
||||
|
|
@ -138,19 +115,13 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
ARGILLA_DATASET_NAME=_credentials_dataset_name,
|
||||
)
|
||||
|
||||
def get_chat_messages(
|
||||
self, payload: StandardLoggingPayload
|
||||
) -> List[Dict[str, Any]]:
|
||||
def get_chat_messages(self, payload: StandardLoggingPayload) -> List[Dict[str, Any]]:
|
||||
payload_messages = payload.get("messages", None)
|
||||
|
||||
if payload_messages is None:
|
||||
raise Exception("No chat messages found in payload.")
|
||||
|
||||
if (
|
||||
isinstance(payload_messages, list)
|
||||
and len(payload_messages) > 0
|
||||
and isinstance(payload_messages[0], dict)
|
||||
):
|
||||
if isinstance(payload_messages, list) and len(payload_messages) > 0 and isinstance(payload_messages[0], dict):
|
||||
return payload_messages
|
||||
elif isinstance(payload_messages, dict):
|
||||
return [payload_messages]
|
||||
|
|
@ -166,20 +137,14 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
if isinstance(response, str):
|
||||
return response
|
||||
elif isinstance(response, dict):
|
||||
return (
|
||||
response.get("choices", [{}])[0].get("message", {}).get("content", "")
|
||||
)
|
||||
return response.get("choices", [{}])[0].get("message", {}).get("content", "")
|
||||
else:
|
||||
raise Exception(f"Invalid response format: {response}")
|
||||
|
||||
def _prepare_log_data(
|
||||
self, kwargs, response_obj, start_time, end_time
|
||||
) -> Optional[ArgillaItem]:
|
||||
def _prepare_log_data(self, kwargs, response_obj, start_time, end_time) -> Optional[ArgillaItem]:
|
||||
try:
|
||||
# Ensure everything in the payload is converted to str
|
||||
payload: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object", None
|
||||
)
|
||||
payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None)
|
||||
|
||||
if payload is None:
|
||||
raise Exception("Error logging request payload. Payload=none.")
|
||||
|
|
@ -220,13 +185,9 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
)
|
||||
|
||||
if response.status_code >= 300:
|
||||
verbose_logger.error(
|
||||
f"Argilla Error: {response.status_code} - {response.text}"
|
||||
)
|
||||
verbose_logger.error(f"Argilla Error: {response.status_code} - {response.text}")
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
f"Batch of {len(self.log_queue)} runs successfully created"
|
||||
)
|
||||
verbose_logger.debug(f"Batch of {len(self.log_queue)} runs successfully created")
|
||||
|
||||
self.log_queue.clear()
|
||||
except Exception:
|
||||
|
|
@ -258,9 +219,7 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
return
|
||||
|
||||
self.log_queue.append(data)
|
||||
verbose_logger.debug(
|
||||
f"Langsmith, event added to queue. Will flush in {self.flush_interval} seconds..."
|
||||
)
|
||||
verbose_logger.debug(f"Langsmith, event added to queue. Will flush in {self.flush_interval} seconds...")
|
||||
|
||||
if len(self.log_queue) >= self.batch_size:
|
||||
self._send_batch()
|
||||
|
|
@ -284,9 +243,7 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
kwargs,
|
||||
response_obj,
|
||||
)
|
||||
payload: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object", None
|
||||
)
|
||||
payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object", None)
|
||||
|
||||
data = self._prepare_log_data(kwargs, response_obj, start_time, end_time)
|
||||
|
||||
|
|
@ -312,18 +269,14 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
if len(self.log_queue) >= self.batch_size:
|
||||
await self.flush_queue()
|
||||
except Exception:
|
||||
verbose_logger.exception(
|
||||
"Argilla Layer Error - error logging async success event."
|
||||
)
|
||||
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()
|
||||
if random_sample > sampling_rate:
|
||||
verbose_logger.info(
|
||||
"Skipping Langsmith logging. Sampling rate={}, random_sample={}".format(
|
||||
sampling_rate, random_sample
|
||||
)
|
||||
"Skipping Langsmith logging. Sampling rate={}, random_sample={}".format(sampling_rate, random_sample)
|
||||
)
|
||||
return # Skip logging
|
||||
verbose_logger.info("Langsmith Failure Event Logging!")
|
||||
|
|
@ -338,9 +291,7 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
if len(self.log_queue) >= self.batch_size:
|
||||
await self.flush_queue()
|
||||
except Exception:
|
||||
verbose_logger.exception(
|
||||
"Langsmith Layer Error - error logging async failure event."
|
||||
)
|
||||
verbose_logger.exception("Langsmith Layer Error - error logging async failure event.")
|
||||
|
||||
async def async_send_batch(self):
|
||||
"""
|
||||
|
|
@ -378,13 +329,9 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
response.raise_for_status()
|
||||
|
||||
if response.status_code >= 300:
|
||||
verbose_logger.error(
|
||||
f"Argilla Error: {response.status_code} - {response.text}"
|
||||
)
|
||||
verbose_logger.error(f"Argilla Error: {response.status_code} - {response.text}")
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
"Batch of %s runs successfully created", len(self.log_queue)
|
||||
)
|
||||
verbose_logger.debug("Batch of %s runs successfully created", len(self.log_queue))
|
||||
except httpx.HTTPStatusError:
|
||||
verbose_logger.exception("Argilla HTTP Error")
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -13,22 +13,16 @@ from .arize_phoenix_prompt_manager import ArizePhoenixPromptManager
|
|||
global_arize_config: Optional[dict] = None
|
||||
|
||||
|
||||
def prompt_initializer(
|
||||
litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec"
|
||||
) -> "CustomPromptManagement":
|
||||
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_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)
|
||||
|
||||
if not api_key or not api_base:
|
||||
raise ValueError(
|
||||
"api_key and api_base are required for Arize Phoenix prompt integration"
|
||||
)
|
||||
raise ValueError("api_key and api_base are required for Arize Phoenix prompt integration")
|
||||
|
||||
try:
|
||||
arize_prompt_manager = ArizePhoenixPromptManager(
|
||||
|
|
@ -36,9 +30,7 @@ def prompt_initializer(
|
|||
"api_key": api_key,
|
||||
"api_base": api_base,
|
||||
"prompt_id": prompt_id,
|
||||
**litellm_params.model_dump(
|
||||
exclude={"api_key", "api_base", "prompt_id"}
|
||||
),
|
||||
**litellm_params.model_dump(exclude={"api_key", "api_base", "prompt_id"}),
|
||||
},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -48,9 +48,7 @@ class ArizeOTELAttributes(BaseLLMObsOTELAttributes):
|
|||
for idx, msg in enumerate(messages):
|
||||
prefix = f"{SpanAttributes.LLM_INPUT_MESSAGES}.{idx}"
|
||||
# Set the role per message.
|
||||
safe_set_attribute(
|
||||
span, f"{prefix}.{MessageAttributes.MESSAGE_ROLE}", msg.get("role")
|
||||
)
|
||||
safe_set_attribute(span, f"{prefix}.{MessageAttributes.MESSAGE_ROLE}", msg.get("role"))
|
||||
# Set the content per message.
|
||||
safe_set_attribute(
|
||||
span,
|
||||
|
|
@ -164,9 +162,7 @@ def _set_audio_outputs(span: "Span", response_obj, audio_attrs, span_attrs):
|
|||
|
||||
audio_transcript = audio_item.get("transcript")
|
||||
if audio_transcript:
|
||||
safe_set_attribute(
|
||||
span, f"{audio_attrs.AUDIO_TRANSCRIPT}.{i}", audio_transcript
|
||||
)
|
||||
safe_set_attribute(span, f"{audio_attrs.AUDIO_TRANSCRIPT}.{i}", audio_transcript)
|
||||
|
||||
|
||||
def _set_embedding_outputs(span: "Span", response_obj, embedding_attrs, span_attrs):
|
||||
|
|
@ -220,9 +216,7 @@ def _set_structured_outputs(span: "Span", response_obj, msg_attrs, span_attrs):
|
|||
message_content = getattr(first_content, "text", "")
|
||||
message_role = getattr(item, "role", "assistant")
|
||||
safe_set_attribute(span, span_attrs.OUTPUT_VALUE, message_content)
|
||||
safe_set_attribute(
|
||||
span, f"{prefix}.{msg_attrs.MESSAGE_CONTENT}", message_content
|
||||
)
|
||||
safe_set_attribute(span, f"{prefix}.{msg_attrs.MESSAGE_CONTENT}", message_content)
|
||||
safe_set_attribute(span, f"{prefix}.{msg_attrs.MESSAGE_ROLE}", message_role)
|
||||
|
||||
|
||||
|
|
@ -253,19 +247,11 @@ def _set_usage_outputs(span: "Span", response_obj, span_attrs):
|
|||
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"
|
||||
)
|
||||
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")
|
||||
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"
|
||||
)
|
||||
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")
|
||||
if prompt_tokens:
|
||||
safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_PROMPT, prompt_tokens)
|
||||
|
||||
|
|
@ -273,9 +259,7 @@ 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"
|
||||
)
|
||||
token_details = _safe_get(usage, "completion_tokens_details") or _safe_get(usage, "output_tokens_details")
|
||||
reasoning_tokens = _safe_get(token_details, "reasoning_tokens")
|
||||
if reasoning_tokens:
|
||||
safe_set_attribute(
|
||||
|
|
@ -291,12 +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 = _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")
|
||||
if cache_read:
|
||||
safe_set_attribute(
|
||||
span,
|
||||
|
|
@ -374,33 +354,24 @@ def _infer_open_inference_span_kind(call_type: Optional[str]) -> str:
|
|||
):
|
||||
return OpenInferenceSpanKindValues.LLM.value
|
||||
|
||||
if any(
|
||||
keyword in lowered
|
||||
for keyword in ("file", "batch", "container", "fine_tuning_job")
|
||||
):
|
||||
if any(keyword in lowered for keyword in ("file", "batch", "container", "fine_tuning_job")):
|
||||
return OpenInferenceSpanKindValues.CHAIN.value
|
||||
|
||||
return OpenInferenceSpanKindValues.UNKNOWN.value
|
||||
|
||||
|
||||
def _set_tool_attributes(
|
||||
span: "Span", optional_tools: Optional[list], metadata_tools: Optional[list]
|
||||
):
|
||||
def _set_tool_attributes(span: "Span", optional_tools: Optional[list], metadata_tools: Optional[list]):
|
||||
"""set tool attributes on span from optional_params or tool call metadata"""
|
||||
if optional_tools:
|
||||
for idx, tool in enumerate(optional_tools):
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
function = (
|
||||
tool.get("function") if isinstance(tool.get("function"), dict) else None
|
||||
)
|
||||
function = tool.get("function") if isinstance(tool.get("function"), dict) else None
|
||||
if not function:
|
||||
continue
|
||||
tool_name = function.get("name")
|
||||
if tool_name:
|
||||
safe_set_attribute(
|
||||
span, f"{SpanAttributes.LLM_TOOLS}.{idx}.name", tool_name
|
||||
)
|
||||
safe_set_attribute(span, f"{SpanAttributes.LLM_TOOLS}.{idx}.name", tool_name)
|
||||
tool_description = function.get("description")
|
||||
if tool_description:
|
||||
safe_set_attribute(
|
||||
|
|
@ -437,9 +408,7 @@ def _set_tool_attributes(
|
|||
)
|
||||
|
||||
|
||||
def set_attributes(
|
||||
span: "Span", kwargs, response_obj, attributes: Type[BaseLLMObsOTELAttributes]
|
||||
):
|
||||
def set_attributes(span: "Span", kwargs, response_obj, attributes: Type[BaseLLMObsOTELAttributes]):
|
||||
"""
|
||||
Populates span with OpenInference-compliant LLM attributes for Arize and Phoenix tracing.
|
||||
"""
|
||||
|
|
@ -458,17 +427,11 @@ def set_attributes(
|
|||
try:
|
||||
optional_params = _sanitize_optional_params(kwargs.get("optional_params"))
|
||||
litellm_params = kwargs.get("litellm_params", {}) or {}
|
||||
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object"
|
||||
)
|
||||
standard_logging_payload: Optional[StandardLoggingPayload] = 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 = standard_logging_payload.get("metadata") if standard_logging_payload else None
|
||||
_set_metadata_attributes(span, metadata, SpanAttributes)
|
||||
|
||||
metadata_tools = _extract_metadata_tools(metadata)
|
||||
|
|
@ -492,19 +455,13 @@ def set_attributes(
|
|||
_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 = 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)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"[Arize/Phoenix] Failed to set OpenInference span attributes: {e}"
|
||||
)
|
||||
verbose_logger.error(f"[Arize/Phoenix] Failed to set OpenInference span attributes: {e}")
|
||||
if hasattr(span, "record_exception"):
|
||||
span.record_exception(e)
|
||||
|
||||
|
|
@ -562,9 +519,7 @@ def _set_request_attributes(
|
|||
if kwargs.get("model"):
|
||||
safe_set_attribute(span, span_attrs.LLM_MODEL_NAME, kwargs.get("model"))
|
||||
|
||||
safe_set_attribute(
|
||||
span, "llm.request.type", standard_logging_payload.get("call_type")
|
||||
)
|
||||
safe_set_attribute(span, "llm.request.type", standard_logging_payload.get("call_type"))
|
||||
safe_set_attribute(
|
||||
span,
|
||||
span_attrs.LLM_PROVIDER,
|
||||
|
|
@ -572,19 +527,13 @@ def _set_request_attributes(
|
|||
)
|
||||
|
||||
if optional_params.get("max_tokens"):
|
||||
safe_set_attribute(
|
||||
span, "llm.request.max_tokens", optional_params.get("max_tokens")
|
||||
)
|
||||
safe_set_attribute(span, "llm.request.max_tokens", optional_params.get("max_tokens"))
|
||||
if optional_params.get("temperature"):
|
||||
safe_set_attribute(
|
||||
span, "llm.request.temperature", optional_params.get("temperature")
|
||||
)
|
||||
safe_set_attribute(span, "llm.request.temperature", optional_params.get("temperature"))
|
||||
if optional_params.get("top_p"):
|
||||
safe_set_attribute(span, "llm.request.top_p", optional_params.get("top_p"))
|
||||
|
||||
safe_set_attribute(
|
||||
span, "llm.is_streaming", str(optional_params.get("stream", False))
|
||||
)
|
||||
safe_set_attribute(span, "llm.is_streaming", str(optional_params.get("stream", False)))
|
||||
|
||||
if optional_params.get("user"):
|
||||
safe_set_attribute(span, "llm.user", optional_params.get("user"))
|
||||
|
|
@ -599,9 +548,7 @@ def _set_model_params(span: "Span", model_params: Optional[dict], span_attrs) ->
|
|||
if not model_params:
|
||||
return
|
||||
|
||||
safe_set_attribute(
|
||||
span, span_attrs.LLM_INVOCATION_PARAMETERS, safe_dumps(model_params)
|
||||
)
|
||||
safe_set_attribute(span, span_attrs.LLM_INVOCATION_PARAMETERS, safe_dumps(model_params))
|
||||
if model_params.get("user"):
|
||||
user_id = model_params.get("user")
|
||||
if user_id is not None:
|
||||
|
|
@ -767,9 +714,7 @@ def _emit_message_tool_calls(span: "Span", prefix: str, message) -> None:
|
|||
continue
|
||||
tc_prefix = f"{prefix}.{MessageAttributes.MESSAGE_TOOL_CALLS}.{tc_idx}"
|
||||
if tc["id"]:
|
||||
safe_set_attribute(
|
||||
span, f"{tc_prefix}.{ToolCallAttributes.TOOL_CALL_ID}", tc["id"]
|
||||
)
|
||||
safe_set_attribute(span, f"{tc_prefix}.{ToolCallAttributes.TOOL_CALL_ID}", tc["id"])
|
||||
fn = tc["function"]
|
||||
if fn["name"]:
|
||||
safe_set_attribute(
|
||||
|
|
@ -862,9 +807,7 @@ def _emit_input_message_extras(span: "Span", prefix: str, message: dict) -> None
|
|||
)
|
||||
|
||||
|
||||
def _set_session_and_user_attrs(
|
||||
span: "Span", kwargs: dict, standard_logging_payload
|
||||
) -> None:
|
||||
def _set_session_and_user_attrs(span: "Span", kwargs: dict, standard_logging_payload) -> None:
|
||||
"""Emit `SESSION_ID` / `USER_ID` / team metadata when source data exists.
|
||||
|
||||
`SESSION_ID` is emitted only when an explicit end-user identifier exists
|
||||
|
|
@ -970,11 +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 = standard_logging_payload.get("call_type") if isinstance(standard_logging_payload, dict) else None
|
||||
if not _is_passthrough_call_type(call_type):
|
||||
return
|
||||
|
||||
|
|
@ -989,18 +928,12 @@ def _maybe_normalize_passthrough(
|
|||
|
||||
# --- INPUT --------------------------------------------------------------
|
||||
additional_args = kwargs.get("additional_args") or {}
|
||||
complete_input_dict = (
|
||||
additional_args.get("complete_input_dict")
|
||||
if isinstance(additional_args, dict)
|
||||
else None
|
||||
)
|
||||
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 = _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs)
|
||||
if not isinstance(parsed_response, dict):
|
||||
return
|
||||
|
||||
|
|
@ -1094,19 +1027,12 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs):
|
|||
candidates = []
|
||||
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
|
||||
):
|
||||
if isinstance(raw_response_obj, dict) and raw_response_obj is not coerced_response_obj:
|
||||
candidates.append(raw_response_obj)
|
||||
|
||||
for candidate in candidates:
|
||||
# StandardPassThroughResponseObject wrapper: {"response": "..."}.
|
||||
if (
|
||||
"response" in candidate
|
||||
and "content" not in candidate
|
||||
and "choices" not in candidate
|
||||
):
|
||||
if "response" in candidate and "content" not in candidate and "choices" not in candidate:
|
||||
inner = candidate.get("response")
|
||||
if isinstance(inner, str):
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -195,20 +195,14 @@ class ArizeLogger(OpenTelemetry):
|
|||
# the suggested param is `arize_space_key`
|
||||
#########################################################
|
||||
if standard_callback_dynamic_params.get("arize_space_id"):
|
||||
dynamic_headers["arize-space-id"] = standard_callback_dynamic_params.get(
|
||||
"arize_space_id"
|
||||
)
|
||||
dynamic_headers["arize-space-id"] = standard_callback_dynamic_params.get("arize_space_id")
|
||||
if standard_callback_dynamic_params.get("arize_space_key"):
|
||||
dynamic_headers["arize-space-id"] = standard_callback_dynamic_params.get(
|
||||
"arize_space_key"
|
||||
)
|
||||
dynamic_headers["arize-space-id"] = standard_callback_dynamic_params.get("arize_space_key")
|
||||
|
||||
#########################################################
|
||||
# `api_key` handling
|
||||
#########################################################
|
||||
if standard_callback_dynamic_params.get("arize_api_key"):
|
||||
dynamic_headers["api_key"] = standard_callback_dynamic_params.get(
|
||||
"arize_api_key"
|
||||
)
|
||||
dynamic_headers["api_key"] = standard_callback_dynamic_params.get("arize_api_key")
|
||||
|
||||
return dynamic_headers
|
||||
|
|
|
|||
|
|
@ -118,9 +118,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
try:
|
||||
provider.force_flush()
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"ArizePhoenixLogger: TracerProvider force_flush failed: %s", e
|
||||
)
|
||||
verbose_logger.debug("ArizePhoenixLogger: TracerProvider force_flush failed: %s", e)
|
||||
|
||||
def _get_litellm_resource_for_project(self, project_name: str):
|
||||
"""
|
||||
|
|
@ -149,9 +147,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
"""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 = TracerProvider(resource=self._get_litellm_resource_for_project(project_name))
|
||||
provider.add_span_processor(self._shared_span_processor)
|
||||
return provider
|
||||
|
||||
|
|
@ -163,9 +159,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
with self._project_providers_lock:
|
||||
if project_name in self._project_providers:
|
||||
self._project_providers.move_to_end(project_name)
|
||||
return self._project_providers[project_name].get_tracer(
|
||||
LITELLM_TRACER_NAME
|
||||
)
|
||||
return self._project_providers[project_name].get_tracer(LITELLM_TRACER_NAME)
|
||||
|
||||
# OTELResourceDetector().detect() is synchronous; build outside the lock so
|
||||
# concurrent requests for other projects are not blocked on cache misses.
|
||||
|
|
@ -174,9 +168,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
with self._project_providers_lock:
|
||||
if project_name in self._project_providers:
|
||||
self._project_providers.move_to_end(project_name)
|
||||
return self._project_providers[project_name].get_tracer(
|
||||
LITELLM_TRACER_NAME
|
||||
)
|
||||
return self._project_providers[project_name].get_tracer(LITELLM_TRACER_NAME)
|
||||
|
||||
if len(self._project_providers) >= _MAX_PROJECT_PROVIDERS:
|
||||
self._project_providers.popitem(last=False)
|
||||
|
|
@ -241,14 +233,10 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
detection to route their telemetry into arbitrary Arize/Phoenix projects.
|
||||
"""
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
return isinstance(litellm_params, dict) and bool(
|
||||
litellm_params.get("proxy_server_request")
|
||||
)
|
||||
return isinstance(litellm_params, dict) and bool(litellm_params.get("proxy_server_request"))
|
||||
|
||||
@staticmethod
|
||||
def _project_from_metadata_dict(
|
||||
metadata: dict, metadata_key: str, *, proxy_mode: bool
|
||||
) -> Optional[str]:
|
||||
def _project_from_metadata_dict(metadata: dict, metadata_key: str, *, proxy_mode: bool) -> Optional[str]:
|
||||
"""
|
||||
Read a Phoenix project field from proxy/SDK metadata.
|
||||
|
||||
|
|
@ -258,25 +246,19 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
"""
|
||||
auth_metadata = metadata.get("user_api_key_auth_metadata")
|
||||
if isinstance(auth_metadata, dict):
|
||||
project = ArizePhoenixLogger._normalize_project_name(
|
||||
auth_metadata.get(metadata_key)
|
||||
)
|
||||
project = ArizePhoenixLogger._normalize_project_name(auth_metadata.get(metadata_key))
|
||||
if project:
|
||||
return project
|
||||
|
||||
if not proxy_mode:
|
||||
return ArizePhoenixLogger._normalize_project_name(
|
||||
metadata.get(metadata_key)
|
||||
)
|
||||
return ArizePhoenixLogger._normalize_project_name(metadata.get(metadata_key))
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _metadata_project_from_kwargs(kwargs: dict, metadata_key: str) -> Optional[str]:
|
||||
proxy_mode = 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
|
||||
)
|
||||
project = ArizePhoenixLogger._project_from_metadata_dict(metadata, metadata_key, proxy_mode=proxy_mode)
|
||||
if project:
|
||||
return project
|
||||
return None
|
||||
|
|
@ -290,21 +272,16 @@ 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 = 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 = ArizePhoenixLogger._metadata_project_from_kwargs(kwargs, "phoenix_project_name")
|
||||
if phoenix_name:
|
||||
return phoenix_name
|
||||
|
||||
env_name = ArizePhoenixLogger._normalize_project_name(
|
||||
os.environ.get("PHOENIX_PROJECT_NAME")
|
||||
or os.environ.get("ARIZE_PROJECT_NAME")
|
||||
os.environ.get("PHOENIX_PROJECT_NAME") or os.environ.get("ARIZE_PROJECT_NAME")
|
||||
)
|
||||
if env_name:
|
||||
return env_name
|
||||
|
|
@ -335,11 +312,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
proxy_server_request = litellm_params.get("proxy_server_request", {}) or {}
|
||||
headers = proxy_server_request.get("headers", {}) or {}
|
||||
|
||||
traceparent_ctx = (
|
||||
self.get_traceparent_from_header(headers=headers)
|
||||
if headers.get("traceparent")
|
||||
else None
|
||||
)
|
||||
traceparent_ctx = self.get_traceparent_from_header(headers=headers) if headers.get("traceparent") else None
|
||||
|
||||
is_proxy_mode = bool(proxy_server_request)
|
||||
|
||||
|
|
@ -347,9 +320,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
start_time_val = kwargs.get("start_time", kwargs.get("api_call_start_time"))
|
||||
parent_span = tracer.start_span(
|
||||
name="litellm_proxy_request",
|
||||
start_time=(
|
||||
self._to_ns(start_time_val) if start_time_val is not None else None
|
||||
),
|
||||
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,
|
||||
)
|
||||
|
|
@ -359,14 +330,10 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
return traceparent_ctx, None
|
||||
|
||||
def _handle_success(self, kwargs, response_obj, start_time, end_time):
|
||||
self._handle_phoenix_trace(
|
||||
kwargs, response_obj, start_time, end_time, success=True
|
||||
)
|
||||
self._handle_phoenix_trace(kwargs, response_obj, start_time, end_time, success=True)
|
||||
|
||||
def _handle_failure(self, kwargs, response_obj, start_time, end_time):
|
||||
self._handle_phoenix_trace(
|
||||
kwargs, response_obj, start_time, end_time, success=False
|
||||
)
|
||||
self._handle_phoenix_trace(kwargs, response_obj, start_time, end_time, success=False)
|
||||
|
||||
def _handle_phoenix_trace(
|
||||
self,
|
||||
|
|
@ -402,9 +369,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
self._record_exception_on_span(span=span, kwargs=kwargs)
|
||||
|
||||
if success:
|
||||
self._maybe_log_raw_request(
|
||||
kwargs, response_obj, start_time, end_time, span
|
||||
)
|
||||
self._maybe_log_raw_request(kwargs, response_obj, start_time, end_time, span)
|
||||
span.end(end_time=self._to_ns(end_time))
|
||||
|
||||
self._create_guardrail_span(kwargs=kwargs, context=ctx)
|
||||
|
|
@ -471,9 +436,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
if api_key is not None:
|
||||
otlp_auth_headers = f"Authorization=Bearer {api_key}"
|
||||
elif "app.phoenix.arize.com" in endpoint:
|
||||
raise ValueError(
|
||||
"PHOENIX_API_KEY must be set when using Phoenix Cloud (app.phoenix.arize.com)."
|
||||
)
|
||||
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"
|
||||
|
||||
|
|
|
|||
|
|
@ -11,9 +11,7 @@ from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
|||
def _sanitize_id(identifier: str) -> str:
|
||||
"""Reject path traversal characters and URL-encode the identifier."""
|
||||
if any(c in identifier for c in ("/", "\\", "#", "?")):
|
||||
raise ValueError(
|
||||
f"Invalid identifier {identifier!r}: contains disallowed characters"
|
||||
)
|
||||
raise ValueError(f"Invalid identifier {identifier!r}: contains disallowed characters")
|
||||
if ".." in identifier:
|
||||
raise ValueError(f"Invalid identifier {identifier!r}: path traversal detected")
|
||||
return urllib.parse.quote(identifier, safe="")
|
||||
|
|
@ -87,17 +85,11 @@ class ArizePhoenixClient:
|
|||
f"Access denied to prompt version '{prompt_version_id}'. Check your Arize Phoenix permissions."
|
||||
)
|
||||
elif response.status_code == 401:
|
||||
raise Exception(
|
||||
"Authentication failed. Check your Arize Phoenix API key and permissions."
|
||||
)
|
||||
raise Exception("Authentication failed. Check your Arize Phoenix API key and permissions.")
|
||||
else:
|
||||
raise Exception(
|
||||
f"Failed to fetch prompt version '{prompt_version_id}': {e}"
|
||||
)
|
||||
raise Exception(f"Failed to fetch prompt version '{prompt_version_id}': {e}")
|
||||
else:
|
||||
raise Exception(
|
||||
f"Error fetching prompt version '{prompt_version_id}': {e}"
|
||||
)
|
||||
raise Exception(f"Error fetching prompt version '{prompt_version_id}': {e}")
|
||||
|
||||
def test_connection(self) -> bool:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -44,9 +44,7 @@ class ArizePhoenixPromptTemplate:
|
|||
self.template_format = metadata.get("template_format", "MUSTACHE")
|
||||
|
||||
def __repr__(self):
|
||||
return (
|
||||
f"ArizePhoenixPromptTemplate(id='{self.template_id}', model='{self.model}')"
|
||||
)
|
||||
return f"ArizePhoenixPromptTemplate(id='{self.template_id}', model='{self.model}')"
|
||||
|
||||
|
||||
class ArizePhoenixTemplateManager:
|
||||
|
|
@ -71,9 +69,7 @@ class ArizePhoenixTemplateManager:
|
|||
self.api_base = api_base
|
||||
self.prompt_id = prompt_id
|
||||
self.prompts: Dict[str, ArizePhoenixPromptTemplate] = {}
|
||||
self.arize_client = ArizePhoenixClient(
|
||||
api_key=self.api_key, api_base=self.api_base
|
||||
)
|
||||
self.arize_client = ArizePhoenixClient(api_key=self.api_key, api_base=self.api_base)
|
||||
|
||||
# Templates fetched from Arize Phoenix come from external workspace
|
||||
# users; in a plain `Environment()` a malicious template could reach
|
||||
|
|
@ -109,13 +105,9 @@ class ArizePhoenixTemplateManager:
|
|||
else:
|
||||
raise ValueError(f"Prompt version '{prompt_version_id}' not found")
|
||||
except Exception as e:
|
||||
raise Exception(
|
||||
f"Failed to load prompt version '{prompt_version_id}' from Arize Phoenix: {e}"
|
||||
)
|
||||
raise Exception(f"Failed to load prompt version '{prompt_version_id}' from Arize Phoenix: {e}")
|
||||
|
||||
def _parse_prompt_data(
|
||||
self, data: Dict[str, Any], prompt_version_id: str
|
||||
) -> ArizePhoenixPromptTemplate:
|
||||
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", [])
|
||||
|
|
@ -154,9 +146,7 @@ class ArizePhoenixTemplateManager:
|
|||
metadata=metadata,
|
||||
)
|
||||
|
||||
def render_template(
|
||||
self, template_id: str, variables: Optional[Dict[str, Any]] = None
|
||||
) -> List[AllMessageValues]:
|
||||
def render_template(self, template_id: str, variables: Optional[Dict[str, Any]] = None) -> List[AllMessageValues]:
|
||||
"""Render a template with the given variables and return formatted messages."""
|
||||
if template_id not in self.prompts:
|
||||
raise ValueError(f"Template '{template_id}' not found")
|
||||
|
|
@ -272,9 +262,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
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 = self.prompt_manager.render_template(prompt_id, prompt_variables or {})
|
||||
|
||||
# Extract metadata
|
||||
metadata = {
|
||||
|
|
@ -317,9 +305,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
|
||||
try:
|
||||
# Get the rendered messages and metadata
|
||||
rendered_messages, prompt_metadata = self.get_prompt_template(
|
||||
prompt_id, prompt_variables
|
||||
)
|
||||
rendered_messages, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables)
|
||||
|
||||
# Merge rendered messages with existing messages
|
||||
if rendered_messages:
|
||||
|
|
@ -353,9 +339,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
# Log error but don't fail the call
|
||||
import litellm
|
||||
|
||||
litellm._logging.verbose_proxy_logger.error(
|
||||
f"Error in Arize Phoenix prompt pre_call_hook: {e}"
|
||||
)
|
||||
litellm._logging.verbose_proxy_logger.error(f"Error in Arize Phoenix prompt pre_call_hook: {e}")
|
||||
return messages, litellm_params
|
||||
|
||||
def get_available_prompts(self) -> List[str]:
|
||||
|
|
@ -408,9 +392,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
self.prompt_manager._load_prompt_from_arize(prompt_id)
|
||||
|
||||
# Get the rendered messages and metadata
|
||||
rendered_messages, prompt_metadata = self.get_prompt_template(
|
||||
prompt_id, prompt_variables
|
||||
)
|
||||
rendered_messages, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables)
|
||||
|
||||
# Extract model from metadata (if specified)
|
||||
template_model = prompt_metadata.get("model")
|
||||
|
|
|
|||
|
|
@ -12,10 +12,7 @@ class AthinaLogger:
|
|||
"athina-api-key": self.athina_api_key,
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
self.athina_logging_url = (
|
||||
os.getenv("ATHINA_BASE_URL", "https://log.athina.ai")
|
||||
+ "/api/v1/log/inference"
|
||||
)
|
||||
self.athina_logging_url = os.getenv("ATHINA_BASE_URL", "https://log.athina.ai") + "/api/v1/log/inference"
|
||||
self.additional_keys = [
|
||||
"environment",
|
||||
"prompt_slug",
|
||||
|
|
@ -42,9 +39,7 @@ class AthinaLogger:
|
|||
if "complete_streaming_response" in kwargs:
|
||||
# Log the completion response in streaming mode
|
||||
completion_response = kwargs["complete_streaming_response"]
|
||||
response_json = (
|
||||
completion_response.model_dump() if completion_response else {}
|
||||
)
|
||||
response_json = completion_response.model_dump() if completion_response else {}
|
||||
else:
|
||||
# Skip logging if the completion response is not available
|
||||
return
|
||||
|
|
@ -56,30 +51,19 @@ class AthinaLogger:
|
|||
"request": kwargs,
|
||||
"response": response_json,
|
||||
"prompt_tokens": response_json.get("usage", {}).get("prompt_tokens"),
|
||||
"completion_tokens": response_json.get("usage", {}).get(
|
||||
"completion_tokens"
|
||||
),
|
||||
"completion_tokens": response_json.get("usage", {}).get("completion_tokens"),
|
||||
"total_tokens": response_json.get("usage", {}).get("total_tokens"),
|
||||
}
|
||||
|
||||
if (
|
||||
type(end_time) is datetime.datetime
|
||||
and type(start_time) is datetime.datetime
|
||||
):
|
||||
data["response_time"] = int(
|
||||
(end_time - start_time).total_seconds() * 1000
|
||||
)
|
||||
if type(end_time) is datetime.datetime and type(start_time) is datetime.datetime:
|
||||
data["response_time"] = int((end_time - start_time).total_seconds() * 1000)
|
||||
|
||||
if "messages" in kwargs:
|
||||
data["prompt"] = kwargs.get("messages", None)
|
||||
|
||||
# Directly add tools or functions if present
|
||||
optional_params = kwargs.get("optional_params", {})
|
||||
data.update(
|
||||
(k, v)
|
||||
for k, v in optional_params.items()
|
||||
if k in ["tools", "functions"]
|
||||
)
|
||||
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", {})
|
||||
|
|
@ -93,13 +77,9 @@ class AthinaLogger:
|
|||
data=json.dumps(data, default=str),
|
||||
)
|
||||
if response.status_code != 200:
|
||||
print_verbose(
|
||||
f"Athina Logger Error - {response.text}, {response.status_code}"
|
||||
)
|
||||
print_verbose(f"Athina Logger Error - {response.text}, {response.status_code}")
|
||||
else:
|
||||
print_verbose(f"Athina Logger Succeeded - {response.text}")
|
||||
except Exception as e:
|
||||
print_verbose(
|
||||
f"Athina Logger Error - {e}, Stack trace: {traceback.format_exc()}"
|
||||
)
|
||||
print_verbose(f"Athina Logger Error - {e}, Stack trace: {traceback.format_exc()}")
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -63,32 +63,16 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
audit_stream_name (str, optional): Stream name from DCR for audit logs.
|
||||
If not provided, audit logs use the standard stream name.
|
||||
"""
|
||||
self.async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.LoggingCallback
|
||||
)
|
||||
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_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 = 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_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 = (
|
||||
client_secret
|
||||
or os.getenv("AZURE_SENTINEL_CLIENT_SECRET")
|
||||
or os.getenv("AZURE_CLIENT_SECRET")
|
||||
client_secret or os.getenv("AZURE_SENTINEL_CLIENT_SECRET") or os.getenv("AZURE_CLIENT_SECRET")
|
||||
)
|
||||
|
||||
if not resolved_dcr_immutable_id:
|
||||
|
|
@ -144,9 +128,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
self.audit_log_queue: List[StandardAuditLogPayload] = []
|
||||
|
||||
@staticmethod
|
||||
def _build_api_endpoint(
|
||||
endpoint: str, dcr_immutable_id: str, stream_name: str
|
||||
) -> str:
|
||||
def _build_api_endpoint(endpoint: str, dcr_immutable_id: str, stream_name: str) -> str:
|
||||
return f"{endpoint.rstrip('/')}/dataCollectionRules/{dcr_immutable_id}/streams/{stream_name}?api-version=2023-01-01"
|
||||
|
||||
async def _get_oauth_token(self) -> str:
|
||||
|
|
@ -157,9 +139,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
Bearer token string
|
||||
"""
|
||||
if (
|
||||
self.oauth_token
|
||||
and self.oauth_token_expires_at
|
||||
and time.time() < self.oauth_token_expires_at - 60
|
||||
self.oauth_token and self.oauth_token_expires_at and time.time() < self.oauth_token_expires_at - 60
|
||||
): # Refresh 60 seconds before expiry
|
||||
return self.oauth_token
|
||||
|
||||
|
|
@ -168,9 +148,7 @@ 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 = f"https://login.microsoftonline.com/{self.tenant_id}/oauth2/v2.0/token"
|
||||
|
||||
token_data = {
|
||||
"client_id": self.client_id,
|
||||
|
|
@ -186,9 +164,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise Exception(
|
||||
f"Failed to get OAuth2 token: {response.status_code} - {response.text}"
|
||||
)
|
||||
raise Exception(f"Failed to get OAuth2 token: {response.status_code} - {response.text}")
|
||||
|
||||
token_response = response.json()
|
||||
self.oauth_token = token_response.get("access_token")
|
||||
|
|
@ -213,15 +189,11 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
Raises a NON Blocking verbose_logger.exception if an error occurs
|
||||
"""
|
||||
try:
|
||||
verbose_logger.debug(
|
||||
"Azure Sentinel: Logging - Enters logging function for model %s", kwargs
|
||||
)
|
||||
verbose_logger.debug("Azure Sentinel: Logging - Enters logging function for model %s", kwargs)
|
||||
standard_logging_payload = kwargs.get("standard_logging_object", None)
|
||||
|
||||
if standard_logging_payload is None:
|
||||
verbose_logger.warning(
|
||||
"Azure Sentinel: standard_logging_object not found in kwargs"
|
||||
)
|
||||
verbose_logger.warning("Azure Sentinel: standard_logging_object not found in kwargs")
|
||||
return
|
||||
|
||||
self.log_queue.append(standard_logging_payload)
|
||||
|
|
@ -230,9 +202,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
await self.async_send_batch()
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"Azure Sentinel Layer Error - {str(e)}\n{traceback.format_exc()}"
|
||||
)
|
||||
verbose_logger.exception(f"Azure Sentinel Layer Error - {str(e)}\n{traceback.format_exc()}")
|
||||
pass
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
|
|
@ -254,9 +224,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
standard_logging_payload = kwargs.get("standard_logging_object", None)
|
||||
|
||||
if standard_logging_payload is None:
|
||||
verbose_logger.warning(
|
||||
"Azure Sentinel: standard_logging_object not found in kwargs"
|
||||
)
|
||||
verbose_logger.warning("Azure Sentinel: standard_logging_object not found in kwargs")
|
||||
return
|
||||
|
||||
self.log_queue.append(standard_logging_payload)
|
||||
|
|
@ -265,14 +233,10 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
await self.async_send_batch()
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"Azure Sentinel Layer Error - {str(e)}\n{traceback.format_exc()}"
|
||||
)
|
||||
verbose_logger.exception(f"Azure Sentinel Layer Error - {str(e)}\n{traceback.format_exc()}")
|
||||
pass
|
||||
|
||||
async def async_log_audit_log_event(
|
||||
self, audit_log: StandardAuditLogPayload
|
||||
) -> None:
|
||||
async def async_log_audit_log_event(self, audit_log: StandardAuditLogPayload) -> None:
|
||||
"""
|
||||
Async log LiteLLM audit log events to Azure Sentinel.
|
||||
|
||||
|
|
@ -293,9 +257,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
await self.async_send_audit_batch()
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"Azure Sentinel Audit Log Layer Error - {str(e)}\n{traceback.format_exc()}"
|
||||
)
|
||||
verbose_logger.exception(f"Azure Sentinel Audit Log Layer Error - {str(e)}\n{traceback.format_exc()}")
|
||||
pass
|
||||
|
||||
async def async_send_batch(self):
|
||||
|
|
@ -331,9 +293,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
if not log_queue:
|
||||
return
|
||||
|
||||
verbose_logger.debug(
|
||||
"Azure Sentinel - about to flush %s %s", len(log_queue), log_type
|
||||
)
|
||||
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()
|
||||
|
|
@ -349,9 +309,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
}
|
||||
|
||||
# Send the request
|
||||
response = await self.async_httpx_client.post(
|
||||
url=api_endpoint, data=body.encode("utf-8"), headers=headers
|
||||
)
|
||||
response = await self.async_httpx_client.post(url=api_endpoint, data=body.encode("utf-8"), headers=headers)
|
||||
|
||||
if response.status_code not in [200, 204]:
|
||||
verbose_logger.error(
|
||||
|
|
@ -359,9 +317,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
response.status_code,
|
||||
response.text,
|
||||
)
|
||||
raise Exception(
|
||||
f"Failed to send logs to Azure Sentinel: {response.status_code} - {response.text}"
|
||||
)
|
||||
raise Exception(f"Failed to send logs to Azure Sentinel: {response.status_code} - {response.text}")
|
||||
|
||||
verbose_logger.debug(
|
||||
"Azure Sentinel: Response from API status_code: %s",
|
||||
|
|
@ -369,9 +325,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"Azure Sentinel Error sending batch API - {str(e)}\n{traceback.format_exc()}"
|
||||
)
|
||||
verbose_logger.exception(f"Azure Sentinel Error sending batch API - {str(e)}\n{traceback.format_exc()}")
|
||||
finally:
|
||||
log_queue.clear()
|
||||
|
||||
|
|
|
|||
|
|
@ -24,42 +24,30 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
**kwargs,
|
||||
):
|
||||
try:
|
||||
verbose_logger.debug(
|
||||
"AzureBlobStorageLogger: in init azure blob storage logger"
|
||||
)
|
||||
verbose_logger.debug("AzureBlobStorageLogger: in init azure blob storage logger")
|
||||
|
||||
# Env Variables used for Azure Storage Authentication
|
||||
self.tenant_id = os.getenv("AZURE_STORAGE_TENANT_ID")
|
||||
self.client_id = os.getenv("AZURE_STORAGE_CLIENT_ID")
|
||||
self.client_secret = os.getenv("AZURE_STORAGE_CLIENT_SECRET")
|
||||
self.azure_storage_account_key: Optional[str] = os.getenv(
|
||||
"AZURE_STORAGE_ACCOUNT_KEY"
|
||||
)
|
||||
self.azure_storage_account_key: Optional[str] = os.getenv("AZURE_STORAGE_ACCOUNT_KEY")
|
||||
|
||||
# Required Env Variables for Azure Storage
|
||||
_azure_storage_account_name = os.getenv("AZURE_STORAGE_ACCOUNT_NAME")
|
||||
if not _azure_storage_account_name:
|
||||
raise ValueError(
|
||||
"Missing required environment variable: 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")
|
||||
if not _azure_storage_file_system:
|
||||
raise ValueError(
|
||||
"Missing required environment variable: AZURE_STORAGE_FILE_SYSTEM"
|
||||
)
|
||||
raise ValueError("Missing required environment variable: AZURE_STORAGE_FILE_SYSTEM")
|
||||
self.azure_storage_file_system: str = _azure_storage_file_system
|
||||
self._service_client = None
|
||||
# Time that the azure service client expires, in order to reset the connection pool and keep it fresh
|
||||
self._service_client_timeout: Optional[float] = None
|
||||
|
||||
# Internal variables used for Token based authentication
|
||||
self.azure_auth_token: Optional[str] = (
|
||||
None # the Azure AD token to use for Azure Storage API requests
|
||||
)
|
||||
self.token_expiry: Optional[datetime] = (
|
||||
None # the expiry time of the currentAzure AD token
|
||||
)
|
||||
self.azure_auth_token: Optional[str] = None # the Azure AD token to use for Azure Storage API requests
|
||||
self.token_expiry: Optional[datetime] = None # the expiry time of the currentAzure AD token
|
||||
|
||||
asyncio.create_task(self.periodic_flush())
|
||||
self.flush_lock = asyncio.Lock()
|
||||
|
|
@ -84,9 +72,7 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
"AzureBlobStorageLogger: Logging - Enters logging function for model %s",
|
||||
kwargs,
|
||||
)
|
||||
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object"
|
||||
)
|
||||
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object")
|
||||
|
||||
if standard_logging_payload is None:
|
||||
raise ValueError("standard_logging_payload is not set")
|
||||
|
|
@ -110,9 +96,7 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
"AzureBlobStorageLogger: Logging - Enters logging function for model %s",
|
||||
kwargs,
|
||||
)
|
||||
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object"
|
||||
)
|
||||
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object")
|
||||
|
||||
if standard_logging_payload is None:
|
||||
raise ValueError("standard_logging_payload is not set")
|
||||
|
|
@ -143,13 +127,9 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
await self.async_upload_payload_to_azure_blob_storage(payload=payload)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"AzureBlobStorageLogger Error sending batch API - {str(e)}"
|
||||
)
|
||||
verbose_logger.exception(f"AzureBlobStorageLogger Error sending batch API - {str(e)}")
|
||||
|
||||
async def async_upload_payload_to_azure_blob_storage(
|
||||
self, payload: StandardLoggingPayload
|
||||
):
|
||||
async def async_upload_payload_to_azure_blob_storage(self, payload: StandardLoggingPayload):
|
||||
"""
|
||||
Uploads the payload to Azure Blob Storage using a 3-step process:
|
||||
1. Create file resource
|
||||
|
|
@ -158,18 +138,12 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
"""
|
||||
try:
|
||||
if self.azure_storage_account_key:
|
||||
await self.upload_to_azure_data_lake_with_azure_account_key(
|
||||
payload=payload
|
||||
)
|
||||
await self.upload_to_azure_data_lake_with_azure_account_key(payload=payload)
|
||||
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
|
||||
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"
|
||||
base_url = f"https://{self.azure_storage_account_name}.dfs.core.windows.net/{self.azure_storage_file_system}/{filename}"
|
||||
|
|
@ -179,9 +153,7 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
await self._append_data(async_client, base_url, json_payload)
|
||||
await self._flush_data(async_client, base_url, len(payload_bytes))
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Successfully uploaded log to Azure Blob Storage: {filename}"
|
||||
)
|
||||
verbose_logger.debug(f"Successfully uploaded log to Azure Blob Storage: {filename}")
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error uploading to Azure Blob Storage: {str(e)}")
|
||||
|
|
@ -203,9 +175,7 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
verbose_logger.exception(f"Error creating file resource: {str(e)}")
|
||||
raise
|
||||
|
||||
async def _append_data(
|
||||
self, client: AsyncHTTPHandler, base_url: str, json_payload: str
|
||||
):
|
||||
async def _append_data(self, client: AsyncHTTPHandler, base_url: str, json_payload: str):
|
||||
"""Helper method to append data to the file"""
|
||||
try:
|
||||
verbose_logger.debug(f"Appending data to file: {base_url}")
|
||||
|
|
@ -234,9 +204,7 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
"Content-Length": "0",
|
||||
"Authorization": f"Bearer {self.azure_auth_token}",
|
||||
}
|
||||
response = await client.patch(
|
||||
f"{base_url}?action=flush&position={position}", headers=headers
|
||||
)
|
||||
response = 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:
|
||||
|
|
@ -282,17 +250,11 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
client_secret is not None,
|
||||
)
|
||||
if tenant_id is None:
|
||||
raise ValueError(
|
||||
"Missing required environment variable: AZURE_STORAGE_TENANT_ID"
|
||||
)
|
||||
raise ValueError("Missing required environment variable: AZURE_STORAGE_TENANT_ID")
|
||||
if client_id is None:
|
||||
raise ValueError(
|
||||
"Missing required environment variable: AZURE_STORAGE_CLIENT_ID"
|
||||
)
|
||||
raise ValueError("Missing required environment variable: AZURE_STORAGE_CLIENT_ID")
|
||||
if client_secret is None:
|
||||
raise ValueError(
|
||||
"Missing required environment variable: AZURE_STORAGE_CLIENT_SECRET"
|
||||
)
|
||||
raise ValueError("Missing required environment variable: AZURE_STORAGE_CLIENT_SECRET")
|
||||
|
||||
token_provider = get_azure_ad_token_from_entra_id(
|
||||
tenant_id=tenant_id,
|
||||
|
|
@ -331,11 +293,7 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
from azure.storage.filedatalake.aio import DataLakeServiceClient
|
||||
|
||||
# expire old clients to recover from connection issues
|
||||
if (
|
||||
self._service_client_timeout
|
||||
and self._service_client
|
||||
and self._service_client_timeout > time.time()
|
||||
):
|
||||
if self._service_client_timeout and self._service_client and self._service_client_timeout > time.time():
|
||||
await self._service_client.close()
|
||||
self._service_client = None
|
||||
if not self._service_client:
|
||||
|
|
@ -346,9 +304,7 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
self._service_client_timeout = time.time() + _DEFAULT_TTL_FOR_HTTPX_CLIENTS
|
||||
return self._service_client
|
||||
|
||||
async def upload_to_azure_data_lake_with_azure_account_key(
|
||||
self, payload: StandardLoggingPayload
|
||||
):
|
||||
async def upload_to_azure_data_lake_with_azure_account_key(self, payload: StandardLoggingPayload):
|
||||
"""
|
||||
Uploads the payload to Azure Data Lake using the Azure SDK
|
||||
|
||||
|
|
@ -359,9 +315,7 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
|
||||
service_client = 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 = service_client.get_file_system_client(file_system=self.azure_storage_file_system)
|
||||
|
||||
try:
|
||||
# Create directory with today's date
|
||||
|
|
@ -391,9 +345,7 @@ class AzureBlobStorageLogger(CustomBatchLogger):
|
|||
# Flush the content to finalize the file
|
||||
await file_client.flush_data(position=len(content), offset=0)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Successfully uploaded and wrote to {today}/{file_name}"
|
||||
)
|
||||
verbose_logger.debug(f"Successfully uploaded and wrote to {today}/{file_name}")
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error occurred: {str(e)}")
|
||||
|
|
|
|||
|
|
@ -29,9 +29,7 @@ def set_global_bitbucket_config(config: dict) -> None:
|
|||
litellm.global_bitbucket_config = config # type: ignore
|
||||
|
||||
|
||||
def prompt_initializer(
|
||||
litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec"
|
||||
) -> "CustomPromptManagement":
|
||||
def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec") -> "CustomPromptManagement":
|
||||
"""
|
||||
Initialize a prompt from a BitBucket repository.
|
||||
"""
|
||||
|
|
@ -39,9 +37,7 @@ def prompt_initializer(
|
|||
prompt_id = getattr(litellm_params, "prompt_id", None)
|
||||
|
||||
if not bitbucket_config:
|
||||
raise ValueError(
|
||||
"bitbucket_config is required for BitBucket prompt integration"
|
||||
)
|
||||
raise ValueError("bitbucket_config is required for BitBucket prompt integration")
|
||||
|
||||
try:
|
||||
bitbucket_prompt_manager = BitBucketPromptManager(
|
||||
|
|
|
|||
|
|
@ -12,15 +12,11 @@ from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
|||
def _sanitize_file_path(file_path: str) -> str:
|
||||
"""Reject path traversal and URL-encode each path segment."""
|
||||
if "#" in file_path or "?" in file_path:
|
||||
raise ValueError(
|
||||
f"Invalid file path {file_path!r}: contains URL special characters"
|
||||
)
|
||||
raise ValueError(f"Invalid file path {file_path!r}: contains URL special characters")
|
||||
parts = file_path.split("/")
|
||||
for part in parts:
|
||||
if part == "..":
|
||||
raise ValueError(
|
||||
f"Invalid file path {file_path!r}: path traversal detected"
|
||||
)
|
||||
raise ValueError(f"Invalid file path {file_path!r}: path traversal detected")
|
||||
return "/".join(urllib.parse.quote(part, safe="") for part in parts)
|
||||
|
||||
|
||||
|
|
@ -115,17 +111,13 @@ class BitBucketClient:
|
|||
f"Access denied to file '{file_path}'. Check your BitBucket permissions for workspace '{self.workspace}' and repository '{self.repository}'."
|
||||
)
|
||||
elif e.response.status_code == 401:
|
||||
raise Exception(
|
||||
"Authentication failed. Check your BitBucket access token and permissions."
|
||||
)
|
||||
raise Exception("Authentication failed. Check your BitBucket access token and permissions.")
|
||||
else:
|
||||
raise Exception(f"Failed to fetch file '{file_path}': {e}")
|
||||
else:
|
||||
raise Exception(f"Error fetching file '{file_path}': {e}")
|
||||
|
||||
def list_files(
|
||||
self, directory_path: str = "", file_extension: str = ".prompt"
|
||||
) -> List[str]:
|
||||
def list_files(self, directory_path: str = "", file_extension: str = ".prompt") -> List[str]:
|
||||
"""
|
||||
List files in a directory with a specific extension.
|
||||
|
||||
|
|
@ -164,9 +156,7 @@ class BitBucketClient:
|
|||
f"Access denied to directory '{directory_path}'. Check your BitBucket permissions for workspace '{self.workspace}' and repository '{self.repository}'."
|
||||
)
|
||||
elif e.response.status_code == 401:
|
||||
raise Exception(
|
||||
"Authentication failed. Check your BitBucket access token and permissions."
|
||||
)
|
||||
raise Exception("Authentication failed. Check your BitBucket access token and permissions.")
|
||||
else:
|
||||
raise Exception(f"Failed to list files in '{directory_path}': {e}")
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -44,9 +44,7 @@ class BitBucketPromptTemplate:
|
|||
self.temperature = metadata.get("temperature")
|
||||
self.max_tokens = metadata.get("max_tokens")
|
||||
self.input_schema = metadata.get("input", {}).get("schema", {})
|
||||
self.optional_params = {
|
||||
k: v for k, v in metadata.items() if k not in ["model", "input", "content"]
|
||||
}
|
||||
self.optional_params = {k: v for k, v in metadata.items() if k not in ["model", "input", "content"]}
|
||||
|
||||
def __repr__(self):
|
||||
return f"BitBucketPromptTemplate(id='{self.template_id}', model='{self.model}')"
|
||||
|
|
@ -101,9 +99,7 @@ class BitBucketTemplateManager:
|
|||
"""Load a specific .prompt file from BitBucket."""
|
||||
try:
|
||||
# Fetch the .prompt file from BitBucket
|
||||
prompt_content = self.bitbucket_client.get_file_content(
|
||||
f"{prompt_id}.prompt"
|
||||
)
|
||||
prompt_content = self.bitbucket_client.get_file_content(f"{prompt_id}.prompt")
|
||||
|
||||
if prompt_content:
|
||||
template = self._parse_prompt_file(prompt_content, prompt_id)
|
||||
|
|
@ -111,9 +107,7 @@ class BitBucketTemplateManager:
|
|||
except Exception as e:
|
||||
raise Exception(f"Failed to load prompt '{prompt_id}' from BitBucket: {e}")
|
||||
|
||||
def _parse_prompt_file(
|
||||
self, content: str, prompt_id: str
|
||||
) -> BitBucketPromptTemplate:
|
||||
def _parse_prompt_file(self, content: str, prompt_id: str) -> BitBucketPromptTemplate:
|
||||
"""Parse a .prompt file content and extract metadata and template."""
|
||||
# Split frontmatter and content
|
||||
if content.startswith("---"):
|
||||
|
|
@ -168,9 +162,7 @@ class BitBucketTemplateManager:
|
|||
result[key] = value.strip("\"'")
|
||||
return result
|
||||
|
||||
def render_template(
|
||||
self, template_id: str, variables: Optional[Dict[str, Any]] = None
|
||||
) -> str:
|
||||
def render_template(self, template_id: str, variables: Optional[Dict[str, Any]] = None) -> str:
|
||||
"""Render a template with the given variables."""
|
||||
if template_id not in self.prompts:
|
||||
raise ValueError(f"Template '{template_id}' not found")
|
||||
|
|
@ -259,9 +251,7 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
raise ValueError(f"Prompt template '{prompt_id}' not found")
|
||||
|
||||
# Render the template
|
||||
rendered_prompt = self.prompt_manager.render_template(
|
||||
prompt_id, prompt_variables or {}
|
||||
)
|
||||
rendered_prompt = self.prompt_manager.render_template(prompt_id, prompt_variables or {})
|
||||
|
||||
# Extract metadata
|
||||
metadata = {
|
||||
|
|
@ -291,9 +281,7 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
|
||||
try:
|
||||
# Get the rendered prompt and metadata
|
||||
rendered_prompt, prompt_metadata = self.get_prompt_template(
|
||||
prompt_id, prompt_variables
|
||||
)
|
||||
rendered_prompt, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables)
|
||||
|
||||
# Parse the rendered prompt into messages
|
||||
parsed_messages = self._parse_prompt_to_messages(rendered_prompt)
|
||||
|
|
@ -332,9 +320,7 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
# Log error but don't fail the call
|
||||
import litellm
|
||||
|
||||
litellm._logging.verbose_proxy_logger.error(
|
||||
f"Error in BitBucket prompt pre_call_hook: {e}"
|
||||
)
|
||||
litellm._logging.verbose_proxy_logger.error(f"Error in BitBucket prompt pre_call_hook: {e}")
|
||||
return messages, litellm_params
|
||||
|
||||
def _parse_prompt_to_messages(self, prompt_content: str) -> List[AllMessageValues]:
|
||||
|
|
@ -389,9 +375,7 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
|
||||
# Add the last message
|
||||
if current_role and current_content:
|
||||
messages.append(
|
||||
{"role": current_role, "content": "\n".join(current_content).strip()}
|
||||
)
|
||||
messages.append({"role": current_role, "content": "\n".join(current_content).strip()})
|
||||
|
||||
# If no role indicators found, treat as a single user message
|
||||
if not messages and prompt_content.strip():
|
||||
|
|
@ -466,9 +450,7 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
self.prompt_manager._load_prompt_from_bitbucket(prompt_id)
|
||||
|
||||
# Get the rendered prompt and metadata
|
||||
rendered_prompt, prompt_metadata = self.get_prompt_template(
|
||||
prompt_id, prompt_variables
|
||||
)
|
||||
rendered_prompt, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables)
|
||||
|
||||
# Convert rendered content to chat messages
|
||||
messages = self._parse_prompt_to_messages(rendered_prompt)
|
||||
|
|
|
|||
|
|
@ -34,16 +34,12 @@ def get_utc_datetime():
|
|||
|
||||
|
||||
class BraintrustLogger(CustomLogger):
|
||||
def __init__(
|
||||
self, api_key: Optional[str] = None, api_base: Optional[str] = None
|
||||
) -> None:
|
||||
def __init__(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> None:
|
||||
super().__init__()
|
||||
self.is_mock_mode = should_use_braintrust_mock()
|
||||
if self.is_mock_mode:
|
||||
create_mock_braintrust_client()
|
||||
verbose_logger.info(
|
||||
"[BRAINTRUST MOCK] Braintrust logger initialized in mock mode"
|
||||
)
|
||||
verbose_logger.info("[BRAINTRUST MOCK] Braintrust logger initialized in mock mode")
|
||||
self.validate_environment(api_key=api_key)
|
||||
self.api_base = api_base or os.getenv("BRAINTRUST_API_BASE") or API_BASE
|
||||
self.default_project_id = None
|
||||
|
|
@ -52,12 +48,8 @@ class BraintrustLogger(CustomLogger):
|
|||
"Authorization": "Bearer " + self.api_key,
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
self._project_id_cache: Dict[
|
||||
str, str
|
||||
] = {} # Cache mapping project names to IDs
|
||||
self.global_braintrust_http_handler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.LoggingCallback
|
||||
)
|
||||
self._project_id_cache: Dict[str, str] = {} # Cache mapping project names to IDs
|
||||
self.global_braintrust_http_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
|
||||
self.global_braintrust_sync_http_handler = HTTPHandler()
|
||||
|
||||
def validate_environment(self, api_key: Optional[str]):
|
||||
|
|
@ -143,23 +135,16 @@ class BraintrustLogger(CustomLogger):
|
|||
output = None
|
||||
choices = []
|
||||
if response_obj is not None and (
|
||||
kwargs.get("call_type", None) == "embedding"
|
||||
or isinstance(response_obj, litellm.EmbeddingResponse)
|
||||
kwargs.get("call_type", None) == "embedding" or isinstance(response_obj, litellm.EmbeddingResponse)
|
||||
):
|
||||
output = None
|
||||
elif response_obj is not None and isinstance(
|
||||
response_obj, litellm.ModelResponse
|
||||
):
|
||||
elif response_obj is not None and isinstance(response_obj, litellm.ModelResponse):
|
||||
output = response_obj["choices"][0]["message"].json()
|
||||
choices = response_obj["choices"]
|
||||
elif response_obj is not None and isinstance(
|
||||
response_obj, litellm.TextCompletionResponse
|
||||
):
|
||||
elif response_obj is not None and isinstance(response_obj, litellm.TextCompletionResponse):
|
||||
output = response_obj.choices[0].text
|
||||
choices = response_obj.choices
|
||||
elif response_obj is not None and isinstance(
|
||||
response_obj, litellm.ImageResponse
|
||||
):
|
||||
elif response_obj is not None and isinstance(response_obj, litellm.ImageResponse):
|
||||
output = response_obj["data"]
|
||||
|
||||
litellm_params = kwargs.get("litellm_params", {}) or {}
|
||||
|
|
@ -169,9 +154,7 @@ class BraintrustLogger(CustomLogger):
|
|||
project_id = dynamic_metadata.get("project_id")
|
||||
if project_id is None:
|
||||
project_name = dynamic_metadata.get("project_name")
|
||||
project_id = (
|
||||
self.get_project_id_sync(project_name) if project_name else None
|
||||
)
|
||||
project_id = self.get_project_id_sync(project_name) if project_name else None
|
||||
|
||||
if project_id is None:
|
||||
if self.default_project_id is None:
|
||||
|
|
@ -206,8 +189,7 @@ class BraintrustLogger(CustomLogger):
|
|||
"completion_tokens": usage_obj.completion_tokens,
|
||||
"total_tokens": usage_obj.total_tokens,
|
||||
"total_cost": cost,
|
||||
"time_to_first_token": end_time.timestamp()
|
||||
- start_time.timestamp(),
|
||||
"time_to_first_token": end_time.timestamp() - start_time.timestamp(),
|
||||
"start": start_time.timestamp(),
|
||||
"end": end_time.timestamp(),
|
||||
}
|
||||
|
|
@ -278,23 +260,16 @@ class BraintrustLogger(CustomLogger):
|
|||
output = None
|
||||
choices = []
|
||||
if response_obj is not None and (
|
||||
kwargs.get("call_type", None) == "embedding"
|
||||
or isinstance(response_obj, litellm.EmbeddingResponse)
|
||||
kwargs.get("call_type", None) == "embedding" or isinstance(response_obj, litellm.EmbeddingResponse)
|
||||
):
|
||||
output = None
|
||||
elif response_obj is not None and isinstance(
|
||||
response_obj, litellm.ModelResponse
|
||||
):
|
||||
elif response_obj is not None and isinstance(response_obj, litellm.ModelResponse):
|
||||
output = response_obj["choices"][0]["message"].json()
|
||||
choices = response_obj["choices"]
|
||||
elif response_obj is not None and isinstance(
|
||||
response_obj, litellm.TextCompletionResponse
|
||||
):
|
||||
elif response_obj is not None and isinstance(response_obj, litellm.TextCompletionResponse):
|
||||
output = response_obj.choices[0].text
|
||||
choices = response_obj.choices
|
||||
elif response_obj is not None and isinstance(
|
||||
response_obj, litellm.ImageResponse
|
||||
):
|
||||
elif response_obj is not None and isinstance(response_obj, litellm.ImageResponse):
|
||||
output = response_obj["data"]
|
||||
|
||||
litellm_params = kwargs.get("litellm_params", {})
|
||||
|
|
@ -304,11 +279,7 @@ class BraintrustLogger(CustomLogger):
|
|||
project_id = dynamic_metadata.get("project_id")
|
||||
if project_id is None:
|
||||
project_name = dynamic_metadata.get("project_name")
|
||||
project_id = (
|
||||
await self.get_project_id_async(project_name)
|
||||
if project_name
|
||||
else None
|
||||
)
|
||||
project_id = await self.get_project_id_async(project_name) if project_name else None
|
||||
|
||||
if project_id is None:
|
||||
if self.default_project_id is None:
|
||||
|
|
@ -350,14 +321,8 @@ class BraintrustLogger(CustomLogger):
|
|||
api_call_start_time = kwargs.get("api_call_start_time")
|
||||
completion_start_time = kwargs.get("completion_start_time")
|
||||
|
||||
if (
|
||||
api_call_start_time is not None
|
||||
and completion_start_time is not None
|
||||
):
|
||||
metrics["time_to_first_token"] = (
|
||||
completion_start_time.timestamp()
|
||||
- api_call_start_time.timestamp()
|
||||
)
|
||||
if api_call_start_time is not None and completion_start_time is not None:
|
||||
metrics["time_to_first_token"] = completion_start_time.timestamp() - api_call_start_time.timestamp()
|
||||
|
||||
# Allow metadata override for span name
|
||||
span_name = dynamic_metadata.get("span_name", "Chat Completion")
|
||||
|
|
|
|||
|
|
@ -156,11 +156,7 @@ def create_mock_braintrust_client():
|
|||
# This is required for async calls to be mocked
|
||||
create_mock_braintrust_factory_client()
|
||||
|
||||
verbose_logger.debug(
|
||||
f"[BRAINTRUST MOCK] Mock latency set to {_MOCK_LATENCY_SECONDS * 1000:.0f}ms"
|
||||
)
|
||||
verbose_logger.debug(
|
||||
"[BRAINTRUST MOCK] Braintrust mock client initialization complete"
|
||||
)
|
||||
verbose_logger.debug(f"[BRAINTRUST MOCK] Mock latency set to {_MOCK_LATENCY_SECONDS * 1000:.0f}ms")
|
||||
verbose_logger.debug("[BRAINTRUST MOCK] Braintrust mock client initialization complete")
|
||||
|
||||
_mocks_initialized = True
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue