fix(coverage): revert ruff UP006/UP045 changes on upstream files

The previous ruff fixes (Dict→dict, Optional→X|None) on 7 upstream files
added ~500 changed lines of pure type-annotation no-ops to our PR diff.
codecov/patch penalised these uncovered lines, dropping patch coverage
to 51.35% (target 61.83%).

Fix: revert these files to exactly match upstream/litellm_internal_staging.
The ruff_strict_gate still passes because the violations exist equally in
both the base and HEAD (total == base_count → no breach).
This commit is contained in:
Yaniv Israel 2026-07-02 15:10:22 +03:00
parent c7665a5227
commit 0b0e3ed8bd
7 changed files with 526 additions and 509 deletions

View file

@ -10,7 +10,7 @@ Supported for both `v1/chat/completions` (via the prompt-management hook) and
"""
import copy
from typing import TYPE_CHECKING, Any, Union, cast
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
@ -39,27 +39,27 @@ class AnthropicCacheControlHook(CustomPromptManagement):
def get_chat_completion_prompt(
self,
model: str,
messages: list[AllMessageValues],
messages: List[AllMessageValues],
non_default_params: dict,
prompt_id: str | None,
prompt_variables: dict | None,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
prompt_spec: PromptSpec | None = None,
prompt_label: str | None = None,
prompt_version: int | None = None,
ignore_prompt_manager_model: bool | None = False,
ignore_prompt_manager_optional_params: bool | None = False,
) -> tuple[str, list[AllMessageValues], dict]:
prompt_spec: Optional[PromptSpec] = None,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
ignore_prompt_manager_model: Optional[bool] = False,
ignore_prompt_manager_optional_params: Optional[bool] = False,
) -> Tuple[str, List[AllMessageValues], dict]:
"""
Apply cache control directives based on specified injection points.
Returns:
- model: str - the model to use
- messages: list[AllMessageValues] - messages with applied cache controls
- messages: List[AllMessageValues] - messages with applied cache controls
- non_default_params: dict - params with any global cache controls
"""
# Extract cache control injection points
injection_points: list[CacheControlInjectionPoint] = non_default_params.pop(
injection_points: List[CacheControlInjectionPoint] = non_default_params.pop(
"cache_control_injection_points", []
)
if not injection_points:
@ -69,8 +69,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
processed_messages = copy.deepcopy(messages)
# Separate message-level and non-message-level injection points
message_points: list[CacheControlMessageInjectionPoint] = []
remaining_points: list[CacheControlInjectionPoint] = []
message_points: List[CacheControlMessageInjectionPoint] = []
remaining_points: List[CacheControlInjectionPoint] = []
for point in injection_points:
if point.get("location") == "message":
message_points.append(cast(CacheControlMessageInjectionPoint, point))
@ -97,10 +97,10 @@ class AnthropicCacheControlHook(CustomPromptManagement):
@staticmethod
def _apply_message_injections(
points: list[CacheControlMessageInjectionPoint],
messages: list[AllMessageValues],
points: List[CacheControlMessageInjectionPoint],
messages: List[AllMessageValues],
max_blocks: int,
) -> list[AllMessageValues]:
) -> List[AllMessageValues]:
"""Apply message-level cache control injection points in order.
Anthropic allows at most ``MAX_CACHE_CONTROL_BLOCKS`` cache_control
@ -149,11 +149,11 @@ class AnthropicCacheControlHook(CustomPromptManagement):
@staticmethod
def _resolve_target_indices(
point: CacheControlMessageInjectionPoint, messages: list[AllMessageValues]
) -> list[int]:
point: CacheControlMessageInjectionPoint, messages: List[AllMessageValues]
) -> List[int]:
"""Resolve which message indices an injection point targets."""
_targetted_index: Union[int, str] | None = point.get("index", None)
targetted_index: int | None = None
_targetted_index: Optional[Union[int, str]] = point.get("index", None)
targetted_index: Optional[int] = None
if isinstance(_targetted_index, str):
try:
targetted_index = int(_targetted_index)
@ -230,10 +230,10 @@ class AnthropicCacheControlHook(CustomPromptManagement):
@staticmethod
def apply_to_anthropic_messages_request(
messages: list[dict],
messages: List[Dict],
system: str | list | None,
injection_points: list[CacheControlInjectionPoint],
) -> tuple[list[dict], str | list | None, list[CacheControlInjectionPoint]]:
injection_points: List[CacheControlInjectionPoint],
) -> Tuple[List[Dict], str | list | None, List[CacheControlInjectionPoint]]:
"""Apply cache control injection for the Anthropic-native v1/messages endpoint.
Returns (messages, system, remaining_non_message_points).
@ -241,12 +241,12 @@ class AnthropicCacheControlHook(CustomPromptManagement):
if not injection_points:
return messages, system, []
processed_messages: list[dict] = copy.deepcopy(messages)
processed_messages: List[Dict] = copy.deepcopy(messages)
processed_system = copy.deepcopy(system) if system is not None else None
message_points: list[CacheControlMessageInjectionPoint] = []
system_points: list[CacheControlMessageInjectionPoint] = []
remaining_points: list[CacheControlInjectionPoint] = []
message_points: List[CacheControlMessageInjectionPoint] = []
system_points: List[CacheControlMessageInjectionPoint] = []
remaining_points: List[CacheControlInjectionPoint] = []
for point in injection_points:
if point.get("location") == "message":
@ -290,7 +290,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
processed_messages = AnthropicCacheControlHook._apply_message_injections(
points=message_points,
messages=cast(list[AllMessageValues], processed_messages),
messages=cast(List[AllMessageValues], processed_messages),
max_blocks=max_blocks - used_blocks,
)
@ -298,10 +298,10 @@ class AnthropicCacheControlHook(CustomPromptManagement):
@staticmethod
def maybe_inject_cache_control(
messages: list[dict],
messages: List[Dict],
system: str | list | None,
kwargs: dict[str, Any],
) -> tuple[list[dict], str | list | None]:
kwargs: Dict[str, Any],
) -> Tuple[List[Dict], str | list | None]:
"""Extract cache_control_injection_points from kwargs and apply if present.
Pops the key from kwargs; if remaining (non-message) points exist they
@ -327,8 +327,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
def should_run_prompt_management(
self,
prompt_id: str | None,
prompt_spec: PromptSpec | None,
prompt_id: Optional[str],
prompt_spec: Optional[PromptSpec],
dynamic_callback_params: StandardCallbackDynamicParams,
) -> bool:
"""Always return False since this is not a true prompt management system."""
@ -336,12 +336,12 @@ class AnthropicCacheControlHook(CustomPromptManagement):
def _compile_prompt_helper(
self,
prompt_id: str | None,
prompt_spec: PromptSpec | None,
prompt_variables: dict | None,
prompt_id: Optional[str],
prompt_spec: Optional[PromptSpec],
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
prompt_label: str | None = None,
prompt_version: int | None = None,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
) -> PromptManagementClient:
"""Not used - this hook only modifies messages, doesn't fetch prompts."""
return PromptManagementClient(
@ -354,12 +354,12 @@ class AnthropicCacheControlHook(CustomPromptManagement):
async def async_compile_prompt_helper(
self,
prompt_id: str | None,
prompt_variables: dict | None,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
prompt_spec: PromptSpec | None = None,
prompt_label: str | None = None,
prompt_version: int | None = None,
prompt_spec: Optional[PromptSpec] = None,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
) -> PromptManagementClient:
"""Not used - this hook only modifies messages, doesn't fetch prompts."""
return self._compile_prompt_helper(
@ -374,19 +374,19 @@ class AnthropicCacheControlHook(CustomPromptManagement):
async def async_get_chat_completion_prompt(
self,
model: str,
messages: list[AllMessageValues],
messages: List[AllMessageValues],
non_default_params: dict,
prompt_id: str | None,
prompt_variables: dict | None,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
litellm_logging_obj: LiteLLMLoggingObj,
prompt_spec: PromptSpec | None = None,
tools: list[dict] | None = None,
prompt_label: str | None = None,
prompt_version: int | None = None,
ignore_prompt_manager_model: bool | None = False,
ignore_prompt_manager_optional_params: bool | None = False,
) -> tuple[str, list[AllMessageValues], dict]:
prompt_spec: Optional[PromptSpec] = None,
tools: Optional[List[Dict]] = None,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
ignore_prompt_manager_model: Optional[bool] = False,
ignore_prompt_manager_optional_params: Optional[bool] = False,
) -> Tuple[str, List[AllMessageValues], dict]:
"""Async version - delegates to sync since no async operations needed."""
return self.get_chat_completion_prompt(
model=model,
@ -403,15 +403,15 @@ class AnthropicCacheControlHook(CustomPromptManagement):
)
@staticmethod
def should_use_anthropic_cache_control_hook(non_default_params: dict) -> bool:
def should_use_anthropic_cache_control_hook(non_default_params: Dict) -> bool:
if non_default_params.get("cache_control_injection_points", None):
return True
return False
@staticmethod
def get_custom_logger_for_anthropic_cache_control_hook(
non_default_params: dict,
) -> CustomLogger | None:
non_default_params: Dict,
) -> Optional[CustomLogger]:
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
)

View file

@ -4,8 +4,11 @@ from typing import (
TYPE_CHECKING,
Any,
ClassVar,
Dict,
List,
Literal,
Optional,
Type,
Union,
get_args,
)
@ -56,7 +59,7 @@ from litellm.exceptions import (
_PRE_CALL_EXECUTED_TOKEN = secrets.token_hex(16)
def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None:
def get_session_id_from_request_data(request_data: Dict[str, Any]) -> Optional[str]:
"""Extract session_id from request data (litellm_session_id or metadata)."""
session_id = request_data.get("litellm_session_id")
if session_id:
@ -81,18 +84,18 @@ class CustomGuardrail(CustomLogger):
def __init__(
self,
guardrail_name: str | None = None,
supported_event_hooks: list[GuardrailEventHooks] | None = None,
event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = None,
guardrail_name: Optional[str] = None,
supported_event_hooks: Optional[List[GuardrailEventHooks]] = None,
event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]] = None,
default_on: bool = False,
mask_request_content: bool = False,
mask_response_content: bool = False,
violation_message_template: str | None = None,
end_session_after_n_fails: int | None = None,
on_violation: str | None = None,
realtime_violation_message: str | None = None,
on_sensitive_data: str | None = None,
sensitive_data_route_to_model: str | None = None,
violation_message_template: Optional[str] = None,
end_session_after_n_fails: Optional[int] = None,
on_violation: Optional[str] = None,
realtime_violation_message: Optional[str] = None,
on_sensitive_data: Optional[str] = None,
sensitive_data_route_to_model: Optional[str] = None,
sticky_session_routing: bool = True,
**kwargs,
):
@ -115,16 +118,16 @@ class CustomGuardrail(CustomLogger):
"""
self.guardrail_name = guardrail_name
self.supported_event_hooks = supported_event_hooks
self.event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = event_hook
self.event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]] = event_hook
self.default_on: bool = default_on
self.mask_request_content: bool = mask_request_content
self.mask_response_content: bool = mask_response_content
self.violation_message_template: str | None = violation_message_template
self.end_session_after_n_fails: int | None = end_session_after_n_fails
self.on_violation: str | None = on_violation
self.realtime_violation_message: str | None = realtime_violation_message
self.on_sensitive_data: str | None = on_sensitive_data
self.sensitive_data_route_to_model: str | None = sensitive_data_route_to_model
self.violation_message_template: Optional[str] = violation_message_template
self.end_session_after_n_fails: Optional[int] = end_session_after_n_fails
self.on_violation: Optional[str] = on_violation
self.realtime_violation_message: Optional[str] = realtime_violation_message
self.on_sensitive_data: Optional[str] = on_sensitive_data
self.sensitive_data_route_to_model: Optional[str] = sensitive_data_route_to_model
self.sticky_session_routing: bool = sticky_session_routing
if supported_event_hooks:
@ -132,13 +135,13 @@ class CustomGuardrail(CustomLogger):
self._validate_event_hook(event_hook, supported_event_hooks)
super().__init__(**kwargs)
def render_violation_message(self, default: str, context: dict[str, Any] | None = None) -> str:
def render_violation_message(self, default: str, context: Optional[Dict[str, Any]] = None) -> str:
"""Return a custom violation message if template is configured."""
if not self.violation_message_template:
return default
format_context: dict[str, Any] = {"default_message": default}
format_context: Dict[str, Any] = {"default_message": default}
if context:
format_context.update(context)
try:
@ -154,8 +157,8 @@ class CustomGuardrail(CustomLogger):
def raise_passthrough_exception(
self,
violation_message: str,
request_data: dict[str, Any],
detection_info: dict[str, Any] | None = None,
request_data: Dict[str, Any],
detection_info: Optional[Dict[str, Any]] = None,
) -> None:
"""
Raise a passthrough exception for guardrail violations.
@ -198,8 +201,8 @@ class CustomGuardrail(CustomLogger):
def raise_sensitive_data_route_exception(
self,
route_to_model: str,
request_data: dict[str, Any],
detection_info: dict[str, Any] | None = None,
request_data: Dict[str, Any],
detection_info: Optional[Dict[str, Any]] = None,
) -> None:
"""
Raise an exception to reroute the request to a different model.
@ -236,7 +239,7 @@ class CustomGuardrail(CustomLogger):
sticky_session_routing=self.sticky_session_routing,
)
def _get_session_id_from_request_data(self, request_data: dict[str, Any]) -> str | None:
def _get_session_id_from_request_data(self, request_data: Dict[str, Any]) -> Optional[str]:
"""Extract session_id from request data."""
return get_session_id_from_request_data(request_data)
@ -249,8 +252,8 @@ class CustomGuardrail(CustomLogger):
def handle_sensitive_data_detection(
self,
request_data: dict[str, Any],
detection_info: dict[str, Any] | None = None,
request_data: Dict[str, Any],
detection_info: Optional[Dict[str, Any]] = None,
) -> None:
"""
Handle sensitive data detection based on guardrail configuration.
@ -292,7 +295,7 @@ class CustomGuardrail(CustomLogger):
)
@staticmethod
def get_config_model() -> type["GuardrailConfigModel"] | None:
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
"""
Returns the config model for the guardrail
@ -302,12 +305,12 @@ class CustomGuardrail(CustomLogger):
def _validate_event_hook(
self,
event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None,
supported_event_hooks: list[GuardrailEventHooks],
event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]],
supported_event_hooks: List[GuardrailEventHooks],
) -> None:
def _validate_event_hook_list_is_in_supported_event_hooks(
event_hook: Union[list[GuardrailEventHooks], list[str]],
supported_event_hooks: list[GuardrailEventHooks],
event_hook: Union[List[GuardrailEventHooks], List[str]],
supported_event_hooks: List[GuardrailEventHooks],
) -> None:
for hook in event_hook:
if isinstance(hook, str):
@ -358,7 +361,7 @@ class CustomGuardrail(CustomLogger):
key_meta = meta.get("user_api_key_metadata") or key_meta
return {**team_meta, **key_meta}
def get_disable_global_guardrail(self, data: dict) -> bool | None:
def get_disable_global_guardrail(self, data: dict) -> Optional[bool]:
"""
Returns True if the global guardrail should be disabled.
@ -367,7 +370,7 @@ class CustomGuardrail(CustomLogger):
"""
return self._get_admin_metadata(data).get("disable_global_guardrails", False)
def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> list[str]:
def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> List[str]:
"""
Returns the list of global guardrail names the team/key has opted out of.
@ -398,7 +401,7 @@ class CustomGuardrail(CustomLogger):
return True
raise
def get_guardrail_from_metadata(self, data: dict) -> Union[list[str], list[dict[str, DynamicGuardrailParams]]]:
def get_guardrail_from_metadata(self, data: dict) -> Union[List[str], List[Dict[str, DynamicGuardrailParams]]]:
"""
Returns the guardrail(s) to be run from the metadata or root
"""
@ -419,7 +422,7 @@ class CustomGuardrail(CustomLogger):
def _guardrail_is_in_requested_guardrails(
self,
requested_guardrails: Union[list[str], list[dict[str, DynamicGuardrailParams]]],
requested_guardrails: Union[List[str], List[Dict[str, DynamicGuardrailParams]]],
) -> bool:
for _guardrail in requested_guardrails:
if isinstance(_guardrail, dict):
@ -431,13 +434,13 @@ class CustomGuardrail(CustomLogger):
return False
def _pre_call_marker(self) -> str | None:
def _pre_call_marker(self) -> Optional[str]:
name = self.guardrail_name
if not name:
return None
return f"{_PRE_CALL_EXECUTED_TOKEN}:{name}"
def mark_pre_call_hook_ran(self, data: dict[str, Any]) -> None:
def mark_pre_call_hook_ran(self, data: Dict[str, Any]) -> None:
"""
Record that this guardrail's ``async_pre_call_hook`` already ran for this
request, so the deployment-level hook does not run it a second time.
@ -462,7 +465,7 @@ class CustomGuardrail(CustomLogger):
return
data["metadata"] = {PRE_CALL_EXECUTED_GUARDRAILS_KEY: [marker]}
def _pre_call_hook_already_ran(self, data: dict[str, Any]) -> bool:
def _pre_call_hook_already_ran(self, data: Dict[str, Any]) -> bool:
marker = self._pre_call_marker()
if marker is None:
return False
@ -474,7 +477,9 @@ class CustomGuardrail(CustomLogger):
return True
return False
async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None:
async def async_pre_call_deployment_hook(
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
) -> Optional[dict]:
from litellm.proxy._types import UserAPIKeyAuth
# should run guardrail
@ -514,8 +519,8 @@ class CustomGuardrail(CustomLogger):
self,
request_data: dict,
response: LLMResponseTypes,
call_type: CallTypes | None,
) -> LLMResponseTypes | None:
call_type: Optional[CallTypes],
) -> Optional[LLMResponseTypes]:
"""
Allow modifying / reviewing the response just after it's received from the deployment.
"""
@ -697,16 +702,16 @@ class CustomGuardrail(CustomLogger):
def add_standard_logging_guardrail_information_to_request_data(
self,
guardrail_json_response: Union[Exception, str, dict, list[dict]],
guardrail_json_response: Union[Exception, str, dict, List[dict]],
request_data: dict,
guardrail_status: GuardrailStatus,
start_time: float | None = None,
end_time: float | None = None,
duration: float | None = None,
masked_entity_count: dict[str, int] | None = None,
guardrail_provider: str | None = None,
event_type: GuardrailEventHooks | None = None,
tracing_detail: GuardrailTracingDetail | None = None,
start_time: Optional[float] = None,
end_time: Optional[float] = None,
duration: Optional[float] = None,
masked_entity_count: Optional[Dict[str, int]] = None,
guardrail_provider: Optional[str] = None,
event_type: Optional[GuardrailEventHooks] = None,
tracing_detail: Optional[GuardrailTracingDetail] = None,
) -> None:
"""
Builds `StandardLoggingGuardrailInformation` and adds it to the request metadata so it can be used for logging to DataDog, Langfuse, etc.
@ -721,7 +726,7 @@ class CustomGuardrail(CustomLogger):
from litellm.types.utils import GuardrailMode
# Use event_type if provided, otherwise fall back to self.event_hook
guardrail_mode: Union[GuardrailEventHooks, GuardrailMode, list[GuardrailEventHooks]]
guardrail_mode: Union[GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks]]
if event_type is not None:
guardrail_mode = event_type
elif isinstance(self.event_hook, Mode):
@ -830,13 +835,13 @@ class CustomGuardrail(CustomLogger):
def _process_response(
self,
response: dict | None,
response: Optional[Dict],
request_data: dict,
start_time: float | None = None,
end_time: float | None = None,
duration: float | None = None,
event_type: GuardrailEventHooks | None = None,
original_inputs: dict | None = None,
start_time: Optional[float] = None,
end_time: Optional[float] = None,
duration: Optional[float] = None,
event_type: Optional[GuardrailEventHooks] = None,
original_inputs: Optional[Dict] = None,
):
"""
Add StandardLoggingGuardrailInformation to the request data
@ -844,7 +849,7 @@ class CustomGuardrail(CustomLogger):
This gets logged on downsteam Langfuse, DataDog, etc.
"""
# Convert None to empty dict to satisfy type requirements
guardrail_response: Union[dict[str, Any], str] = {} if response is None else response
guardrail_response: Union[Dict[str, Any], str] = {} if response is None else response
# For apply_guardrail functions in custom_code_guardrail scenario,
# simplify the logged response to "allow", "deny", or "mask"
@ -900,10 +905,10 @@ class CustomGuardrail(CustomLogger):
self,
e: Exception,
request_data: dict,
start_time: float | None = None,
end_time: float | None = None,
duration: float | None = None,
event_type: GuardrailEventHooks | None = None,
start_time: Optional[float] = None,
end_time: Optional[float] = None,
duration: Optional[float] = None,
event_type: Optional[GuardrailEventHooks] = None,
):
"""
Add StandardLoggingGuardrailInformation to the request data
@ -930,7 +935,7 @@ class CustomGuardrail(CustomLogger):
)
raise e
def _inputs_were_modified(self, original_inputs: dict, response: dict) -> bool:
def _inputs_were_modified(self, original_inputs: Dict, response: Dict) -> bool:
"""
Compare original inputs with response to determine if content was modified.
@ -975,8 +980,8 @@ class CustomGuardrail(CustomLogger):
setattr(self, key, value)
def get_guardrails_messages_for_call_type(
self, call_type: CallTypes, data: dict | None = None
) -> list[AllMessageValues] | None:
self, call_type: CallTypes, data: Optional[dict] = None
) -> Optional[List[AllMessageValues]]:
"""
Returns the messages for the given call type and data
"""
@ -1014,7 +1019,7 @@ class CustomGuardrail(CustomLogger):
input=input_data,
responses_api_request=data,
)
return cast(list[AllMessageValues], messages)
return cast(List[AllMessageValues], messages)
return None
@ -1078,7 +1083,7 @@ def log_guardrail_information(func):
def _infer_event_type_from_function_name(
func_name: str,
) -> GuardrailEventHooks | None:
) -> Optional[GuardrailEventHooks]:
"""Infer the actual event type from the function name"""
if func_name == "async_pre_call_hook":
return GuardrailEventHooks.pre_call

View file

@ -10,7 +10,7 @@ identical metrics. The attribute cardinality filter is reused from v1 by import
from dataclasses import dataclass
from datetime import datetime
from typing import Any, Mapping
from typing import Any, FrozenSet, Mapping, Optional
from opentelemetry.metrics import Histogram, Meter
@ -81,11 +81,11 @@ class GenAIMetricRecorder:
survives.
"""
def __init__(self, metrics: GenAIMetrics, callback_name: str | None = None) -> None:
def __init__(self, metrics: GenAIMetrics, callback_name: Optional[str] = None) -> None:
self._metrics = metrics
self._callback_name = callback_name
self._include: frozenset[str] | None = None
self._exclude: frozenset[str] | None = None
self._include: Optional[FrozenSet[str]] = None
self._exclude: Optional[FrozenSet[str]] = None
self._filter_resolved = False
def record(

View file

@ -4,7 +4,11 @@ import time
from typing import (
TYPE_CHECKING,
Any,
Dict,
List,
NoReturn,
Optional,
Tuple,
Union,
cast,
)
@ -146,8 +150,8 @@ def _basic_sanitize_anthropic_tool_name(name: str) -> str:
def _build_anthropic_tool_name_maps(
original_names: list[str],
) -> tuple[dict[str, str], dict[str, str]]:
original_names: List[str],
) -> Tuple[Dict[str, str], Dict[str, str]]:
"""Build (forward, reverse) tool-name maps for a single request.
forward[original] = sanitized -- only present when name was rewritten
@ -167,7 +171,7 @@ def _build_anthropic_tool_name_maps(
seen gets the disambiguating suffix. Callers should preserve the
caller's tool order (we do).
"""
forward: dict[str, str] = {}
forward: Dict[str, str] = {}
used: set = set()
# First pass: reserve slots for names that are already valid so they
@ -208,7 +212,7 @@ def _build_anthropic_tool_name_maps(
return forward, reverse
REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT: dict[str, str] = {
REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT: Dict[str, str] = {
"low": "low",
"minimal": "low",
"medium": "medium",
@ -235,23 +239,23 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
to pass metadata to anthropic, it's {"user_id": "any-relevant-information"}
"""
max_tokens: int | None = None
stop_sequences: list | None = None
temperature: int | None = None
top_p: int | None = None
top_k: int | None = None
metadata: dict | None = None
system: str | None = None
max_tokens: Optional[int] = None
stop_sequences: Optional[list] = None
temperature: Optional[int] = None
top_p: Optional[int] = None
top_k: Optional[int] = None
metadata: Optional[dict] = None
system: Optional[str] = None
def __init__(
self,
max_tokens: int | None = None,
stop_sequences: list | None = None,
temperature: int | None = None,
top_p: int | None = None,
top_k: int | None = None,
metadata: dict | None = None,
system: str | None = None,
max_tokens: Optional[int] = None,
stop_sequences: Optional[list] = None,
temperature: Optional[int] = None,
top_p: Optional[int] = None,
top_k: Optional[int] = None,
metadata: Optional[dict] = None,
system: Optional[str] = None,
) -> None:
locals_ = locals().copy()
for key, value in locals_.items():
@ -259,11 +263,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
setattr(self.__class__, key, value)
@property
def custom_llm_provider(self) -> str | None:
def custom_llm_provider(self) -> Optional[str]:
return "anthropic"
@classmethod
def get_config(cls, *, model: str | None = None):
def get_config(cls, *, model: Optional[str] = None):
config = super().get_config()
# anthropic requires a default value for max_tokens
@ -273,7 +277,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return config
@staticmethod
def get_max_tokens_for_model(model: str | None = None) -> int:
def get_max_tokens_for_model(model: Optional[str] = None) -> int:
"""
Get the max output tokens for a given model.
Falls back to DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS (configurable via env var) if model is not found.
@ -290,7 +294,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def convert_tool_use_to_openai_format(
anthropic_tool_content: dict[str, Any],
anthropic_tool_content: Dict[str, Any],
index: int,
) -> ChatCompletionToolCallChunk:
"""
@ -315,7 +319,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
# Include caller information if present (for programmatic tool calling)
if "caller" in anthropic_tool_content:
tool_call["caller"] = cast(dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item]
tool_call["caller"] = cast(Dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item]
return tool_call
@staticmethod
@ -336,7 +340,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return AnthropicConfig._supports_model_capability(model, f"supports_{level}_reasoning_effort")
@staticmethod
def _validate_effort_for_model(model: str, effort: str | None) -> str | None:
def _validate_effort_for_model(model: str, effort: Optional[str]) -> Optional[str]:
"""Return ``None`` if ``effort`` is allowed on ``model``, else an error message."""
if effort == "max" and not (
AnthropicConfig._is_adaptive_thinking_model(model) or AnthropicConfig._supports_effort_level(model, "max")
@ -363,7 +367,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
@staticmethod
def _model_supports_speed_param(model: str, custom_llm_provider: str | None = None) -> bool:
def _model_supports_speed_param(model: str, custom_llm_provider: Optional[str] = None) -> bool:
"""Whether the model accepts Anthropic's ``speed`` parameter (fast mode).
Fast mode is direct Anthropic API-only (not Bedrock, Vertex, or Azure).
@ -380,7 +384,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
model: str,
optional_params: dict,
drop_params: bool,
custom_llm_provider: str | None = None,
custom_llm_provider: Optional[str] = None,
) -> None:
if "speed" not in optional_params:
return
@ -459,7 +463,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return params
@staticmethod
def filter_anthropic_output_schema(schema: dict[str, Any]) -> dict[str, Any]:
def filter_anthropic_output_schema(schema: Dict[str, Any]) -> Dict[str, Any]:
"""
Filter out unsupported fields from JSON schema for Anthropic's output_format API.
@ -515,7 +519,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if field in schema:
constraint_descriptions.append(constraint_labels[field].format(schema[field]))
result: dict[str, Any] = {}
result: Dict[str, Any] = {}
# Update description with removed constraint info
if constraint_descriptions:
@ -555,7 +559,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return result
def get_json_schema_from_pydantic_object(self, response_format: Union[Any, dict, None]) -> dict | None:
def get_json_schema_from_pydantic_object(self, response_format: Union[Any, Dict, None]) -> Optional[dict]:
return type_to_response_format_param(
response_format, ref_template="/$defs/{model}"
) # Relevant issue: https://github.com/BerriAI/litellm/issues/7755
@ -570,10 +574,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _map_tool_choice(
self,
tool_choice: str | None,
parallel_tool_use: bool | None,
) -> AnthropicMessagesToolChoice | None:
_tool_choice: AnthropicMessagesToolChoice | None = None
tool_choice: Optional[str],
parallel_tool_use: Optional[bool],
) -> Optional[AnthropicMessagesToolChoice]:
_tool_choice: Optional[AnthropicMessagesToolChoice] = None
if tool_choice == "auto":
_tool_choice = AnthropicMessagesToolChoice(
type="auto",
@ -614,9 +618,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _map_tool_helper(
self,
tool: ChatCompletionToolParam,
) -> tuple[AllAnthropicToolsValues | None, AnthropicMcpServerTool | None]:
returned_tool: AllAnthropicToolsValues | None = None
mcp_server: AnthropicMcpServerTool | None = None
) -> Tuple[Optional[AllAnthropicToolsValues], Optional[AnthropicMcpServerTool]]:
returned_tool: Optional[AllAnthropicToolsValues] = None
mcp_server: Optional[AnthropicMcpServerTool] = None
if tool["type"] == "function" or tool["type"] == "custom":
_input_schema: dict = tool["function"].get(
@ -667,8 +671,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if "parameters" not in tool["function"]:
raise ValueError("Missing required parameter: parameters")
_display_width_px: int | None = tool["function"]["parameters"].get("display_width_px")
_display_height_px: int | None = tool["function"]["parameters"].get("display_height_px")
_display_width_px: Optional[int] = tool["function"]["parameters"].get("display_width_px")
_display_height_px: Optional[int] = tool["function"]["parameters"].get("display_height_px")
if _display_width_px is None or _display_height_px is None:
raise ValueError("Missing required parameter: display_width_px or display_height_px")
@ -829,14 +833,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
from litellm.types.llms.anthropic import AnthropicMcpServerToolConfiguration
allowed_tools = tool.get("allowed_tools", None)
tool_configuration: AnthropicMcpServerToolConfiguration | None = None
tool_configuration: Optional[AnthropicMcpServerToolConfiguration] = None
if allowed_tools is not None:
tool_configuration = AnthropicMcpServerToolConfiguration(
allowed_tools=tool.get("allowed_tools", None),
)
headers = tool.get("headers", {})
authorization_token: str | None = None
authorization_token: Optional[str] = None
if headers is not None:
bearer_token = headers.get("Authorization", None)
if bearer_token is not None:
@ -856,8 +860,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _map_tools(
self,
tools: list,
) -> tuple[list[AllAnthropicToolsValues], list[AnthropicMcpServerTool]]:
tools: List,
) -> Tuple[List[AllAnthropicToolsValues], List[AnthropicMcpServerTool]]:
anthropic_tools = []
mcp_servers = []
for tool in tools:
@ -902,9 +906,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def _rewrite_tool_names_in_messages(
messages: list[AllMessageValues],
name_forward_map: dict[str, str],
) -> list[AllMessageValues]:
messages: List[AllMessageValues],
name_forward_map: Dict[str, str],
) -> List[AllMessageValues]:
"""Return a copy of `messages` with tool_call/function_call names
rewritten using the per-request forward map.
@ -915,7 +919,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
"""
if not name_forward_map:
return messages
new_messages: list[AllMessageValues] = []
new_messages: List[AllMessageValues] = []
for msg in messages:
if not isinstance(msg, dict):
new_messages.append(msg)
@ -953,8 +957,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def _build_request_tool_name_maps(
tools: list,
) -> tuple[dict[str, str], dict[str, str]]:
tools: List,
) -> Tuple[Dict[str, str], Dict[str, str]]:
"""Build the (forward, reverse) tool-name maps for an OpenAI tools list.
Operates on **OpenAI-format** tool dicts (pre-``_map_tools``). The
@ -968,7 +972,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
the original name out of either ``{"function": {"name": ...}}``
(legacy OpenAI shape) or ``{"name": ...}`` (rare top-level shape).
"""
original_names: list[str] = []
original_names: List[str] = []
for tool in tools or []:
if not isinstance(tool, dict):
continue
@ -981,8 +985,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def _sanitize_tool_names_in_request(
optional_params: dict[str, Any],
) -> tuple[dict[str, str], dict[str, str]]:
optional_params: Dict[str, Any],
) -> Tuple[Dict[str, str], Dict[str, str]]:
"""Sanitize ``optional_params['tools']`` and ``optional_params['tool_choice']``
in place so every name matches Anthropic's ``^[a-zA-Z0-9_-]{1,128}$``.
@ -1006,7 +1010,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
# Order matters: the first occurrence wins the canonical slot;
# later collisions get numeric suffixes (see
# ``_build_anthropic_tool_name_maps``).
original_names: list[str] = []
original_names: List[str] = []
for t in tools:
if not isinstance(t, dict):
continue
@ -1028,7 +1032,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
# so a caller reusing the same tool list/dicts across requests
# doesn't see its inputs permanently rewritten (which would also
# drop the original key from `forward` on the next request).
new_tools: list[Any] = []
new_tools: List[Any] = []
for t in tools:
if (
isinstance(t, dict)
@ -1054,7 +1058,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return forward, reverse
def _detect_tool_search_tools(self, tools: list | None) -> bool:
def _detect_tool_search_tools(self, tools: Optional[List]) -> bool:
"""Check if tool search tools are present in the tools list."""
if not tools:
return False
@ -1068,7 +1072,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return True
return False
def _separate_deferred_tools(self, tools: list) -> tuple[list, list]:
def _separate_deferred_tools(self, tools: List) -> Tuple[List, List]:
"""
Separate tools into deferred and non-deferred lists.
@ -1088,9 +1092,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _expand_tool_references(
self,
content: list,
deferred_tools: list,
) -> list:
content: List,
deferred_tools: List,
) -> List:
"""
Expand tool_reference blocks to full tool definitions.
@ -1131,8 +1135,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return expanded_content
def _map_stop_sequences(self, stop: Union[str, list[str]] | None) -> list[str] | None:
new_stop: list[str] | None = None
def _map_stop_sequences(self, stop: Optional[Union[str, List[str]]]) -> Optional[List[str]]:
new_stop: Optional[List[str]] = None
if isinstance(stop, str):
if (
stop.isspace() and litellm.drop_params is True
@ -1153,10 +1157,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def _map_reasoning_effort(
reasoning_effort: Union[REASONING_EFFORT, str] | None,
reasoning_effort: Optional[Union[REASONING_EFFORT, str]],
model: str,
llm_provider: str = "anthropic",
) -> AnthropicThinkingParam | None:
) -> Optional[AnthropicThinkingParam]:
if reasoning_effort is None or reasoning_effort == "none":
return None
if AnthropicConfig._is_adaptive_thinking_model(model):
@ -1207,10 +1211,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
llm_provider=llm_provider,
)
def _extract_json_schema_from_response_format(self, value: dict | None) -> dict | None:
def _extract_json_schema_from_response_format(self, value: Optional[dict]) -> Optional[dict]:
if value is None:
return None
json_schema: dict | None = None
json_schema: Optional[dict] = None
if "response_schema" in value:
json_schema = value["response_schema"]
elif "json_schema" in value:
@ -1218,8 +1222,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return json_schema
def map_response_format_to_anthropic_output_format(self, value: dict | None) -> AnthropicOutputSchema | None:
json_schema: dict | None = self._extract_json_schema_from_response_format(value)
def map_response_format_to_anthropic_output_format(self, value: Optional[dict]) -> Optional[AnthropicOutputSchema]:
json_schema: Optional[dict] = self._extract_json_schema_from_response_format(value)
if json_schema is None:
return None
@ -1245,13 +1249,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
def map_response_format_to_anthropic_tool(
self, value: dict | None, optional_params: dict, is_thinking_enabled: bool
) -> AnthropicMessagesTool | None:
self, value: Optional[dict], optional_params: dict, is_thinking_enabled: bool
) -> Optional[AnthropicMessagesTool]:
ignore_response_format_types = ["text"]
if value is None or value["type"] in ignore_response_format_types: # value is a no-op
return None
json_schema: dict | None = self._extract_json_schema_from_response_format(value)
json_schema: Optional[dict] = self._extract_json_schema_from_response_format(value)
if json_schema is None:
return None
"""
@ -1296,8 +1300,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def map_openai_context_management_to_anthropic(
context_management: Union[list[dict[str, Any]], dict[str, Any]],
) -> dict[str, Any] | None:
context_management: Union[List[Dict[str, Any]], Dict[str, Any]],
) -> Optional[Dict[str, Any]]:
"""
OpenAI format: [{"type": "compaction", "compact_threshold": 200000}]
Anthropic format: {
@ -1328,7 +1332,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
entry_type = entry.get("type")
if entry_type == "compaction":
anthropic_edit: dict[str, Any] = {"type": "compact_20260112"}
anthropic_edit: Dict[str, Any] = {"type": "compact_20260112"}
compact_threshold = entry.get("compact_threshold")
# Rewrite to 'trigger' with correct nesting if threshold exists
if compact_threshold is not None and isinstance(compact_threshold, (int, float)):
@ -1384,7 +1388,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if mcp_servers:
optional_params["mcp_servers"] = mcp_servers
elif param == "tool_choice" or param == "parallel_tool_calls":
_tool_choice: AnthropicMessagesToolChoice | None = self._map_tool_choice(
_tool_choice: Optional[AnthropicMessagesToolChoice] = self._map_tool_choice(
tool_choice=non_default_params.get("tool_choice"),
parallel_tool_use=non_default_params.get("parallel_tool_calls"),
)
@ -1515,7 +1519,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _create_json_tool_call_for_response_format(
self,
json_schema: dict | None = None,
json_schema: Optional[dict] = None,
) -> AnthropicMessagesTool:
"""
Handles creating a tool call for getting responses in JSON format.
@ -1550,7 +1554,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
"""
return False
def translate_system_message(self, messages: list[AllMessageValues]) -> list[AnthropicSystemMessageContent]:
def translate_system_message(self, messages: List[AllMessageValues]) -> List[AnthropicSystemMessageContent]:
"""
Translate system message to anthropic format.
@ -1558,7 +1562,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
When should_strip_billing_metadata() is True, x-anthropic-billing-header system blocks are dropped.
"""
system_prompt_indices = []
anthropic_system_message_list: list[AnthropicSystemMessageContent] = []
anthropic_system_message_list: List[AnthropicSystemMessageContent] = []
for idx, message in enumerate(messages):
if message["role"] == "system":
system_prompt_indices.append(idx)
@ -1608,9 +1612,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def add_code_execution_tool(
self,
messages: list[AllAnthropicMessageValues],
tools: list[Union[AllAnthropicToolsValues, dict]],
) -> list[Union[AllAnthropicToolsValues, dict]]:
messages: List[AllAnthropicMessageValues],
tools: List[Union[AllAnthropicToolsValues, Dict]],
) -> List[Union[AllAnthropicToolsValues, Dict]]:
"""if 'container_upload' in messages, add code_execution tool"""
add_code_execution_tool = False
for message in messages:
@ -1725,7 +1729,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def transform_request(
self,
model: str,
messages: list[AllMessageValues],
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
@ -1827,7 +1831,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
## Add code_execution tool if container_upload is in messages
_tools = (
cast(
list[Union[AllAnthropicToolsValues, dict]] | None,
Optional[List[Union[AllAnthropicToolsValues, Dict]]],
optional_params.get("tools"),
)
or []
@ -1925,12 +1929,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _resolve_json_mode_non_streaming(
self,
json_mode: bool | None,
tool_calls: list[ChatCompletionToolCallChunk],
) -> tuple[
LitellmMessage | None,
list[ChatCompletionToolCallChunk],
str | None,
json_mode: Optional[bool],
tool_calls: List[ChatCompletionToolCallChunk],
) -> Tuple[
Optional[LitellmMessage],
List[ChatCompletionToolCallChunk],
Optional[str],
]:
"""Strip internal response_format tool calls; merge payload into content when mixed with user tools."""
if json_mode is not True or not tool_calls:
@ -1951,30 +1955,30 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
first_json = tool_calls[json_indices[0]]
json_msg = AnthropicConfig._convert_tool_response_to_message([first_json])
extra_content: str | None = json_msg.content if json_msg is not None else None
extra_content: Optional[str] = json_msg.content if json_msg is not None else None
filtered_tools = [t for i, t in enumerate(tool_calls) if i not in json_indices]
return None, filtered_tools, extra_content
def extract_response_content(
self, completion_response: dict
) -> tuple[
) -> Tuple[
str,
list[Any] | None,
list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None,
str | None,
list[ChatCompletionToolCallChunk],
list[Any] | None,
list[Any] | None,
list[Any] | None,
Optional[List[Any]],
Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]],
Optional[str],
List[ChatCompletionToolCallChunk],
Optional[List[Any]],
Optional[List[Any]],
Optional[List[Any]],
]:
text_content = ""
citations: list[Any] | None = None
thinking_blocks: list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None
reasoning_content: str | None = None
tool_calls: list[ChatCompletionToolCallChunk] = []
web_search_results: list[Any] | None = None
tool_results: list[Any] | None = None
compaction_blocks: list[Any] | None = None
citations: Optional[List[Any]] = None
thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]] = None
reasoning_content: Optional[str] = None
tool_calls: List[ChatCompletionToolCallChunk] = []
web_search_results: Optional[List[Any]] = None
tool_results: Optional[List[Any]] = None
compaction_blocks: Optional[List[Any]] = None
for idx, content in enumerate(completion_response["content"]):
if content["type"] == "text":
text_content += content["text"]
@ -2037,7 +2041,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if thinking_blocks is not None:
reasoning_content = ""
for block in thinking_blocks:
thinking_content = cast(str | None, block.get("thinking"))
thinking_content = cast(Optional[str], block.get("thinking"))
if thinking_content is not None:
reasoning_content += thinking_content
@ -2055,9 +2059,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def calculate_usage(
self,
usage_object: dict,
reasoning_content: str | None,
completion_response: dict | None = None,
speed: str | None = None,
reasoning_content: Optional[str],
completion_response: Optional[dict] = None,
speed: Optional[str] = None,
) -> Usage:
# NOTE: Sometimes the usage object has None set explicitly for token counts, meaning .get() & key access returns None, and we need to account for this
raw_prompt_tokens = usage_object.get("input_tokens", 0) or 0
@ -2067,10 +2071,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
_usage = usage_object
cache_creation_input_tokens: int = 0
cache_read_input_tokens: int = 0
cache_creation_token_details: CacheCreationTokenDetails | None = None
web_search_requests: int | None = None
tool_search_requests: int | None = None
inference_geo: str | None = None
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
web_search_requests: Optional[int] = None
tool_search_requests: Optional[int] = None
inference_geo: Optional[str] = None
if "inference_geo" in _usage and _usage["inference_geo"] is not None:
inference_geo = _usage["inference_geo"]
service_tier = cast(
@ -2078,7 +2082,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
_usage.get("service_tier"),
)
iterations: list[Any] | None = _usage.get("iterations")
iterations: Optional[List[Any]] = _usage.get("iterations")
if iterations:
prompt_tokens = sum(it.get("input_tokens", 0) or 0 for it in iterations)
completion_tokens = sum(it.get("output_tokens", 0) or 0 for it in iterations)
@ -2164,8 +2168,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
return usage
def _build_code_by_id_map(self, tool_calls: list[ChatCompletionToolCallChunk]) -> dict[str, str]:
code_by_id: dict[str, str] = {}
def _build_code_by_id_map(self, tool_calls: List[ChatCompletionToolCallChunk]) -> Dict[str, str]:
code_by_id: Dict[str, str] = {}
for tc in tool_calls:
try:
args = json.loads(tc.get("function", {}).get("arguments", "{}"))
@ -2179,10 +2183,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _build_code_interpreter_results(
self,
tool_results: list[Any],
code_by_id: dict[str, str],
container_id: str | None,
) -> list[OutputCodeInterpreterCall]:
tool_results: List[Any],
code_by_id: Dict[str, str],
container_id: Optional[str],
) -> List[OutputCodeInterpreterCall]:
code_interpreter_results = []
for tr in tool_results:
if tr.get("type") != "bash_code_execution_tool_result":
@ -2205,14 +2209,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _build_provider_specific_fields(
self,
completion_response: dict,
citations: list[Any] | None,
thinking_blocks: list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None,
web_search_results: list[Any] | None,
tool_results: list[Any] | None,
compaction_blocks: list[Any] | None,
tool_calls: list[ChatCompletionToolCallChunk],
) -> dict[str, Any]:
provider_specific_fields: dict[str, Any] = {
citations: Optional[List[Any]],
thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]],
web_search_results: Optional[List[Any]],
tool_results: Optional[List[Any]],
compaction_blocks: Optional[List[Any]],
tool_calls: List[ChatCompletionToolCallChunk],
) -> Dict[str, Any]:
provider_specific_fields: Dict[str, Any] = {
"citations": citations,
"thinking_blocks": thinking_blocks,
}
@ -2249,12 +2253,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
completion_response: dict,
raw_response: httpx.Response,
model_response: ModelResponse,
json_mode: bool | None = None,
prefix_prompt: str | None = None,
speed: str | None = None,
tool_name_reverse_map: dict[str, str] | None = None,
json_mode: Optional[bool] = None,
prefix_prompt: Optional[str] = None,
speed: Optional[str] = None,
tool_name_reverse_map: Optional[Dict[str, str]] = None,
):
_hidden_params: dict = {}
_hidden_params: Dict = {}
_hidden_params["additional_headers"] = process_anthropic_headers(dict(raw_response.headers))
if "error" in completion_response:
response_headers = getattr(raw_response, "headers", None)
@ -2350,7 +2354,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
model_response._hidden_params = _hidden_params
return model_response
def get_prefix_prompt(self, messages: list[AllMessageValues]) -> str | None:
def get_prefix_prompt(self, messages: List[AllMessageValues]) -> Optional[str]:
"""
Get the prefix prompt from the messages.
@ -2375,13 +2379,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LoggingClass,
request_data: dict,
messages: list[AllMessageValues],
optional_params: dict,
request_data: Dict,
messages: List[AllMessageValues],
optional_params: Dict,
litellm_params: dict,
encoding: Any,
api_key: str | None = None,
json_mode: bool | None = None,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ModelResponse:
## LOGGING
logging_obj.post_call(
@ -2404,7 +2408,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
prefix_prompt = self.get_prefix_prompt(messages=messages)
speed = optional_params.get("speed")
tool_name_reverse_map: dict[str, str] | None = None
tool_name_reverse_map: Optional[Dict[str, str]] = None
if isinstance(litellm_params, dict):
_candidate = litellm_params.get(ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY)
if isinstance(_candidate, dict):
@ -2423,14 +2427,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def _convert_tool_response_to_message(
tool_calls: list[ChatCompletionToolCallChunk],
) -> LitellmMessage | None:
tool_calls: List[ChatCompletionToolCallChunk],
) -> Optional[LitellmMessage]:
"""
In JSON mode, Anthropic API returns JSON schema as a tool call, we need to convert it to a message to follow the OpenAI format
"""
## HANDLE JSON MODE - anthropic returns single function call
json_mode_content_str: str | None = tool_calls[0]["function"].get("arguments")
json_mode_content_str: Optional[str] = tool_calls[0]["function"].get("arguments")
try:
if json_mode_content_str is not None:
args = json.loads(json_mode_content_str)
@ -2448,7 +2452,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return None
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers]
) -> BaseLLMException:
return AnthropicError(
status_code=status_code,

View file

@ -2,7 +2,7 @@
Translates from OpenAI's `/v1/chat/completions` to ModelScope's `/v1/chat/completions`
"""
from typing import Any, Coroutine, Literal, Union, cast, overload
from typing import Any, Coroutine, Literal, Optional, Tuple, Union, cast, overload
from typing_extensions import override
@ -59,8 +59,8 @@ class ModelScopeChatConfig(OpenAIGPTConfig):
return super()._transform_messages(messages=messages, model=model, is_async=False)
def _get_openai_compatible_provider_info(
self, api_base: str | None, api_key: str | None
) -> tuple[str | None, str | None]:
self, api_base: Optional[str], api_key: Optional[str]
) -> Tuple[Optional[str], Optional[str]]:
api_base = api_base or get_secret_str("MODELSCOPE_API_BASE") or self.DEFAULT_BASE_URL # type: ignore
dynamic_api_key = api_key or get_secret_str("MODELSCOPE_API_KEY")
return api_base, dynamic_api_key
@ -68,12 +68,12 @@ class ModelScopeChatConfig(OpenAIGPTConfig):
@override
def get_complete_url(
self,
api_base: str | None,
api_key: str | None,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: bool | None = None,
stream: Optional[bool] = None,
) -> str:
"""
If api_base is not provided, use the default ModelScope /chat/completions endpoint.

View file

@ -9,9 +9,13 @@ from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
List,
Literal,
Mapping,
Optional,
Tuple,
Type,
Union,
cast,
)
@ -131,7 +135,7 @@ class VertexAIBaseConfig:
optional_params[mapped_params[param]] = value
return optional_params
def get_eu_regions(self) -> list[str]:
def get_eu_regions(self) -> List[str]:
"""
Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions
"""
@ -148,7 +152,7 @@ class VertexAIBaseConfig:
"europe-west9",
]
def get_us_regions(self) -> list[str]:
def get_us_regions(self) -> List[str]:
"""
Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions
"""
@ -194,29 +198,29 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
Note: Please make sure to modify the default parameters as required for your use case.
"""
temperature: float | None = None
max_output_tokens: int | None = None
top_p: float | None = None
top_k: int | None = None
response_mime_type: str | None = None
candidate_count: int | None = None
stop_sequences: list | None = None
frequency_penalty: float | None = None
presence_penalty: float | None = None
seed: int | None = None
temperature: Optional[float] = None
max_output_tokens: Optional[int] = None
top_p: Optional[float] = None
top_k: Optional[int] = None
response_mime_type: Optional[str] = None
candidate_count: Optional[int] = None
stop_sequences: Optional[list] = None
frequency_penalty: Optional[float] = None
presence_penalty: Optional[float] = None
seed: Optional[int] = None
def __init__(
self,
temperature: float | None = None,
max_output_tokens: int | None = None,
top_p: float | None = None,
top_k: int | None = None,
response_mime_type: str | None = None,
candidate_count: int | None = None,
stop_sequences: list | None = None,
frequency_penalty: float | None = None,
presence_penalty: float | None = None,
seed: int | None = None,
temperature: Optional[float] = None,
max_output_tokens: Optional[int] = None,
top_p: Optional[float] = None,
top_k: Optional[int] = None,
response_mime_type: Optional[str] = None,
candidate_count: Optional[int] = None,
stop_sequences: Optional[list] = None,
frequency_penalty: Optional[float] = None,
presence_penalty: Optional[float] = None,
seed: Optional[int] = None,
) -> None:
locals_ = locals().copy()
for key, value in locals_.items():
@ -228,8 +232,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return super().get_config()
def get_json_schema_from_pydantic_object(
self, response_format: Union[type["BaseModel"], dict] | None
) -> dict | None:
self, response_format: Optional[Union[Type["BaseModel"], dict]]
) -> Optional[dict]:
"""
Override to use Pydantic's model_json_schema() instead of OpenAI's
to_strict_json_schema().
@ -285,7 +289,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return False
@staticmethod
def _forward_gemini_function_call_id(model: str, custom_llm_provider: str | None = None) -> bool:
def _forward_gemini_function_call_id(model: str, custom_llm_provider: Optional[str] = None) -> bool:
"""
Whether to include `id` on function_call / function_response parts.
@ -305,7 +309,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return False
return True
def get_supported_openai_params(self, model: str) -> list[str]:
def get_supported_openai_params(self, model: str) -> List[str]:
supported_params = [
"temperature",
"top_p",
@ -340,7 +344,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
supported_params.append("thinking")
return supported_params
def map_tool_choice_values(self, model: str, tool_choice: Union[str, dict]) -> ToolConfig | None:
def map_tool_choice_values(self, model: str, tool_choice: Union[str, dict]) -> Optional[ToolConfig]:
if tool_choice == "none":
return ToolConfig(functionCallingConfig=FunctionCallingConfig(mode="NONE"))
elif tool_choice == "required":
@ -466,7 +470,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return transformed_config
def _extract_google_maps_retrieval_config(self, google_maps_config: dict) -> tuple[dict, dict | None]:
def _extract_google_maps_retrieval_config(self, google_maps_config: dict) -> Tuple[dict, Optional[dict]]:
"""
Extract location configuration from googleMaps tool for Vertex AI toolConfig.
@ -504,7 +508,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return cleaned_config, retrieval_config
def get_tool_value(self, tool: dict, tool_name: str) -> dict | None:
def get_tool_value(self, tool: dict, tool_name: str) -> Optional[dict]:
"""
Helper function to get tool value handling both camelCase and underscore_case variants
@ -529,10 +533,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _resolve_search_tool_conflict(
gtool_func_declarations: list,
googleSearch: dict | None,
googleSearchRetrieval: dict | None,
enterpriseWebSearch: dict | None,
urlContext: dict | None,
googleSearch: Optional[dict],
googleSearchRetrieval: Optional[dict],
enterpriseWebSearch: Optional[dict],
urlContext: Optional[dict],
optional_params: dict,
) -> tuple:
"""
@ -576,7 +580,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return googleSearch, googleSearchRetrieval, enterpriseWebSearch, urlContext
def _map_function(self, value: list[dict], optional_params: dict) -> list[Tools]:
def _map_function(self, value: List[dict], optional_params: dict) -> List[Tools]:
"""
Map OpenAI-style tools/functions to Vertex AI format.
@ -592,21 +596,21 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
googleMaps tools contain location data
"""
gtool_func_declarations = []
googleSearch: dict | None = None
googleSearchRetrieval: dict | None = None
enterpriseWebSearch: dict | None = None
urlContext: dict | None = None
code_execution: dict | None = None
googleMaps: dict | None = None
google_maps_retrieval_config: dict | None = None
computerUse: dict | None = None
googleSearch: Optional[dict] = None
googleSearchRetrieval: Optional[dict] = None
enterpriseWebSearch: Optional[dict] = None
urlContext: Optional[dict] = None
code_execution: Optional[dict] = None
googleMaps: Optional[dict] = None
google_maps_retrieval_config: Optional[dict] = None
computerUse: Optional[dict] = None
# remove 'additionalProperties' from tools
value = _remove_additional_properties(value)
# remove 'strict' from tools
value = _remove_strict_from_schema(value)
for tool in value:
openai_function_object: ChatCompletionToolParamFunctionChunk | None = None
openai_function_object: Optional[ChatCompletionToolParamFunctionChunk] = None
if "function" in tool: # tools list
_openai_function_object = ChatCompletionToolParamFunctionChunk( # type: ignore
**tool["function"]
@ -696,7 +700,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
# Build list of Tool objects - each Tool should contain exactly one type
# per Vertex AI API spec: "A Tool object should contain exactly one type of Tool"
_tools_list: list[Tools] = []
_tools_list: List[Tools] = []
(
googleSearch,
@ -814,7 +818,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _map_reasoning_effort_to_thinking_budget(
reasoning_effort: str,
model: str | None = None,
model: Optional[str] = None,
) -> GeminiThinkingConfig:
if reasoning_effort == "minimal":
# Use model-specific minimum thinking budget or fallback
@ -863,7 +867,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _map_reasoning_effort_to_thinking_level(
reasoning_effort: str,
model: str | None = None,
model: Optional[str] = None,
) -> GeminiThinkingConfig:
"""
Map reasoning_effort to thinking_level for Gemini 3+ models.
@ -909,12 +913,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
raise ValueError(f"Invalid reasoning effort: {reasoning_effort}")
@staticmethod
def _is_thinking_budget_zero(thinking_budget: int | None) -> bool:
def _is_thinking_budget_zero(thinking_budget: Optional[int]) -> bool:
return thinking_budget is not None and thinking_budget == 0
@staticmethod
def _validate_thinking_config_conflicts(
optional_params: dict,
optional_params: Dict,
param_name: str,
param_description: str = "thinking_budget",
) -> None:
@ -935,7 +939,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _validate_thinking_level_conflicts(
optional_params: dict,
optional_params: Dict,
) -> None:
"""
Validate that thinking_level and thinking_budget are not both specified.
@ -955,7 +959,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _map_thinking_param(
thinking_param: AnthropicThinkingParam,
model: str | None = None,
model: Optional[str] = None,
) -> GeminiThinkingConfig:
thinking_enabled = thinking_param.get("type") == "enabled"
thinking_budget = thinking_param.get("budget_tokens")
@ -1060,8 +1064,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _apply_include_server_side_tool_invocations(
non_default_params: dict,
optional_params: dict,
non_default_params: Dict,
optional_params: Dict,
) -> None:
"""
Set include_server_side_tool_invocations before tools are mapped.
@ -1080,11 +1084,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
non_default_params: Dict,
optional_params: Dict,
model: str,
drop_params: bool,
) -> dict:
) -> Dict:
self._apply_include_server_side_tool_invocations(non_default_params, optional_params)
gemini_sampling_params_warned: bool = False
for param, value in non_default_params.items():
@ -1176,7 +1180,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
elif param == "reasoning_effort":
# Extract effort value - handle both string and dict formats
# Dict format comes from OpenAI Agents SDK: {"effort": "high", "summary": "auto"}
effort_value: str | None = None
effort_value: Optional[str] = None
if isinstance(value, str):
effort_value = value
elif isinstance(value, dict):
@ -1252,7 +1256,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
optional_params[mapped_params[param]] = value
return optional_params
def get_eu_regions(self) -> list[str]:
def get_eu_regions(self) -> List[str]:
"""
Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions
"""
@ -1289,7 +1293,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return model
@staticmethod
def _is_model_gemini_spec_model(model: str | None) -> bool:
def _is_model_gemini_spec_model(model: Optional[str]) -> bool:
"""
Returns true if user is trying to call custom model in `/gemini` request/response format
"""
@ -1312,7 +1316,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return model.split("/")[-1]
return model
def get_flagged_finish_reasons(self) -> dict[str, str]:
def get_flagged_finish_reasons(self) -> Dict[str, str]:
"""
Return Dictionary of finish reasons which indicate response was flagged
@ -1349,7 +1353,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
@staticmethod
def get_finish_reason_mapping() -> dict[str, OpenAIChatCompletionFinishReason]:
def get_finish_reason_mapping() -> Dict[str, OpenAIChatCompletionFinishReason]:
"""
Return Dictionary of Gemini/Vertex AI finish reasons and their
OpenAI-compatible mappings.
@ -1368,9 +1372,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
return exception_string
def get_assistant_content_message(self, parts: list[HttpxPartType]) -> tuple[str | None, str | None]:
content_str: str | None = None
reasoning_content_str: str | None = None
def get_assistant_content_message(self, parts: List[HttpxPartType]) -> Tuple[Optional[str], Optional[str]]:
content_str: Optional[str] = None
reasoning_content_str: Optional[str] = None
for part in parts:
_content_str = ""
@ -1409,7 +1413,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return content_str, reasoning_content_str
def _extract_thinking_blocks_from_parts(self, parts: list[HttpxPartType]) -> list[ChatCompletionThinkingBlock]:
def _extract_thinking_blocks_from_parts(self, parts: List[HttpxPartType]) -> List[ChatCompletionThinkingBlock]:
"""Extract thinking blocks from parts if present.
Per Google's docs (https://ai.google.dev/gemini-api/docs/thinking):
@ -1418,7 +1422,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
it does NOT indicate that the content is thinking (a part can have
thoughtSignature without thought: true, e.g., function calls)
"""
thinking_blocks: list[ChatCompletionThinkingBlock] = []
thinking_blocks: List[ChatCompletionThinkingBlock] = []
for part in parts:
if part.get("thought") is True:
thinking_text = part.get("text", "")
@ -1432,7 +1436,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
thinking_blocks.append(block)
return thinking_blocks
def _extract_thought_signatures_from_parts(self, parts: list[HttpxPartType]) -> list[str] | None:
def _extract_thought_signatures_from_parts(self, parts: List[HttpxPartType]) -> Optional[List[str]]:
"""Extract thoughtSignature values from parts.
Per Google's docs, thoughtSignature is returned for multi-turn context preservation
@ -1442,7 +1446,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
Returns:
List of thoughtSignature strings if any are found, None otherwise
"""
signatures: list[str] = []
signatures: List[str] = []
for part in parts:
signature = part.get("thoughtSignature")
if signature is not None:
@ -1451,8 +1455,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _extract_server_side_tool_invocations(
parts: list[HttpxPartType],
) -> list[dict[str, Any]] | None:
parts: List[HttpxPartType],
) -> Optional[List[Dict[str, Any]]]:
"""Extract server-side tool invocations (toolCall/toolResponse) from parts.
These are returned by Gemini when context circulation is enabled
@ -1463,15 +1467,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
Returns:
List of server-side invocation dicts if any found, None otherwise.
"""
invocations: list[dict[str, Any]] = []
invocations: List[Dict[str, Any]] = []
# Index toolCalls by id so we can pair them with responses
tool_calls_by_id: dict[str, dict[str, Any]] = {}
tool_responses_by_id: dict[str, dict[str, Any]] = {}
tool_calls_by_id: Dict[str, Dict[str, Any]] = {}
tool_responses_by_id: Dict[str, Dict[str, Any]] = {}
for part in parts:
if "toolCall" in part:
tc = part["toolCall"]
entry: dict[str, Any] = {
entry: Dict[str, Any] = {
"tool_type": tc.get("toolType"),
"id": tc.get("id"),
"args": tc.get("args"),
@ -1511,9 +1515,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return invocations if invocations else None
def _extract_image_response_from_parts(self, parts: list[HttpxPartType]) -> list[ImageURLListItem] | None:
def _extract_image_response_from_parts(self, parts: List[HttpxPartType]) -> Optional[List[ImageURLListItem]]:
"""Extract image response from parts if present"""
images: list[ImageURLListItem] = []
images: List[ImageURLListItem] = []
for part in parts:
if "inlineData" in part:
inline_data = part.get("inlineData", {})
@ -1531,7 +1535,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
return images
def _extract_audio_response_from_parts(self, parts: list[HttpxPartType]) -> ChatCompletionAudioResponse | None:
def _extract_audio_response_from_parts(self, parts: List[HttpxPartType]) -> Optional[ChatCompletionAudioResponse]:
"""Extract audio response from parts if present"""
for part in parts:
if "text" in part:
@ -1569,16 +1573,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _transform_parts(
parts: list[HttpxPartType],
parts: List[HttpxPartType],
cumulative_tool_call_idx: int,
is_function_call: bool | None,
) -> tuple[
ChatCompletionToolCallFunctionChunk | None,
list[ChatCompletionToolCallChunk] | None,
is_function_call: Optional[bool],
) -> Tuple[
Optional[ChatCompletionToolCallFunctionChunk],
Optional[List[ChatCompletionToolCallChunk]],
int,
]:
function: ChatCompletionToolCallFunctionChunk | None = None
_tools: list[ChatCompletionToolCallChunk] = []
function: Optional[ChatCompletionToolCallFunctionChunk] = None
_tools: List[ChatCompletionToolCallChunk] = []
for part in parts:
if "functionCall" in part:
_function_chunk: ChatCompletionToolCallFunctionChunk = {
@ -1593,7 +1597,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
gemini_call_id = part["functionCall"].get("id")
if is_function_call is True:
function_dict: dict[str, Any] = dict(_function_chunk)
function_dict: Dict[str, Any] = dict(_function_chunk)
if thought_signature:
if "provider_specific_fields" not in function_dict:
function_dict["provider_specific_fields"] = {}
@ -1622,22 +1626,22 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
_tools.append(_tool_response_chunk)
cumulative_tool_call_idx += 1
if len(_tools) == 0:
tools: list[ChatCompletionToolCallChunk] | None = None
tools: Optional[List[ChatCompletionToolCallChunk]] = None
else:
tools = _tools
return function, tools, cumulative_tool_call_idx
@staticmethod
def _transform_logprobs(
logprobs_result: LogprobsResult | None,
) -> ChoiceLogprobs | None:
logprobs_result: Optional[LogprobsResult],
) -> Optional[ChoiceLogprobs]:
if logprobs_result is None:
return None
if "chosenCandidates" not in logprobs_result:
return None
logprobs_list: list[ChatCompletionTokenLogprob] = []
logprobs_list: List[ChatCompletionTokenLogprob] = []
for index, candidate in enumerate(logprobs_result["chosenCandidates"]):
top_logprobs: list[TopLogprob] = []
top_logprobs: List[TopLogprob] = []
if "topCandidates" in logprobs_result and index < len(logprobs_result["topCandidates"]):
top_candidates_for_index = logprobs_result["topCandidates"][index]["candidates"]
@ -1744,16 +1748,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
) -> Usage:
if completion_response is not None and "usageMetadata" not in completion_response:
raise ValueError(f"usageMetadata not found in completion_response. Got={completion_response}")
cached_tokens: int | None = None
cached_tokens: Optional[int] = None
# Separate variables for prompt tokens by modality
prompt_audio_tokens: int | None = None
prompt_image_tokens: int | None = None
prompt_text_tokens: int | None = None
prompt_video_tokens: int | None = None
prompt_tokens_details: PromptTokensDetailsWrapper | None = None
reasoning_tokens: int | None = None
response_tokens: int | None = None
response_tokens_details: CompletionTokensDetailsWrapper | None = None
prompt_audio_tokens: Optional[int] = None
prompt_image_tokens: Optional[int] = None
prompt_text_tokens: Optional[int] = None
prompt_video_tokens: Optional[int] = None
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None
reasoning_tokens: Optional[int] = None
response_tokens: Optional[int] = None
response_tokens_details: Optional[CompletionTokensDetailsWrapper] = None
usage_metadata = completion_response["usageMetadata"]
def _get_token_count(detail: Mapping[str, Any]) -> int:
@ -1832,10 +1836,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
## Parse cacheTokensDetails (breakdown of cached tokens by modality)
## When explicit caching is used, Gemini provides this field to show which modalities were cached
cached_text_tokens: int | None = None
cached_audio_tokens: int | None = None
cached_image_tokens: int | None = None
cached_video_tokens: int | None = None
cached_text_tokens: Optional[int] = None
cached_audio_tokens: Optional[int] = None
cached_image_tokens: Optional[int] = None
cached_video_tokens: Optional[int] = None
if "cacheTokensDetails" in usage_metadata:
for detail in usage_metadata["cacheTokensDetails"]:
@ -1908,8 +1912,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _check_finish_reason(
chat_completion_message: ChatCompletionResponseMessage | None,
finish_reason: str | None,
chat_completion_message: Optional[ChatCompletionResponseMessage],
finish_reason: Optional[str],
) -> OpenAIChatCompletionFinishReason:
from litellm.litellm_core_utils.core_helpers import map_finish_reason
@ -1925,7 +1929,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _check_prompt_level_content_filter(
processed_chunk: GenerateContentResponseBody,
response_id: str | None,
response_id: Optional[str],
) -> Optional["ModelResponseStream"]:
"""
Check if prompt is blocked due to content filtering at the prompt level.
@ -1969,8 +1973,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return None
@staticmethod
def _calculate_web_search_requests(grounding_metadata: list[dict]) -> int | None:
web_search_requests: int | None = None
def _calculate_web_search_requests(grounding_metadata: List[dict]) -> Optional[int]:
web_search_requests: Optional[int] = None
if grounding_metadata and isinstance(grounding_metadata, list) and len(grounding_metadata) > 0:
for grounding_metadata_item in grounding_metadata:
@ -1986,10 +1990,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
chat_completion_message: ChatCompletionResponseMessage,
candidate: Candidates,
idx: int,
tools: list[ChatCompletionToolCallChunk] | None,
functions: ChatCompletionToolCallFunctionChunk | None,
chat_completion_logprobs: ChoiceLogprobs | None,
image_response: list[ImageURLListItem] | None,
tools: Optional[List[ChatCompletionToolCallChunk]],
functions: Optional[ChatCompletionToolCallFunctionChunk],
chat_completion_logprobs: Optional[ChoiceLogprobs],
image_response: Optional[List[ImageURLListItem]],
) -> StreamingChoices:
"""
Helper method to create a streaming choice object for Vertex AI
@ -2021,7 +2025,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _extract_candidate_metadata(
candidate: Candidates,
) -> tuple[list[dict], list[dict], list, list]:
) -> Tuple[List[dict], List[dict], List, List]:
"""
Extract metadata from a single candidate response.
@ -2031,10 +2035,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
safety_ratings: List
citation_metadata: List
"""
grounding_metadata: list[dict] = []
url_context_metadata: list[dict] = []
safety_ratings: list = []
citation_metadata: list = []
grounding_metadata: List[dict] = []
url_context_metadata: List[dict] = []
safety_ratings: List = []
citation_metadata: List = []
if "groundingMetadata" in candidate:
if isinstance(candidate["groundingMetadata"], list):
@ -2079,10 +2083,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _set_stream_metadata_on_response(
model_response: Any,
grounding_metadata: list[dict],
url_context_metadata: list[dict],
safety_ratings: list[dict],
citation_metadata: list[dict],
grounding_metadata: List[dict],
url_context_metadata: List[dict],
safety_ratings: List[dict],
citation_metadata: List[dict],
) -> None:
setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore
if grounding_metadata:
@ -2102,10 +2106,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
def apply_assembled_streaming_response_metadata(
self,
response: ModelResponse,
chunks: list[Any],
chunks: List[Any],
) -> None:
for field_name in VERTEX_AI_PROVIDER_METADATA_FIELDS:
merged: list[Any] = []
merged: List[Any] = []
for chunk in chunks:
value = VertexGeminiConfig._get_stream_chunk_attr(chunk, field_name)
if not value:
@ -2120,14 +2124,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _convert_grounding_metadata_to_annotations(
grounding_metadata: list[dict],
content_text: str | None,
) -> list[ChatCompletionAnnotation]:
grounding_metadata: List[dict],
content_text: Optional[str],
) -> List[ChatCompletionAnnotation]:
"""
Convert Vertex AI grounding metadata to OpenAI-style annotations.
"""
annotations: list[ChatCompletionAnnotation] = []
annotations: List[ChatCompletionAnnotation] = []
for metadata in grounding_metadata:
# Extract groundingSupports - these map text segments to sources
@ -2135,7 +2139,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
grounding_chunks = metadata.get("groundingChunks", [])
# Build a map of chunk indices to web URIs
chunk_to_uri_map: dict[int, dict[str, str]] = {}
chunk_to_uri_map: Dict[int, Dict[str, str]] = {}
for idx, chunk in enumerate(grounding_chunks):
if "web" in chunk:
web_data = chunk["web"]
@ -2175,11 +2179,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _process_candidates(
_candidates: list[Candidates],
_candidates: List[Candidates],
model_response: Union[ModelResponse, "ModelResponseStream"],
standard_optional_params: dict,
cumulative_tool_call_index: int = 0,
) -> tuple[list[dict], list[dict], list, list, int]:
) -> Tuple[List[dict], List[dict], List, List, int]:
"""
Helper method to process candidates and extract metadata
@ -2195,19 +2199,19 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
from litellm.types.utils import ModelResponseStream
grounding_metadata: list[dict] = []
url_context_metadata: list[dict] = []
image_response: list[ImageURLListItem] | None = None
safety_ratings: list = []
citation_metadata: list = []
grounding_metadata: List[dict] = []
url_context_metadata: List[dict] = []
image_response: Optional[List[ImageURLListItem]] = None
safety_ratings: List = []
citation_metadata: List = []
chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"}
chat_completion_logprobs: ChoiceLogprobs | None = None
tools: list[ChatCompletionToolCallChunk] | None = []
functions: ChatCompletionToolCallFunctionChunk | None = None
thinking_blocks: list[ChatCompletionThinkingBlock] | None = None
reasoning_content: str | None = None
thought_signatures: Any | None = None
server_side_tool_invocations: list[dict[str, Any]] | None = None
chat_completion_logprobs: Optional[ChoiceLogprobs] = None
tools: Optional[List[ChatCompletionToolCallChunk]] = []
functions: Optional[ChatCompletionToolCallFunctionChunk] = None
thinking_blocks: Optional[List[ChatCompletionThinkingBlock]] = None
reasoning_content: Optional[str] = None
thought_signatures: Optional[Any] = None
server_side_tool_invocations: Optional[List[Dict[str, Any]]] = None
for idx, candidate in enumerate(_candidates):
if "content" not in candidate:
@ -2254,11 +2258,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
if audio_response is not None:
cast(dict[str, Any], chat_completion_message)["audio"] = audio_response
cast(Dict[str, Any], chat_completion_message)["audio"] = audio_response
chat_completion_message["content"] = None # OpenAI spec
if image_response is not None:
# Handle image response - combine with text content into structured format
cast(dict[str, Any], chat_completion_message)["images"] = image_response
cast(Dict[str, Any], chat_completion_message)["images"] = image_response
if content is not None:
chat_completion_message["content"] = content
@ -2360,13 +2364,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LoggingClass,
request_data: dict,
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
request_data: Dict,
messages: List[AllMessageValues],
optional_params: Dict,
litellm_params: Dict,
encoding: Any,
api_key: str | None = None,
json_mode: bool | None = None,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ModelResponse:
## LOGGING
logging_obj.post_call(
@ -2433,11 +2437,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
response_id = completion_response.get("responseId")
if response_id:
model_response.id = response_id
url_context_metadata: list[dict] = []
url_context_metadata: List[dict] = []
try:
grounding_metadata: list[dict] = []
safety_ratings: list[dict] = []
citation_metadata: list[dict] = []
grounding_metadata: List[dict] = []
safety_ratings: List[dict] = []
citation_metadata: List[dict] = []
if _candidates:
(
grounding_metadata,
@ -2501,10 +2505,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
def _transform_messages(
self,
messages: list[AllMessageValues],
model: str | None = None,
litellm_params: dict | None = None,
) -> list[ContentType]:
messages: List[AllMessageValues],
model: Optional[str] = None,
litellm_params: Optional[dict] = None,
) -> List[ContentType]:
return _gemini_convert_messages_with_history(
messages=messages,
model=model,
@ -2513,30 +2517,30 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers]
) -> BaseLLMException:
return VertexAIError(message=error_message, status_code=status_code, headers=headers)
def transform_request(
self,
model: str,
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
messages: List[AllMessageValues],
optional_params: Dict,
litellm_params: Dict,
headers: Dict,
) -> Dict:
raise NotImplementedError("Vertex AI has a custom implementation of transform_request. Needs sync + async.")
def validate_environment(
self,
headers: dict | None,
headers: Optional[Dict],
model: str,
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Union[str, dict] | None = None,
api_base: str | None = None,
) -> dict:
messages: List[AllMessageValues],
optional_params: Dict,
litellm_params: Dict,
api_key: Optional[Union[str, Dict]] = None,
api_base: Optional[str] = None,
) -> Dict:
default_headers = {
"Content-Type": "application/json",
}
@ -2551,8 +2555,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
async def make_call(
client: AsyncHTTPHandler | None, # module-level client
gemini_client: AsyncHTTPHandler | None, # if passed by user
client: Optional[AsyncHTTPHandler], # module-level client
gemini_client: Optional[AsyncHTTPHandler], # if passed by user
api_base: str,
headers: dict,
data: str,
@ -2603,8 +2607,8 @@ async def make_call(
def make_sync_call(
client: HTTPHandler | None, # module-level client
gemini_client: HTTPHandler | None, # if passed by user
client: Optional[HTTPHandler], # module-level client
gemini_client: Optional[HTTPHandler], # if passed by user
api_base: str,
headers: dict,
data: str,
@ -2659,20 +2663,20 @@ class VertexLLM(VertexBase):
model_response: ModelResponse,
print_verbose: Callable,
data: dict,
timeout: Union[float, httpx.Timeout] | None,
timeout: Optional[Union[float, httpx.Timeout]],
encoding,
logging_obj,
stream,
optional_params: dict,
litellm_params: dict,
logger_fn=None,
api_base: str | None = None,
client: AsyncHTTPHandler | None = None,
vertex_project: str | None = None,
vertex_location: str | None = None,
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None,
gemini_api_key: str | None = None,
extra_headers: dict | None = None,
api_base: Optional[str] = None,
client: Optional[AsyncHTTPHandler] = None,
vertex_project: Optional[str] = None,
vertex_location: Optional[str] = None,
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None,
gemini_api_key: Optional[str] = None,
extra_headers: Optional[dict] = None,
) -> CustomStreamWrapper:
should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params)
@ -2755,20 +2759,20 @@ class VertexLLM(VertexBase):
custom_llm_provider: Literal[
"vertex_ai", "vertex_ai_beta", "gemini"
], # if it's vertex_ai or gemini (google ai studio)
timeout: Union[float, httpx.Timeout] | None,
timeout: Optional[Union[float, httpx.Timeout]],
encoding,
logging_obj,
stream,
optional_params: dict,
litellm_params: dict,
logger_fn=None,
api_base: str | None = None,
client: AsyncHTTPHandler | None = None,
vertex_project: str | None = None,
vertex_location: str | None = None,
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None,
gemini_api_key: str | None = None,
extra_headers: dict | None = None,
api_base: Optional[str] = None,
client: Optional[AsyncHTTPHandler] = None,
vertex_project: Optional[str] = None,
vertex_location: Optional[str] = None,
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None,
gemini_api_key: Optional[str] = None,
extra_headers: Optional[dict] = None,
) -> Union[ModelResponse, CustomStreamWrapper]:
should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params)
@ -2877,18 +2881,18 @@ class VertexLLM(VertexBase):
logging_obj,
optional_params: dict,
acompletion: bool,
timeout: Union[float, httpx.Timeout] | None,
vertex_project: str | None,
vertex_location: str | None,
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None,
gemini_api_key: str | None,
timeout: Optional[Union[float, httpx.Timeout]],
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES],
gemini_api_key: Optional[str],
litellm_params: dict,
logger_fn=None,
extra_headers: dict | None = None,
client: Union[AsyncHTTPHandler, HTTPHandler] | None = None,
api_base: str | None = None,
extra_headers: Optional[dict] = None,
client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None,
api_base: Optional[str] = None,
) -> Union[ModelResponse, CustomStreamWrapper]:
stream: bool | None = optional_params.pop("stream", None) # type: ignore
stream: Optional[bool] = optional_params.pop("stream", None) # type: ignore
transform_request_params = {
"gemini_api_key": gemini_api_key,
@ -3076,7 +3080,7 @@ class ModelResponseIterator:
streaming_response,
sync_stream: bool,
logging_obj: LoggingClass,
response_headers: dict[str, str] | None = None,
response_headers: Optional[Dict[str, str]] = None,
response: httpx.Response | None = None,
):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
@ -3121,9 +3125,9 @@ class ModelResponseIterator:
def _apply_stream_candidates(
self,
_candidates: list[Candidates],
_candidates: List[Candidates],
model_response: Any,
) -> tuple[list[dict], list[dict], list[dict], list[dict]]:
) -> Tuple[List[dict], List[dict], List[dict], List[dict]]:
(
grounding_metadata,
url_context_metadata,
@ -3202,8 +3206,8 @@ class ModelResponseIterator:
self,
processed_chunk: Any,
model_response: Any,
grounding_metadata: list[dict],
) -> Usage | None:
grounding_metadata: List[dict],
) -> Optional[Usage]:
if "usageMetadata" not in processed_chunk:
return None
@ -3250,12 +3254,12 @@ class ModelResponseIterator:
if blocked_response is not None:
model_response = blocked_response
grounding_metadata: list[dict] = []
url_context_metadata: list[dict] = []
safety_ratings: list[dict] = []
citation_metadata: list[dict] = []
grounding_metadata: List[dict] = []
url_context_metadata: List[dict] = []
safety_ratings: List[dict] = []
citation_metadata: List[dict] = []
_candidates: list[Candidates] | None = processed_chunk.get("candidates")
_candidates: Optional[List[Candidates]] = processed_chunk.get("candidates")
if _candidates:
(
grounding_metadata,

View file

@ -3,7 +3,7 @@
import importlib
import os
from datetime import datetime, timezone
from typing import Any, Literal, Optional, cast
from typing import Any, Dict, List, Literal, Optional, Set, Type, cast
from pydantic import ValidationError
@ -54,7 +54,7 @@ guardrail_initializer_registry = {
SupportedGuardrailIntegrations.LLM_AS_A_JUDGE.value: initialize_llm_as_a_judge,
}
guardrail_class_registry: dict[str, type[CustomGuardrail]] = {
guardrail_class_registry: Dict[str, Type[CustomGuardrail]] = {
SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail
}
@ -220,7 +220,7 @@ class GuardrailRegistry:
###########################################################
########### In memory management helpers for guardrails ###########
############################################################
def get_initialized_guardrail_callback(self, guardrail_name: str) -> CustomGuardrail | None:
def get_initialized_guardrail_callback(self, guardrail_name: str) -> Optional[CustomGuardrail]:
"""
Returns the initialized guardrail callback for a given guardrail name
"""
@ -314,7 +314,7 @@ class GuardrailRegistry:
@staticmethod
async def get_all_guardrails_from_db(
prisma_client: PrismaClient,
) -> list[Guardrail]:
) -> List[Guardrail]:
"""
Get all active guardrails from the database.
Only rows with status == "active" are returned (pending_review and rejected are excluded).
@ -325,7 +325,7 @@ class GuardrailRegistry:
order={"created_at": "desc"},
)
guardrails: list[Guardrail] = []
guardrails: List[Guardrail] = []
for guardrail in guardrails_from_db:
guardrails.append(Guardrail(**(dict(guardrail)))) # type: ignore
@ -333,7 +333,7 @@ class GuardrailRegistry:
except Exception as e:
raise Exception(f"Error getting guardrails from DB: {str(e)}")
async def get_guardrail_by_id_from_db(self, guardrail_id: str, prisma_client: PrismaClient) -> Guardrail | None:
async def get_guardrail_by_id_from_db(self, guardrail_id: str, prisma_client: PrismaClient) -> Optional[Guardrail]:
"""
Get a guardrail by its ID from the database
"""
@ -349,7 +349,9 @@ class GuardrailRegistry:
except Exception as e:
raise Exception(f"Error getting guardrail from DB: {str(e)}")
async def get_guardrail_by_name_from_db(self, guardrail_name: str, prisma_client: PrismaClient) -> Guardrail | None:
async def get_guardrail_by_name_from_db(
self, guardrail_name: str, prisma_client: PrismaClient
) -> Optional[Guardrail]:
"""
Get a guardrail by its name from the database
"""
@ -372,17 +374,17 @@ class InMemoryGuardrailHandler:
"""
def __init__(self):
self.IN_MEMORY_GUARDRAILS: dict[str, Guardrail] = {}
self.IN_MEMORY_GUARDRAILS: Dict[str, Guardrail] = {}
"""
Guardrail id to Guardrail object mapping
"""
self.guardrail_id_to_custom_guardrail: dict[str, CustomGuardrail | None] = {}
self.guardrail_id_to_custom_guardrail: Dict[str, Optional[CustomGuardrail]] = {}
"""
Guardrail id to CustomGuardrail object mapping
"""
self._sources: dict[str, Literal["db", "config"]] = {}
self._sources: Dict[str, Literal["db", "config"]] = {}
"""
Guardrail id to provenance marker. "db" entries are reconciled against
the DB on each polling tick; "config" entries are owned by proxy_config.yaml
@ -392,10 +394,10 @@ class InMemoryGuardrailHandler:
def initialize_guardrail(
self,
guardrail: Guardrail,
config_file_path: str | None = None,
config_file_path: Optional[str] = None,
llm_router: Optional["Router"] = None,
source: Literal["db", "config"] = "config",
) -> Guardrail | None:
) -> Optional[Guardrail]:
"""
Initialize a guardrail from a dictionary and add it to the litellm callback manager
@ -411,7 +413,7 @@ class InMemoryGuardrailHandler:
self._sources[guardrail_id] = source
return self.IN_MEMORY_GUARDRAILS[guardrail_id]
custom_guardrail_callback: CustomGuardrail | None = None
custom_guardrail_callback: Optional[CustomGuardrail] = None
litellm_params_data = guardrail["litellm_params"]
verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data)
@ -487,11 +489,11 @@ class InMemoryGuardrailHandler:
def initialize_custom_guardrail(
self,
guardrail: dict,
guardrail: Dict,
guardrail_type: str,
litellm_params: LitellmParams,
config_file_path: str | None = None,
) -> CustomGuardrail | None:
config_file_path: Optional[str] = None,
) -> Optional[CustomGuardrail]:
"""
Initialize a Custom Guardrail from a python file or module path
@ -576,25 +578,25 @@ class InMemoryGuardrailHandler:
litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback)
def list_in_memory_guardrails(self) -> list[Guardrail]:
def list_in_memory_guardrails(self) -> List[Guardrail]:
"""
List all guardrails in memory
"""
return list(self.IN_MEMORY_GUARDRAILS.values())
def get_guardrail_by_id(self, guardrail_id: str) -> Guardrail | None:
def get_guardrail_by_id(self, guardrail_id: str) -> Optional[Guardrail]:
"""
Get a guardrail by its ID from memory
"""
return self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
def get_source(self, guardrail_id: str) -> Literal["db", "config"] | None:
def get_source(self, guardrail_id: str) -> Optional[Literal["db", "config"]]:
"""
Return the provenance of an in-memory guardrail.
"""
return self._sources.get(guardrail_id)
def reconcile_db_guardrails(self, db_guardrail_ids: set[str]) -> list[str]:
def reconcile_db_guardrails(self, db_guardrail_ids: Set[str]) -> List[str]:
"""
Drop in-memory entries that originated from the DB but are no longer
present in db_guardrail_ids. Config-loaded guardrails are never touched.
@ -617,8 +619,8 @@ class InMemoryGuardrailHandler:
@staticmethod
def _normalize_litellm_params_for_comparison(
params: Any | None,
) -> dict[str, Any] | None:
params: Optional[Any],
) -> Optional[Dict[str, Any]]:
"""
Render litellm_params to a canonical dict so an in-memory LitellmParams and
the raw dict loaded from the DB compare equal when they describe the same
@ -682,9 +684,9 @@ class InMemoryGuardrailHandler:
def reinitialize_guardrail(
self,
guardrail: Guardrail,
config_file_path: str | None = None,
config_file_path: Optional[str] = None,
source: Literal["db", "config"] = "config",
) -> Guardrail | None:
) -> Optional[Guardrail]:
"""
Force re-initialization of a guardrail even if it exists in memory.
Removes old callback from litellm.callbacks and creates fresh instance.
@ -701,7 +703,9 @@ class InMemoryGuardrailHandler:
# Initialize fresh (will add new callback to litellm.callbacks)
return self.initialize_guardrail(guardrail=guardrail, config_file_path=config_file_path, source=source)
def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None:
def sync_guardrail_from_db(
self, guardrail: Guardrail, config_file_path: Optional[str] = None
) -> Optional[Guardrail]:
"""
Sync a guardrail from DB - initializes if new, re-initializes if changed.
This is the method to call during DB polling.