fix(lint): fix UP006/UP045/F401 in changed files; add depth guard to check_any_discipline

- Replace Dict/List/Optional/Tuple typing imports with built-in equivalents
  (UP006, UP045) across files touched in this PR diff, then clean up
  the now-unused typing imports (F401).
- Add _MAX_CONTAINS_ANY_DEPTH guard to check_any_discipline.contains_any()
  to prevent RecursionError on deeply-nested mypy types.
This commit is contained in:
Yaniv Israel 2026-06-17 18:29:29 +03:00
parent be3e2706b0
commit 6c6ed854b1
11 changed files with 606 additions and 582 deletions

37
litellm/build-and-push.sh Normal file
View file

@ -0,0 +1,37 @@
#!/bin/bash
set -e
# Configuration
ECR_REPO="420049223852.dkr.ecr.eu-central-1.amazonaws.com/litellm-deepkeep"
TAG="${1:-dev-$(date +%Y%m%d-%H%M%S)}"
FULL_IMAGE="${ECR_REPO}:${TAG}"
echo "================================================"
echo "Building LiteLLM Docker image with UI"
echo "Image: ${FULL_IMAGE}"
echo "================================================"
# Login to ECR
echo "Logging in to ECR..."
aws ecr get-login-password --region eu-central-1 | docker login --username AWS --password-stdin 420049223852.dkr.ecr.eu-central-1.amazonaws.com
# Build the image
echo "Building Docker image..."
docker build -t "${FULL_IMAGE}" -f Dockerfile .
# Push the image
echo "Pushing image to ECR..."
docker push "${FULL_IMAGE}"
echo "================================================"
echo "✅ Image pushed successfully!"
echo ""
echo "To use in Tilt, update your tilt_config.yaniv-v2.yaml:"
echo ""
echo "litellm:"
echo " image:"
echo " repository: ${ECR_REPO}"
echo " tag: ${TAG}"
echo ""
echo "Or run: tilt trigger litellm"
echo "================================================"

View file

@ -0,0 +1,22 @@
model_list:
- model_name: fake-openai-endpoint
litellm_params:
model: openai/fake-model
api_key: fake-key
api_base: https://exampleopenaiendpoint-production.up.railway.app/
general_settings:
master_key: sk-1234
litellm_settings:
drop_params: True
telemetry: False
guardrails:
- guardrail_name: deepkeep-firewall
litellm_params:
guardrail: deepkeep
mode: [pre_call, post_call]
api_key: "1D-VOZhxhsgG2o8pnbDcXaDmzLOWSwZpgkGZPDyVQBk"
api_base: "http://localhost:8081/api"
deepkeep_firewall_id: "063e4e5ba5c2be55"

View file

@ -7,7 +7,7 @@ Users can define
"""
import copy
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
from typing import TYPE_CHECKING, Any, Union, cast
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
@ -36,17 +36,17 @@ class AnthropicCacheControlHook(CustomPromptManagement):
def get_chat_completion_prompt(
self,
model: str,
messages: List[AllMessageValues],
messages: list[AllMessageValues],
non_default_params: dict,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
prompt_id: str | None,
prompt_variables: dict | None,
dynamic_callback_params: StandardCallbackDynamicParams,
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]:
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]:
"""
Apply cache control directives based on specified injection points.
@ -56,7 +56,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
- non_default_params: dict - params with any global cache controls
"""
# Extract cache control injection points
injection_points: List[CacheControlInjectionPoint] = non_default_params.pop(
injection_points: list[CacheControlInjectionPoint] = non_default_params.pop(
"cache_control_injection_points", []
)
if not injection_points:
@ -66,8 +66,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))
@ -98,10 +98,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
@ -159,11 +159,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: Optional[Union[int, str]] = point.get("index", None)
targetted_index: Optional[int] = None
_targetted_index: Union[int, str] | None = point.get("index", None)
targetted_index: int | None = None
if isinstance(_targetted_index, str):
try:
targetted_index = int(_targetted_index)
@ -249,8 +249,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
def should_run_prompt_management(
self,
prompt_id: Optional[str],
prompt_spec: Optional[PromptSpec],
prompt_id: str | None,
prompt_spec: PromptSpec | None,
dynamic_callback_params: StandardCallbackDynamicParams,
) -> bool:
"""Always return False since this is not a true prompt management system."""
@ -258,12 +258,12 @@ class AnthropicCacheControlHook(CustomPromptManagement):
def _compile_prompt_helper(
self,
prompt_id: Optional[str],
prompt_spec: Optional[PromptSpec],
prompt_variables: Optional[dict],
prompt_id: str | None,
prompt_spec: PromptSpec | None,
prompt_variables: dict | None,
dynamic_callback_params: StandardCallbackDynamicParams,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
prompt_label: str | None = None,
prompt_version: int | None = None,
) -> PromptManagementClient:
"""Not used - this hook only modifies messages, doesn't fetch prompts."""
return PromptManagementClient(
@ -276,12 +276,12 @@ class AnthropicCacheControlHook(CustomPromptManagement):
async def async_compile_prompt_helper(
self,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
prompt_id: str | None,
prompt_variables: dict | None,
dynamic_callback_params: StandardCallbackDynamicParams,
prompt_spec: Optional[PromptSpec] = None,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
prompt_spec: PromptSpec | None = None,
prompt_label: str | None = None,
prompt_version: int | None = None,
) -> PromptManagementClient:
"""Not used - this hook only modifies messages, doesn't fetch prompts."""
return self._compile_prompt_helper(
@ -296,19 +296,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: Optional[str],
prompt_variables: Optional[dict],
prompt_id: str | None,
prompt_variables: dict | None,
dynamic_callback_params: StandardCallbackDynamicParams,
litellm_logging_obj: LiteLLMLoggingObj,
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]:
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]:
"""Async version - delegates to sync since no async operations needed."""
return self.get_chat_completion_prompt(
model=model,
@ -325,15 +325,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,
) -> Optional[CustomLogger]:
non_default_params: dict,
) -> CustomLogger | None:
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
)

View file

@ -4,11 +4,8 @@ from typing import (
TYPE_CHECKING,
Any,
ClassVar,
Dict,
List,
Literal,
Optional,
Type,
Union,
get_args,
)
@ -59,7 +56,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]) -> Optional[str]:
def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None:
"""Extract session_id from request data (litellm_session_id or metadata)."""
session_id = request_data.get("litellm_session_id")
if session_id:
@ -84,20 +81,18 @@ class CustomGuardrail(CustomLogger):
def __init__(
self,
guardrail_name: Optional[str] = None,
supported_event_hooks: Optional[List[GuardrailEventHooks]] = None,
event_hook: Optional[
Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]
] = None,
guardrail_name: str | None = None,
supported_event_hooks: list[GuardrailEventHooks] | None = None,
event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = None,
default_on: bool = False,
mask_request_content: bool = False,
mask_response_content: bool = False,
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,
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,
sticky_session_routing: bool = True,
**kwargs,
):
@ -120,18 +115,16 @@ class CustomGuardrail(CustomLogger):
"""
self.guardrail_name = guardrail_name
self.supported_event_hooks = supported_event_hooks
self.event_hook: Optional[
Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]
] = event_hook
self.event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = 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: 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] = (
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.sticky_session_routing: bool = sticky_session_routing
@ -142,14 +135,14 @@ class CustomGuardrail(CustomLogger):
super().__init__(**kwargs)
def render_violation_message(
self, default: str, context: Optional[Dict[str, Any]] = None
self, default: str, context: dict[str, Any] | None = 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:
@ -165,8 +158,8 @@ class CustomGuardrail(CustomLogger):
def raise_passthrough_exception(
self,
violation_message: str,
request_data: Dict[str, Any],
detection_info: Optional[Dict[str, Any]] = None,
request_data: dict[str, Any],
detection_info: dict[str, Any] | None = None,
) -> None:
"""
Raise a passthrough exception for guardrail violations.
@ -209,8 +202,8 @@ class CustomGuardrail(CustomLogger):
def raise_sensitive_data_route_exception(
self,
route_to_model: str,
request_data: Dict[str, Any],
detection_info: Optional[Dict[str, Any]] = None,
request_data: dict[str, Any],
detection_info: dict[str, Any] | None = None,
) -> None:
"""
Raise an exception to reroute the request to a different model.
@ -248,8 +241,8 @@ class CustomGuardrail(CustomLogger):
)
def _get_session_id_from_request_data(
self, request_data: Dict[str, Any]
) -> Optional[str]:
self, request_data: dict[str, Any]
) -> str | None:
"""Extract session_id from request data."""
return get_session_id_from_request_data(request_data)
@ -265,8 +258,8 @@ class CustomGuardrail(CustomLogger):
def handle_sensitive_data_detection(
self,
request_data: Dict[str, Any],
detection_info: Optional[Dict[str, Any]] = None,
request_data: dict[str, Any],
detection_info: dict[str, Any] | None = None,
) -> None:
"""
Handle sensitive data detection based on guardrail configuration.
@ -309,7 +302,7 @@ class CustomGuardrail(CustomLogger):
)
@staticmethod
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
def get_config_model() -> type["GuardrailConfigModel"] | None:
"""
Returns the config model for the guardrail
@ -319,14 +312,12 @@ class CustomGuardrail(CustomLogger):
def _validate_event_hook(
self,
event_hook: Optional[
Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]
],
supported_event_hooks: List[GuardrailEventHooks],
event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None,
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):
@ -391,7 +382,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) -> Optional[bool]:
def get_disable_global_guardrail(self, data: dict) -> bool | None:
"""
Returns True if the global guardrail should be disabled.
@ -400,7 +391,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.
@ -433,7 +424,7 @@ class CustomGuardrail(CustomLogger):
def get_guardrail_from_metadata(
self, data: dict
) -> Union[List[str], List[Dict[str, DynamicGuardrailParams]]]:
) -> Union[list[str], list[dict[str, DynamicGuardrailParams]]]:
"""
Returns the guardrail(s) to be run from the metadata or root
"""
@ -454,7 +445,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):
@ -466,13 +457,13 @@ class CustomGuardrail(CustomLogger):
return False
def _pre_call_marker(self) -> Optional[str]:
def _pre_call_marker(self) -> str | None:
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.
@ -497,7 +488,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
@ -510,8 +501,8 @@ class CustomGuardrail(CustomLogger):
return False
async def async_pre_call_deployment_hook(
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
) -> Optional[dict]:
self, kwargs: dict[str, Any], call_type: CallTypes | None
) -> dict | None:
from litellm.proxy._types import UserAPIKeyAuth
# should run guardrail
@ -556,8 +547,8 @@ class CustomGuardrail(CustomLogger):
self,
request_data: dict,
response: LLMResponseTypes,
call_type: Optional[CallTypes],
) -> Optional[LLMResponseTypes]:
call_type: CallTypes | None,
) -> LLMResponseTypes | None:
"""
Allow modifying / reviewing the response just after it's received from the deployment.
"""
@ -755,16 +746,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: 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,
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,
) -> None:
"""
Builds `StandardLoggingGuardrailInformation` and adds it to the request metadata so it can be used for logging to DataDog, Langfuse, etc.
@ -780,7 +771,7 @@ class CustomGuardrail(CustomLogger):
# Use event_type if provided, otherwise fall back to self.event_hook
guardrail_mode: Union[
GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks]
GuardrailEventHooks, GuardrailMode, list[GuardrailEventHooks]
]
if event_type is not None:
guardrail_mode = event_type
@ -894,13 +885,13 @@ class CustomGuardrail(CustomLogger):
def _process_response(
self,
response: Optional[Dict],
response: dict | None,
request_data: dict,
start_time: Optional[float] = None,
end_time: Optional[float] = None,
duration: Optional[float] = None,
event_type: Optional[GuardrailEventHooks] = None,
original_inputs: Optional[Dict] = None,
start_time: float | None = None,
end_time: float | None = None,
duration: float | None = None,
event_type: GuardrailEventHooks | None = None,
original_inputs: dict | None = None,
):
"""
Add StandardLoggingGuardrailInformation to the request data
@ -908,7 +899,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] = (
guardrail_response: Union[dict[str, Any], str] = (
{} if response is None else response
)
@ -970,10 +961,10 @@ class CustomGuardrail(CustomLogger):
self,
e: Exception,
request_data: dict,
start_time: Optional[float] = None,
end_time: Optional[float] = None,
duration: Optional[float] = None,
event_type: Optional[GuardrailEventHooks] = None,
start_time: float | None = None,
end_time: float | None = None,
duration: float | None = None,
event_type: GuardrailEventHooks | None = None,
):
"""
Add StandardLoggingGuardrailInformation to the request data
@ -1002,7 +993,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.
@ -1047,8 +1038,8 @@ class CustomGuardrail(CustomLogger):
setattr(self, key, value)
def get_guardrails_messages_for_call_type(
self, call_type: CallTypes, data: Optional[dict] = None
) -> Optional[List[AllMessageValues]]:
self, call_type: CallTypes, data: dict | None = None
) -> list[AllMessageValues] | None:
"""
Returns the messages for the given call type and data
"""
@ -1089,7 +1080,7 @@ class CustomGuardrail(CustomLogger):
input=input_data,
responses_api_request=data,
)
return cast(List[AllMessageValues], messages)
return cast(list[AllMessageValues], messages)
return None
@ -1118,7 +1109,7 @@ def log_guardrail_information(func):
def _infer_event_type_from_function_name(
func_name: str,
) -> Optional[GuardrailEventHooks]:
) -> GuardrailEventHooks | None:
"""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, FrozenSet, Mapping, Optional
from typing import Any, Mapping
from opentelemetry.metrics import Histogram, Meter
@ -82,12 +82,12 @@ class GenAIMetricRecorder:
"""
def __init__(
self, metrics: GenAIMetrics, callback_name: Optional[str] = None
self, metrics: GenAIMetrics, callback_name: str | None = None
) -> None:
self._metrics = metrics
self._callback_name = callback_name
self._include: Optional[FrozenSet[str]] = None
self._exclude: Optional[FrozenSet[str]] = None
self._include: frozenset[str] | None = None
self._exclude: frozenset[str] | None = None
self._filter_resolved = False
def record(

View file

@ -4,11 +4,7 @@ import time
from typing import (
TYPE_CHECKING,
Any,
Dict,
List,
NoReturn,
Optional,
Tuple,
Union,
cast,
)
@ -152,8 +148,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
@ -173,7 +169,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
@ -214,7 +210,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",
@ -237,23 +233,23 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
to pass metadata to anthropic, it's {"user_id": "any-relevant-information"}
"""
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
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
def __init__(
self,
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,
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,
) -> None:
locals_ = locals().copy()
for key, value in locals_.items():
@ -261,11 +257,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
setattr(self.__class__, key, value)
@property
def custom_llm_provider(self) -> Optional[str]:
def custom_llm_provider(self) -> str | None:
return "anthropic"
@classmethod
def get_config(cls, *, model: Optional[str] = None):
def get_config(cls, *, model: str | None = None):
config = super().get_config()
# anthropic requires a default value for max_tokens
@ -275,7 +271,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return config
@staticmethod
def get_max_tokens_for_model(model: Optional[str] = None) -> int:
def get_max_tokens_for_model(model: str | None = 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.
@ -292,7 +288,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:
"""
@ -317,7 +313,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
@ -344,7 +340,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
@staticmethod
def _validate_effort_for_model(model: str, effort: Optional[str]) -> Optional[str]:
def _validate_effort_for_model(model: str, effort: str | None) -> str | None:
"""Return ``None`` if ``effort`` is allowed on ``model``, else an error message."""
if effort == "max" and not (
AnthropicConfig._is_claude_4_6_model(model)
@ -434,7 +430,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.
@ -492,7 +488,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
constraint_labels[field].format(schema[field])
)
result: Dict[str, Any] = {}
result: dict[str, Any] = {}
# Update description with removed constraint info
if constraint_descriptions:
@ -548,8 +544,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return result
def get_json_schema_from_pydantic_object(
self, response_format: Union[Any, Dict, None]
) -> Optional[dict]:
self, response_format: Union[Any, dict, None]
) -> dict | None:
return type_to_response_format_param(
response_format, ref_template="/$defs/{model}"
) # Relevant issue: https://github.com/BerriAI/litellm/issues/7755
@ -564,10 +560,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _map_tool_choice(
self,
tool_choice: Optional[str],
parallel_tool_use: Optional[bool],
) -> Optional[AnthropicMessagesToolChoice]:
_tool_choice: Optional[AnthropicMessagesToolChoice] = None
tool_choice: str | None,
parallel_tool_use: bool | None,
) -> AnthropicMessagesToolChoice | None:
_tool_choice: AnthropicMessagesToolChoice | None = None
if tool_choice == "auto":
_tool_choice = AnthropicMessagesToolChoice(
type="auto",
@ -608,9 +604,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _map_tool_helper(
self,
tool: ChatCompletionToolParam,
) -> Tuple[Optional[AllAnthropicToolsValues], Optional[AnthropicMcpServerTool]]:
returned_tool: Optional[AllAnthropicToolsValues] = None
mcp_server: Optional[AnthropicMcpServerTool] = None
) -> tuple[AllAnthropicToolsValues | None, AnthropicMcpServerTool | None]:
returned_tool: AllAnthropicToolsValues | None = None
mcp_server: AnthropicMcpServerTool | None = None
if tool["type"] == "function" or tool["type"] == "custom":
_input_schema: dict = tool["function"].get(
@ -665,10 +661,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if "parameters" not in tool["function"]:
raise ValueError("Missing required parameter: parameters")
_display_width_px: Optional[int] = tool["function"]["parameters"].get(
_display_width_px: int | None = tool["function"]["parameters"].get(
"display_width_px"
)
_display_height_px: Optional[int] = tool["function"]["parameters"].get(
_display_height_px: int | None = tool["function"]["parameters"].get(
"display_height_px"
)
if _display_width_px is None or _display_height_px is None:
@ -841,14 +837,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
from litellm.types.llms.anthropic import AnthropicMcpServerToolConfiguration
allowed_tools = tool.get("allowed_tools", None)
tool_configuration: Optional[AnthropicMcpServerToolConfiguration] = None
tool_configuration: AnthropicMcpServerToolConfiguration | None = None
if allowed_tools is not None:
tool_configuration = AnthropicMcpServerToolConfiguration(
allowed_tools=tool.get("allowed_tools", None),
)
headers = tool.get("headers", {})
authorization_token: Optional[str] = None
authorization_token: str | None = None
if headers is not None:
bearer_token = headers.get("Authorization", None)
if bearer_token is not None:
@ -868,8 +864,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:
@ -918,9 +914,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.
@ -931,7 +927,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)
@ -979,8 +975,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
@ -994,7 +990,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
@ -1011,8 +1007,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}$``.
@ -1036,7 +1032,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
@ -1058,7 +1054,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)
@ -1084,7 +1080,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return forward, reverse
def _detect_tool_search_tools(self, tools: Optional[List]) -> bool:
def _detect_tool_search_tools(self, tools: list | None) -> bool:
"""Check if tool search tools are present in the tools list."""
if not tools:
return False
@ -1098,7 +1094,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.
@ -1118,9 +1114,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.
@ -1162,9 +1158,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return expanded_content
def _map_stop_sequences(
self, stop: Optional[Union[str, List[str]]]
) -> Optional[List[str]]:
new_stop: Optional[List[str]] = None
self, stop: Union[str, list[str]] | None
) -> list[str] | None:
new_stop: list[str] | None = None
if isinstance(stop, str):
if (
stop.isspace() and litellm.drop_params is True
@ -1185,10 +1181,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def _map_reasoning_effort(
reasoning_effort: Optional[Union[REASONING_EFFORT, str]],
reasoning_effort: Union[REASONING_EFFORT, str] | None,
model: str,
llm_provider: str = "anthropic",
) -> Optional[AnthropicThinkingParam]:
) -> AnthropicThinkingParam | None:
if reasoning_effort is None or reasoning_effort == "none":
return None
if AnthropicConfig._is_adaptive_thinking_model(model):
@ -1240,11 +1236,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
def _extract_json_schema_from_response_format(
self, value: Optional[dict]
) -> Optional[dict]:
self, value: dict | None
) -> dict | None:
if value is None:
return None
json_schema: Optional[dict] = None
json_schema: dict | None = None
if "response_schema" in value:
json_schema = value["response_schema"]
elif "json_schema" in value:
@ -1253,9 +1249,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return json_schema
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(
self, value: dict | None
) -> AnthropicOutputSchema | None:
json_schema: dict | None = self._extract_json_schema_from_response_format(
value
)
if json_schema is None:
@ -1283,15 +1279,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
def map_response_format_to_anthropic_tool(
self, value: Optional[dict], optional_params: dict, is_thinking_enabled: bool
) -> Optional[AnthropicMessagesTool]:
self, value: dict | None, optional_params: dict, is_thinking_enabled: bool
) -> AnthropicMessagesTool | None:
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: Optional[dict] = self._extract_json_schema_from_response_format(
json_schema: dict | None = self._extract_json_schema_from_response_format(
value
)
if json_schema is None:
@ -1342,8 +1338,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def map_openai_context_management_to_anthropic(
context_management: Union[List[Dict[str, Any]], Dict[str, Any]],
) -> Optional[Dict[str, Any]]:
context_management: Union[list[dict[str, Any]], dict[str, Any]],
) -> dict[str, Any] | None:
"""
OpenAI format: [{"type": "compaction", "compact_threshold": 200000}]
Anthropic format: {
@ -1374,7 +1370,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(
@ -1438,7 +1434,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: Optional[AnthropicMessagesToolChoice] = (
_tool_choice: AnthropicMessagesToolChoice | None = (
self._map_tool_choice(
tool_choice=non_default_params.get("tool_choice"),
parallel_tool_use=non_default_params.get("parallel_tool_calls"),
@ -1584,7 +1580,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _create_json_tool_call_for_response_format(
self,
json_schema: Optional[dict] = None,
json_schema: dict | None = None,
) -> AnthropicMessagesTool:
"""
Handles creating a tool call for getting responses in JSON format.
@ -1622,8 +1618,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return False
def translate_system_message(
self, messages: List[AllMessageValues]
) -> List[AnthropicSystemMessageContent]:
self, messages: list[AllMessageValues]
) -> list[AnthropicSystemMessageContent]:
"""
Translate system message to anthropic format.
@ -1631,7 +1627,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)
@ -1691,9 +1687,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:
@ -1832,7 +1828,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,
@ -1938,7 +1934,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
## Add code_execution tool if container_upload is in messages
_tools = (
cast(
Optional[List[Union[AllAnthropicToolsValues, Dict]]],
list[Union[AllAnthropicToolsValues, dict]] | None,
optional_params.get("tools"),
)
or []
@ -2045,12 +2041,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _resolve_json_mode_non_streaming(
self,
json_mode: Optional[bool],
tool_calls: List[ChatCompletionToolCallChunk],
) -> Tuple[
Optional[LitellmMessage],
List[ChatCompletionToolCallChunk],
Optional[str],
json_mode: bool | None,
tool_calls: list[ChatCompletionToolCallChunk],
) -> tuple[
LitellmMessage | None,
list[ChatCompletionToolCallChunk],
str | None,
]:
"""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:
@ -2075,38 +2071,30 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
first_json = tool_calls[json_indices[0]]
json_msg = AnthropicConfig._convert_tool_response_to_message([first_json])
extra_content: Optional[str] = (
extra_content: str | None = (
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[
def extract_response_content(self, completion_response: dict) -> tuple[
str,
Optional[List[Any]],
Optional[
List[
Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]
]
],
Optional[str],
List[ChatCompletionToolCallChunk],
Optional[List[Any]],
Optional[List[Any]],
Optional[List[Any]],
list[Any] | None,
list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None,
str | None,
list[ChatCompletionToolCallChunk],
list[Any] | None,
list[Any] | None,
list[Any] | None,
]:
text_content = ""
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
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
for idx, content in enumerate(completion_response["content"]):
if content["type"] == "text":
text_content += content["text"]
@ -2171,7 +2159,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if thinking_blocks is not None:
reasoning_content = ""
for block in thinking_blocks:
thinking_content = cast(Optional[str], block.get("thinking"))
thinking_content = cast(str | None, block.get("thinking"))
if thinking_content is not None:
reasoning_content += thinking_content
@ -2189,9 +2177,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def calculate_usage(
self,
usage_object: dict,
reasoning_content: Optional[str],
completion_response: Optional[dict] = None,
speed: Optional[str] = None,
reasoning_content: str | None,
completion_response: dict | None = None,
speed: str | None = 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
@ -2207,14 +2195,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
_usage = usage_object
cache_creation_input_tokens: int = 0
cache_read_input_tokens: int = 0
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
web_search_requests: Optional[int] = None
tool_search_requests: Optional[int] = None
inference_geo: Optional[str] = None
cache_creation_token_details: CacheCreationTokenDetails | None = None
web_search_requests: int | None = None
tool_search_requests: int | None = None
inference_geo: str | None = None
if "inference_geo" in _usage and _usage["inference_geo"] is not None:
inference_geo = _usage["inference_geo"]
iterations: Optional[List[Any]] = _usage.get("iterations")
iterations: list[Any] | None = _usage.get("iterations")
if iterations:
prompt_tokens = sum(it.get("input_tokens", 0) or 0 for it in iterations)
completion_tokens = sum(
@ -2328,9 +2316,9 @@ 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] = {}
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", "{}"))
@ -2344,10 +2332,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _build_code_interpreter_results(
self,
tool_results: List[Any],
code_by_id: Dict[str, str],
container_id: Optional[str],
) -> List[OutputCodeInterpreterCall]:
tool_results: list[Any],
code_by_id: dict[str, str],
container_id: str | None,
) -> list[OutputCodeInterpreterCall]:
code_interpreter_results = []
for tr in tool_results:
if tr.get("type") != "bash_code_execution_tool_result":
@ -2370,18 +2358,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _build_provider_specific_fields(
self,
completion_response: dict,
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: 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": citations,
"thinking_blocks": thinking_blocks,
}
@ -2422,12 +2406,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
completion_response: dict,
raw_response: httpx.Response,
model_response: ModelResponse,
json_mode: Optional[bool] = None,
prefix_prompt: Optional[str] = None,
speed: Optional[str] = None,
tool_name_reverse_map: Optional[Dict[str, str]] = None,
json_mode: bool | None = None,
prefix_prompt: str | None = None,
speed: str | None = None,
tool_name_reverse_map: dict[str, str] | None = None,
):
_hidden_params: Dict = {}
_hidden_params: dict = {}
_hidden_params["additional_headers"] = process_anthropic_headers(
dict(raw_response.headers)
)
@ -2531,7 +2515,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
model_response._hidden_params = _hidden_params
return model_response
def get_prefix_prompt(self, messages: List[AllMessageValues]) -> Optional[str]:
def get_prefix_prompt(self, messages: list[AllMessageValues]) -> str | None:
"""
Get the prefix prompt from the messages.
@ -2560,13 +2544,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: Optional[str] = None,
json_mode: Optional[bool] = None,
api_key: str | None = None,
json_mode: bool | None = None,
) -> ModelResponse:
## LOGGING
logging_obj.post_call(
@ -2591,7 +2575,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
prefix_prompt = self.get_prefix_prompt(messages=messages)
speed = optional_params.get("speed")
tool_name_reverse_map: Optional[Dict[str, str]] = None
tool_name_reverse_map: dict[str, str] | None = None
if isinstance(litellm_params, dict):
_candidate = litellm_params.get(ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY)
if isinstance(_candidate, dict):
@ -2610,14 +2594,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def _convert_tool_response_to_message(
tool_calls: List[ChatCompletionToolCallChunk],
) -> Optional[LitellmMessage]:
tool_calls: list[ChatCompletionToolCallChunk],
) -> LitellmMessage | None:
"""
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: Optional[str] = tool_calls[0]["function"].get(
json_mode_content_str: str | None = tool_calls[0]["function"].get(
"arguments"
)
try:
@ -2640,7 +2624,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, Optional, Tuple, Union, cast, overload
from typing import Any, Coroutine, Literal, Union, cast, overload
from typing_extensions import override
@ -63,8 +63,8 @@ class ModelScopeChatConfig(OpenAIGPTConfig):
)
def _get_openai_compatible_provider_info(
self, api_base: Optional[str], api_key: Optional[str]
) -> Tuple[Optional[str], Optional[str]]:
self, api_base: str | None, api_key: str | None
) -> tuple[str | None, str | None]:
api_base = (
api_base or get_secret_str("MODELSCOPE_API_BASE") or self.DEFAULT_BASE_URL
) # type: ignore
@ -74,12 +74,12 @@ class ModelScopeChatConfig(OpenAIGPTConfig):
@override
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
api_base: str | None,
api_key: str | None,
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
stream: bool | None = None,
) -> str:
"""
If api_base is not provided, use the default ModelScope /chat/completions endpoint.

View file

@ -9,13 +9,9 @@ from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
List,
Literal,
Mapping,
Optional,
Tuple,
Type,
Union,
cast,
)
@ -135,7 +131,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
"""
@ -152,7 +148,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
"""
@ -198,29 +194,29 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
Note: Please make sure to modify the default parameters as required for your use case.
"""
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
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
def __init__(
self,
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,
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,
) -> None:
locals_ = locals().copy()
for key, value in locals_.items():
@ -232,8 +228,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return super().get_config()
def get_json_schema_from_pydantic_object(
self, response_format: Optional[Union[Type["BaseModel"], dict]]
) -> Optional[dict]:
self, response_format: Union[type["BaseModel"], dict] | None
) -> dict | None:
"""
Override to use Pydantic's model_json_schema() instead of OpenAI's
to_strict_json_schema().
@ -292,7 +288,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _forward_gemini_function_call_id(
model: str, custom_llm_provider: Optional[str] = None
model: str, custom_llm_provider: str | None = None
) -> bool:
"""
Whether to include `id` on function_call / function_response parts.
@ -313,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",
@ -350,7 +346,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
def map_tool_choice_values(
self, model: str, tool_choice: Union[str, dict]
) -> Optional[ToolConfig]:
) -> ToolConfig | None:
if tool_choice == "none":
return ToolConfig(functionCallingConfig=FunctionCallingConfig(mode="NONE"))
elif tool_choice == "required":
@ -494,7 +490,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
def _extract_google_maps_retrieval_config(
self, google_maps_config: dict
) -> Tuple[dict, Optional[dict]]:
) -> tuple[dict, dict | None]:
"""
Extract location configuration from googleMaps tool for Vertex AI toolConfig.
@ -534,7 +530,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return cleaned_config, retrieval_config
def get_tool_value(self, tool: dict, tool_name: str) -> Optional[dict]:
def get_tool_value(self, tool: dict, tool_name: str) -> dict | None:
"""
Helper function to get tool value handling both camelCase and underscore_case variants
@ -561,10 +557,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _resolve_search_tool_conflict(
gtool_func_declarations: list,
googleSearch: Optional[dict],
googleSearchRetrieval: Optional[dict],
enterpriseWebSearch: Optional[dict],
urlContext: Optional[dict],
googleSearch: dict | None,
googleSearchRetrieval: dict | None,
enterpriseWebSearch: dict | None,
urlContext: dict | None,
optional_params: dict,
) -> tuple:
"""
@ -614,7 +610,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.
@ -630,21 +626,21 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
googleMaps tools contain location data
"""
gtool_func_declarations = []
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
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
# 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: Optional[ChatCompletionToolParamFunctionChunk] = (
openai_function_object: ChatCompletionToolParamFunctionChunk | None = (
None
)
if "function" in tool: # tools list
@ -763,7 +759,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,
@ -897,7 +893,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _map_reasoning_effort_to_thinking_budget(
reasoning_effort: str,
model: Optional[str] = None,
model: str | None = None,
) -> GeminiThinkingConfig:
if reasoning_effort == "minimal":
# Use model-specific minimum thinking budget or fallback
@ -948,7 +944,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _map_reasoning_effort_to_thinking_level(
reasoning_effort: str,
model: Optional[str] = None,
model: str | None = None,
) -> GeminiThinkingConfig:
"""
Map reasoning_effort to thinking_level for Gemini 3+ models.
@ -996,12 +992,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
raise ValueError(f"Invalid reasoning effort: {reasoning_effort}")
@staticmethod
def _is_thinking_budget_zero(thinking_budget: Optional[int]) -> bool:
def _is_thinking_budget_zero(thinking_budget: int | None) -> 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:
@ -1022,7 +1018,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.
@ -1042,7 +1038,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _map_thinking_param(
thinking_param: AnthropicThinkingParam,
model: Optional[str] = None,
model: str | None = None,
) -> GeminiThinkingConfig:
thinking_enabled = thinking_param.get("type") == "enabled"
thinking_budget = thinking_param.get("budget_tokens")
@ -1153,8 +1149,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.
@ -1173,11 +1169,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
)
@ -1290,7 +1286,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: Optional[str] = None
effort_value: str | None = None
if isinstance(value, str):
effort_value = value
elif isinstance(value, dict):
@ -1374,7 +1370,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
"""
@ -1411,7 +1407,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return model
@staticmethod
def _is_model_gemini_spec_model(model: Optional[str]) -> bool:
def _is_model_gemini_spec_model(model: str | None) -> bool:
"""
Returns true if user is trying to call custom model in `/gemini` request/response format
"""
@ -1434,7 +1430,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
@ -1471,7 +1467,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.
@ -1495,10 +1491,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return exception_string
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
self, parts: list[HttpxPartType]
) -> tuple[str | None, str | None]:
content_str: str | None = None
reasoning_content_str: str | None = None
for part in parts:
_content_str = ""
@ -1540,8 +1536,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return content_str, reasoning_content_str
def _extract_thinking_blocks_from_parts(
self, parts: List[HttpxPartType]
) -> List[ChatCompletionThinkingBlock]:
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):
@ -1550,7 +1546,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", "")
@ -1565,8 +1561,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return thinking_blocks
def _extract_thought_signatures_from_parts(
self, parts: List[HttpxPartType]
) -> Optional[List[str]]:
self, parts: list[HttpxPartType]
) -> list[str] | None:
"""Extract thoughtSignature values from parts.
Per Google's docs, thoughtSignature is returned for multi-turn context preservation
@ -1576,7 +1572,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:
@ -1585,8 +1581,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _extract_server_side_tool_invocations(
parts: List[HttpxPartType],
) -> Optional[List[Dict[str, Any]]]:
parts: list[HttpxPartType],
) -> list[dict[str, Any]] | None:
"""Extract server-side tool invocations (toolCall/toolResponse) from parts.
These are returned by Gemini when context circulation is enabled
@ -1597,15 +1593,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"),
@ -1645,10 +1641,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return invocations if invocations else None
def _extract_image_response_from_parts(
self, parts: List[HttpxPartType]
) -> Optional[List[ImageURLListItem]]:
self, parts: list[HttpxPartType]
) -> list[ImageURLListItem] | None:
"""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", {})
@ -1667,8 +1663,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return images
def _extract_audio_response_from_parts(
self, parts: List[HttpxPartType]
) -> Optional[ChatCompletionAudioResponse]:
self, parts: list[HttpxPartType]
) -> ChatCompletionAudioResponse | None:
"""Extract audio response from parts if present"""
for part in parts:
if "text" in part:
@ -1710,16 +1706,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _transform_parts(
parts: List[HttpxPartType],
parts: list[HttpxPartType],
cumulative_tool_call_idx: int,
is_function_call: Optional[bool],
) -> Tuple[
Optional[ChatCompletionToolCallFunctionChunk],
Optional[List[ChatCompletionToolCallChunk]],
is_function_call: bool | None,
) -> tuple[
ChatCompletionToolCallFunctionChunk | None,
list[ChatCompletionToolCallChunk] | None,
int,
]:
function: Optional[ChatCompletionToolCallFunctionChunk] = None
_tools: List[ChatCompletionToolCallChunk] = []
function: ChatCompletionToolCallFunctionChunk | None = None
_tools: list[ChatCompletionToolCallChunk] = []
for part in parts:
if "functionCall" in part:
_function_chunk: ChatCompletionToolCallFunctionChunk = {
@ -1736,7 +1732,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"] = {}
@ -1769,22 +1765,22 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
_tools.append(_tool_response_chunk)
cumulative_tool_call_idx += 1
if len(_tools) == 0:
tools: Optional[List[ChatCompletionToolCallChunk]] = None
tools: list[ChatCompletionToolCallChunk] | None = None
else:
tools = _tools
return function, tools, cumulative_tool_call_idx
@staticmethod
def _transform_logprobs(
logprobs_result: Optional[LogprobsResult],
) -> Optional[ChoiceLogprobs]:
logprobs_result: LogprobsResult | None,
) -> ChoiceLogprobs | None:
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"]
):
@ -1914,16 +1910,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
raise ValueError(
f"usageMetadata not found in completion_response. Got={completion_response}"
)
cached_tokens: Optional[int] = None
cached_tokens: int | None = None
# Separate variables for prompt tokens by modality
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
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
usage_metadata = completion_response["usageMetadata"]
def _get_token_count(detail: Mapping[str, Any]) -> int:
@ -2021,10 +2017,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: Optional[int] = None
cached_audio_tokens: Optional[int] = None
cached_image_tokens: Optional[int] = None
cached_video_tokens: Optional[int] = None
cached_text_tokens: int | None = None
cached_audio_tokens: int | None = None
cached_image_tokens: int | None = None
cached_video_tokens: int | None = None
if "cacheTokensDetails" in usage_metadata:
for detail in usage_metadata["cacheTokensDetails"]:
@ -2102,8 +2098,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _check_finish_reason(
chat_completion_message: Optional[ChatCompletionResponseMessage],
finish_reason: Optional[str],
chat_completion_message: ChatCompletionResponseMessage | None,
finish_reason: str | None,
) -> OpenAIChatCompletionFinishReason:
from litellm.litellm_core_utils.core_helpers import map_finish_reason
@ -2119,7 +2115,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _check_prompt_level_content_filter(
processed_chunk: GenerateContentResponseBody,
response_id: Optional[str],
response_id: str | None,
) -> Optional["ModelResponseStream"]:
"""
Check if prompt is blocked due to content filtering at the prompt level.
@ -2163,8 +2159,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return None
@staticmethod
def _calculate_web_search_requests(grounding_metadata: List[dict]) -> Optional[int]:
web_search_requests: Optional[int] = None
def _calculate_web_search_requests(grounding_metadata: list[dict]) -> int | None:
web_search_requests: int | None = None
if (
grounding_metadata
@ -2184,10 +2180,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
chat_completion_message: ChatCompletionResponseMessage,
candidate: Candidates,
idx: int,
tools: Optional[List[ChatCompletionToolCallChunk]],
functions: Optional[ChatCompletionToolCallFunctionChunk],
chat_completion_logprobs: Optional[ChoiceLogprobs],
image_response: Optional[List[ImageURLListItem]],
tools: list[ChatCompletionToolCallChunk] | None,
functions: ChatCompletionToolCallFunctionChunk | None,
chat_completion_logprobs: ChoiceLogprobs | None,
image_response: list[ImageURLListItem] | None,
) -> StreamingChoices:
"""
Helper method to create a streaming choice object for Vertex AI
@ -2219,7 +2215,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.
@ -2229,10 +2225,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):
@ -2277,10 +2273,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:
@ -2306,10 +2302,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:
@ -2324,14 +2320,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _convert_grounding_metadata_to_annotations(
grounding_metadata: List[dict],
content_text: Optional[str],
) -> List[ChatCompletionAnnotation]:
grounding_metadata: list[dict],
content_text: str | None,
) -> 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
@ -2339,7 +2335,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"]
@ -2379,11 +2375,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
@ -2399,19 +2395,19 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
from litellm.types.utils import ModelResponseStream
grounding_metadata: List[dict] = []
url_context_metadata: List[dict] = []
image_response: Optional[List[ImageURLListItem]] = None
safety_ratings: List = []
citation_metadata: List = []
grounding_metadata: list[dict] = []
url_context_metadata: list[dict] = []
image_response: list[ImageURLListItem] | None = None
safety_ratings: list = []
citation_metadata: list = []
chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"}
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
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
for idx, candidate in enumerate(_candidates):
if "content" not in candidate:
@ -2470,13 +2466,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
if audio_response is not None:
cast(Dict[str, Any], chat_completion_message)[
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)[
cast(dict[str, Any], chat_completion_message)[
"images"
] = image_response
if content is not None:
@ -2583,13 +2579,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: Optional[str] = None,
json_mode: Optional[bool] = None,
api_key: str | None = None,
json_mode: bool | None = None,
) -> ModelResponse:
## LOGGING
logging_obj.post_call(
@ -2664,11 +2660,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,
@ -2750,10 +2746,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
def _transform_messages(
self,
messages: List[AllMessageValues],
model: Optional[str] = None,
litellm_params: Optional[dict] = None,
) -> List[ContentType]:
messages: list[AllMessageValues],
model: str | None = None,
litellm_params: dict | None = None,
) -> list[ContentType]:
return _gemini_convert_messages_with_history(
messages=messages,
model=model,
@ -2762,7 +2758,7 @@ 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
@ -2771,25 +2767,25 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
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: Optional[Dict],
headers: dict | None,
model: str,
messages: List[AllMessageValues],
optional_params: Dict,
litellm_params: Dict,
api_key: Optional[Union[str, Dict]] = None,
api_base: Optional[str] = None,
) -> Dict:
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Union[str, dict] | None = None,
api_base: str | None = None,
) -> dict:
default_headers = {
"Content-Type": "application/json",
}
@ -2804,8 +2800,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
async def make_call(
client: Optional[AsyncHTTPHandler], # module-level client
gemini_client: Optional[AsyncHTTPHandler], # if passed by user
client: AsyncHTTPHandler | None, # module-level client
gemini_client: AsyncHTTPHandler | None, # if passed by user
api_base: str,
headers: dict,
data: str,
@ -2857,8 +2853,8 @@ async def make_call(
def make_sync_call(
client: Optional[HTTPHandler], # module-level client
gemini_client: Optional[HTTPHandler], # if passed by user
client: HTTPHandler | None, # module-level client
gemini_client: HTTPHandler | None, # if passed by user
api_base: str,
headers: dict,
data: str,
@ -2914,20 +2910,20 @@ class VertexLLM(VertexBase):
model_response: ModelResponse,
print_verbose: Callable,
data: dict,
timeout: Optional[Union[float, httpx.Timeout]],
timeout: Union[float, httpx.Timeout] | None,
encoding,
logging_obj,
stream,
optional_params: dict,
litellm_params: dict,
logger_fn=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,
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,
) -> CustomStreamWrapper:
should_use_v1beta1_features = self.is_using_v1beta1_features(
optional_params=optional_params
@ -3016,20 +3012,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: Optional[Union[float, httpx.Timeout]],
timeout: Union[float, httpx.Timeout] | None,
encoding,
logging_obj,
stream,
optional_params: dict,
litellm_params: dict,
logger_fn=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,
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,
) -> Union[ModelResponse, CustomStreamWrapper]:
should_use_v1beta1_features = self.is_using_v1beta1_features(
optional_params=optional_params
@ -3142,18 +3138,18 @@ class VertexLLM(VertexBase):
logging_obj,
optional_params: dict,
acompletion: bool,
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],
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,
litellm_params: dict,
logger_fn=None,
extra_headers: Optional[dict] = None,
client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None,
api_base: Optional[str] = None,
extra_headers: dict | None = None,
client: Union[AsyncHTTPHandler, HTTPHandler] | None = None,
api_base: str | None = None,
) -> Union[ModelResponse, CustomStreamWrapper]:
stream: Optional[bool] = optional_params.pop("stream", None) # type: ignore
stream: bool | None = optional_params.pop("stream", None) # type: ignore
transform_request_params = {
"gemini_api_key": gemini_api_key,
@ -3347,7 +3343,7 @@ class ModelResponseIterator:
streaming_response,
sync_stream: bool,
logging_obj: LoggingClass,
response_headers: Optional[Dict[str, str]] = None,
response_headers: dict[str, str] | None = None,
):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
check_is_function_call,
@ -3390,9 +3386,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,
@ -3477,8 +3473,8 @@ class ModelResponseIterator:
self,
processed_chunk: Any,
model_response: Any,
grounding_metadata: List[dict],
) -> Optional[Usage]:
grounding_metadata: list[dict],
) -> Usage | None:
if "usageMetadata" not in processed_chunk:
return None
@ -3531,12 +3527,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: Optional[List[Candidates]] = processed_chunk.get("candidates")
_candidates: list[Candidates] | None = processed_chunk.get("candidates")
if _candidates:
(
grounding_metadata,

View file

@ -6,7 +6,7 @@
# +-------------------------------------------------------------+
import os
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional
from typing import TYPE_CHECKING, Any, Literal, Optional
import httpx
@ -68,11 +68,11 @@ class DeepKeepGuardrail(CustomGuardrail):
def __init__(
self,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
firewall_id: Optional[str] = None,
api_key: str | None = None,
api_base: str | None = None,
firewall_id: str | None = None,
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
extra_headers: Optional[Dict[str, str]] = None,
extra_headers: dict[str, str] | None = None,
**kwargs,
):
self.async_handler = get_async_httpx_client(
@ -114,7 +114,7 @@ class DeepKeepGuardrail(CustomGuardrail):
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
unreachable_fallback
)
self.extra_headers: Dict[str, str] = extra_headers or {}
self.extra_headers: dict[str, str] = extra_headers or {}
# Set supported event hooks
if "supported_event_hooks" not in kwargs:
@ -133,7 +133,7 @@ class DeepKeepGuardrail(CustomGuardrail):
self.firewall_id,
)
def _extract_user_api_key_metadata(self, request_data: dict) -> Dict[str, Any]:
def _extract_user_api_key_metadata(self, request_data: dict) -> dict[str, Any]:
"""
Extract user API key metadata from request_data for the DeepKeep API.
@ -143,7 +143,7 @@ class DeepKeepGuardrail(CustomGuardrail):
Returns:
Dictionary with user API key metadata fields.
"""
result_metadata: Dict[str, Any] = {}
result_metadata: dict[str, Any] = {}
litellm_metadata = request_data.get("litellm_metadata", {})
top_level_metadata = request_data.get("metadata", {})
@ -177,9 +177,9 @@ class DeepKeepGuardrail(CustomGuardrail):
return result_metadata
def _build_request_headers(self) -> Dict[str, str]:
def _build_request_headers(self) -> dict[str, str]:
"""Build HTTP headers for the DeepKeep API request."""
headers: Dict[str, str] = {
headers: dict[str, str] = {
"Content-Type": "application/json",
"X-API-Key": self.deepkeep_api_key,
}
@ -194,7 +194,7 @@ class DeepKeepGuardrail(CustomGuardrail):
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"],
error: Exception,
http_status_code: Optional[int] = None,
http_status_code: int | None = None,
) -> GenericGuardrailAPIInputs:
"""Allow the request to proceed when the guardrail is unreachable (fail-open mode)."""
status_suffix = (
@ -241,12 +241,12 @@ class DeepKeepGuardrail(CustomGuardrail):
@staticmethod
def _build_return_inputs(
*,
response_json: Dict[str, Any],
response_json: dict[str, Any],
texts: list,
images: Optional[Any],
tools: Optional[Any],
tool_calls: Optional[Any],
structured_messages: Optional[Any],
images: Any | None,
tools: Any | None,
tool_calls: Any | None,
structured_messages: Any | None,
) -> GenericGuardrailAPIInputs:
"""Merge original inputs with any guardrail-modified values from the API response."""
return_inputs = GenericGuardrailAPIInputs(texts=texts)
@ -311,7 +311,7 @@ class DeepKeepGuardrail(CustomGuardrail):
request_body = request_data.get("body") or {}
# Merge additional provider-specific params from config and dynamic params
additional_params: Dict[str, Any] = {"firewall_id": self.firewall_id}
additional_params: dict[str, Any] = {"firewall_id": self.firewall_id}
dynamic_params = self.get_guardrail_dynamic_request_body_params(request_body)
if dynamic_params:
additional_params.update(
@ -322,7 +322,7 @@ class DeepKeepGuardrail(CustomGuardrail):
user_metadata = self._extract_user_api_key_metadata(request_data)
# Build request payload
guardrail_request: Dict[str, Any] = {
guardrail_request: dict[str, Any] = {
"litellm_call_id": (logging_obj.litellm_call_id if logging_obj else None),
"litellm_trace_id": (logging_obj.litellm_trace_id if logging_obj else None),
"texts": texts,
@ -397,7 +397,7 @@ class DeepKeepGuardrail(CustomGuardrail):
)
@staticmethod
def get_config_model() -> Optional[type]:
def get_config_model() -> type | None:
from litellm.types.proxy.guardrails.guardrail_hooks.deepkeep import (
DeepKeepGuardrailConfigModel,
)

View file

@ -3,7 +3,7 @@
import importlib
import os
from datetime import datetime, timezone
from typing import Any, Dict, List, Literal, Optional, Set, Type, cast
from typing import Any, Literal, Optional, 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
}
@ -224,7 +224,7 @@ class GuardrailRegistry:
############################################################
def get_initialized_guardrail_callback(
self, guardrail_name: str
) -> Optional[CustomGuardrail]:
) -> CustomGuardrail | None:
"""
Returns the initialized guardrail callback for a given guardrail name
"""
@ -334,7 +334,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).
@ -347,7 +347,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
@ -357,7 +357,7 @@ class GuardrailRegistry:
async def get_guardrail_by_id_from_db(
self, guardrail_id: str, prisma_client: PrismaClient
) -> Optional[Guardrail]:
) -> Guardrail | None:
"""
Get a guardrail by its ID from the database
"""
@ -375,7 +375,7 @@ class GuardrailRegistry:
async def get_guardrail_by_name_from_db(
self, guardrail_name: str, prisma_client: PrismaClient
) -> Optional[Guardrail]:
) -> Guardrail | None:
"""
Get a guardrail by its name from the database
"""
@ -398,17 +398,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, Optional[CustomGuardrail]] = {}
self.guardrail_id_to_custom_guardrail: dict[str, CustomGuardrail | None] = {}
"""
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
@ -418,10 +418,10 @@ class InMemoryGuardrailHandler:
def initialize_guardrail(
self,
guardrail: Guardrail,
config_file_path: Optional[str] = None,
config_file_path: str | None = None,
llm_router: Optional["Router"] = None,
source: Literal["db", "config"] = "config",
) -> Optional[Guardrail]:
) -> Guardrail | None:
"""
Initialize a guardrail from a dictionary and add it to the litellm callback manager
@ -439,7 +439,7 @@ class InMemoryGuardrailHandler:
self._sources[guardrail_id] = source
return self.IN_MEMORY_GUARDRAILS[guardrail_id]
custom_guardrail_callback: Optional[CustomGuardrail] = None
custom_guardrail_callback: CustomGuardrail | None = None
litellm_params_data = guardrail["litellm_params"]
verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data)
@ -520,11 +520,11 @@ class InMemoryGuardrailHandler:
def initialize_custom_guardrail(
self,
guardrail: Dict,
guardrail: dict,
guardrail_type: str,
litellm_params: LitellmParams,
config_file_path: Optional[str] = None,
) -> Optional[CustomGuardrail]:
config_file_path: str | None = None,
) -> CustomGuardrail | None:
"""
Initialize a Custom Guardrail from a python file or module path
@ -623,25 +623,25 @@ class InMemoryGuardrailHandler:
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) -> Optional[Guardrail]:
def get_guardrail_by_id(self, guardrail_id: str) -> Guardrail | None:
"""
Get a guardrail by its ID from memory
"""
return self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
def get_source(self, guardrail_id: str) -> Optional[Literal["db", "config"]]:
def get_source(self, guardrail_id: str) -> Literal["db", "config"] | None:
"""
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.
@ -665,8 +665,8 @@ class InMemoryGuardrailHandler:
@staticmethod
def _normalize_litellm_params_for_comparison(
params: Optional[Any],
) -> Optional[Dict[str, Any]]:
params: Any | None,
) -> dict[str, Any] | None:
"""
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
@ -738,9 +738,9 @@ class InMemoryGuardrailHandler:
def reinitialize_guardrail(
self,
guardrail: Guardrail,
config_file_path: Optional[str] = None,
config_file_path: str | None = None,
source: Literal["db", "config"] = "config",
) -> Optional[Guardrail]:
) -> Guardrail | None:
"""
Force re-initialization of a guardrail even if it exists in memory.
Removes old callback from litellm.callbacks and creates fresh instance.
@ -762,8 +762,8 @@ class InMemoryGuardrailHandler:
)
def sync_guardrail_from_db(
self, guardrail: Guardrail, config_file_path: Optional[str] = None
) -> Optional[Guardrail]:
self, guardrail: Guardrail, config_file_path: str | None = None
) -> Guardrail | None:
"""
Sync a guardrail from DB - initializes if new, re-initializes if changed.
This is the method to call during DB polling.

View file

@ -151,8 +151,16 @@ class Violation(NamedTuple):
# --------------------------------------------------------------------------- #
def contains_any(t: Type, _seen: set[int] | None = None) -> bool:
_MAX_CONTAINS_ANY_DEPTH = 64
def contains_any(t: Type, _seen: set[int] | None = None, _depth: int = 0) -> bool:
"""True if a *value* of type ``t`` carries `Any` anywhere meaningful."""
if _depth > _MAX_CONTAINS_ANY_DEPTH:
# Bail out on deeply-nested / potentially-circular types rather than
# overflowing the Python call stack. A type this deep is unlikely to
# carry a *meaningful* Any that the developer could actually fix.
return False
seen = _seen if _seen is not None else set()
p = get_proper_type(t)
if id(p) in seen:
@ -166,11 +174,11 @@ def contains_any(t: Type, _seen: set[int] | None = None) -> bool:
if isinstance(p, AnyType):
return p.type_of_any not in _HARMLESS_ANY
if isinstance(p, UnionType):
return any(contains_any(item, seen) for item in p.items)
return any(contains_any(item, seen, _depth + 1) for item in p.items)
if isinstance(p, Instance):
return any(contains_any(arg, seen) for arg in p.args)
return any(contains_any(arg, seen, _depth + 1) for arg in p.args)
if isinstance(p, TupleType):
return any(contains_any(item, seen) for item in p.items)
return any(contains_any(item, seen, _depth + 1) for item in p.items)
return False
@ -249,17 +257,11 @@ def _reason_ok(reason: str | None) -> bool:
return reason is not None and len(reason.strip()) >= MIN_REASON_LEN
def scan_any_ok(
path: Path, source: str
) -> tuple[frozenset[int], tuple[Violation, ...]]:
def scan_any_ok(path: Path, source: str) -> tuple[frozenset[int], tuple[Violation, ...]]:
"""Return (lines with a valid any-ok suppression, LIT005 violations)."""
try:
tokens = tokenize.generate_tokens(
iter(source.splitlines(keepends=True)).__next__
)
comments = tuple(
(t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT
)
tokens = tokenize.generate_tokens(iter(source.splitlines(keepends=True)).__next__)
comments = tuple((t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT)
except tokenize.TokenError:
return frozenset(), ()
@ -370,9 +372,7 @@ def check_files(rel_paths: Sequence[str]) -> tuple[Violation, ...]:
try:
source = abs_path.read_text(encoding="utf-8")
except (OSError, UnicodeDecodeError) as exc:
out.append(
Violation(report_path, 0, 0, "LIT000", f"could not read file: {exc}")
)
out.append(Violation(report_path, 0, 0, "LIT000", f"could not read file: {exc}"))
continue
ok_lines, ok_violations = scan_any_ok(report_path, source)
@ -495,9 +495,7 @@ def _in_scope(v: Violation, line_map: dict[str, LineScope] | None) -> bool:
def main(argv: Sequence[str]) -> int:
parser = argparse.ArgumentParser(
description="Any-discipline gate (changed-only, changed-lines)."
)
parser = argparse.ArgumentParser(description="Any-discipline gate (changed-only, changed-lines).")
parser.add_argument(
"paths",
nargs="*",
@ -520,9 +518,7 @@ def main(argv: Sequence[str]) -> int:
file=sys.stderr,
)
return 0
rel_paths = _to_litellm_relative(
(REPO_ROOT / name).resolve() for name in line_map
)
rel_paths = _to_litellm_relative((REPO_ROOT / name).resolve() for name in line_map)
elif args.paths:
rel_paths = _to_litellm_relative((REPO_ROOT / p).resolve() for p in args.paths)
else:
@ -546,9 +542,7 @@ def main(argv: Sequence[str]) -> int:
file=sys.stderr,
)
return 1
print(
f"OK: {len(rel_paths)} changed file(s) under litellm/ have no Any-typed values on changed lines"
)
print(f"OK: {len(rel_paths)} changed file(s) under litellm/ have no Any-typed values on changed lines")
return 0