mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(coverage): revert ruff UP006/UP045 changes on upstream files
The previous ruff fixes (Dict→dict, Optional→X|None) on 7 upstream files added ~500 changed lines of pure type-annotation no-ops to our PR diff. codecov/patch penalised these uncovered lines, dropping patch coverage to 51.35% (target 61.83%). Fix: revert these files to exactly match upstream/litellm_internal_staging. The ruff_strict_gate still passes because the violations exist equally in both the base and HEAD (total == base_count → no breach).
This commit is contained in:
parent
c7665a5227
commit
0b0e3ed8bd
7 changed files with 526 additions and 509 deletions
|
|
@ -10,7 +10,7 @@ Supported for both `v1/chat/completions` (via the prompt-management hook) and
|
|||
"""
|
||||
|
||||
import copy
|
||||
from typing import TYPE_CHECKING, Any, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -39,27 +39,27 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
def get_chat_completion_prompt(
|
||||
self,
|
||||
model: str,
|
||||
messages: list[AllMessageValues],
|
||||
messages: List[AllMessageValues],
|
||||
non_default_params: dict,
|
||||
prompt_id: str | None,
|
||||
prompt_variables: dict | None,
|
||||
prompt_id: Optional[str],
|
||||
prompt_variables: Optional[dict],
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
prompt_spec: PromptSpec | None = None,
|
||||
prompt_label: str | None = None,
|
||||
prompt_version: int | None = None,
|
||||
ignore_prompt_manager_model: bool | None = False,
|
||||
ignore_prompt_manager_optional_params: bool | None = False,
|
||||
) -> tuple[str, list[AllMessageValues], dict]:
|
||||
prompt_spec: Optional[PromptSpec] = None,
|
||||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
ignore_prompt_manager_model: Optional[bool] = False,
|
||||
ignore_prompt_manager_optional_params: Optional[bool] = False,
|
||||
) -> Tuple[str, List[AllMessageValues], dict]:
|
||||
"""
|
||||
Apply cache control directives based on specified injection points.
|
||||
|
||||
Returns:
|
||||
- model: str - the model to use
|
||||
- messages: list[AllMessageValues] - messages with applied cache controls
|
||||
- messages: List[AllMessageValues] - messages with applied cache controls
|
||||
- non_default_params: dict - params with any global cache controls
|
||||
"""
|
||||
# Extract cache control injection points
|
||||
injection_points: list[CacheControlInjectionPoint] = non_default_params.pop(
|
||||
injection_points: List[CacheControlInjectionPoint] = non_default_params.pop(
|
||||
"cache_control_injection_points", []
|
||||
)
|
||||
if not injection_points:
|
||||
|
|
@ -69,8 +69,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
processed_messages = copy.deepcopy(messages)
|
||||
|
||||
# Separate message-level and non-message-level injection points
|
||||
message_points: list[CacheControlMessageInjectionPoint] = []
|
||||
remaining_points: list[CacheControlInjectionPoint] = []
|
||||
message_points: List[CacheControlMessageInjectionPoint] = []
|
||||
remaining_points: List[CacheControlInjectionPoint] = []
|
||||
for point in injection_points:
|
||||
if point.get("location") == "message":
|
||||
message_points.append(cast(CacheControlMessageInjectionPoint, point))
|
||||
|
|
@ -97,10 +97,10 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
@staticmethod
|
||||
def _apply_message_injections(
|
||||
points: list[CacheControlMessageInjectionPoint],
|
||||
messages: list[AllMessageValues],
|
||||
points: List[CacheControlMessageInjectionPoint],
|
||||
messages: List[AllMessageValues],
|
||||
max_blocks: int,
|
||||
) -> list[AllMessageValues]:
|
||||
) -> List[AllMessageValues]:
|
||||
"""Apply message-level cache control injection points in order.
|
||||
|
||||
Anthropic allows at most ``MAX_CACHE_CONTROL_BLOCKS`` cache_control
|
||||
|
|
@ -149,11 +149,11 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
@staticmethod
|
||||
def _resolve_target_indices(
|
||||
point: CacheControlMessageInjectionPoint, messages: list[AllMessageValues]
|
||||
) -> list[int]:
|
||||
point: CacheControlMessageInjectionPoint, messages: List[AllMessageValues]
|
||||
) -> List[int]:
|
||||
"""Resolve which message indices an injection point targets."""
|
||||
_targetted_index: Union[int, str] | None = point.get("index", None)
|
||||
targetted_index: int | None = None
|
||||
_targetted_index: Optional[Union[int, str]] = point.get("index", None)
|
||||
targetted_index: Optional[int] = None
|
||||
if isinstance(_targetted_index, str):
|
||||
try:
|
||||
targetted_index = int(_targetted_index)
|
||||
|
|
@ -230,10 +230,10 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
@staticmethod
|
||||
def apply_to_anthropic_messages_request(
|
||||
messages: list[dict],
|
||||
messages: List[Dict],
|
||||
system: str | list | None,
|
||||
injection_points: list[CacheControlInjectionPoint],
|
||||
) -> tuple[list[dict], str | list | None, list[CacheControlInjectionPoint]]:
|
||||
injection_points: List[CacheControlInjectionPoint],
|
||||
) -> Tuple[List[Dict], str | list | None, List[CacheControlInjectionPoint]]:
|
||||
"""Apply cache control injection for the Anthropic-native v1/messages endpoint.
|
||||
|
||||
Returns (messages, system, remaining_non_message_points).
|
||||
|
|
@ -241,12 +241,12 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
if not injection_points:
|
||||
return messages, system, []
|
||||
|
||||
processed_messages: list[dict] = copy.deepcopy(messages)
|
||||
processed_messages: List[Dict] = copy.deepcopy(messages)
|
||||
processed_system = copy.deepcopy(system) if system is not None else None
|
||||
|
||||
message_points: list[CacheControlMessageInjectionPoint] = []
|
||||
system_points: list[CacheControlMessageInjectionPoint] = []
|
||||
remaining_points: list[CacheControlInjectionPoint] = []
|
||||
message_points: List[CacheControlMessageInjectionPoint] = []
|
||||
system_points: List[CacheControlMessageInjectionPoint] = []
|
||||
remaining_points: List[CacheControlInjectionPoint] = []
|
||||
|
||||
for point in injection_points:
|
||||
if point.get("location") == "message":
|
||||
|
|
@ -290,7 +290,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
processed_messages = AnthropicCacheControlHook._apply_message_injections(
|
||||
points=message_points,
|
||||
messages=cast(list[AllMessageValues], processed_messages),
|
||||
messages=cast(List[AllMessageValues], processed_messages),
|
||||
max_blocks=max_blocks - used_blocks,
|
||||
)
|
||||
|
||||
|
|
@ -298,10 +298,10 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
@staticmethod
|
||||
def maybe_inject_cache_control(
|
||||
messages: list[dict],
|
||||
messages: List[Dict],
|
||||
system: str | list | None,
|
||||
kwargs: dict[str, Any],
|
||||
) -> tuple[list[dict], str | list | None]:
|
||||
kwargs: Dict[str, Any],
|
||||
) -> Tuple[List[Dict], str | list | None]:
|
||||
"""Extract cache_control_injection_points from kwargs and apply if present.
|
||||
|
||||
Pops the key from kwargs; if remaining (non-message) points exist they
|
||||
|
|
@ -327,8 +327,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
def should_run_prompt_management(
|
||||
self,
|
||||
prompt_id: str | None,
|
||||
prompt_spec: PromptSpec | None,
|
||||
prompt_id: Optional[str],
|
||||
prompt_spec: Optional[PromptSpec],
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
) -> bool:
|
||||
"""Always return False since this is not a true prompt management system."""
|
||||
|
|
@ -336,12 +336,12 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
def _compile_prompt_helper(
|
||||
self,
|
||||
prompt_id: str | None,
|
||||
prompt_spec: PromptSpec | None,
|
||||
prompt_variables: dict | None,
|
||||
prompt_id: Optional[str],
|
||||
prompt_spec: Optional[PromptSpec],
|
||||
prompt_variables: Optional[dict],
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
prompt_label: str | None = None,
|
||||
prompt_version: int | None = None,
|
||||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
) -> PromptManagementClient:
|
||||
"""Not used - this hook only modifies messages, doesn't fetch prompts."""
|
||||
return PromptManagementClient(
|
||||
|
|
@ -354,12 +354,12 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
async def async_compile_prompt_helper(
|
||||
self,
|
||||
prompt_id: str | None,
|
||||
prompt_variables: dict | None,
|
||||
prompt_id: Optional[str],
|
||||
prompt_variables: Optional[dict],
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
prompt_spec: PromptSpec | None = None,
|
||||
prompt_label: str | None = None,
|
||||
prompt_version: int | None = None,
|
||||
prompt_spec: Optional[PromptSpec] = None,
|
||||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
) -> PromptManagementClient:
|
||||
"""Not used - this hook only modifies messages, doesn't fetch prompts."""
|
||||
return self._compile_prompt_helper(
|
||||
|
|
@ -374,19 +374,19 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
async def async_get_chat_completion_prompt(
|
||||
self,
|
||||
model: str,
|
||||
messages: list[AllMessageValues],
|
||||
messages: List[AllMessageValues],
|
||||
non_default_params: dict,
|
||||
prompt_id: str | None,
|
||||
prompt_variables: dict | None,
|
||||
prompt_id: Optional[str],
|
||||
prompt_variables: Optional[dict],
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
prompt_spec: PromptSpec | None = None,
|
||||
tools: list[dict] | None = None,
|
||||
prompt_label: str | None = None,
|
||||
prompt_version: int | None = None,
|
||||
ignore_prompt_manager_model: bool | None = False,
|
||||
ignore_prompt_manager_optional_params: bool | None = False,
|
||||
) -> tuple[str, list[AllMessageValues], dict]:
|
||||
prompt_spec: Optional[PromptSpec] = None,
|
||||
tools: Optional[List[Dict]] = None,
|
||||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
ignore_prompt_manager_model: Optional[bool] = False,
|
||||
ignore_prompt_manager_optional_params: Optional[bool] = False,
|
||||
) -> Tuple[str, List[AllMessageValues], dict]:
|
||||
"""Async version - delegates to sync since no async operations needed."""
|
||||
return self.get_chat_completion_prompt(
|
||||
model=model,
|
||||
|
|
@ -403,15 +403,15 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def should_use_anthropic_cache_control_hook(non_default_params: dict) -> bool:
|
||||
def should_use_anthropic_cache_control_hook(non_default_params: Dict) -> bool:
|
||||
if non_default_params.get("cache_control_injection_points", None):
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def get_custom_logger_for_anthropic_cache_control_hook(
|
||||
non_default_params: dict,
|
||||
) -> CustomLogger | None:
|
||||
non_default_params: Dict,
|
||||
) -> Optional[CustomLogger]:
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
_init_custom_logger_compatible_class,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,8 +4,11 @@ from typing import (
|
|||
TYPE_CHECKING,
|
||||
Any,
|
||||
ClassVar,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Type,
|
||||
Union,
|
||||
get_args,
|
||||
)
|
||||
|
|
@ -56,7 +59,7 @@ from litellm.exceptions import (
|
|||
_PRE_CALL_EXECUTED_TOKEN = secrets.token_hex(16)
|
||||
|
||||
|
||||
def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None:
|
||||
def get_session_id_from_request_data(request_data: Dict[str, Any]) -> Optional[str]:
|
||||
"""Extract session_id from request data (litellm_session_id or metadata)."""
|
||||
session_id = request_data.get("litellm_session_id")
|
||||
if session_id:
|
||||
|
|
@ -81,18 +84,18 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
guardrail_name: str | None = None,
|
||||
supported_event_hooks: list[GuardrailEventHooks] | None = None,
|
||||
event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = None,
|
||||
guardrail_name: Optional[str] = None,
|
||||
supported_event_hooks: Optional[List[GuardrailEventHooks]] = None,
|
||||
event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]] = None,
|
||||
default_on: bool = False,
|
||||
mask_request_content: bool = False,
|
||||
mask_response_content: bool = False,
|
||||
violation_message_template: str | None = None,
|
||||
end_session_after_n_fails: int | None = None,
|
||||
on_violation: str | None = None,
|
||||
realtime_violation_message: str | None = None,
|
||||
on_sensitive_data: str | None = None,
|
||||
sensitive_data_route_to_model: str | None = None,
|
||||
violation_message_template: Optional[str] = None,
|
||||
end_session_after_n_fails: Optional[int] = None,
|
||||
on_violation: Optional[str] = None,
|
||||
realtime_violation_message: Optional[str] = None,
|
||||
on_sensitive_data: Optional[str] = None,
|
||||
sensitive_data_route_to_model: Optional[str] = None,
|
||||
sticky_session_routing: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
|
|
@ -115,16 +118,16 @@ class CustomGuardrail(CustomLogger):
|
|||
"""
|
||||
self.guardrail_name = guardrail_name
|
||||
self.supported_event_hooks = supported_event_hooks
|
||||
self.event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = event_hook
|
||||
self.event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]] = event_hook
|
||||
self.default_on: bool = default_on
|
||||
self.mask_request_content: bool = mask_request_content
|
||||
self.mask_response_content: bool = mask_response_content
|
||||
self.violation_message_template: str | None = violation_message_template
|
||||
self.end_session_after_n_fails: int | None = end_session_after_n_fails
|
||||
self.on_violation: str | None = on_violation
|
||||
self.realtime_violation_message: str | None = realtime_violation_message
|
||||
self.on_sensitive_data: str | None = on_sensitive_data
|
||||
self.sensitive_data_route_to_model: str | None = sensitive_data_route_to_model
|
||||
self.violation_message_template: Optional[str] = violation_message_template
|
||||
self.end_session_after_n_fails: Optional[int] = end_session_after_n_fails
|
||||
self.on_violation: Optional[str] = on_violation
|
||||
self.realtime_violation_message: Optional[str] = realtime_violation_message
|
||||
self.on_sensitive_data: Optional[str] = on_sensitive_data
|
||||
self.sensitive_data_route_to_model: Optional[str] = sensitive_data_route_to_model
|
||||
self.sticky_session_routing: bool = sticky_session_routing
|
||||
|
||||
if supported_event_hooks:
|
||||
|
|
@ -132,13 +135,13 @@ class CustomGuardrail(CustomLogger):
|
|||
self._validate_event_hook(event_hook, supported_event_hooks)
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def render_violation_message(self, default: str, context: dict[str, Any] | None = None) -> str:
|
||||
def render_violation_message(self, default: str, context: Optional[Dict[str, Any]] = None) -> str:
|
||||
"""Return a custom violation message if template is configured."""
|
||||
|
||||
if not self.violation_message_template:
|
||||
return default
|
||||
|
||||
format_context: dict[str, Any] = {"default_message": default}
|
||||
format_context: Dict[str, Any] = {"default_message": default}
|
||||
if context:
|
||||
format_context.update(context)
|
||||
try:
|
||||
|
|
@ -154,8 +157,8 @@ class CustomGuardrail(CustomLogger):
|
|||
def raise_passthrough_exception(
|
||||
self,
|
||||
violation_message: str,
|
||||
request_data: dict[str, Any],
|
||||
detection_info: dict[str, Any] | None = None,
|
||||
request_data: Dict[str, Any],
|
||||
detection_info: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Raise a passthrough exception for guardrail violations.
|
||||
|
|
@ -198,8 +201,8 @@ class CustomGuardrail(CustomLogger):
|
|||
def raise_sensitive_data_route_exception(
|
||||
self,
|
||||
route_to_model: str,
|
||||
request_data: dict[str, Any],
|
||||
detection_info: dict[str, Any] | None = None,
|
||||
request_data: Dict[str, Any],
|
||||
detection_info: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Raise an exception to reroute the request to a different model.
|
||||
|
|
@ -236,7 +239,7 @@ class CustomGuardrail(CustomLogger):
|
|||
sticky_session_routing=self.sticky_session_routing,
|
||||
)
|
||||
|
||||
def _get_session_id_from_request_data(self, request_data: dict[str, Any]) -> str | None:
|
||||
def _get_session_id_from_request_data(self, request_data: Dict[str, Any]) -> Optional[str]:
|
||||
"""Extract session_id from request data."""
|
||||
return get_session_id_from_request_data(request_data)
|
||||
|
||||
|
|
@ -249,8 +252,8 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def handle_sensitive_data_detection(
|
||||
self,
|
||||
request_data: dict[str, Any],
|
||||
detection_info: dict[str, Any] | None = None,
|
||||
request_data: Dict[str, Any],
|
||||
detection_info: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Handle sensitive data detection based on guardrail configuration.
|
||||
|
|
@ -292,7 +295,7 @@ class CustomGuardrail(CustomLogger):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> type["GuardrailConfigModel"] | None:
|
||||
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
|
||||
"""
|
||||
Returns the config model for the guardrail
|
||||
|
||||
|
|
@ -302,12 +305,12 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def _validate_event_hook(
|
||||
self,
|
||||
event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None,
|
||||
supported_event_hooks: list[GuardrailEventHooks],
|
||||
event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]],
|
||||
supported_event_hooks: List[GuardrailEventHooks],
|
||||
) -> None:
|
||||
def _validate_event_hook_list_is_in_supported_event_hooks(
|
||||
event_hook: Union[list[GuardrailEventHooks], list[str]],
|
||||
supported_event_hooks: list[GuardrailEventHooks],
|
||||
event_hook: Union[List[GuardrailEventHooks], List[str]],
|
||||
supported_event_hooks: List[GuardrailEventHooks],
|
||||
) -> None:
|
||||
for hook in event_hook:
|
||||
if isinstance(hook, str):
|
||||
|
|
@ -358,7 +361,7 @@ class CustomGuardrail(CustomLogger):
|
|||
key_meta = meta.get("user_api_key_metadata") or key_meta
|
||||
return {**team_meta, **key_meta}
|
||||
|
||||
def get_disable_global_guardrail(self, data: dict) -> bool | None:
|
||||
def get_disable_global_guardrail(self, data: dict) -> Optional[bool]:
|
||||
"""
|
||||
Returns True if the global guardrail should be disabled.
|
||||
|
||||
|
|
@ -367,7 +370,7 @@ class CustomGuardrail(CustomLogger):
|
|||
"""
|
||||
return self._get_admin_metadata(data).get("disable_global_guardrails", False)
|
||||
|
||||
def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> list[str]:
|
||||
def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> List[str]:
|
||||
"""
|
||||
Returns the list of global guardrail names the team/key has opted out of.
|
||||
|
||||
|
|
@ -398,7 +401,7 @@ class CustomGuardrail(CustomLogger):
|
|||
return True
|
||||
raise
|
||||
|
||||
def get_guardrail_from_metadata(self, data: dict) -> Union[list[str], list[dict[str, DynamicGuardrailParams]]]:
|
||||
def get_guardrail_from_metadata(self, data: dict) -> Union[List[str], List[Dict[str, DynamicGuardrailParams]]]:
|
||||
"""
|
||||
Returns the guardrail(s) to be run from the metadata or root
|
||||
"""
|
||||
|
|
@ -419,7 +422,7 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def _guardrail_is_in_requested_guardrails(
|
||||
self,
|
||||
requested_guardrails: Union[list[str], list[dict[str, DynamicGuardrailParams]]],
|
||||
requested_guardrails: Union[List[str], List[Dict[str, DynamicGuardrailParams]]],
|
||||
) -> bool:
|
||||
for _guardrail in requested_guardrails:
|
||||
if isinstance(_guardrail, dict):
|
||||
|
|
@ -431,13 +434,13 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
return False
|
||||
|
||||
def _pre_call_marker(self) -> str | None:
|
||||
def _pre_call_marker(self) -> Optional[str]:
|
||||
name = self.guardrail_name
|
||||
if not name:
|
||||
return None
|
||||
return f"{_PRE_CALL_EXECUTED_TOKEN}:{name}"
|
||||
|
||||
def mark_pre_call_hook_ran(self, data: dict[str, Any]) -> None:
|
||||
def mark_pre_call_hook_ran(self, data: Dict[str, Any]) -> None:
|
||||
"""
|
||||
Record that this guardrail's ``async_pre_call_hook`` already ran for this
|
||||
request, so the deployment-level hook does not run it a second time.
|
||||
|
|
@ -462,7 +465,7 @@ class CustomGuardrail(CustomLogger):
|
|||
return
|
||||
data["metadata"] = {PRE_CALL_EXECUTED_GUARDRAILS_KEY: [marker]}
|
||||
|
||||
def _pre_call_hook_already_ran(self, data: dict[str, Any]) -> bool:
|
||||
def _pre_call_hook_already_ran(self, data: Dict[str, Any]) -> bool:
|
||||
marker = self._pre_call_marker()
|
||||
if marker is None:
|
||||
return False
|
||||
|
|
@ -474,7 +477,9 @@ class CustomGuardrail(CustomLogger):
|
|||
return True
|
||||
return False
|
||||
|
||||
async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None:
|
||||
async def async_pre_call_deployment_hook(
|
||||
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
|
||||
) -> Optional[dict]:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
# should run guardrail
|
||||
|
|
@ -514,8 +519,8 @@ class CustomGuardrail(CustomLogger):
|
|||
self,
|
||||
request_data: dict,
|
||||
response: LLMResponseTypes,
|
||||
call_type: CallTypes | None,
|
||||
) -> LLMResponseTypes | None:
|
||||
call_type: Optional[CallTypes],
|
||||
) -> Optional[LLMResponseTypes]:
|
||||
"""
|
||||
Allow modifying / reviewing the response just after it's received from the deployment.
|
||||
"""
|
||||
|
|
@ -697,16 +702,16 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def add_standard_logging_guardrail_information_to_request_data(
|
||||
self,
|
||||
guardrail_json_response: Union[Exception, str, dict, list[dict]],
|
||||
guardrail_json_response: Union[Exception, str, dict, List[dict]],
|
||||
request_data: dict,
|
||||
guardrail_status: GuardrailStatus,
|
||||
start_time: float | None = None,
|
||||
end_time: float | None = None,
|
||||
duration: float | None = None,
|
||||
masked_entity_count: dict[str, int] | None = None,
|
||||
guardrail_provider: str | None = None,
|
||||
event_type: GuardrailEventHooks | None = None,
|
||||
tracing_detail: GuardrailTracingDetail | None = None,
|
||||
start_time: Optional[float] = None,
|
||||
end_time: Optional[float] = None,
|
||||
duration: Optional[float] = None,
|
||||
masked_entity_count: Optional[Dict[str, int]] = None,
|
||||
guardrail_provider: Optional[str] = None,
|
||||
event_type: Optional[GuardrailEventHooks] = None,
|
||||
tracing_detail: Optional[GuardrailTracingDetail] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Builds `StandardLoggingGuardrailInformation` and adds it to the request metadata so it can be used for logging to DataDog, Langfuse, etc.
|
||||
|
|
@ -721,7 +726,7 @@ class CustomGuardrail(CustomLogger):
|
|||
from litellm.types.utils import GuardrailMode
|
||||
|
||||
# Use event_type if provided, otherwise fall back to self.event_hook
|
||||
guardrail_mode: Union[GuardrailEventHooks, GuardrailMode, list[GuardrailEventHooks]]
|
||||
guardrail_mode: Union[GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks]]
|
||||
if event_type is not None:
|
||||
guardrail_mode = event_type
|
||||
elif isinstance(self.event_hook, Mode):
|
||||
|
|
@ -830,13 +835,13 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def _process_response(
|
||||
self,
|
||||
response: dict | None,
|
||||
response: Optional[Dict],
|
||||
request_data: dict,
|
||||
start_time: float | None = None,
|
||||
end_time: float | None = None,
|
||||
duration: float | None = None,
|
||||
event_type: GuardrailEventHooks | None = None,
|
||||
original_inputs: dict | None = None,
|
||||
start_time: Optional[float] = None,
|
||||
end_time: Optional[float] = None,
|
||||
duration: Optional[float] = None,
|
||||
event_type: Optional[GuardrailEventHooks] = None,
|
||||
original_inputs: Optional[Dict] = None,
|
||||
):
|
||||
"""
|
||||
Add StandardLoggingGuardrailInformation to the request data
|
||||
|
|
@ -844,7 +849,7 @@ class CustomGuardrail(CustomLogger):
|
|||
This gets logged on downsteam Langfuse, DataDog, etc.
|
||||
"""
|
||||
# Convert None to empty dict to satisfy type requirements
|
||||
guardrail_response: Union[dict[str, Any], str] = {} if response is None else response
|
||||
guardrail_response: Union[Dict[str, Any], str] = {} if response is None else response
|
||||
|
||||
# For apply_guardrail functions in custom_code_guardrail scenario,
|
||||
# simplify the logged response to "allow", "deny", or "mask"
|
||||
|
|
@ -900,10 +905,10 @@ class CustomGuardrail(CustomLogger):
|
|||
self,
|
||||
e: Exception,
|
||||
request_data: dict,
|
||||
start_time: float | None = None,
|
||||
end_time: float | None = None,
|
||||
duration: float | None = None,
|
||||
event_type: GuardrailEventHooks | None = None,
|
||||
start_time: Optional[float] = None,
|
||||
end_time: Optional[float] = None,
|
||||
duration: Optional[float] = None,
|
||||
event_type: Optional[GuardrailEventHooks] = None,
|
||||
):
|
||||
"""
|
||||
Add StandardLoggingGuardrailInformation to the request data
|
||||
|
|
@ -930,7 +935,7 @@ class CustomGuardrail(CustomLogger):
|
|||
)
|
||||
raise e
|
||||
|
||||
def _inputs_were_modified(self, original_inputs: dict, response: dict) -> bool:
|
||||
def _inputs_were_modified(self, original_inputs: Dict, response: Dict) -> bool:
|
||||
"""
|
||||
Compare original inputs with response to determine if content was modified.
|
||||
|
||||
|
|
@ -975,8 +980,8 @@ class CustomGuardrail(CustomLogger):
|
|||
setattr(self, key, value)
|
||||
|
||||
def get_guardrails_messages_for_call_type(
|
||||
self, call_type: CallTypes, data: dict | None = None
|
||||
) -> list[AllMessageValues] | None:
|
||||
self, call_type: CallTypes, data: Optional[dict] = None
|
||||
) -> Optional[List[AllMessageValues]]:
|
||||
"""
|
||||
Returns the messages for the given call type and data
|
||||
"""
|
||||
|
|
@ -1014,7 +1019,7 @@ class CustomGuardrail(CustomLogger):
|
|||
input=input_data,
|
||||
responses_api_request=data,
|
||||
)
|
||||
return cast(list[AllMessageValues], messages)
|
||||
return cast(List[AllMessageValues], messages)
|
||||
return None
|
||||
|
||||
|
||||
|
|
@ -1078,7 +1083,7 @@ def log_guardrail_information(func):
|
|||
|
||||
def _infer_event_type_from_function_name(
|
||||
func_name: str,
|
||||
) -> GuardrailEventHooks | None:
|
||||
) -> Optional[GuardrailEventHooks]:
|
||||
"""Infer the actual event type from the function name"""
|
||||
if func_name == "async_pre_call_hook":
|
||||
return GuardrailEventHooks.pre_call
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ identical metrics. The attribute cardinality filter is reused from v1 by import
|
|||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any, Mapping
|
||||
from typing import Any, FrozenSet, Mapping, Optional
|
||||
|
||||
from opentelemetry.metrics import Histogram, Meter
|
||||
|
||||
|
|
@ -81,11 +81,11 @@ class GenAIMetricRecorder:
|
|||
survives.
|
||||
"""
|
||||
|
||||
def __init__(self, metrics: GenAIMetrics, callback_name: str | None = None) -> None:
|
||||
def __init__(self, metrics: GenAIMetrics, callback_name: Optional[str] = None) -> None:
|
||||
self._metrics = metrics
|
||||
self._callback_name = callback_name
|
||||
self._include: frozenset[str] | None = None
|
||||
self._exclude: frozenset[str] | None = None
|
||||
self._include: Optional[FrozenSet[str]] = None
|
||||
self._exclude: Optional[FrozenSet[str]] = None
|
||||
self._filter_resolved = False
|
||||
|
||||
def record(
|
||||
|
|
|
|||
|
|
@ -4,7 +4,11 @@ import time
|
|||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Dict,
|
||||
List,
|
||||
NoReturn,
|
||||
Optional,
|
||||
Tuple,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
|
@ -146,8 +150,8 @@ def _basic_sanitize_anthropic_tool_name(name: str) -> str:
|
|||
|
||||
|
||||
def _build_anthropic_tool_name_maps(
|
||||
original_names: list[str],
|
||||
) -> tuple[dict[str, str], dict[str, str]]:
|
||||
original_names: List[str],
|
||||
) -> Tuple[Dict[str, str], Dict[str, str]]:
|
||||
"""Build (forward, reverse) tool-name maps for a single request.
|
||||
|
||||
forward[original] = sanitized -- only present when name was rewritten
|
||||
|
|
@ -167,7 +171,7 @@ def _build_anthropic_tool_name_maps(
|
|||
seen gets the disambiguating suffix. Callers should preserve the
|
||||
caller's tool order (we do).
|
||||
"""
|
||||
forward: dict[str, str] = {}
|
||||
forward: Dict[str, str] = {}
|
||||
used: set = set()
|
||||
|
||||
# First pass: reserve slots for names that are already valid so they
|
||||
|
|
@ -208,7 +212,7 @@ def _build_anthropic_tool_name_maps(
|
|||
return forward, reverse
|
||||
|
||||
|
||||
REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT: dict[str, str] = {
|
||||
REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT: Dict[str, str] = {
|
||||
"low": "low",
|
||||
"minimal": "low",
|
||||
"medium": "medium",
|
||||
|
|
@ -235,23 +239,23 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
to pass metadata to anthropic, it's {"user_id": "any-relevant-information"}
|
||||
"""
|
||||
|
||||
max_tokens: int | None = None
|
||||
stop_sequences: list | None = None
|
||||
temperature: int | None = None
|
||||
top_p: int | None = None
|
||||
top_k: int | None = None
|
||||
metadata: dict | None = None
|
||||
system: str | None = None
|
||||
max_tokens: Optional[int] = None
|
||||
stop_sequences: Optional[list] = None
|
||||
temperature: Optional[int] = None
|
||||
top_p: Optional[int] = None
|
||||
top_k: Optional[int] = None
|
||||
metadata: Optional[dict] = None
|
||||
system: Optional[str] = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_tokens: int | None = None,
|
||||
stop_sequences: list | None = None,
|
||||
temperature: int | None = None,
|
||||
top_p: int | None = None,
|
||||
top_k: int | None = None,
|
||||
metadata: dict | None = None,
|
||||
system: str | None = None,
|
||||
max_tokens: Optional[int] = None,
|
||||
stop_sequences: Optional[list] = None,
|
||||
temperature: Optional[int] = None,
|
||||
top_p: Optional[int] = None,
|
||||
top_k: Optional[int] = None,
|
||||
metadata: Optional[dict] = None,
|
||||
system: Optional[str] = None,
|
||||
) -> None:
|
||||
locals_ = locals().copy()
|
||||
for key, value in locals_.items():
|
||||
|
|
@ -259,11 +263,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
setattr(self.__class__, key, value)
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> str | None:
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "anthropic"
|
||||
|
||||
@classmethod
|
||||
def get_config(cls, *, model: str | None = None):
|
||||
def get_config(cls, *, model: Optional[str] = None):
|
||||
config = super().get_config()
|
||||
|
||||
# anthropic requires a default value for max_tokens
|
||||
|
|
@ -273,7 +277,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return config
|
||||
|
||||
@staticmethod
|
||||
def get_max_tokens_for_model(model: str | None = None) -> int:
|
||||
def get_max_tokens_for_model(model: Optional[str] = None) -> int:
|
||||
"""
|
||||
Get the max output tokens for a given model.
|
||||
Falls back to DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS (configurable via env var) if model is not found.
|
||||
|
|
@ -290,7 +294,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def convert_tool_use_to_openai_format(
|
||||
anthropic_tool_content: dict[str, Any],
|
||||
anthropic_tool_content: Dict[str, Any],
|
||||
index: int,
|
||||
) -> ChatCompletionToolCallChunk:
|
||||
"""
|
||||
|
|
@ -315,7 +319,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
)
|
||||
# Include caller information if present (for programmatic tool calling)
|
||||
if "caller" in anthropic_tool_content:
|
||||
tool_call["caller"] = cast(dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item]
|
||||
tool_call["caller"] = cast(Dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item]
|
||||
return tool_call
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -336,7 +340,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return AnthropicConfig._supports_model_capability(model, f"supports_{level}_reasoning_effort")
|
||||
|
||||
@staticmethod
|
||||
def _validate_effort_for_model(model: str, effort: str | None) -> str | None:
|
||||
def _validate_effort_for_model(model: str, effort: Optional[str]) -> Optional[str]:
|
||||
"""Return ``None`` if ``effort`` is allowed on ``model``, else an error message."""
|
||||
if effort == "max" and not (
|
||||
AnthropicConfig._is_adaptive_thinking_model(model) or AnthropicConfig._supports_effort_level(model, "max")
|
||||
|
|
@ -363,7 +367,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _model_supports_speed_param(model: str, custom_llm_provider: str | None = None) -> bool:
|
||||
def _model_supports_speed_param(model: str, custom_llm_provider: Optional[str] = None) -> bool:
|
||||
"""Whether the model accepts Anthropic's ``speed`` parameter (fast mode).
|
||||
|
||||
Fast mode is direct Anthropic API-only (not Bedrock, Vertex, or Azure).
|
||||
|
|
@ -380,7 +384,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
model: str,
|
||||
optional_params: dict,
|
||||
drop_params: bool,
|
||||
custom_llm_provider: str | None = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
) -> None:
|
||||
if "speed" not in optional_params:
|
||||
return
|
||||
|
|
@ -459,7 +463,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return params
|
||||
|
||||
@staticmethod
|
||||
def filter_anthropic_output_schema(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
def filter_anthropic_output_schema(schema: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Filter out unsupported fields from JSON schema for Anthropic's output_format API.
|
||||
|
||||
|
|
@ -515,7 +519,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
if field in schema:
|
||||
constraint_descriptions.append(constraint_labels[field].format(schema[field]))
|
||||
|
||||
result: dict[str, Any] = {}
|
||||
result: Dict[str, Any] = {}
|
||||
|
||||
# Update description with removed constraint info
|
||||
if constraint_descriptions:
|
||||
|
|
@ -555,7 +559,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
return result
|
||||
|
||||
def get_json_schema_from_pydantic_object(self, response_format: Union[Any, dict, None]) -> dict | None:
|
||||
def get_json_schema_from_pydantic_object(self, response_format: Union[Any, Dict, None]) -> Optional[dict]:
|
||||
return type_to_response_format_param(
|
||||
response_format, ref_template="/$defs/{model}"
|
||||
) # Relevant issue: https://github.com/BerriAI/litellm/issues/7755
|
||||
|
|
@ -570,10 +574,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _map_tool_choice(
|
||||
self,
|
||||
tool_choice: str | None,
|
||||
parallel_tool_use: bool | None,
|
||||
) -> AnthropicMessagesToolChoice | None:
|
||||
_tool_choice: AnthropicMessagesToolChoice | None = None
|
||||
tool_choice: Optional[str],
|
||||
parallel_tool_use: Optional[bool],
|
||||
) -> Optional[AnthropicMessagesToolChoice]:
|
||||
_tool_choice: Optional[AnthropicMessagesToolChoice] = None
|
||||
if tool_choice == "auto":
|
||||
_tool_choice = AnthropicMessagesToolChoice(
|
||||
type="auto",
|
||||
|
|
@ -614,9 +618,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def _map_tool_helper(
|
||||
self,
|
||||
tool: ChatCompletionToolParam,
|
||||
) -> tuple[AllAnthropicToolsValues | None, AnthropicMcpServerTool | None]:
|
||||
returned_tool: AllAnthropicToolsValues | None = None
|
||||
mcp_server: AnthropicMcpServerTool | None = None
|
||||
) -> Tuple[Optional[AllAnthropicToolsValues], Optional[AnthropicMcpServerTool]]:
|
||||
returned_tool: Optional[AllAnthropicToolsValues] = None
|
||||
mcp_server: Optional[AnthropicMcpServerTool] = None
|
||||
|
||||
if tool["type"] == "function" or tool["type"] == "custom":
|
||||
_input_schema: dict = tool["function"].get(
|
||||
|
|
@ -667,8 +671,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
if "parameters" not in tool["function"]:
|
||||
raise ValueError("Missing required parameter: parameters")
|
||||
|
||||
_display_width_px: int | None = tool["function"]["parameters"].get("display_width_px")
|
||||
_display_height_px: int | None = tool["function"]["parameters"].get("display_height_px")
|
||||
_display_width_px: Optional[int] = tool["function"]["parameters"].get("display_width_px")
|
||||
_display_height_px: Optional[int] = tool["function"]["parameters"].get("display_height_px")
|
||||
if _display_width_px is None or _display_height_px is None:
|
||||
raise ValueError("Missing required parameter: display_width_px or display_height_px")
|
||||
|
||||
|
|
@ -829,14 +833,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
from litellm.types.llms.anthropic import AnthropicMcpServerToolConfiguration
|
||||
|
||||
allowed_tools = tool.get("allowed_tools", None)
|
||||
tool_configuration: AnthropicMcpServerToolConfiguration | None = None
|
||||
tool_configuration: Optional[AnthropicMcpServerToolConfiguration] = None
|
||||
if allowed_tools is not None:
|
||||
tool_configuration = AnthropicMcpServerToolConfiguration(
|
||||
allowed_tools=tool.get("allowed_tools", None),
|
||||
)
|
||||
|
||||
headers = tool.get("headers", {})
|
||||
authorization_token: str | None = None
|
||||
authorization_token: Optional[str] = None
|
||||
if headers is not None:
|
||||
bearer_token = headers.get("Authorization", None)
|
||||
if bearer_token is not None:
|
||||
|
|
@ -856,8 +860,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _map_tools(
|
||||
self,
|
||||
tools: list,
|
||||
) -> tuple[list[AllAnthropicToolsValues], list[AnthropicMcpServerTool]]:
|
||||
tools: List,
|
||||
) -> Tuple[List[AllAnthropicToolsValues], List[AnthropicMcpServerTool]]:
|
||||
anthropic_tools = []
|
||||
mcp_servers = []
|
||||
for tool in tools:
|
||||
|
|
@ -902,9 +906,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _rewrite_tool_names_in_messages(
|
||||
messages: list[AllMessageValues],
|
||||
name_forward_map: dict[str, str],
|
||||
) -> list[AllMessageValues]:
|
||||
messages: List[AllMessageValues],
|
||||
name_forward_map: Dict[str, str],
|
||||
) -> List[AllMessageValues]:
|
||||
"""Return a copy of `messages` with tool_call/function_call names
|
||||
rewritten using the per-request forward map.
|
||||
|
||||
|
|
@ -915,7 +919,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
"""
|
||||
if not name_forward_map:
|
||||
return messages
|
||||
new_messages: list[AllMessageValues] = []
|
||||
new_messages: List[AllMessageValues] = []
|
||||
for msg in messages:
|
||||
if not isinstance(msg, dict):
|
||||
new_messages.append(msg)
|
||||
|
|
@ -953,8 +957,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _build_request_tool_name_maps(
|
||||
tools: list,
|
||||
) -> tuple[dict[str, str], dict[str, str]]:
|
||||
tools: List,
|
||||
) -> Tuple[Dict[str, str], Dict[str, str]]:
|
||||
"""Build the (forward, reverse) tool-name maps for an OpenAI tools list.
|
||||
|
||||
Operates on **OpenAI-format** tool dicts (pre-``_map_tools``). The
|
||||
|
|
@ -968,7 +972,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
the original name out of either ``{"function": {"name": ...}}``
|
||||
(legacy OpenAI shape) or ``{"name": ...}`` (rare top-level shape).
|
||||
"""
|
||||
original_names: list[str] = []
|
||||
original_names: List[str] = []
|
||||
for tool in tools or []:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
|
|
@ -981,8 +985,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _sanitize_tool_names_in_request(
|
||||
optional_params: dict[str, Any],
|
||||
) -> tuple[dict[str, str], dict[str, str]]:
|
||||
optional_params: Dict[str, Any],
|
||||
) -> Tuple[Dict[str, str], Dict[str, str]]:
|
||||
"""Sanitize ``optional_params['tools']`` and ``optional_params['tool_choice']``
|
||||
in place so every name matches Anthropic's ``^[a-zA-Z0-9_-]{1,128}$``.
|
||||
|
||||
|
|
@ -1006,7 +1010,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
# Order matters: the first occurrence wins the canonical slot;
|
||||
# later collisions get numeric suffixes (see
|
||||
# ``_build_anthropic_tool_name_maps``).
|
||||
original_names: list[str] = []
|
||||
original_names: List[str] = []
|
||||
for t in tools:
|
||||
if not isinstance(t, dict):
|
||||
continue
|
||||
|
|
@ -1028,7 +1032,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
# so a caller reusing the same tool list/dicts across requests
|
||||
# doesn't see its inputs permanently rewritten (which would also
|
||||
# drop the original key from `forward` on the next request).
|
||||
new_tools: list[Any] = []
|
||||
new_tools: List[Any] = []
|
||||
for t in tools:
|
||||
if (
|
||||
isinstance(t, dict)
|
||||
|
|
@ -1054,7 +1058,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
return forward, reverse
|
||||
|
||||
def _detect_tool_search_tools(self, tools: list | None) -> bool:
|
||||
def _detect_tool_search_tools(self, tools: Optional[List]) -> bool:
|
||||
"""Check if tool search tools are present in the tools list."""
|
||||
if not tools:
|
||||
return False
|
||||
|
|
@ -1068,7 +1072,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return True
|
||||
return False
|
||||
|
||||
def _separate_deferred_tools(self, tools: list) -> tuple[list, list]:
|
||||
def _separate_deferred_tools(self, tools: List) -> Tuple[List, List]:
|
||||
"""
|
||||
Separate tools into deferred and non-deferred lists.
|
||||
|
||||
|
|
@ -1088,9 +1092,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _expand_tool_references(
|
||||
self,
|
||||
content: list,
|
||||
deferred_tools: list,
|
||||
) -> list:
|
||||
content: List,
|
||||
deferred_tools: List,
|
||||
) -> List:
|
||||
"""
|
||||
Expand tool_reference blocks to full tool definitions.
|
||||
|
||||
|
|
@ -1131,8 +1135,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
return expanded_content
|
||||
|
||||
def _map_stop_sequences(self, stop: Union[str, list[str]] | None) -> list[str] | None:
|
||||
new_stop: list[str] | None = None
|
||||
def _map_stop_sequences(self, stop: Optional[Union[str, List[str]]]) -> Optional[List[str]]:
|
||||
new_stop: Optional[List[str]] = None
|
||||
if isinstance(stop, str):
|
||||
if (
|
||||
stop.isspace() and litellm.drop_params is True
|
||||
|
|
@ -1153,10 +1157,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _map_reasoning_effort(
|
||||
reasoning_effort: Union[REASONING_EFFORT, str] | None,
|
||||
reasoning_effort: Optional[Union[REASONING_EFFORT, str]],
|
||||
model: str,
|
||||
llm_provider: str = "anthropic",
|
||||
) -> AnthropicThinkingParam | None:
|
||||
) -> Optional[AnthropicThinkingParam]:
|
||||
if reasoning_effort is None or reasoning_effort == "none":
|
||||
return None
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model):
|
||||
|
|
@ -1207,10 +1211,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
def _extract_json_schema_from_response_format(self, value: dict | None) -> dict | None:
|
||||
def _extract_json_schema_from_response_format(self, value: Optional[dict]) -> Optional[dict]:
|
||||
if value is None:
|
||||
return None
|
||||
json_schema: dict | None = None
|
||||
json_schema: Optional[dict] = None
|
||||
if "response_schema" in value:
|
||||
json_schema = value["response_schema"]
|
||||
elif "json_schema" in value:
|
||||
|
|
@ -1218,8 +1222,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
return json_schema
|
||||
|
||||
def map_response_format_to_anthropic_output_format(self, value: dict | None) -> AnthropicOutputSchema | None:
|
||||
json_schema: dict | None = self._extract_json_schema_from_response_format(value)
|
||||
def map_response_format_to_anthropic_output_format(self, value: Optional[dict]) -> Optional[AnthropicOutputSchema]:
|
||||
json_schema: Optional[dict] = self._extract_json_schema_from_response_format(value)
|
||||
if json_schema is None:
|
||||
return None
|
||||
|
||||
|
|
@ -1245,13 +1249,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
)
|
||||
|
||||
def map_response_format_to_anthropic_tool(
|
||||
self, value: dict | None, optional_params: dict, is_thinking_enabled: bool
|
||||
) -> AnthropicMessagesTool | None:
|
||||
self, value: Optional[dict], optional_params: dict, is_thinking_enabled: bool
|
||||
) -> Optional[AnthropicMessagesTool]:
|
||||
ignore_response_format_types = ["text"]
|
||||
if value is None or value["type"] in ignore_response_format_types: # value is a no-op
|
||||
return None
|
||||
|
||||
json_schema: dict | None = self._extract_json_schema_from_response_format(value)
|
||||
json_schema: Optional[dict] = self._extract_json_schema_from_response_format(value)
|
||||
if json_schema is None:
|
||||
return None
|
||||
"""
|
||||
|
|
@ -1296,8 +1300,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def map_openai_context_management_to_anthropic(
|
||||
context_management: Union[list[dict[str, Any]], dict[str, Any]],
|
||||
) -> dict[str, Any] | None:
|
||||
context_management: Union[List[Dict[str, Any]], Dict[str, Any]],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
OpenAI format: [{"type": "compaction", "compact_threshold": 200000}]
|
||||
Anthropic format: {
|
||||
|
|
@ -1328,7 +1332,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
entry_type = entry.get("type")
|
||||
if entry_type == "compaction":
|
||||
anthropic_edit: dict[str, Any] = {"type": "compact_20260112"}
|
||||
anthropic_edit: Dict[str, Any] = {"type": "compact_20260112"}
|
||||
compact_threshold = entry.get("compact_threshold")
|
||||
# Rewrite to 'trigger' with correct nesting if threshold exists
|
||||
if compact_threshold is not None and isinstance(compact_threshold, (int, float)):
|
||||
|
|
@ -1384,7 +1388,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
if mcp_servers:
|
||||
optional_params["mcp_servers"] = mcp_servers
|
||||
elif param == "tool_choice" or param == "parallel_tool_calls":
|
||||
_tool_choice: AnthropicMessagesToolChoice | None = self._map_tool_choice(
|
||||
_tool_choice: Optional[AnthropicMessagesToolChoice] = self._map_tool_choice(
|
||||
tool_choice=non_default_params.get("tool_choice"),
|
||||
parallel_tool_use=non_default_params.get("parallel_tool_calls"),
|
||||
)
|
||||
|
|
@ -1515,7 +1519,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _create_json_tool_call_for_response_format(
|
||||
self,
|
||||
json_schema: dict | None = None,
|
||||
json_schema: Optional[dict] = None,
|
||||
) -> AnthropicMessagesTool:
|
||||
"""
|
||||
Handles creating a tool call for getting responses in JSON format.
|
||||
|
|
@ -1550,7 +1554,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
"""
|
||||
return False
|
||||
|
||||
def translate_system_message(self, messages: list[AllMessageValues]) -> list[AnthropicSystemMessageContent]:
|
||||
def translate_system_message(self, messages: List[AllMessageValues]) -> List[AnthropicSystemMessageContent]:
|
||||
"""
|
||||
Translate system message to anthropic format.
|
||||
|
||||
|
|
@ -1558,7 +1562,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
When should_strip_billing_metadata() is True, x-anthropic-billing-header system blocks are dropped.
|
||||
"""
|
||||
system_prompt_indices = []
|
||||
anthropic_system_message_list: list[AnthropicSystemMessageContent] = []
|
||||
anthropic_system_message_list: List[AnthropicSystemMessageContent] = []
|
||||
for idx, message in enumerate(messages):
|
||||
if message["role"] == "system":
|
||||
system_prompt_indices.append(idx)
|
||||
|
|
@ -1608,9 +1612,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def add_code_execution_tool(
|
||||
self,
|
||||
messages: list[AllAnthropicMessageValues],
|
||||
tools: list[Union[AllAnthropicToolsValues, dict]],
|
||||
) -> list[Union[AllAnthropicToolsValues, dict]]:
|
||||
messages: List[AllAnthropicMessageValues],
|
||||
tools: List[Union[AllAnthropicToolsValues, Dict]],
|
||||
) -> List[Union[AllAnthropicToolsValues, Dict]]:
|
||||
"""if 'container_upload' in messages, add code_execution tool"""
|
||||
add_code_execution_tool = False
|
||||
for message in messages:
|
||||
|
|
@ -1725,7 +1729,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: list[AllMessageValues],
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
|
|
@ -1827,7 +1831,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
## Add code_execution tool if container_upload is in messages
|
||||
_tools = (
|
||||
cast(
|
||||
list[Union[AllAnthropicToolsValues, dict]] | None,
|
||||
Optional[List[Union[AllAnthropicToolsValues, Dict]]],
|
||||
optional_params.get("tools"),
|
||||
)
|
||||
or []
|
||||
|
|
@ -1925,12 +1929,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _resolve_json_mode_non_streaming(
|
||||
self,
|
||||
json_mode: bool | None,
|
||||
tool_calls: list[ChatCompletionToolCallChunk],
|
||||
) -> tuple[
|
||||
LitellmMessage | None,
|
||||
list[ChatCompletionToolCallChunk],
|
||||
str | None,
|
||||
json_mode: Optional[bool],
|
||||
tool_calls: List[ChatCompletionToolCallChunk],
|
||||
) -> Tuple[
|
||||
Optional[LitellmMessage],
|
||||
List[ChatCompletionToolCallChunk],
|
||||
Optional[str],
|
||||
]:
|
||||
"""Strip internal response_format tool calls; merge payload into content when mixed with user tools."""
|
||||
if json_mode is not True or not tool_calls:
|
||||
|
|
@ -1951,30 +1955,30 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
first_json = tool_calls[json_indices[0]]
|
||||
json_msg = AnthropicConfig._convert_tool_response_to_message([first_json])
|
||||
extra_content: str | None = json_msg.content if json_msg is not None else None
|
||||
extra_content: Optional[str] = json_msg.content if json_msg is not None else None
|
||||
filtered_tools = [t for i, t in enumerate(tool_calls) if i not in json_indices]
|
||||
return None, filtered_tools, extra_content
|
||||
|
||||
def extract_response_content(
|
||||
self, completion_response: dict
|
||||
) -> tuple[
|
||||
) -> Tuple[
|
||||
str,
|
||||
list[Any] | None,
|
||||
list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None,
|
||||
str | None,
|
||||
list[ChatCompletionToolCallChunk],
|
||||
list[Any] | None,
|
||||
list[Any] | None,
|
||||
list[Any] | None,
|
||||
Optional[List[Any]],
|
||||
Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]],
|
||||
Optional[str],
|
||||
List[ChatCompletionToolCallChunk],
|
||||
Optional[List[Any]],
|
||||
Optional[List[Any]],
|
||||
Optional[List[Any]],
|
||||
]:
|
||||
text_content = ""
|
||||
citations: list[Any] | None = None
|
||||
thinking_blocks: list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None
|
||||
reasoning_content: str | None = None
|
||||
tool_calls: list[ChatCompletionToolCallChunk] = []
|
||||
web_search_results: list[Any] | None = None
|
||||
tool_results: list[Any] | None = None
|
||||
compaction_blocks: list[Any] | None = None
|
||||
citations: Optional[List[Any]] = None
|
||||
thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]] = None
|
||||
reasoning_content: Optional[str] = None
|
||||
tool_calls: List[ChatCompletionToolCallChunk] = []
|
||||
web_search_results: Optional[List[Any]] = None
|
||||
tool_results: Optional[List[Any]] = None
|
||||
compaction_blocks: Optional[List[Any]] = None
|
||||
for idx, content in enumerate(completion_response["content"]):
|
||||
if content["type"] == "text":
|
||||
text_content += content["text"]
|
||||
|
|
@ -2037,7 +2041,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
if thinking_blocks is not None:
|
||||
reasoning_content = ""
|
||||
for block in thinking_blocks:
|
||||
thinking_content = cast(str | None, block.get("thinking"))
|
||||
thinking_content = cast(Optional[str], block.get("thinking"))
|
||||
if thinking_content is not None:
|
||||
reasoning_content += thinking_content
|
||||
|
||||
|
|
@ -2055,9 +2059,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def calculate_usage(
|
||||
self,
|
||||
usage_object: dict,
|
||||
reasoning_content: str | None,
|
||||
completion_response: dict | None = None,
|
||||
speed: str | None = None,
|
||||
reasoning_content: Optional[str],
|
||||
completion_response: Optional[dict] = None,
|
||||
speed: Optional[str] = None,
|
||||
) -> Usage:
|
||||
# NOTE: Sometimes the usage object has None set explicitly for token counts, meaning .get() & key access returns None, and we need to account for this
|
||||
raw_prompt_tokens = usage_object.get("input_tokens", 0) or 0
|
||||
|
|
@ -2067,10 +2071,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
_usage = usage_object
|
||||
cache_creation_input_tokens: int = 0
|
||||
cache_read_input_tokens: int = 0
|
||||
cache_creation_token_details: CacheCreationTokenDetails | None = None
|
||||
web_search_requests: int | None = None
|
||||
tool_search_requests: int | None = None
|
||||
inference_geo: str | None = None
|
||||
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
|
||||
web_search_requests: Optional[int] = None
|
||||
tool_search_requests: Optional[int] = None
|
||||
inference_geo: Optional[str] = None
|
||||
if "inference_geo" in _usage and _usage["inference_geo"] is not None:
|
||||
inference_geo = _usage["inference_geo"]
|
||||
service_tier = cast(
|
||||
|
|
@ -2078,7 +2082,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
_usage.get("service_tier"),
|
||||
)
|
||||
|
||||
iterations: list[Any] | None = _usage.get("iterations")
|
||||
iterations: Optional[List[Any]] = _usage.get("iterations")
|
||||
if iterations:
|
||||
prompt_tokens = sum(it.get("input_tokens", 0) or 0 for it in iterations)
|
||||
completion_tokens = sum(it.get("output_tokens", 0) or 0 for it in iterations)
|
||||
|
|
@ -2164,8 +2168,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
)
|
||||
return usage
|
||||
|
||||
def _build_code_by_id_map(self, tool_calls: list[ChatCompletionToolCallChunk]) -> dict[str, str]:
|
||||
code_by_id: dict[str, str] = {}
|
||||
def _build_code_by_id_map(self, tool_calls: List[ChatCompletionToolCallChunk]) -> Dict[str, str]:
|
||||
code_by_id: Dict[str, str] = {}
|
||||
for tc in tool_calls:
|
||||
try:
|
||||
args = json.loads(tc.get("function", {}).get("arguments", "{}"))
|
||||
|
|
@ -2179,10 +2183,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _build_code_interpreter_results(
|
||||
self,
|
||||
tool_results: list[Any],
|
||||
code_by_id: dict[str, str],
|
||||
container_id: str | None,
|
||||
) -> list[OutputCodeInterpreterCall]:
|
||||
tool_results: List[Any],
|
||||
code_by_id: Dict[str, str],
|
||||
container_id: Optional[str],
|
||||
) -> List[OutputCodeInterpreterCall]:
|
||||
code_interpreter_results = []
|
||||
for tr in tool_results:
|
||||
if tr.get("type") != "bash_code_execution_tool_result":
|
||||
|
|
@ -2205,14 +2209,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def _build_provider_specific_fields(
|
||||
self,
|
||||
completion_response: dict,
|
||||
citations: list[Any] | None,
|
||||
thinking_blocks: list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None,
|
||||
web_search_results: list[Any] | None,
|
||||
tool_results: list[Any] | None,
|
||||
compaction_blocks: list[Any] | None,
|
||||
tool_calls: list[ChatCompletionToolCallChunk],
|
||||
) -> dict[str, Any]:
|
||||
provider_specific_fields: dict[str, Any] = {
|
||||
citations: Optional[List[Any]],
|
||||
thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]],
|
||||
web_search_results: Optional[List[Any]],
|
||||
tool_results: Optional[List[Any]],
|
||||
compaction_blocks: Optional[List[Any]],
|
||||
tool_calls: List[ChatCompletionToolCallChunk],
|
||||
) -> Dict[str, Any]:
|
||||
provider_specific_fields: Dict[str, Any] = {
|
||||
"citations": citations,
|
||||
"thinking_blocks": thinking_blocks,
|
||||
}
|
||||
|
|
@ -2249,12 +2253,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
completion_response: dict,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
json_mode: bool | None = None,
|
||||
prefix_prompt: str | None = None,
|
||||
speed: str | None = None,
|
||||
tool_name_reverse_map: dict[str, str] | None = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
prefix_prompt: Optional[str] = None,
|
||||
speed: Optional[str] = None,
|
||||
tool_name_reverse_map: Optional[Dict[str, str]] = None,
|
||||
):
|
||||
_hidden_params: dict = {}
|
||||
_hidden_params: Dict = {}
|
||||
_hidden_params["additional_headers"] = process_anthropic_headers(dict(raw_response.headers))
|
||||
if "error" in completion_response:
|
||||
response_headers = getattr(raw_response, "headers", None)
|
||||
|
|
@ -2350,7 +2354,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
model_response._hidden_params = _hidden_params
|
||||
return model_response
|
||||
|
||||
def get_prefix_prompt(self, messages: list[AllMessageValues]) -> str | None:
|
||||
def get_prefix_prompt(self, messages: List[AllMessageValues]) -> Optional[str]:
|
||||
"""
|
||||
Get the prefix prompt from the messages.
|
||||
|
||||
|
|
@ -2375,13 +2379,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: LoggingClass,
|
||||
request_data: dict,
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
request_data: Dict,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: Dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ModelResponse:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -2404,7 +2408,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
prefix_prompt = self.get_prefix_prompt(messages=messages)
|
||||
speed = optional_params.get("speed")
|
||||
tool_name_reverse_map: dict[str, str] | None = None
|
||||
tool_name_reverse_map: Optional[Dict[str, str]] = None
|
||||
if isinstance(litellm_params, dict):
|
||||
_candidate = litellm_params.get(ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY)
|
||||
if isinstance(_candidate, dict):
|
||||
|
|
@ -2423,14 +2427,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _convert_tool_response_to_message(
|
||||
tool_calls: list[ChatCompletionToolCallChunk],
|
||||
) -> LitellmMessage | None:
|
||||
tool_calls: List[ChatCompletionToolCallChunk],
|
||||
) -> Optional[LitellmMessage]:
|
||||
"""
|
||||
In JSON mode, Anthropic API returns JSON schema as a tool call, we need to convert it to a message to follow the OpenAI format
|
||||
|
||||
"""
|
||||
## HANDLE JSON MODE - anthropic returns single function call
|
||||
json_mode_content_str: str | None = tool_calls[0]["function"].get("arguments")
|
||||
json_mode_content_str: Optional[str] = tool_calls[0]["function"].get("arguments")
|
||||
try:
|
||||
if json_mode_content_str is not None:
|
||||
args = json.loads(json_mode_content_str)
|
||||
|
|
@ -2448,7 +2452,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return None
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
return AnthropicError(
|
||||
status_code=status_code,
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
Translates from OpenAI's `/v1/chat/completions` to ModelScope's `/v1/chat/completions`
|
||||
"""
|
||||
|
||||
from typing import Any, Coroutine, Literal, Union, cast, overload
|
||||
from typing import Any, Coroutine, Literal, Optional, Tuple, Union, cast, overload
|
||||
|
||||
from typing_extensions import override
|
||||
|
||||
|
|
@ -59,8 +59,8 @@ class ModelScopeChatConfig(OpenAIGPTConfig):
|
|||
return super()._transform_messages(messages=messages, model=model, is_async=False)
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self, api_base: str | None, api_key: str | None
|
||||
) -> tuple[str | None, str | None]:
|
||||
self, api_base: Optional[str], api_key: Optional[str]
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
api_base = api_base or get_secret_str("MODELSCOPE_API_BASE") or self.DEFAULT_BASE_URL # type: ignore
|
||||
dynamic_api_key = api_key or get_secret_str("MODELSCOPE_API_KEY")
|
||||
return api_base, dynamic_api_key
|
||||
|
|
@ -68,12 +68,12 @@ class ModelScopeChatConfig(OpenAIGPTConfig):
|
|||
@override
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: bool | None = None,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
If api_base is not provided, use the default ModelScope /chat/completions endpoint.
|
||||
|
|
|
|||
|
|
@ -9,9 +9,13 @@ from typing import (
|
|||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Tuple,
|
||||
Type,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
|
@ -131,7 +135,7 @@ class VertexAIBaseConfig:
|
|||
optional_params[mapped_params[param]] = value
|
||||
return optional_params
|
||||
|
||||
def get_eu_regions(self) -> list[str]:
|
||||
def get_eu_regions(self) -> List[str]:
|
||||
"""
|
||||
Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions
|
||||
"""
|
||||
|
|
@ -148,7 +152,7 @@ class VertexAIBaseConfig:
|
|||
"europe-west9",
|
||||
]
|
||||
|
||||
def get_us_regions(self) -> list[str]:
|
||||
def get_us_regions(self) -> List[str]:
|
||||
"""
|
||||
Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions
|
||||
"""
|
||||
|
|
@ -194,29 +198,29 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
Note: Please make sure to modify the default parameters as required for your use case.
|
||||
"""
|
||||
|
||||
temperature: float | None = None
|
||||
max_output_tokens: int | None = None
|
||||
top_p: float | None = None
|
||||
top_k: int | None = None
|
||||
response_mime_type: str | None = None
|
||||
candidate_count: int | None = None
|
||||
stop_sequences: list | None = None
|
||||
frequency_penalty: float | None = None
|
||||
presence_penalty: float | None = None
|
||||
seed: int | None = None
|
||||
temperature: Optional[float] = None
|
||||
max_output_tokens: Optional[int] = None
|
||||
top_p: Optional[float] = None
|
||||
top_k: Optional[int] = None
|
||||
response_mime_type: Optional[str] = None
|
||||
candidate_count: Optional[int] = None
|
||||
stop_sequences: Optional[list] = None
|
||||
frequency_penalty: Optional[float] = None
|
||||
presence_penalty: Optional[float] = None
|
||||
seed: Optional[int] = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
temperature: float | None = None,
|
||||
max_output_tokens: int | None = None,
|
||||
top_p: float | None = None,
|
||||
top_k: int | None = None,
|
||||
response_mime_type: str | None = None,
|
||||
candidate_count: int | None = None,
|
||||
stop_sequences: list | None = None,
|
||||
frequency_penalty: float | None = None,
|
||||
presence_penalty: float | None = None,
|
||||
seed: int | None = None,
|
||||
temperature: Optional[float] = None,
|
||||
max_output_tokens: Optional[int] = None,
|
||||
top_p: Optional[float] = None,
|
||||
top_k: Optional[int] = None,
|
||||
response_mime_type: Optional[str] = None,
|
||||
candidate_count: Optional[int] = None,
|
||||
stop_sequences: Optional[list] = None,
|
||||
frequency_penalty: Optional[float] = None,
|
||||
presence_penalty: Optional[float] = None,
|
||||
seed: Optional[int] = None,
|
||||
) -> None:
|
||||
locals_ = locals().copy()
|
||||
for key, value in locals_.items():
|
||||
|
|
@ -228,8 +232,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return super().get_config()
|
||||
|
||||
def get_json_schema_from_pydantic_object(
|
||||
self, response_format: Union[type["BaseModel"], dict] | None
|
||||
) -> dict | None:
|
||||
self, response_format: Optional[Union[Type["BaseModel"], dict]]
|
||||
) -> Optional[dict]:
|
||||
"""
|
||||
Override to use Pydantic's model_json_schema() instead of OpenAI's
|
||||
to_strict_json_schema().
|
||||
|
|
@ -285,7 +289,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return False
|
||||
|
||||
@staticmethod
|
||||
def _forward_gemini_function_call_id(model: str, custom_llm_provider: str | None = None) -> bool:
|
||||
def _forward_gemini_function_call_id(model: str, custom_llm_provider: Optional[str] = None) -> bool:
|
||||
"""
|
||||
Whether to include `id` on function_call / function_response parts.
|
||||
|
||||
|
|
@ -305,7 +309,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return False
|
||||
return True
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list[str]:
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
supported_params = [
|
||||
"temperature",
|
||||
"top_p",
|
||||
|
|
@ -340,7 +344,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
supported_params.append("thinking")
|
||||
return supported_params
|
||||
|
||||
def map_tool_choice_values(self, model: str, tool_choice: Union[str, dict]) -> ToolConfig | None:
|
||||
def map_tool_choice_values(self, model: str, tool_choice: Union[str, dict]) -> Optional[ToolConfig]:
|
||||
if tool_choice == "none":
|
||||
return ToolConfig(functionCallingConfig=FunctionCallingConfig(mode="NONE"))
|
||||
elif tool_choice == "required":
|
||||
|
|
@ -466,7 +470,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
return transformed_config
|
||||
|
||||
def _extract_google_maps_retrieval_config(self, google_maps_config: dict) -> tuple[dict, dict | None]:
|
||||
def _extract_google_maps_retrieval_config(self, google_maps_config: dict) -> Tuple[dict, Optional[dict]]:
|
||||
"""
|
||||
Extract location configuration from googleMaps tool for Vertex AI toolConfig.
|
||||
|
||||
|
|
@ -504,7 +508,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
return cleaned_config, retrieval_config
|
||||
|
||||
def get_tool_value(self, tool: dict, tool_name: str) -> dict | None:
|
||||
def get_tool_value(self, tool: dict, tool_name: str) -> Optional[dict]:
|
||||
"""
|
||||
Helper function to get tool value handling both camelCase and underscore_case variants
|
||||
|
||||
|
|
@ -529,10 +533,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _resolve_search_tool_conflict(
|
||||
gtool_func_declarations: list,
|
||||
googleSearch: dict | None,
|
||||
googleSearchRetrieval: dict | None,
|
||||
enterpriseWebSearch: dict | None,
|
||||
urlContext: dict | None,
|
||||
googleSearch: Optional[dict],
|
||||
googleSearchRetrieval: Optional[dict],
|
||||
enterpriseWebSearch: Optional[dict],
|
||||
urlContext: Optional[dict],
|
||||
optional_params: dict,
|
||||
) -> tuple:
|
||||
"""
|
||||
|
|
@ -576,7 +580,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
return googleSearch, googleSearchRetrieval, enterpriseWebSearch, urlContext
|
||||
|
||||
def _map_function(self, value: list[dict], optional_params: dict) -> list[Tools]:
|
||||
def _map_function(self, value: List[dict], optional_params: dict) -> List[Tools]:
|
||||
"""
|
||||
Map OpenAI-style tools/functions to Vertex AI format.
|
||||
|
||||
|
|
@ -592,21 +596,21 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
googleMaps tools contain location data
|
||||
"""
|
||||
gtool_func_declarations = []
|
||||
googleSearch: dict | None = None
|
||||
googleSearchRetrieval: dict | None = None
|
||||
enterpriseWebSearch: dict | None = None
|
||||
urlContext: dict | None = None
|
||||
code_execution: dict | None = None
|
||||
googleMaps: dict | None = None
|
||||
google_maps_retrieval_config: dict | None = None
|
||||
computerUse: dict | None = None
|
||||
googleSearch: Optional[dict] = None
|
||||
googleSearchRetrieval: Optional[dict] = None
|
||||
enterpriseWebSearch: Optional[dict] = None
|
||||
urlContext: Optional[dict] = None
|
||||
code_execution: Optional[dict] = None
|
||||
googleMaps: Optional[dict] = None
|
||||
google_maps_retrieval_config: Optional[dict] = None
|
||||
computerUse: Optional[dict] = None
|
||||
# remove 'additionalProperties' from tools
|
||||
value = _remove_additional_properties(value)
|
||||
# remove 'strict' from tools
|
||||
value = _remove_strict_from_schema(value)
|
||||
|
||||
for tool in value:
|
||||
openai_function_object: ChatCompletionToolParamFunctionChunk | None = None
|
||||
openai_function_object: Optional[ChatCompletionToolParamFunctionChunk] = None
|
||||
if "function" in tool: # tools list
|
||||
_openai_function_object = ChatCompletionToolParamFunctionChunk( # type: ignore
|
||||
**tool["function"]
|
||||
|
|
@ -696,7 +700,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
# Build list of Tool objects - each Tool should contain exactly one type
|
||||
# per Vertex AI API spec: "A Tool object should contain exactly one type of Tool"
|
||||
_tools_list: list[Tools] = []
|
||||
_tools_list: List[Tools] = []
|
||||
|
||||
(
|
||||
googleSearch,
|
||||
|
|
@ -814,7 +818,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _map_reasoning_effort_to_thinking_budget(
|
||||
reasoning_effort: str,
|
||||
model: str | None = None,
|
||||
model: Optional[str] = None,
|
||||
) -> GeminiThinkingConfig:
|
||||
if reasoning_effort == "minimal":
|
||||
# Use model-specific minimum thinking budget or fallback
|
||||
|
|
@ -863,7 +867,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _map_reasoning_effort_to_thinking_level(
|
||||
reasoning_effort: str,
|
||||
model: str | None = None,
|
||||
model: Optional[str] = None,
|
||||
) -> GeminiThinkingConfig:
|
||||
"""
|
||||
Map reasoning_effort to thinking_level for Gemini 3+ models.
|
||||
|
|
@ -909,12 +913,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
raise ValueError(f"Invalid reasoning effort: {reasoning_effort}")
|
||||
|
||||
@staticmethod
|
||||
def _is_thinking_budget_zero(thinking_budget: int | None) -> bool:
|
||||
def _is_thinking_budget_zero(thinking_budget: Optional[int]) -> bool:
|
||||
return thinking_budget is not None and thinking_budget == 0
|
||||
|
||||
@staticmethod
|
||||
def _validate_thinking_config_conflicts(
|
||||
optional_params: dict,
|
||||
optional_params: Dict,
|
||||
param_name: str,
|
||||
param_description: str = "thinking_budget",
|
||||
) -> None:
|
||||
|
|
@ -935,7 +939,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _validate_thinking_level_conflicts(
|
||||
optional_params: dict,
|
||||
optional_params: Dict,
|
||||
) -> None:
|
||||
"""
|
||||
Validate that thinking_level and thinking_budget are not both specified.
|
||||
|
|
@ -955,7 +959,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _map_thinking_param(
|
||||
thinking_param: AnthropicThinkingParam,
|
||||
model: str | None = None,
|
||||
model: Optional[str] = None,
|
||||
) -> GeminiThinkingConfig:
|
||||
thinking_enabled = thinking_param.get("type") == "enabled"
|
||||
thinking_budget = thinking_param.get("budget_tokens")
|
||||
|
|
@ -1060,8 +1064,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _apply_include_server_side_tool_invocations(
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
non_default_params: Dict,
|
||||
optional_params: Dict,
|
||||
) -> None:
|
||||
"""
|
||||
Set include_server_side_tool_invocations before tools are mapped.
|
||||
|
|
@ -1080,11 +1084,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
non_default_params: Dict,
|
||||
optional_params: Dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
) -> Dict:
|
||||
self._apply_include_server_side_tool_invocations(non_default_params, optional_params)
|
||||
gemini_sampling_params_warned: bool = False
|
||||
for param, value in non_default_params.items():
|
||||
|
|
@ -1176,7 +1180,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
elif param == "reasoning_effort":
|
||||
# Extract effort value - handle both string and dict formats
|
||||
# Dict format comes from OpenAI Agents SDK: {"effort": "high", "summary": "auto"}
|
||||
effort_value: str | None = None
|
||||
effort_value: Optional[str] = None
|
||||
if isinstance(value, str):
|
||||
effort_value = value
|
||||
elif isinstance(value, dict):
|
||||
|
|
@ -1252,7 +1256,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
optional_params[mapped_params[param]] = value
|
||||
return optional_params
|
||||
|
||||
def get_eu_regions(self) -> list[str]:
|
||||
def get_eu_regions(self) -> List[str]:
|
||||
"""
|
||||
Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions
|
||||
"""
|
||||
|
|
@ -1289,7 +1293,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return model
|
||||
|
||||
@staticmethod
|
||||
def _is_model_gemini_spec_model(model: str | None) -> bool:
|
||||
def _is_model_gemini_spec_model(model: Optional[str]) -> bool:
|
||||
"""
|
||||
Returns true if user is trying to call custom model in `/gemini` request/response format
|
||||
"""
|
||||
|
|
@ -1312,7 +1316,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return model.split("/")[-1]
|
||||
return model
|
||||
|
||||
def get_flagged_finish_reasons(self) -> dict[str, str]:
|
||||
def get_flagged_finish_reasons(self) -> Dict[str, str]:
|
||||
"""
|
||||
Return Dictionary of finish reasons which indicate response was flagged
|
||||
|
||||
|
|
@ -1349,7 +1353,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def get_finish_reason_mapping() -> dict[str, OpenAIChatCompletionFinishReason]:
|
||||
def get_finish_reason_mapping() -> Dict[str, OpenAIChatCompletionFinishReason]:
|
||||
"""
|
||||
Return Dictionary of Gemini/Vertex AI finish reasons and their
|
||||
OpenAI-compatible mappings.
|
||||
|
|
@ -1368,9 +1372,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
return exception_string
|
||||
|
||||
def get_assistant_content_message(self, parts: list[HttpxPartType]) -> tuple[str | None, str | None]:
|
||||
content_str: str | None = None
|
||||
reasoning_content_str: str | None = None
|
||||
def get_assistant_content_message(self, parts: List[HttpxPartType]) -> Tuple[Optional[str], Optional[str]]:
|
||||
content_str: Optional[str] = None
|
||||
reasoning_content_str: Optional[str] = None
|
||||
|
||||
for part in parts:
|
||||
_content_str = ""
|
||||
|
|
@ -1409,7 +1413,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
return content_str, reasoning_content_str
|
||||
|
||||
def _extract_thinking_blocks_from_parts(self, parts: list[HttpxPartType]) -> list[ChatCompletionThinkingBlock]:
|
||||
def _extract_thinking_blocks_from_parts(self, parts: List[HttpxPartType]) -> List[ChatCompletionThinkingBlock]:
|
||||
"""Extract thinking blocks from parts if present.
|
||||
|
||||
Per Google's docs (https://ai.google.dev/gemini-api/docs/thinking):
|
||||
|
|
@ -1418,7 +1422,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
it does NOT indicate that the content is thinking (a part can have
|
||||
thoughtSignature without thought: true, e.g., function calls)
|
||||
"""
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock] = []
|
||||
thinking_blocks: List[ChatCompletionThinkingBlock] = []
|
||||
for part in parts:
|
||||
if part.get("thought") is True:
|
||||
thinking_text = part.get("text", "")
|
||||
|
|
@ -1432,7 +1436,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
thinking_blocks.append(block)
|
||||
return thinking_blocks
|
||||
|
||||
def _extract_thought_signatures_from_parts(self, parts: list[HttpxPartType]) -> list[str] | None:
|
||||
def _extract_thought_signatures_from_parts(self, parts: List[HttpxPartType]) -> Optional[List[str]]:
|
||||
"""Extract thoughtSignature values from parts.
|
||||
|
||||
Per Google's docs, thoughtSignature is returned for multi-turn context preservation
|
||||
|
|
@ -1442,7 +1446,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
Returns:
|
||||
List of thoughtSignature strings if any are found, None otherwise
|
||||
"""
|
||||
signatures: list[str] = []
|
||||
signatures: List[str] = []
|
||||
for part in parts:
|
||||
signature = part.get("thoughtSignature")
|
||||
if signature is not None:
|
||||
|
|
@ -1451,8 +1455,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _extract_server_side_tool_invocations(
|
||||
parts: list[HttpxPartType],
|
||||
) -> list[dict[str, Any]] | None:
|
||||
parts: List[HttpxPartType],
|
||||
) -> Optional[List[Dict[str, Any]]]:
|
||||
"""Extract server-side tool invocations (toolCall/toolResponse) from parts.
|
||||
|
||||
These are returned by Gemini when context circulation is enabled
|
||||
|
|
@ -1463,15 +1467,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
Returns:
|
||||
List of server-side invocation dicts if any found, None otherwise.
|
||||
"""
|
||||
invocations: list[dict[str, Any]] = []
|
||||
invocations: List[Dict[str, Any]] = []
|
||||
# Index toolCalls by id so we can pair them with responses
|
||||
tool_calls_by_id: dict[str, dict[str, Any]] = {}
|
||||
tool_responses_by_id: dict[str, dict[str, Any]] = {}
|
||||
tool_calls_by_id: Dict[str, Dict[str, Any]] = {}
|
||||
tool_responses_by_id: Dict[str, Dict[str, Any]] = {}
|
||||
|
||||
for part in parts:
|
||||
if "toolCall" in part:
|
||||
tc = part["toolCall"]
|
||||
entry: dict[str, Any] = {
|
||||
entry: Dict[str, Any] = {
|
||||
"tool_type": tc.get("toolType"),
|
||||
"id": tc.get("id"),
|
||||
"args": tc.get("args"),
|
||||
|
|
@ -1511,9 +1515,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
return invocations if invocations else None
|
||||
|
||||
def _extract_image_response_from_parts(self, parts: list[HttpxPartType]) -> list[ImageURLListItem] | None:
|
||||
def _extract_image_response_from_parts(self, parts: List[HttpxPartType]) -> Optional[List[ImageURLListItem]]:
|
||||
"""Extract image response from parts if present"""
|
||||
images: list[ImageURLListItem] = []
|
||||
images: List[ImageURLListItem] = []
|
||||
for part in parts:
|
||||
if "inlineData" in part:
|
||||
inline_data = part.get("inlineData", {})
|
||||
|
|
@ -1531,7 +1535,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
return images
|
||||
|
||||
def _extract_audio_response_from_parts(self, parts: list[HttpxPartType]) -> ChatCompletionAudioResponse | None:
|
||||
def _extract_audio_response_from_parts(self, parts: List[HttpxPartType]) -> Optional[ChatCompletionAudioResponse]:
|
||||
"""Extract audio response from parts if present"""
|
||||
for part in parts:
|
||||
if "text" in part:
|
||||
|
|
@ -1569,16 +1573,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _transform_parts(
|
||||
parts: list[HttpxPartType],
|
||||
parts: List[HttpxPartType],
|
||||
cumulative_tool_call_idx: int,
|
||||
is_function_call: bool | None,
|
||||
) -> tuple[
|
||||
ChatCompletionToolCallFunctionChunk | None,
|
||||
list[ChatCompletionToolCallChunk] | None,
|
||||
is_function_call: Optional[bool],
|
||||
) -> Tuple[
|
||||
Optional[ChatCompletionToolCallFunctionChunk],
|
||||
Optional[List[ChatCompletionToolCallChunk]],
|
||||
int,
|
||||
]:
|
||||
function: ChatCompletionToolCallFunctionChunk | None = None
|
||||
_tools: list[ChatCompletionToolCallChunk] = []
|
||||
function: Optional[ChatCompletionToolCallFunctionChunk] = None
|
||||
_tools: List[ChatCompletionToolCallChunk] = []
|
||||
for part in parts:
|
||||
if "functionCall" in part:
|
||||
_function_chunk: ChatCompletionToolCallFunctionChunk = {
|
||||
|
|
@ -1593,7 +1597,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
gemini_call_id = part["functionCall"].get("id")
|
||||
|
||||
if is_function_call is True:
|
||||
function_dict: dict[str, Any] = dict(_function_chunk)
|
||||
function_dict: Dict[str, Any] = dict(_function_chunk)
|
||||
if thought_signature:
|
||||
if "provider_specific_fields" not in function_dict:
|
||||
function_dict["provider_specific_fields"] = {}
|
||||
|
|
@ -1622,22 +1626,22 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
_tools.append(_tool_response_chunk)
|
||||
cumulative_tool_call_idx += 1
|
||||
if len(_tools) == 0:
|
||||
tools: list[ChatCompletionToolCallChunk] | None = None
|
||||
tools: Optional[List[ChatCompletionToolCallChunk]] = None
|
||||
else:
|
||||
tools = _tools
|
||||
return function, tools, cumulative_tool_call_idx
|
||||
|
||||
@staticmethod
|
||||
def _transform_logprobs(
|
||||
logprobs_result: LogprobsResult | None,
|
||||
) -> ChoiceLogprobs | None:
|
||||
logprobs_result: Optional[LogprobsResult],
|
||||
) -> Optional[ChoiceLogprobs]:
|
||||
if logprobs_result is None:
|
||||
return None
|
||||
if "chosenCandidates" not in logprobs_result:
|
||||
return None
|
||||
logprobs_list: list[ChatCompletionTokenLogprob] = []
|
||||
logprobs_list: List[ChatCompletionTokenLogprob] = []
|
||||
for index, candidate in enumerate(logprobs_result["chosenCandidates"]):
|
||||
top_logprobs: list[TopLogprob] = []
|
||||
top_logprobs: List[TopLogprob] = []
|
||||
if "topCandidates" in logprobs_result and index < len(logprobs_result["topCandidates"]):
|
||||
top_candidates_for_index = logprobs_result["topCandidates"][index]["candidates"]
|
||||
|
||||
|
|
@ -1744,16 +1748,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
) -> Usage:
|
||||
if completion_response is not None and "usageMetadata" not in completion_response:
|
||||
raise ValueError(f"usageMetadata not found in completion_response. Got={completion_response}")
|
||||
cached_tokens: int | None = None
|
||||
cached_tokens: Optional[int] = None
|
||||
# Separate variables for prompt tokens by modality
|
||||
prompt_audio_tokens: int | None = None
|
||||
prompt_image_tokens: int | None = None
|
||||
prompt_text_tokens: int | None = None
|
||||
prompt_video_tokens: int | None = None
|
||||
prompt_tokens_details: PromptTokensDetailsWrapper | None = None
|
||||
reasoning_tokens: int | None = None
|
||||
response_tokens: int | None = None
|
||||
response_tokens_details: CompletionTokensDetailsWrapper | None = None
|
||||
prompt_audio_tokens: Optional[int] = None
|
||||
prompt_image_tokens: Optional[int] = None
|
||||
prompt_text_tokens: Optional[int] = None
|
||||
prompt_video_tokens: Optional[int] = None
|
||||
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None
|
||||
reasoning_tokens: Optional[int] = None
|
||||
response_tokens: Optional[int] = None
|
||||
response_tokens_details: Optional[CompletionTokensDetailsWrapper] = None
|
||||
usage_metadata = completion_response["usageMetadata"]
|
||||
|
||||
def _get_token_count(detail: Mapping[str, Any]) -> int:
|
||||
|
|
@ -1832,10 +1836,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
## Parse cacheTokensDetails (breakdown of cached tokens by modality)
|
||||
## When explicit caching is used, Gemini provides this field to show which modalities were cached
|
||||
cached_text_tokens: int | None = None
|
||||
cached_audio_tokens: int | None = None
|
||||
cached_image_tokens: int | None = None
|
||||
cached_video_tokens: int | None = None
|
||||
cached_text_tokens: Optional[int] = None
|
||||
cached_audio_tokens: Optional[int] = None
|
||||
cached_image_tokens: Optional[int] = None
|
||||
cached_video_tokens: Optional[int] = None
|
||||
|
||||
if "cacheTokensDetails" in usage_metadata:
|
||||
for detail in usage_metadata["cacheTokensDetails"]:
|
||||
|
|
@ -1908,8 +1912,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _check_finish_reason(
|
||||
chat_completion_message: ChatCompletionResponseMessage | None,
|
||||
finish_reason: str | None,
|
||||
chat_completion_message: Optional[ChatCompletionResponseMessage],
|
||||
finish_reason: Optional[str],
|
||||
) -> OpenAIChatCompletionFinishReason:
|
||||
from litellm.litellm_core_utils.core_helpers import map_finish_reason
|
||||
|
||||
|
|
@ -1925,7 +1929,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _check_prompt_level_content_filter(
|
||||
processed_chunk: GenerateContentResponseBody,
|
||||
response_id: str | None,
|
||||
response_id: Optional[str],
|
||||
) -> Optional["ModelResponseStream"]:
|
||||
"""
|
||||
Check if prompt is blocked due to content filtering at the prompt level.
|
||||
|
|
@ -1969,8 +1973,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
def _calculate_web_search_requests(grounding_metadata: list[dict]) -> int | None:
|
||||
web_search_requests: int | None = None
|
||||
def _calculate_web_search_requests(grounding_metadata: List[dict]) -> Optional[int]:
|
||||
web_search_requests: Optional[int] = None
|
||||
|
||||
if grounding_metadata and isinstance(grounding_metadata, list) and len(grounding_metadata) > 0:
|
||||
for grounding_metadata_item in grounding_metadata:
|
||||
|
|
@ -1986,10 +1990,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
chat_completion_message: ChatCompletionResponseMessage,
|
||||
candidate: Candidates,
|
||||
idx: int,
|
||||
tools: list[ChatCompletionToolCallChunk] | None,
|
||||
functions: ChatCompletionToolCallFunctionChunk | None,
|
||||
chat_completion_logprobs: ChoiceLogprobs | None,
|
||||
image_response: list[ImageURLListItem] | None,
|
||||
tools: Optional[List[ChatCompletionToolCallChunk]],
|
||||
functions: Optional[ChatCompletionToolCallFunctionChunk],
|
||||
chat_completion_logprobs: Optional[ChoiceLogprobs],
|
||||
image_response: Optional[List[ImageURLListItem]],
|
||||
) -> StreamingChoices:
|
||||
"""
|
||||
Helper method to create a streaming choice object for Vertex AI
|
||||
|
|
@ -2021,7 +2025,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _extract_candidate_metadata(
|
||||
candidate: Candidates,
|
||||
) -> tuple[list[dict], list[dict], list, list]:
|
||||
) -> Tuple[List[dict], List[dict], List, List]:
|
||||
"""
|
||||
Extract metadata from a single candidate response.
|
||||
|
||||
|
|
@ -2031,10 +2035,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
safety_ratings: List
|
||||
citation_metadata: List
|
||||
"""
|
||||
grounding_metadata: list[dict] = []
|
||||
url_context_metadata: list[dict] = []
|
||||
safety_ratings: list = []
|
||||
citation_metadata: list = []
|
||||
grounding_metadata: List[dict] = []
|
||||
url_context_metadata: List[dict] = []
|
||||
safety_ratings: List = []
|
||||
citation_metadata: List = []
|
||||
|
||||
if "groundingMetadata" in candidate:
|
||||
if isinstance(candidate["groundingMetadata"], list):
|
||||
|
|
@ -2079,10 +2083,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _set_stream_metadata_on_response(
|
||||
model_response: Any,
|
||||
grounding_metadata: list[dict],
|
||||
url_context_metadata: list[dict],
|
||||
safety_ratings: list[dict],
|
||||
citation_metadata: list[dict],
|
||||
grounding_metadata: List[dict],
|
||||
url_context_metadata: List[dict],
|
||||
safety_ratings: List[dict],
|
||||
citation_metadata: List[dict],
|
||||
) -> None:
|
||||
setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore
|
||||
if grounding_metadata:
|
||||
|
|
@ -2102,10 +2106,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
def apply_assembled_streaming_response_metadata(
|
||||
self,
|
||||
response: ModelResponse,
|
||||
chunks: list[Any],
|
||||
chunks: List[Any],
|
||||
) -> None:
|
||||
for field_name in VERTEX_AI_PROVIDER_METADATA_FIELDS:
|
||||
merged: list[Any] = []
|
||||
merged: List[Any] = []
|
||||
for chunk in chunks:
|
||||
value = VertexGeminiConfig._get_stream_chunk_attr(chunk, field_name)
|
||||
if not value:
|
||||
|
|
@ -2120,14 +2124,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _convert_grounding_metadata_to_annotations(
|
||||
grounding_metadata: list[dict],
|
||||
content_text: str | None,
|
||||
) -> list[ChatCompletionAnnotation]:
|
||||
grounding_metadata: List[dict],
|
||||
content_text: Optional[str],
|
||||
) -> List[ChatCompletionAnnotation]:
|
||||
"""
|
||||
Convert Vertex AI grounding metadata to OpenAI-style annotations.
|
||||
"""
|
||||
|
||||
annotations: list[ChatCompletionAnnotation] = []
|
||||
annotations: List[ChatCompletionAnnotation] = []
|
||||
|
||||
for metadata in grounding_metadata:
|
||||
# Extract groundingSupports - these map text segments to sources
|
||||
|
|
@ -2135,7 +2139,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
grounding_chunks = metadata.get("groundingChunks", [])
|
||||
|
||||
# Build a map of chunk indices to web URIs
|
||||
chunk_to_uri_map: dict[int, dict[str, str]] = {}
|
||||
chunk_to_uri_map: Dict[int, Dict[str, str]] = {}
|
||||
for idx, chunk in enumerate(grounding_chunks):
|
||||
if "web" in chunk:
|
||||
web_data = chunk["web"]
|
||||
|
|
@ -2175,11 +2179,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _process_candidates(
|
||||
_candidates: list[Candidates],
|
||||
_candidates: List[Candidates],
|
||||
model_response: Union[ModelResponse, "ModelResponseStream"],
|
||||
standard_optional_params: dict,
|
||||
cumulative_tool_call_index: int = 0,
|
||||
) -> tuple[list[dict], list[dict], list, list, int]:
|
||||
) -> Tuple[List[dict], List[dict], List, List, int]:
|
||||
"""
|
||||
Helper method to process candidates and extract metadata
|
||||
|
||||
|
|
@ -2195,19 +2199,19 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
from litellm.types.utils import ModelResponseStream
|
||||
|
||||
grounding_metadata: list[dict] = []
|
||||
url_context_metadata: list[dict] = []
|
||||
image_response: list[ImageURLListItem] | None = None
|
||||
safety_ratings: list = []
|
||||
citation_metadata: list = []
|
||||
grounding_metadata: List[dict] = []
|
||||
url_context_metadata: List[dict] = []
|
||||
image_response: Optional[List[ImageURLListItem]] = None
|
||||
safety_ratings: List = []
|
||||
citation_metadata: List = []
|
||||
chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"}
|
||||
chat_completion_logprobs: ChoiceLogprobs | None = None
|
||||
tools: list[ChatCompletionToolCallChunk] | None = []
|
||||
functions: ChatCompletionToolCallFunctionChunk | None = None
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock] | None = None
|
||||
reasoning_content: str | None = None
|
||||
thought_signatures: Any | None = None
|
||||
server_side_tool_invocations: list[dict[str, Any]] | None = None
|
||||
chat_completion_logprobs: Optional[ChoiceLogprobs] = None
|
||||
tools: Optional[List[ChatCompletionToolCallChunk]] = []
|
||||
functions: Optional[ChatCompletionToolCallFunctionChunk] = None
|
||||
thinking_blocks: Optional[List[ChatCompletionThinkingBlock]] = None
|
||||
reasoning_content: Optional[str] = None
|
||||
thought_signatures: Optional[Any] = None
|
||||
server_side_tool_invocations: Optional[List[Dict[str, Any]]] = None
|
||||
|
||||
for idx, candidate in enumerate(_candidates):
|
||||
if "content" not in candidate:
|
||||
|
|
@ -2254,11 +2258,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
|
||||
if audio_response is not None:
|
||||
cast(dict[str, Any], chat_completion_message)["audio"] = audio_response
|
||||
cast(Dict[str, Any], chat_completion_message)["audio"] = audio_response
|
||||
chat_completion_message["content"] = None # OpenAI spec
|
||||
if image_response is not None:
|
||||
# Handle image response - combine with text content into structured format
|
||||
cast(dict[str, Any], chat_completion_message)["images"] = image_response
|
||||
cast(Dict[str, Any], chat_completion_message)["images"] = image_response
|
||||
if content is not None:
|
||||
chat_completion_message["content"] = content
|
||||
|
||||
|
|
@ -2360,13 +2364,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: LoggingClass,
|
||||
request_data: dict,
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
request_data: Dict,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: Dict,
|
||||
litellm_params: Dict,
|
||||
encoding: Any,
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ModelResponse:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -2433,11 +2437,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
response_id = completion_response.get("responseId")
|
||||
if response_id:
|
||||
model_response.id = response_id
|
||||
url_context_metadata: list[dict] = []
|
||||
url_context_metadata: List[dict] = []
|
||||
try:
|
||||
grounding_metadata: list[dict] = []
|
||||
safety_ratings: list[dict] = []
|
||||
citation_metadata: list[dict] = []
|
||||
grounding_metadata: List[dict] = []
|
||||
safety_ratings: List[dict] = []
|
||||
citation_metadata: List[dict] = []
|
||||
if _candidates:
|
||||
(
|
||||
grounding_metadata,
|
||||
|
|
@ -2501,10 +2505,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
def _transform_messages(
|
||||
self,
|
||||
messages: list[AllMessageValues],
|
||||
model: str | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
) -> list[ContentType]:
|
||||
messages: List[AllMessageValues],
|
||||
model: Optional[str] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> List[ContentType]:
|
||||
return _gemini_convert_messages_with_history(
|
||||
messages=messages,
|
||||
model=model,
|
||||
|
|
@ -2513,30 +2517,30 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
return VertexAIError(message=error_message, status_code=status_code, headers=headers)
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: Dict,
|
||||
litellm_params: Dict,
|
||||
headers: Dict,
|
||||
) -> Dict:
|
||||
raise NotImplementedError("Vertex AI has a custom implementation of transform_request. Needs sync + async.")
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict | None,
|
||||
headers: Optional[Dict],
|
||||
model: str,
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Union[str, dict] | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict:
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: Dict,
|
||||
litellm_params: Dict,
|
||||
api_key: Optional[Union[str, Dict]] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> Dict:
|
||||
default_headers = {
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
|
@ -2551,8 +2555,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
|
||||
async def make_call(
|
||||
client: AsyncHTTPHandler | None, # module-level client
|
||||
gemini_client: AsyncHTTPHandler | None, # if passed by user
|
||||
client: Optional[AsyncHTTPHandler], # module-level client
|
||||
gemini_client: Optional[AsyncHTTPHandler], # if passed by user
|
||||
api_base: str,
|
||||
headers: dict,
|
||||
data: str,
|
||||
|
|
@ -2603,8 +2607,8 @@ async def make_call(
|
|||
|
||||
|
||||
def make_sync_call(
|
||||
client: HTTPHandler | None, # module-level client
|
||||
gemini_client: HTTPHandler | None, # if passed by user
|
||||
client: Optional[HTTPHandler], # module-level client
|
||||
gemini_client: Optional[HTTPHandler], # if passed by user
|
||||
api_base: str,
|
||||
headers: dict,
|
||||
data: str,
|
||||
|
|
@ -2659,20 +2663,20 @@ class VertexLLM(VertexBase):
|
|||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
data: dict,
|
||||
timeout: Union[float, httpx.Timeout] | None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
encoding,
|
||||
logging_obj,
|
||||
stream,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
logger_fn=None,
|
||||
api_base: str | None = None,
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
vertex_project: str | None = None,
|
||||
vertex_location: str | None = None,
|
||||
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None,
|
||||
gemini_api_key: str | None = None,
|
||||
extra_headers: dict | None = None,
|
||||
api_base: Optional[str] = None,
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
vertex_project: Optional[str] = None,
|
||||
vertex_location: Optional[str] = None,
|
||||
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None,
|
||||
gemini_api_key: Optional[str] = None,
|
||||
extra_headers: Optional[dict] = None,
|
||||
) -> CustomStreamWrapper:
|
||||
should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params)
|
||||
|
||||
|
|
@ -2755,20 +2759,20 @@ class VertexLLM(VertexBase):
|
|||
custom_llm_provider: Literal[
|
||||
"vertex_ai", "vertex_ai_beta", "gemini"
|
||||
], # if it's vertex_ai or gemini (google ai studio)
|
||||
timeout: Union[float, httpx.Timeout] | None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
encoding,
|
||||
logging_obj,
|
||||
stream,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
logger_fn=None,
|
||||
api_base: str | None = None,
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
vertex_project: str | None = None,
|
||||
vertex_location: str | None = None,
|
||||
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None,
|
||||
gemini_api_key: str | None = None,
|
||||
extra_headers: dict | None = None,
|
||||
api_base: Optional[str] = None,
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
vertex_project: Optional[str] = None,
|
||||
vertex_location: Optional[str] = None,
|
||||
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None,
|
||||
gemini_api_key: Optional[str] = None,
|
||||
extra_headers: Optional[dict] = None,
|
||||
) -> Union[ModelResponse, CustomStreamWrapper]:
|
||||
should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params)
|
||||
|
||||
|
|
@ -2877,18 +2881,18 @@ class VertexLLM(VertexBase):
|
|||
logging_obj,
|
||||
optional_params: dict,
|
||||
acompletion: bool,
|
||||
timeout: Union[float, httpx.Timeout] | None,
|
||||
vertex_project: str | None,
|
||||
vertex_location: str | None,
|
||||
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None,
|
||||
gemini_api_key: str | None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
gemini_api_key: Optional[str],
|
||||
litellm_params: dict,
|
||||
logger_fn=None,
|
||||
extra_headers: dict | None = None,
|
||||
client: Union[AsyncHTTPHandler, HTTPHandler] | None = None,
|
||||
api_base: str | None = None,
|
||||
extra_headers: Optional[dict] = None,
|
||||
client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> Union[ModelResponse, CustomStreamWrapper]:
|
||||
stream: bool | None = optional_params.pop("stream", None) # type: ignore
|
||||
stream: Optional[bool] = optional_params.pop("stream", None) # type: ignore
|
||||
|
||||
transform_request_params = {
|
||||
"gemini_api_key": gemini_api_key,
|
||||
|
|
@ -3076,7 +3080,7 @@ class ModelResponseIterator:
|
|||
streaming_response,
|
||||
sync_stream: bool,
|
||||
logging_obj: LoggingClass,
|
||||
response_headers: dict[str, str] | None = None,
|
||||
response_headers: Optional[Dict[str, str]] = None,
|
||||
response: httpx.Response | None = None,
|
||||
):
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
|
|
@ -3121,9 +3125,9 @@ class ModelResponseIterator:
|
|||
|
||||
def _apply_stream_candidates(
|
||||
self,
|
||||
_candidates: list[Candidates],
|
||||
_candidates: List[Candidates],
|
||||
model_response: Any,
|
||||
) -> tuple[list[dict], list[dict], list[dict], list[dict]]:
|
||||
) -> Tuple[List[dict], List[dict], List[dict], List[dict]]:
|
||||
(
|
||||
grounding_metadata,
|
||||
url_context_metadata,
|
||||
|
|
@ -3202,8 +3206,8 @@ class ModelResponseIterator:
|
|||
self,
|
||||
processed_chunk: Any,
|
||||
model_response: Any,
|
||||
grounding_metadata: list[dict],
|
||||
) -> Usage | None:
|
||||
grounding_metadata: List[dict],
|
||||
) -> Optional[Usage]:
|
||||
if "usageMetadata" not in processed_chunk:
|
||||
return None
|
||||
|
||||
|
|
@ -3250,12 +3254,12 @@ class ModelResponseIterator:
|
|||
if blocked_response is not None:
|
||||
model_response = blocked_response
|
||||
|
||||
grounding_metadata: list[dict] = []
|
||||
url_context_metadata: list[dict] = []
|
||||
safety_ratings: list[dict] = []
|
||||
citation_metadata: list[dict] = []
|
||||
grounding_metadata: List[dict] = []
|
||||
url_context_metadata: List[dict] = []
|
||||
safety_ratings: List[dict] = []
|
||||
citation_metadata: List[dict] = []
|
||||
|
||||
_candidates: list[Candidates] | None = processed_chunk.get("candidates")
|
||||
_candidates: Optional[List[Candidates]] = processed_chunk.get("candidates")
|
||||
if _candidates:
|
||||
(
|
||||
grounding_metadata,
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
import importlib
|
||||
import os
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Literal, Optional, cast
|
||||
from typing import Any, Dict, List, Literal, Optional, Set, Type, cast
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
|
|
@ -54,7 +54,7 @@ guardrail_initializer_registry = {
|
|||
SupportedGuardrailIntegrations.LLM_AS_A_JUDGE.value: initialize_llm_as_a_judge,
|
||||
}
|
||||
|
||||
guardrail_class_registry: dict[str, type[CustomGuardrail]] = {
|
||||
guardrail_class_registry: Dict[str, Type[CustomGuardrail]] = {
|
||||
SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail
|
||||
}
|
||||
|
||||
|
|
@ -220,7 +220,7 @@ class GuardrailRegistry:
|
|||
###########################################################
|
||||
########### In memory management helpers for guardrails ###########
|
||||
############################################################
|
||||
def get_initialized_guardrail_callback(self, guardrail_name: str) -> CustomGuardrail | None:
|
||||
def get_initialized_guardrail_callback(self, guardrail_name: str) -> Optional[CustomGuardrail]:
|
||||
"""
|
||||
Returns the initialized guardrail callback for a given guardrail name
|
||||
"""
|
||||
|
|
@ -314,7 +314,7 @@ class GuardrailRegistry:
|
|||
@staticmethod
|
||||
async def get_all_guardrails_from_db(
|
||||
prisma_client: PrismaClient,
|
||||
) -> list[Guardrail]:
|
||||
) -> List[Guardrail]:
|
||||
"""
|
||||
Get all active guardrails from the database.
|
||||
Only rows with status == "active" are returned (pending_review and rejected are excluded).
|
||||
|
|
@ -325,7 +325,7 @@ class GuardrailRegistry:
|
|||
order={"created_at": "desc"},
|
||||
)
|
||||
|
||||
guardrails: list[Guardrail] = []
|
||||
guardrails: List[Guardrail] = []
|
||||
for guardrail in guardrails_from_db:
|
||||
guardrails.append(Guardrail(**(dict(guardrail)))) # type: ignore
|
||||
|
||||
|
|
@ -333,7 +333,7 @@ class GuardrailRegistry:
|
|||
except Exception as e:
|
||||
raise Exception(f"Error getting guardrails from DB: {str(e)}")
|
||||
|
||||
async def get_guardrail_by_id_from_db(self, guardrail_id: str, prisma_client: PrismaClient) -> Guardrail | None:
|
||||
async def get_guardrail_by_id_from_db(self, guardrail_id: str, prisma_client: PrismaClient) -> Optional[Guardrail]:
|
||||
"""
|
||||
Get a guardrail by its ID from the database
|
||||
"""
|
||||
|
|
@ -349,7 +349,9 @@ class GuardrailRegistry:
|
|||
except Exception as e:
|
||||
raise Exception(f"Error getting guardrail from DB: {str(e)}")
|
||||
|
||||
async def get_guardrail_by_name_from_db(self, guardrail_name: str, prisma_client: PrismaClient) -> Guardrail | None:
|
||||
async def get_guardrail_by_name_from_db(
|
||||
self, guardrail_name: str, prisma_client: PrismaClient
|
||||
) -> Optional[Guardrail]:
|
||||
"""
|
||||
Get a guardrail by its name from the database
|
||||
"""
|
||||
|
|
@ -372,17 +374,17 @@ class InMemoryGuardrailHandler:
|
|||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.IN_MEMORY_GUARDRAILS: dict[str, Guardrail] = {}
|
||||
self.IN_MEMORY_GUARDRAILS: Dict[str, Guardrail] = {}
|
||||
"""
|
||||
Guardrail id to Guardrail object mapping
|
||||
"""
|
||||
|
||||
self.guardrail_id_to_custom_guardrail: dict[str, CustomGuardrail | None] = {}
|
||||
self.guardrail_id_to_custom_guardrail: Dict[str, Optional[CustomGuardrail]] = {}
|
||||
"""
|
||||
Guardrail id to CustomGuardrail object mapping
|
||||
"""
|
||||
|
||||
self._sources: dict[str, Literal["db", "config"]] = {}
|
||||
self._sources: Dict[str, Literal["db", "config"]] = {}
|
||||
"""
|
||||
Guardrail id to provenance marker. "db" entries are reconciled against
|
||||
the DB on each polling tick; "config" entries are owned by proxy_config.yaml
|
||||
|
|
@ -392,10 +394,10 @@ class InMemoryGuardrailHandler:
|
|||
def initialize_guardrail(
|
||||
self,
|
||||
guardrail: Guardrail,
|
||||
config_file_path: str | None = None,
|
||||
config_file_path: Optional[str] = None,
|
||||
llm_router: Optional["Router"] = None,
|
||||
source: Literal["db", "config"] = "config",
|
||||
) -> Guardrail | None:
|
||||
) -> Optional[Guardrail]:
|
||||
"""
|
||||
Initialize a guardrail from a dictionary and add it to the litellm callback manager
|
||||
|
||||
|
|
@ -411,7 +413,7 @@ class InMemoryGuardrailHandler:
|
|||
self._sources[guardrail_id] = source
|
||||
return self.IN_MEMORY_GUARDRAILS[guardrail_id]
|
||||
|
||||
custom_guardrail_callback: CustomGuardrail | None = None
|
||||
custom_guardrail_callback: Optional[CustomGuardrail] = None
|
||||
litellm_params_data = guardrail["litellm_params"]
|
||||
verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data)
|
||||
|
||||
|
|
@ -487,11 +489,11 @@ class InMemoryGuardrailHandler:
|
|||
|
||||
def initialize_custom_guardrail(
|
||||
self,
|
||||
guardrail: dict,
|
||||
guardrail: Dict,
|
||||
guardrail_type: str,
|
||||
litellm_params: LitellmParams,
|
||||
config_file_path: str | None = None,
|
||||
) -> CustomGuardrail | None:
|
||||
config_file_path: Optional[str] = None,
|
||||
) -> Optional[CustomGuardrail]:
|
||||
"""
|
||||
Initialize a Custom Guardrail from a python file or module path
|
||||
|
||||
|
|
@ -576,25 +578,25 @@ class InMemoryGuardrailHandler:
|
|||
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback)
|
||||
|
||||
def list_in_memory_guardrails(self) -> list[Guardrail]:
|
||||
def list_in_memory_guardrails(self) -> List[Guardrail]:
|
||||
"""
|
||||
List all guardrails in memory
|
||||
"""
|
||||
return list(self.IN_MEMORY_GUARDRAILS.values())
|
||||
|
||||
def get_guardrail_by_id(self, guardrail_id: str) -> Guardrail | None:
|
||||
def get_guardrail_by_id(self, guardrail_id: str) -> Optional[Guardrail]:
|
||||
"""
|
||||
Get a guardrail by its ID from memory
|
||||
"""
|
||||
return self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
|
||||
|
||||
def get_source(self, guardrail_id: str) -> Literal["db", "config"] | None:
|
||||
def get_source(self, guardrail_id: str) -> Optional[Literal["db", "config"]]:
|
||||
"""
|
||||
Return the provenance of an in-memory guardrail.
|
||||
"""
|
||||
return self._sources.get(guardrail_id)
|
||||
|
||||
def reconcile_db_guardrails(self, db_guardrail_ids: set[str]) -> list[str]:
|
||||
def reconcile_db_guardrails(self, db_guardrail_ids: Set[str]) -> List[str]:
|
||||
"""
|
||||
Drop in-memory entries that originated from the DB but are no longer
|
||||
present in db_guardrail_ids. Config-loaded guardrails are never touched.
|
||||
|
|
@ -617,8 +619,8 @@ class InMemoryGuardrailHandler:
|
|||
|
||||
@staticmethod
|
||||
def _normalize_litellm_params_for_comparison(
|
||||
params: Any | None,
|
||||
) -> dict[str, Any] | None:
|
||||
params: Optional[Any],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Render litellm_params to a canonical dict so an in-memory LitellmParams and
|
||||
the raw dict loaded from the DB compare equal when they describe the same
|
||||
|
|
@ -682,9 +684,9 @@ class InMemoryGuardrailHandler:
|
|||
def reinitialize_guardrail(
|
||||
self,
|
||||
guardrail: Guardrail,
|
||||
config_file_path: str | None = None,
|
||||
config_file_path: Optional[str] = None,
|
||||
source: Literal["db", "config"] = "config",
|
||||
) -> Guardrail | None:
|
||||
) -> Optional[Guardrail]:
|
||||
"""
|
||||
Force re-initialization of a guardrail even if it exists in memory.
|
||||
Removes old callback from litellm.callbacks and creates fresh instance.
|
||||
|
|
@ -701,7 +703,9 @@ class InMemoryGuardrailHandler:
|
|||
# Initialize fresh (will add new callback to litellm.callbacks)
|
||||
return self.initialize_guardrail(guardrail=guardrail, config_file_path=config_file_path, source=source)
|
||||
|
||||
def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None:
|
||||
def sync_guardrail_from_db(
|
||||
self, guardrail: Guardrail, config_file_path: Optional[str] = None
|
||||
) -> Optional[Guardrail]:
|
||||
"""
|
||||
Sync a guardrail from DB - initializes if new, re-initializes if changed.
|
||||
This is the method to call during DB polling.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue