mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(lint): fix UP006/UP045/F401 in changed files; add depth guard to check_any_discipline
- Replace Dict/List/Optional/Tuple typing imports with built-in equivalents (UP006, UP045) across files touched in this PR diff, then clean up the now-unused typing imports (F401). - Add _MAX_CONTAINS_ANY_DEPTH guard to check_any_discipline.contains_any() to prevent RecursionError on deeply-nested mypy types.
This commit is contained in:
parent
be3e2706b0
commit
6c6ed854b1
11 changed files with 606 additions and 582 deletions
37
litellm/build-and-push.sh
Normal file
37
litellm/build-and-push.sh
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
#!/bin/bash
|
||||
set -e
|
||||
|
||||
# Configuration
|
||||
ECR_REPO="420049223852.dkr.ecr.eu-central-1.amazonaws.com/litellm-deepkeep"
|
||||
TAG="${1:-dev-$(date +%Y%m%d-%H%M%S)}"
|
||||
FULL_IMAGE="${ECR_REPO}:${TAG}"
|
||||
|
||||
echo "================================================"
|
||||
echo "Building LiteLLM Docker image with UI"
|
||||
echo "Image: ${FULL_IMAGE}"
|
||||
echo "================================================"
|
||||
|
||||
# Login to ECR
|
||||
echo "Logging in to ECR..."
|
||||
aws ecr get-login-password --region eu-central-1 | docker login --username AWS --password-stdin 420049223852.dkr.ecr.eu-central-1.amazonaws.com
|
||||
|
||||
# Build the image
|
||||
echo "Building Docker image..."
|
||||
docker build -t "${FULL_IMAGE}" -f Dockerfile .
|
||||
|
||||
# Push the image
|
||||
echo "Pushing image to ECR..."
|
||||
docker push "${FULL_IMAGE}"
|
||||
|
||||
echo "================================================"
|
||||
echo "✅ Image pushed successfully!"
|
||||
echo ""
|
||||
echo "To use in Tilt, update your tilt_config.yaniv-v2.yaml:"
|
||||
echo ""
|
||||
echo "litellm:"
|
||||
echo " image:"
|
||||
echo " repository: ${ECR_REPO}"
|
||||
echo " tag: ${TAG}"
|
||||
echo ""
|
||||
echo "Or run: tilt trigger litellm"
|
||||
echo "================================================"
|
||||
22
litellm/deepkeep_tilt_config.yaml
Normal file
22
litellm/deepkeep_tilt_config.yaml
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
model_list:
|
||||
- model_name: fake-openai-endpoint
|
||||
litellm_params:
|
||||
model: openai/fake-model
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
|
||||
litellm_settings:
|
||||
drop_params: True
|
||||
telemetry: False
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: deepkeep-firewall
|
||||
litellm_params:
|
||||
guardrail: deepkeep
|
||||
mode: [pre_call, post_call]
|
||||
api_key: "1D-VOZhxhsgG2o8pnbDcXaDmzLOWSwZpgkGZPDyVQBk"
|
||||
api_base: "http://localhost:8081/api"
|
||||
deepkeep_firewall_id: "063e4e5ba5c2be55"
|
||||
|
|
@ -7,7 +7,7 @@ Users can define
|
|||
"""
|
||||
|
||||
import copy
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Union, cast
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -36,17 +36,17 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
def get_chat_completion_prompt(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
messages: list[AllMessageValues],
|
||||
non_default_params: dict,
|
||||
prompt_id: Optional[str],
|
||||
prompt_variables: Optional[dict],
|
||||
prompt_id: str | None,
|
||||
prompt_variables: dict | None,
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
prompt_spec: Optional[PromptSpec] = None,
|
||||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
ignore_prompt_manager_model: Optional[bool] = False,
|
||||
ignore_prompt_manager_optional_params: Optional[bool] = False,
|
||||
) -> Tuple[str, List[AllMessageValues], dict]:
|
||||
prompt_spec: PromptSpec | None = None,
|
||||
prompt_label: str | None = None,
|
||||
prompt_version: int | None = None,
|
||||
ignore_prompt_manager_model: bool | None = False,
|
||||
ignore_prompt_manager_optional_params: bool | None = False,
|
||||
) -> tuple[str, list[AllMessageValues], dict]:
|
||||
"""
|
||||
Apply cache control directives based on specified injection points.
|
||||
|
||||
|
|
@ -56,7 +56,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
- non_default_params: dict - params with any global cache controls
|
||||
"""
|
||||
# Extract cache control injection points
|
||||
injection_points: List[CacheControlInjectionPoint] = non_default_params.pop(
|
||||
injection_points: list[CacheControlInjectionPoint] = non_default_params.pop(
|
||||
"cache_control_injection_points", []
|
||||
)
|
||||
if not injection_points:
|
||||
|
|
@ -66,8 +66,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
processed_messages = copy.deepcopy(messages)
|
||||
|
||||
# Separate message-level and non-message-level injection points
|
||||
message_points: List[CacheControlMessageInjectionPoint] = []
|
||||
remaining_points: List[CacheControlInjectionPoint] = []
|
||||
message_points: list[CacheControlMessageInjectionPoint] = []
|
||||
remaining_points: list[CacheControlInjectionPoint] = []
|
||||
for point in injection_points:
|
||||
if point.get("location") == "message":
|
||||
message_points.append(cast(CacheControlMessageInjectionPoint, point))
|
||||
|
|
@ -98,10 +98,10 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
@staticmethod
|
||||
def _apply_message_injections(
|
||||
points: List[CacheControlMessageInjectionPoint],
|
||||
messages: List[AllMessageValues],
|
||||
points: list[CacheControlMessageInjectionPoint],
|
||||
messages: list[AllMessageValues],
|
||||
max_blocks: int,
|
||||
) -> List[AllMessageValues]:
|
||||
) -> list[AllMessageValues]:
|
||||
"""Apply message-level cache control injection points in order.
|
||||
|
||||
Anthropic allows at most ``MAX_CACHE_CONTROL_BLOCKS`` cache_control
|
||||
|
|
@ -159,11 +159,11 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
@staticmethod
|
||||
def _resolve_target_indices(
|
||||
point: CacheControlMessageInjectionPoint, messages: List[AllMessageValues]
|
||||
) -> List[int]:
|
||||
point: CacheControlMessageInjectionPoint, messages: list[AllMessageValues]
|
||||
) -> list[int]:
|
||||
"""Resolve which message indices an injection point targets."""
|
||||
_targetted_index: Optional[Union[int, str]] = point.get("index", None)
|
||||
targetted_index: Optional[int] = None
|
||||
_targetted_index: Union[int, str] | None = point.get("index", None)
|
||||
targetted_index: int | None = None
|
||||
if isinstance(_targetted_index, str):
|
||||
try:
|
||||
targetted_index = int(_targetted_index)
|
||||
|
|
@ -249,8 +249,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
def should_run_prompt_management(
|
||||
self,
|
||||
prompt_id: Optional[str],
|
||||
prompt_spec: Optional[PromptSpec],
|
||||
prompt_id: str | None,
|
||||
prompt_spec: PromptSpec | None,
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
) -> bool:
|
||||
"""Always return False since this is not a true prompt management system."""
|
||||
|
|
@ -258,12 +258,12 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
def _compile_prompt_helper(
|
||||
self,
|
||||
prompt_id: Optional[str],
|
||||
prompt_spec: Optional[PromptSpec],
|
||||
prompt_variables: Optional[dict],
|
||||
prompt_id: str | None,
|
||||
prompt_spec: PromptSpec | None,
|
||||
prompt_variables: dict | None,
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
prompt_label: str | None = None,
|
||||
prompt_version: int | None = None,
|
||||
) -> PromptManagementClient:
|
||||
"""Not used - this hook only modifies messages, doesn't fetch prompts."""
|
||||
return PromptManagementClient(
|
||||
|
|
@ -276,12 +276,12 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
async def async_compile_prompt_helper(
|
||||
self,
|
||||
prompt_id: Optional[str],
|
||||
prompt_variables: Optional[dict],
|
||||
prompt_id: str | None,
|
||||
prompt_variables: dict | None,
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
prompt_spec: Optional[PromptSpec] = None,
|
||||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
prompt_spec: PromptSpec | None = None,
|
||||
prompt_label: str | None = None,
|
||||
prompt_version: int | None = None,
|
||||
) -> PromptManagementClient:
|
||||
"""Not used - this hook only modifies messages, doesn't fetch prompts."""
|
||||
return self._compile_prompt_helper(
|
||||
|
|
@ -296,19 +296,19 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
async def async_get_chat_completion_prompt(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
messages: list[AllMessageValues],
|
||||
non_default_params: dict,
|
||||
prompt_id: Optional[str],
|
||||
prompt_variables: Optional[dict],
|
||||
prompt_id: str | None,
|
||||
prompt_variables: dict | None,
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
prompt_spec: Optional[PromptSpec] = None,
|
||||
tools: Optional[List[Dict]] = None,
|
||||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
ignore_prompt_manager_model: Optional[bool] = False,
|
||||
ignore_prompt_manager_optional_params: Optional[bool] = False,
|
||||
) -> Tuple[str, List[AllMessageValues], dict]:
|
||||
prompt_spec: PromptSpec | None = None,
|
||||
tools: list[dict] | None = None,
|
||||
prompt_label: str | None = None,
|
||||
prompt_version: int | None = None,
|
||||
ignore_prompt_manager_model: bool | None = False,
|
||||
ignore_prompt_manager_optional_params: bool | None = False,
|
||||
) -> tuple[str, list[AllMessageValues], dict]:
|
||||
"""Async version - delegates to sync since no async operations needed."""
|
||||
return self.get_chat_completion_prompt(
|
||||
model=model,
|
||||
|
|
@ -325,15 +325,15 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def should_use_anthropic_cache_control_hook(non_default_params: Dict) -> bool:
|
||||
def should_use_anthropic_cache_control_hook(non_default_params: dict) -> bool:
|
||||
if non_default_params.get("cache_control_injection_points", None):
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def get_custom_logger_for_anthropic_cache_control_hook(
|
||||
non_default_params: Dict,
|
||||
) -> Optional[CustomLogger]:
|
||||
non_default_params: dict,
|
||||
) -> CustomLogger | None:
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
_init_custom_logger_compatible_class,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,11 +4,8 @@ from typing import (
|
|||
TYPE_CHECKING,
|
||||
Any,
|
||||
ClassVar,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Type,
|
||||
Union,
|
||||
get_args,
|
||||
)
|
||||
|
|
@ -59,7 +56,7 @@ from litellm.exceptions import (
|
|||
_PRE_CALL_EXECUTED_TOKEN = secrets.token_hex(16)
|
||||
|
||||
|
||||
def get_session_id_from_request_data(request_data: Dict[str, Any]) -> Optional[str]:
|
||||
def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None:
|
||||
"""Extract session_id from request data (litellm_session_id or metadata)."""
|
||||
session_id = request_data.get("litellm_session_id")
|
||||
if session_id:
|
||||
|
|
@ -84,20 +81,18 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
guardrail_name: Optional[str] = None,
|
||||
supported_event_hooks: Optional[List[GuardrailEventHooks]] = None,
|
||||
event_hook: Optional[
|
||||
Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]
|
||||
] = None,
|
||||
guardrail_name: str | None = None,
|
||||
supported_event_hooks: list[GuardrailEventHooks] | None = None,
|
||||
event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = None,
|
||||
default_on: bool = False,
|
||||
mask_request_content: bool = False,
|
||||
mask_response_content: bool = False,
|
||||
violation_message_template: Optional[str] = None,
|
||||
end_session_after_n_fails: Optional[int] = None,
|
||||
on_violation: Optional[str] = None,
|
||||
realtime_violation_message: Optional[str] = None,
|
||||
on_sensitive_data: Optional[str] = None,
|
||||
sensitive_data_route_to_model: Optional[str] = None,
|
||||
violation_message_template: str | None = None,
|
||||
end_session_after_n_fails: int | None = None,
|
||||
on_violation: str | None = None,
|
||||
realtime_violation_message: str | None = None,
|
||||
on_sensitive_data: str | None = None,
|
||||
sensitive_data_route_to_model: str | None = None,
|
||||
sticky_session_routing: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
|
|
@ -120,18 +115,16 @@ class CustomGuardrail(CustomLogger):
|
|||
"""
|
||||
self.guardrail_name = guardrail_name
|
||||
self.supported_event_hooks = supported_event_hooks
|
||||
self.event_hook: Optional[
|
||||
Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]
|
||||
] = event_hook
|
||||
self.event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = event_hook
|
||||
self.default_on: bool = default_on
|
||||
self.mask_request_content: bool = mask_request_content
|
||||
self.mask_response_content: bool = mask_response_content
|
||||
self.violation_message_template: Optional[str] = violation_message_template
|
||||
self.end_session_after_n_fails: Optional[int] = end_session_after_n_fails
|
||||
self.on_violation: Optional[str] = on_violation
|
||||
self.realtime_violation_message: Optional[str] = realtime_violation_message
|
||||
self.on_sensitive_data: Optional[str] = on_sensitive_data
|
||||
self.sensitive_data_route_to_model: Optional[str] = (
|
||||
self.violation_message_template: str | None = violation_message_template
|
||||
self.end_session_after_n_fails: int | None = end_session_after_n_fails
|
||||
self.on_violation: str | None = on_violation
|
||||
self.realtime_violation_message: str | None = realtime_violation_message
|
||||
self.on_sensitive_data: str | None = on_sensitive_data
|
||||
self.sensitive_data_route_to_model: str | None = (
|
||||
sensitive_data_route_to_model
|
||||
)
|
||||
self.sticky_session_routing: bool = sticky_session_routing
|
||||
|
|
@ -142,14 +135,14 @@ class CustomGuardrail(CustomLogger):
|
|||
super().__init__(**kwargs)
|
||||
|
||||
def render_violation_message(
|
||||
self, default: str, context: Optional[Dict[str, Any]] = None
|
||||
self, default: str, context: dict[str, Any] | None = None
|
||||
) -> str:
|
||||
"""Return a custom violation message if template is configured."""
|
||||
|
||||
if not self.violation_message_template:
|
||||
return default
|
||||
|
||||
format_context: Dict[str, Any] = {"default_message": default}
|
||||
format_context: dict[str, Any] = {"default_message": default}
|
||||
if context:
|
||||
format_context.update(context)
|
||||
try:
|
||||
|
|
@ -165,8 +158,8 @@ class CustomGuardrail(CustomLogger):
|
|||
def raise_passthrough_exception(
|
||||
self,
|
||||
violation_message: str,
|
||||
request_data: Dict[str, Any],
|
||||
detection_info: Optional[Dict[str, Any]] = None,
|
||||
request_data: dict[str, Any],
|
||||
detection_info: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Raise a passthrough exception for guardrail violations.
|
||||
|
|
@ -209,8 +202,8 @@ class CustomGuardrail(CustomLogger):
|
|||
def raise_sensitive_data_route_exception(
|
||||
self,
|
||||
route_to_model: str,
|
||||
request_data: Dict[str, Any],
|
||||
detection_info: Optional[Dict[str, Any]] = None,
|
||||
request_data: dict[str, Any],
|
||||
detection_info: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Raise an exception to reroute the request to a different model.
|
||||
|
|
@ -248,8 +241,8 @@ class CustomGuardrail(CustomLogger):
|
|||
)
|
||||
|
||||
def _get_session_id_from_request_data(
|
||||
self, request_data: Dict[str, Any]
|
||||
) -> Optional[str]:
|
||||
self, request_data: dict[str, Any]
|
||||
) -> str | None:
|
||||
"""Extract session_id from request data."""
|
||||
return get_session_id_from_request_data(request_data)
|
||||
|
||||
|
|
@ -265,8 +258,8 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def handle_sensitive_data_detection(
|
||||
self,
|
||||
request_data: Dict[str, Any],
|
||||
detection_info: Optional[Dict[str, Any]] = None,
|
||||
request_data: dict[str, Any],
|
||||
detection_info: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Handle sensitive data detection based on guardrail configuration.
|
||||
|
|
@ -309,7 +302,7 @@ class CustomGuardrail(CustomLogger):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
|
||||
def get_config_model() -> type["GuardrailConfigModel"] | None:
|
||||
"""
|
||||
Returns the config model for the guardrail
|
||||
|
||||
|
|
@ -319,14 +312,12 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def _validate_event_hook(
|
||||
self,
|
||||
event_hook: Optional[
|
||||
Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]
|
||||
],
|
||||
supported_event_hooks: List[GuardrailEventHooks],
|
||||
event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None,
|
||||
supported_event_hooks: list[GuardrailEventHooks],
|
||||
) -> None:
|
||||
def _validate_event_hook_list_is_in_supported_event_hooks(
|
||||
event_hook: Union[List[GuardrailEventHooks], List[str]],
|
||||
supported_event_hooks: List[GuardrailEventHooks],
|
||||
event_hook: Union[list[GuardrailEventHooks], list[str]],
|
||||
supported_event_hooks: list[GuardrailEventHooks],
|
||||
) -> None:
|
||||
for hook in event_hook:
|
||||
if isinstance(hook, str):
|
||||
|
|
@ -391,7 +382,7 @@ class CustomGuardrail(CustomLogger):
|
|||
key_meta = meta.get("user_api_key_metadata") or key_meta
|
||||
return {**team_meta, **key_meta}
|
||||
|
||||
def get_disable_global_guardrail(self, data: dict) -> Optional[bool]:
|
||||
def get_disable_global_guardrail(self, data: dict) -> bool | None:
|
||||
"""
|
||||
Returns True if the global guardrail should be disabled.
|
||||
|
||||
|
|
@ -400,7 +391,7 @@ class CustomGuardrail(CustomLogger):
|
|||
"""
|
||||
return self._get_admin_metadata(data).get("disable_global_guardrails", False)
|
||||
|
||||
def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> List[str]:
|
||||
def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> list[str]:
|
||||
"""
|
||||
Returns the list of global guardrail names the team/key has opted out of.
|
||||
|
||||
|
|
@ -433,7 +424,7 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def get_guardrail_from_metadata(
|
||||
self, data: dict
|
||||
) -> Union[List[str], List[Dict[str, DynamicGuardrailParams]]]:
|
||||
) -> Union[list[str], list[dict[str, DynamicGuardrailParams]]]:
|
||||
"""
|
||||
Returns the guardrail(s) to be run from the metadata or root
|
||||
"""
|
||||
|
|
@ -454,7 +445,7 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def _guardrail_is_in_requested_guardrails(
|
||||
self,
|
||||
requested_guardrails: Union[List[str], List[Dict[str, DynamicGuardrailParams]]],
|
||||
requested_guardrails: Union[list[str], list[dict[str, DynamicGuardrailParams]]],
|
||||
) -> bool:
|
||||
for _guardrail in requested_guardrails:
|
||||
if isinstance(_guardrail, dict):
|
||||
|
|
@ -466,13 +457,13 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
return False
|
||||
|
||||
def _pre_call_marker(self) -> Optional[str]:
|
||||
def _pre_call_marker(self) -> str | None:
|
||||
name = self.guardrail_name
|
||||
if not name:
|
||||
return None
|
||||
return f"{_PRE_CALL_EXECUTED_TOKEN}:{name}"
|
||||
|
||||
def mark_pre_call_hook_ran(self, data: Dict[str, Any]) -> None:
|
||||
def mark_pre_call_hook_ran(self, data: dict[str, Any]) -> None:
|
||||
"""
|
||||
Record that this guardrail's ``async_pre_call_hook`` already ran for this
|
||||
request, so the deployment-level hook does not run it a second time.
|
||||
|
|
@ -497,7 +488,7 @@ class CustomGuardrail(CustomLogger):
|
|||
return
|
||||
data["metadata"] = {PRE_CALL_EXECUTED_GUARDRAILS_KEY: [marker]}
|
||||
|
||||
def _pre_call_hook_already_ran(self, data: Dict[str, Any]) -> bool:
|
||||
def _pre_call_hook_already_ran(self, data: dict[str, Any]) -> bool:
|
||||
marker = self._pre_call_marker()
|
||||
if marker is None:
|
||||
return False
|
||||
|
|
@ -510,8 +501,8 @@ class CustomGuardrail(CustomLogger):
|
|||
return False
|
||||
|
||||
async def async_pre_call_deployment_hook(
|
||||
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
|
||||
) -> Optional[dict]:
|
||||
self, kwargs: dict[str, Any], call_type: CallTypes | None
|
||||
) -> dict | None:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
# should run guardrail
|
||||
|
|
@ -556,8 +547,8 @@ class CustomGuardrail(CustomLogger):
|
|||
self,
|
||||
request_data: dict,
|
||||
response: LLMResponseTypes,
|
||||
call_type: Optional[CallTypes],
|
||||
) -> Optional[LLMResponseTypes]:
|
||||
call_type: CallTypes | None,
|
||||
) -> LLMResponseTypes | None:
|
||||
"""
|
||||
Allow modifying / reviewing the response just after it's received from the deployment.
|
||||
"""
|
||||
|
|
@ -755,16 +746,16 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def add_standard_logging_guardrail_information_to_request_data(
|
||||
self,
|
||||
guardrail_json_response: Union[Exception, str, dict, List[dict]],
|
||||
guardrail_json_response: Union[Exception, str, dict, list[dict]],
|
||||
request_data: dict,
|
||||
guardrail_status: GuardrailStatus,
|
||||
start_time: Optional[float] = None,
|
||||
end_time: Optional[float] = None,
|
||||
duration: Optional[float] = None,
|
||||
masked_entity_count: Optional[Dict[str, int]] = None,
|
||||
guardrail_provider: Optional[str] = None,
|
||||
event_type: Optional[GuardrailEventHooks] = None,
|
||||
tracing_detail: Optional[GuardrailTracingDetail] = None,
|
||||
start_time: float | None = None,
|
||||
end_time: float | None = None,
|
||||
duration: float | None = None,
|
||||
masked_entity_count: dict[str, int] | None = None,
|
||||
guardrail_provider: str | None = None,
|
||||
event_type: GuardrailEventHooks | None = None,
|
||||
tracing_detail: GuardrailTracingDetail | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Builds `StandardLoggingGuardrailInformation` and adds it to the request metadata so it can be used for logging to DataDog, Langfuse, etc.
|
||||
|
|
@ -780,7 +771,7 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
# Use event_type if provided, otherwise fall back to self.event_hook
|
||||
guardrail_mode: Union[
|
||||
GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks]
|
||||
GuardrailEventHooks, GuardrailMode, list[GuardrailEventHooks]
|
||||
]
|
||||
if event_type is not None:
|
||||
guardrail_mode = event_type
|
||||
|
|
@ -894,13 +885,13 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def _process_response(
|
||||
self,
|
||||
response: Optional[Dict],
|
||||
response: dict | None,
|
||||
request_data: dict,
|
||||
start_time: Optional[float] = None,
|
||||
end_time: Optional[float] = None,
|
||||
duration: Optional[float] = None,
|
||||
event_type: Optional[GuardrailEventHooks] = None,
|
||||
original_inputs: Optional[Dict] = None,
|
||||
start_time: float | None = None,
|
||||
end_time: float | None = None,
|
||||
duration: float | None = None,
|
||||
event_type: GuardrailEventHooks | None = None,
|
||||
original_inputs: dict | None = None,
|
||||
):
|
||||
"""
|
||||
Add StandardLoggingGuardrailInformation to the request data
|
||||
|
|
@ -908,7 +899,7 @@ class CustomGuardrail(CustomLogger):
|
|||
This gets logged on downsteam Langfuse, DataDog, etc.
|
||||
"""
|
||||
# Convert None to empty dict to satisfy type requirements
|
||||
guardrail_response: Union[Dict[str, Any], str] = (
|
||||
guardrail_response: Union[dict[str, Any], str] = (
|
||||
{} if response is None else response
|
||||
)
|
||||
|
||||
|
|
@ -970,10 +961,10 @@ class CustomGuardrail(CustomLogger):
|
|||
self,
|
||||
e: Exception,
|
||||
request_data: dict,
|
||||
start_time: Optional[float] = None,
|
||||
end_time: Optional[float] = None,
|
||||
duration: Optional[float] = None,
|
||||
event_type: Optional[GuardrailEventHooks] = None,
|
||||
start_time: float | None = None,
|
||||
end_time: float | None = None,
|
||||
duration: float | None = None,
|
||||
event_type: GuardrailEventHooks | None = None,
|
||||
):
|
||||
"""
|
||||
Add StandardLoggingGuardrailInformation to the request data
|
||||
|
|
@ -1002,7 +993,7 @@ class CustomGuardrail(CustomLogger):
|
|||
)
|
||||
raise e
|
||||
|
||||
def _inputs_were_modified(self, original_inputs: Dict, response: Dict) -> bool:
|
||||
def _inputs_were_modified(self, original_inputs: dict, response: dict) -> bool:
|
||||
"""
|
||||
Compare original inputs with response to determine if content was modified.
|
||||
|
||||
|
|
@ -1047,8 +1038,8 @@ class CustomGuardrail(CustomLogger):
|
|||
setattr(self, key, value)
|
||||
|
||||
def get_guardrails_messages_for_call_type(
|
||||
self, call_type: CallTypes, data: Optional[dict] = None
|
||||
) -> Optional[List[AllMessageValues]]:
|
||||
self, call_type: CallTypes, data: dict | None = None
|
||||
) -> list[AllMessageValues] | None:
|
||||
"""
|
||||
Returns the messages for the given call type and data
|
||||
"""
|
||||
|
|
@ -1089,7 +1080,7 @@ class CustomGuardrail(CustomLogger):
|
|||
input=input_data,
|
||||
responses_api_request=data,
|
||||
)
|
||||
return cast(List[AllMessageValues], messages)
|
||||
return cast(list[AllMessageValues], messages)
|
||||
return None
|
||||
|
||||
|
||||
|
|
@ -1118,7 +1109,7 @@ def log_guardrail_information(func):
|
|||
|
||||
def _infer_event_type_from_function_name(
|
||||
func_name: str,
|
||||
) -> Optional[GuardrailEventHooks]:
|
||||
) -> GuardrailEventHooks | None:
|
||||
"""Infer the actual event type from the function name"""
|
||||
if func_name == "async_pre_call_hook":
|
||||
return GuardrailEventHooks.pre_call
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ identical metrics. The attribute cardinality filter is reused from v1 by import
|
|||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any, FrozenSet, Mapping, Optional
|
||||
from typing import Any, Mapping
|
||||
|
||||
from opentelemetry.metrics import Histogram, Meter
|
||||
|
||||
|
|
@ -82,12 +82,12 @@ class GenAIMetricRecorder:
|
|||
"""
|
||||
|
||||
def __init__(
|
||||
self, metrics: GenAIMetrics, callback_name: Optional[str] = None
|
||||
self, metrics: GenAIMetrics, callback_name: str | None = None
|
||||
) -> None:
|
||||
self._metrics = metrics
|
||||
self._callback_name = callback_name
|
||||
self._include: Optional[FrozenSet[str]] = None
|
||||
self._exclude: Optional[FrozenSet[str]] = None
|
||||
self._include: frozenset[str] | None = None
|
||||
self._exclude: frozenset[str] | None = None
|
||||
self._filter_resolved = False
|
||||
|
||||
def record(
|
||||
|
|
|
|||
|
|
@ -4,11 +4,7 @@ import time
|
|||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Dict,
|
||||
List,
|
||||
NoReturn,
|
||||
Optional,
|
||||
Tuple,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
|
@ -152,8 +148,8 @@ def _basic_sanitize_anthropic_tool_name(name: str) -> str:
|
|||
|
||||
|
||||
def _build_anthropic_tool_name_maps(
|
||||
original_names: List[str],
|
||||
) -> Tuple[Dict[str, str], Dict[str, str]]:
|
||||
original_names: list[str],
|
||||
) -> tuple[dict[str, str], dict[str, str]]:
|
||||
"""Build (forward, reverse) tool-name maps for a single request.
|
||||
|
||||
forward[original] = sanitized -- only present when name was rewritten
|
||||
|
|
@ -173,7 +169,7 @@ def _build_anthropic_tool_name_maps(
|
|||
seen gets the disambiguating suffix. Callers should preserve the
|
||||
caller's tool order (we do).
|
||||
"""
|
||||
forward: Dict[str, str] = {}
|
||||
forward: dict[str, str] = {}
|
||||
used: set = set()
|
||||
|
||||
# First pass: reserve slots for names that are already valid so they
|
||||
|
|
@ -214,7 +210,7 @@ def _build_anthropic_tool_name_maps(
|
|||
return forward, reverse
|
||||
|
||||
|
||||
REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT: Dict[str, str] = {
|
||||
REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT: dict[str, str] = {
|
||||
"low": "low",
|
||||
"minimal": "low",
|
||||
"medium": "medium",
|
||||
|
|
@ -237,23 +233,23 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
to pass metadata to anthropic, it's {"user_id": "any-relevant-information"}
|
||||
"""
|
||||
|
||||
max_tokens: Optional[int] = None
|
||||
stop_sequences: Optional[list] = None
|
||||
temperature: Optional[int] = None
|
||||
top_p: Optional[int] = None
|
||||
top_k: Optional[int] = None
|
||||
metadata: Optional[dict] = None
|
||||
system: Optional[str] = None
|
||||
max_tokens: int | None = None
|
||||
stop_sequences: list | None = None
|
||||
temperature: int | None = None
|
||||
top_p: int | None = None
|
||||
top_k: int | None = None
|
||||
metadata: dict | None = None
|
||||
system: str | None = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_tokens: Optional[int] = None,
|
||||
stop_sequences: Optional[list] = None,
|
||||
temperature: Optional[int] = None,
|
||||
top_p: Optional[int] = None,
|
||||
top_k: Optional[int] = None,
|
||||
metadata: Optional[dict] = None,
|
||||
system: Optional[str] = None,
|
||||
max_tokens: int | None = None,
|
||||
stop_sequences: list | None = None,
|
||||
temperature: int | None = None,
|
||||
top_p: int | None = None,
|
||||
top_k: int | None = None,
|
||||
metadata: dict | None = None,
|
||||
system: str | None = None,
|
||||
) -> None:
|
||||
locals_ = locals().copy()
|
||||
for key, value in locals_.items():
|
||||
|
|
@ -261,11 +257,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
setattr(self.__class__, key, value)
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
def custom_llm_provider(self) -> str | None:
|
||||
return "anthropic"
|
||||
|
||||
@classmethod
|
||||
def get_config(cls, *, model: Optional[str] = None):
|
||||
def get_config(cls, *, model: str | None = None):
|
||||
config = super().get_config()
|
||||
|
||||
# anthropic requires a default value for max_tokens
|
||||
|
|
@ -275,7 +271,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return config
|
||||
|
||||
@staticmethod
|
||||
def get_max_tokens_for_model(model: Optional[str] = None) -> int:
|
||||
def get_max_tokens_for_model(model: str | None = None) -> int:
|
||||
"""
|
||||
Get the max output tokens for a given model.
|
||||
Falls back to DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS (configurable via env var) if model is not found.
|
||||
|
|
@ -292,7 +288,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def convert_tool_use_to_openai_format(
|
||||
anthropic_tool_content: Dict[str, Any],
|
||||
anthropic_tool_content: dict[str, Any],
|
||||
index: int,
|
||||
) -> ChatCompletionToolCallChunk:
|
||||
"""
|
||||
|
|
@ -317,7 +313,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
)
|
||||
# Include caller information if present (for programmatic tool calling)
|
||||
if "caller" in anthropic_tool_content:
|
||||
tool_call["caller"] = cast(Dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item]
|
||||
tool_call["caller"] = cast(dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item]
|
||||
return tool_call
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -344,7 +340,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _validate_effort_for_model(model: str, effort: Optional[str]) -> Optional[str]:
|
||||
def _validate_effort_for_model(model: str, effort: str | None) -> str | None:
|
||||
"""Return ``None`` if ``effort`` is allowed on ``model``, else an error message."""
|
||||
if effort == "max" and not (
|
||||
AnthropicConfig._is_claude_4_6_model(model)
|
||||
|
|
@ -434,7 +430,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return params
|
||||
|
||||
@staticmethod
|
||||
def filter_anthropic_output_schema(schema: Dict[str, Any]) -> Dict[str, Any]:
|
||||
def filter_anthropic_output_schema(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
Filter out unsupported fields from JSON schema for Anthropic's output_format API.
|
||||
|
||||
|
|
@ -492,7 +488,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
constraint_labels[field].format(schema[field])
|
||||
)
|
||||
|
||||
result: Dict[str, Any] = {}
|
||||
result: dict[str, Any] = {}
|
||||
|
||||
# Update description with removed constraint info
|
||||
if constraint_descriptions:
|
||||
|
|
@ -548,8 +544,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return result
|
||||
|
||||
def get_json_schema_from_pydantic_object(
|
||||
self, response_format: Union[Any, Dict, None]
|
||||
) -> Optional[dict]:
|
||||
self, response_format: Union[Any, dict, None]
|
||||
) -> dict | None:
|
||||
return type_to_response_format_param(
|
||||
response_format, ref_template="/$defs/{model}"
|
||||
) # Relevant issue: https://github.com/BerriAI/litellm/issues/7755
|
||||
|
|
@ -564,10 +560,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _map_tool_choice(
|
||||
self,
|
||||
tool_choice: Optional[str],
|
||||
parallel_tool_use: Optional[bool],
|
||||
) -> Optional[AnthropicMessagesToolChoice]:
|
||||
_tool_choice: Optional[AnthropicMessagesToolChoice] = None
|
||||
tool_choice: str | None,
|
||||
parallel_tool_use: bool | None,
|
||||
) -> AnthropicMessagesToolChoice | None:
|
||||
_tool_choice: AnthropicMessagesToolChoice | None = None
|
||||
if tool_choice == "auto":
|
||||
_tool_choice = AnthropicMessagesToolChoice(
|
||||
type="auto",
|
||||
|
|
@ -608,9 +604,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def _map_tool_helper(
|
||||
self,
|
||||
tool: ChatCompletionToolParam,
|
||||
) -> Tuple[Optional[AllAnthropicToolsValues], Optional[AnthropicMcpServerTool]]:
|
||||
returned_tool: Optional[AllAnthropicToolsValues] = None
|
||||
mcp_server: Optional[AnthropicMcpServerTool] = None
|
||||
) -> tuple[AllAnthropicToolsValues | None, AnthropicMcpServerTool | None]:
|
||||
returned_tool: AllAnthropicToolsValues | None = None
|
||||
mcp_server: AnthropicMcpServerTool | None = None
|
||||
|
||||
if tool["type"] == "function" or tool["type"] == "custom":
|
||||
_input_schema: dict = tool["function"].get(
|
||||
|
|
@ -665,10 +661,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
if "parameters" not in tool["function"]:
|
||||
raise ValueError("Missing required parameter: parameters")
|
||||
|
||||
_display_width_px: Optional[int] = tool["function"]["parameters"].get(
|
||||
_display_width_px: int | None = tool["function"]["parameters"].get(
|
||||
"display_width_px"
|
||||
)
|
||||
_display_height_px: Optional[int] = tool["function"]["parameters"].get(
|
||||
_display_height_px: int | None = tool["function"]["parameters"].get(
|
||||
"display_height_px"
|
||||
)
|
||||
if _display_width_px is None or _display_height_px is None:
|
||||
|
|
@ -841,14 +837,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
from litellm.types.llms.anthropic import AnthropicMcpServerToolConfiguration
|
||||
|
||||
allowed_tools = tool.get("allowed_tools", None)
|
||||
tool_configuration: Optional[AnthropicMcpServerToolConfiguration] = None
|
||||
tool_configuration: AnthropicMcpServerToolConfiguration | None = None
|
||||
if allowed_tools is not None:
|
||||
tool_configuration = AnthropicMcpServerToolConfiguration(
|
||||
allowed_tools=tool.get("allowed_tools", None),
|
||||
)
|
||||
|
||||
headers = tool.get("headers", {})
|
||||
authorization_token: Optional[str] = None
|
||||
authorization_token: str | None = None
|
||||
if headers is not None:
|
||||
bearer_token = headers.get("Authorization", None)
|
||||
if bearer_token is not None:
|
||||
|
|
@ -868,8 +864,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _map_tools(
|
||||
self,
|
||||
tools: List,
|
||||
) -> Tuple[List[AllAnthropicToolsValues], List[AnthropicMcpServerTool]]:
|
||||
tools: list,
|
||||
) -> tuple[list[AllAnthropicToolsValues], list[AnthropicMcpServerTool]]:
|
||||
anthropic_tools = []
|
||||
mcp_servers = []
|
||||
for tool in tools:
|
||||
|
|
@ -918,9 +914,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _rewrite_tool_names_in_messages(
|
||||
messages: List[AllMessageValues],
|
||||
name_forward_map: Dict[str, str],
|
||||
) -> List[AllMessageValues]:
|
||||
messages: list[AllMessageValues],
|
||||
name_forward_map: dict[str, str],
|
||||
) -> list[AllMessageValues]:
|
||||
"""Return a copy of `messages` with tool_call/function_call names
|
||||
rewritten using the per-request forward map.
|
||||
|
||||
|
|
@ -931,7 +927,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
"""
|
||||
if not name_forward_map:
|
||||
return messages
|
||||
new_messages: List[AllMessageValues] = []
|
||||
new_messages: list[AllMessageValues] = []
|
||||
for msg in messages:
|
||||
if not isinstance(msg, dict):
|
||||
new_messages.append(msg)
|
||||
|
|
@ -979,8 +975,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _build_request_tool_name_maps(
|
||||
tools: List,
|
||||
) -> Tuple[Dict[str, str], Dict[str, str]]:
|
||||
tools: list,
|
||||
) -> tuple[dict[str, str], dict[str, str]]:
|
||||
"""Build the (forward, reverse) tool-name maps for an OpenAI tools list.
|
||||
|
||||
Operates on **OpenAI-format** tool dicts (pre-``_map_tools``). The
|
||||
|
|
@ -994,7 +990,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
the original name out of either ``{"function": {"name": ...}}``
|
||||
(legacy OpenAI shape) or ``{"name": ...}`` (rare top-level shape).
|
||||
"""
|
||||
original_names: List[str] = []
|
||||
original_names: list[str] = []
|
||||
for tool in tools or []:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
|
|
@ -1011,8 +1007,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _sanitize_tool_names_in_request(
|
||||
optional_params: Dict[str, Any],
|
||||
) -> Tuple[Dict[str, str], Dict[str, str]]:
|
||||
optional_params: dict[str, Any],
|
||||
) -> tuple[dict[str, str], dict[str, str]]:
|
||||
"""Sanitize ``optional_params['tools']`` and ``optional_params['tool_choice']``
|
||||
in place so every name matches Anthropic's ``^[a-zA-Z0-9_-]{1,128}$``.
|
||||
|
||||
|
|
@ -1036,7 +1032,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
# Order matters: the first occurrence wins the canonical slot;
|
||||
# later collisions get numeric suffixes (see
|
||||
# ``_build_anthropic_tool_name_maps``).
|
||||
original_names: List[str] = []
|
||||
original_names: list[str] = []
|
||||
for t in tools:
|
||||
if not isinstance(t, dict):
|
||||
continue
|
||||
|
|
@ -1058,7 +1054,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
# so a caller reusing the same tool list/dicts across requests
|
||||
# doesn't see its inputs permanently rewritten (which would also
|
||||
# drop the original key from `forward` on the next request).
|
||||
new_tools: List[Any] = []
|
||||
new_tools: list[Any] = []
|
||||
for t in tools:
|
||||
if (
|
||||
isinstance(t, dict)
|
||||
|
|
@ -1084,7 +1080,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
return forward, reverse
|
||||
|
||||
def _detect_tool_search_tools(self, tools: Optional[List]) -> bool:
|
||||
def _detect_tool_search_tools(self, tools: list | None) -> bool:
|
||||
"""Check if tool search tools are present in the tools list."""
|
||||
if not tools:
|
||||
return False
|
||||
|
|
@ -1098,7 +1094,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return True
|
||||
return False
|
||||
|
||||
def _separate_deferred_tools(self, tools: List) -> Tuple[List, List]:
|
||||
def _separate_deferred_tools(self, tools: list) -> tuple[list, list]:
|
||||
"""
|
||||
Separate tools into deferred and non-deferred lists.
|
||||
|
||||
|
|
@ -1118,9 +1114,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _expand_tool_references(
|
||||
self,
|
||||
content: List,
|
||||
deferred_tools: List,
|
||||
) -> List:
|
||||
content: list,
|
||||
deferred_tools: list,
|
||||
) -> list:
|
||||
"""
|
||||
Expand tool_reference blocks to full tool definitions.
|
||||
|
||||
|
|
@ -1162,9 +1158,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return expanded_content
|
||||
|
||||
def _map_stop_sequences(
|
||||
self, stop: Optional[Union[str, List[str]]]
|
||||
) -> Optional[List[str]]:
|
||||
new_stop: Optional[List[str]] = None
|
||||
self, stop: Union[str, list[str]] | None
|
||||
) -> list[str] | None:
|
||||
new_stop: list[str] | None = None
|
||||
if isinstance(stop, str):
|
||||
if (
|
||||
stop.isspace() and litellm.drop_params is True
|
||||
|
|
@ -1185,10 +1181,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _map_reasoning_effort(
|
||||
reasoning_effort: Optional[Union[REASONING_EFFORT, str]],
|
||||
reasoning_effort: Union[REASONING_EFFORT, str] | None,
|
||||
model: str,
|
||||
llm_provider: str = "anthropic",
|
||||
) -> Optional[AnthropicThinkingParam]:
|
||||
) -> AnthropicThinkingParam | None:
|
||||
if reasoning_effort is None or reasoning_effort == "none":
|
||||
return None
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model):
|
||||
|
|
@ -1240,11 +1236,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
)
|
||||
|
||||
def _extract_json_schema_from_response_format(
|
||||
self, value: Optional[dict]
|
||||
) -> Optional[dict]:
|
||||
self, value: dict | None
|
||||
) -> dict | None:
|
||||
if value is None:
|
||||
return None
|
||||
json_schema: Optional[dict] = None
|
||||
json_schema: dict | None = None
|
||||
if "response_schema" in value:
|
||||
json_schema = value["response_schema"]
|
||||
elif "json_schema" in value:
|
||||
|
|
@ -1253,9 +1249,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return json_schema
|
||||
|
||||
def map_response_format_to_anthropic_output_format(
|
||||
self, value: Optional[dict]
|
||||
) -> Optional[AnthropicOutputSchema]:
|
||||
json_schema: Optional[dict] = self._extract_json_schema_from_response_format(
|
||||
self, value: dict | None
|
||||
) -> AnthropicOutputSchema | None:
|
||||
json_schema: dict | None = self._extract_json_schema_from_response_format(
|
||||
value
|
||||
)
|
||||
if json_schema is None:
|
||||
|
|
@ -1283,15 +1279,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
)
|
||||
|
||||
def map_response_format_to_anthropic_tool(
|
||||
self, value: Optional[dict], optional_params: dict, is_thinking_enabled: bool
|
||||
) -> Optional[AnthropicMessagesTool]:
|
||||
self, value: dict | None, optional_params: dict, is_thinking_enabled: bool
|
||||
) -> AnthropicMessagesTool | None:
|
||||
ignore_response_format_types = ["text"]
|
||||
if (
|
||||
value is None or value["type"] in ignore_response_format_types
|
||||
): # value is a no-op
|
||||
return None
|
||||
|
||||
json_schema: Optional[dict] = self._extract_json_schema_from_response_format(
|
||||
json_schema: dict | None = self._extract_json_schema_from_response_format(
|
||||
value
|
||||
)
|
||||
if json_schema is None:
|
||||
|
|
@ -1342,8 +1338,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def map_openai_context_management_to_anthropic(
|
||||
context_management: Union[List[Dict[str, Any]], Dict[str, Any]],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
context_management: Union[list[dict[str, Any]], dict[str, Any]],
|
||||
) -> dict[str, Any] | None:
|
||||
"""
|
||||
OpenAI format: [{"type": "compaction", "compact_threshold": 200000}]
|
||||
Anthropic format: {
|
||||
|
|
@ -1374,7 +1370,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
entry_type = entry.get("type")
|
||||
if entry_type == "compaction":
|
||||
anthropic_edit: Dict[str, Any] = {"type": "compact_20260112"}
|
||||
anthropic_edit: dict[str, Any] = {"type": "compact_20260112"}
|
||||
compact_threshold = entry.get("compact_threshold")
|
||||
# Rewrite to 'trigger' with correct nesting if threshold exists
|
||||
if compact_threshold is not None and isinstance(
|
||||
|
|
@ -1438,7 +1434,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
if mcp_servers:
|
||||
optional_params["mcp_servers"] = mcp_servers
|
||||
elif param == "tool_choice" or param == "parallel_tool_calls":
|
||||
_tool_choice: Optional[AnthropicMessagesToolChoice] = (
|
||||
_tool_choice: AnthropicMessagesToolChoice | None = (
|
||||
self._map_tool_choice(
|
||||
tool_choice=non_default_params.get("tool_choice"),
|
||||
parallel_tool_use=non_default_params.get("parallel_tool_calls"),
|
||||
|
|
@ -1584,7 +1580,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _create_json_tool_call_for_response_format(
|
||||
self,
|
||||
json_schema: Optional[dict] = None,
|
||||
json_schema: dict | None = None,
|
||||
) -> AnthropicMessagesTool:
|
||||
"""
|
||||
Handles creating a tool call for getting responses in JSON format.
|
||||
|
|
@ -1622,8 +1618,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return False
|
||||
|
||||
def translate_system_message(
|
||||
self, messages: List[AllMessageValues]
|
||||
) -> List[AnthropicSystemMessageContent]:
|
||||
self, messages: list[AllMessageValues]
|
||||
) -> list[AnthropicSystemMessageContent]:
|
||||
"""
|
||||
Translate system message to anthropic format.
|
||||
|
||||
|
|
@ -1631,7 +1627,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
When should_strip_billing_metadata() is True, x-anthropic-billing-header system blocks are dropped.
|
||||
"""
|
||||
system_prompt_indices = []
|
||||
anthropic_system_message_list: List[AnthropicSystemMessageContent] = []
|
||||
anthropic_system_message_list: list[AnthropicSystemMessageContent] = []
|
||||
for idx, message in enumerate(messages):
|
||||
if message["role"] == "system":
|
||||
system_prompt_indices.append(idx)
|
||||
|
|
@ -1691,9 +1687,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def add_code_execution_tool(
|
||||
self,
|
||||
messages: List[AllAnthropicMessageValues],
|
||||
tools: List[Union[AllAnthropicToolsValues, Dict]],
|
||||
) -> List[Union[AllAnthropicToolsValues, Dict]]:
|
||||
messages: list[AllAnthropicMessageValues],
|
||||
tools: list[Union[AllAnthropicToolsValues, dict]],
|
||||
) -> list[Union[AllAnthropicToolsValues, dict]]:
|
||||
"""if 'container_upload' in messages, add code_execution tool"""
|
||||
add_code_execution_tool = False
|
||||
for message in messages:
|
||||
|
|
@ -1832,7 +1828,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
|
|
@ -1938,7 +1934,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
## Add code_execution tool if container_upload is in messages
|
||||
_tools = (
|
||||
cast(
|
||||
Optional[List[Union[AllAnthropicToolsValues, Dict]]],
|
||||
list[Union[AllAnthropicToolsValues, dict]] | None,
|
||||
optional_params.get("tools"),
|
||||
)
|
||||
or []
|
||||
|
|
@ -2045,12 +2041,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _resolve_json_mode_non_streaming(
|
||||
self,
|
||||
json_mode: Optional[bool],
|
||||
tool_calls: List[ChatCompletionToolCallChunk],
|
||||
) -> Tuple[
|
||||
Optional[LitellmMessage],
|
||||
List[ChatCompletionToolCallChunk],
|
||||
Optional[str],
|
||||
json_mode: bool | None,
|
||||
tool_calls: list[ChatCompletionToolCallChunk],
|
||||
) -> tuple[
|
||||
LitellmMessage | None,
|
||||
list[ChatCompletionToolCallChunk],
|
||||
str | None,
|
||||
]:
|
||||
"""Strip internal response_format tool calls; merge payload into content when mixed with user tools."""
|
||||
if json_mode is not True or not tool_calls:
|
||||
|
|
@ -2075,38 +2071,30 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
first_json = tool_calls[json_indices[0]]
|
||||
json_msg = AnthropicConfig._convert_tool_response_to_message([first_json])
|
||||
extra_content: Optional[str] = (
|
||||
extra_content: str | None = (
|
||||
json_msg.content if json_msg is not None else None
|
||||
)
|
||||
filtered_tools = [t for i, t in enumerate(tool_calls) if i not in json_indices]
|
||||
return None, filtered_tools, extra_content
|
||||
|
||||
def extract_response_content(self, completion_response: dict) -> Tuple[
|
||||
def extract_response_content(self, completion_response: dict) -> tuple[
|
||||
str,
|
||||
Optional[List[Any]],
|
||||
Optional[
|
||||
List[
|
||||
Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]
|
||||
]
|
||||
],
|
||||
Optional[str],
|
||||
List[ChatCompletionToolCallChunk],
|
||||
Optional[List[Any]],
|
||||
Optional[List[Any]],
|
||||
Optional[List[Any]],
|
||||
list[Any] | None,
|
||||
list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None,
|
||||
str | None,
|
||||
list[ChatCompletionToolCallChunk],
|
||||
list[Any] | None,
|
||||
list[Any] | None,
|
||||
list[Any] | None,
|
||||
]:
|
||||
text_content = ""
|
||||
citations: Optional[List[Any]] = None
|
||||
thinking_blocks: Optional[
|
||||
List[
|
||||
Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]
|
||||
]
|
||||
] = None
|
||||
reasoning_content: Optional[str] = None
|
||||
tool_calls: List[ChatCompletionToolCallChunk] = []
|
||||
web_search_results: Optional[List[Any]] = None
|
||||
tool_results: Optional[List[Any]] = None
|
||||
compaction_blocks: Optional[List[Any]] = None
|
||||
citations: list[Any] | None = None
|
||||
thinking_blocks: list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None
|
||||
reasoning_content: str | None = None
|
||||
tool_calls: list[ChatCompletionToolCallChunk] = []
|
||||
web_search_results: list[Any] | None = None
|
||||
tool_results: list[Any] | None = None
|
||||
compaction_blocks: list[Any] | None = None
|
||||
for idx, content in enumerate(completion_response["content"]):
|
||||
if content["type"] == "text":
|
||||
text_content += content["text"]
|
||||
|
|
@ -2171,7 +2159,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
if thinking_blocks is not None:
|
||||
reasoning_content = ""
|
||||
for block in thinking_blocks:
|
||||
thinking_content = cast(Optional[str], block.get("thinking"))
|
||||
thinking_content = cast(str | None, block.get("thinking"))
|
||||
if thinking_content is not None:
|
||||
reasoning_content += thinking_content
|
||||
|
||||
|
|
@ -2189,9 +2177,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def calculate_usage(
|
||||
self,
|
||||
usage_object: dict,
|
||||
reasoning_content: Optional[str],
|
||||
completion_response: Optional[dict] = None,
|
||||
speed: Optional[str] = None,
|
||||
reasoning_content: str | None,
|
||||
completion_response: dict | None = None,
|
||||
speed: str | None = None,
|
||||
) -> Usage:
|
||||
# NOTE: Sometimes the usage object has None set explicitly for token counts, meaning .get() & key access returns None, and we need to account for this
|
||||
raw_prompt_tokens = usage_object.get("input_tokens", 0) or 0
|
||||
|
|
@ -2207,14 +2195,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
_usage = usage_object
|
||||
cache_creation_input_tokens: int = 0
|
||||
cache_read_input_tokens: int = 0
|
||||
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
|
||||
web_search_requests: Optional[int] = None
|
||||
tool_search_requests: Optional[int] = None
|
||||
inference_geo: Optional[str] = None
|
||||
cache_creation_token_details: CacheCreationTokenDetails | None = None
|
||||
web_search_requests: int | None = None
|
||||
tool_search_requests: int | None = None
|
||||
inference_geo: str | None = None
|
||||
if "inference_geo" in _usage and _usage["inference_geo"] is not None:
|
||||
inference_geo = _usage["inference_geo"]
|
||||
|
||||
iterations: Optional[List[Any]] = _usage.get("iterations")
|
||||
iterations: list[Any] | None = _usage.get("iterations")
|
||||
if iterations:
|
||||
prompt_tokens = sum(it.get("input_tokens", 0) or 0 for it in iterations)
|
||||
completion_tokens = sum(
|
||||
|
|
@ -2328,9 +2316,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return usage
|
||||
|
||||
def _build_code_by_id_map(
|
||||
self, tool_calls: List[ChatCompletionToolCallChunk]
|
||||
) -> Dict[str, str]:
|
||||
code_by_id: Dict[str, str] = {}
|
||||
self, tool_calls: list[ChatCompletionToolCallChunk]
|
||||
) -> dict[str, str]:
|
||||
code_by_id: dict[str, str] = {}
|
||||
for tc in tool_calls:
|
||||
try:
|
||||
args = json.loads(tc.get("function", {}).get("arguments", "{}"))
|
||||
|
|
@ -2344,10 +2332,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _build_code_interpreter_results(
|
||||
self,
|
||||
tool_results: List[Any],
|
||||
code_by_id: Dict[str, str],
|
||||
container_id: Optional[str],
|
||||
) -> List[OutputCodeInterpreterCall]:
|
||||
tool_results: list[Any],
|
||||
code_by_id: dict[str, str],
|
||||
container_id: str | None,
|
||||
) -> list[OutputCodeInterpreterCall]:
|
||||
code_interpreter_results = []
|
||||
for tr in tool_results:
|
||||
if tr.get("type") != "bash_code_execution_tool_result":
|
||||
|
|
@ -2370,18 +2358,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def _build_provider_specific_fields(
|
||||
self,
|
||||
completion_response: dict,
|
||||
citations: Optional[List[Any]],
|
||||
thinking_blocks: Optional[
|
||||
List[
|
||||
Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]
|
||||
]
|
||||
],
|
||||
web_search_results: Optional[List[Any]],
|
||||
tool_results: Optional[List[Any]],
|
||||
compaction_blocks: Optional[List[Any]],
|
||||
tool_calls: List[ChatCompletionToolCallChunk],
|
||||
) -> Dict[str, Any]:
|
||||
provider_specific_fields: Dict[str, Any] = {
|
||||
citations: list[Any] | None,
|
||||
thinking_blocks: list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None,
|
||||
web_search_results: list[Any] | None,
|
||||
tool_results: list[Any] | None,
|
||||
compaction_blocks: list[Any] | None,
|
||||
tool_calls: list[ChatCompletionToolCallChunk],
|
||||
) -> dict[str, Any]:
|
||||
provider_specific_fields: dict[str, Any] = {
|
||||
"citations": citations,
|
||||
"thinking_blocks": thinking_blocks,
|
||||
}
|
||||
|
|
@ -2422,12 +2406,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
completion_response: dict,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
json_mode: Optional[bool] = None,
|
||||
prefix_prompt: Optional[str] = None,
|
||||
speed: Optional[str] = None,
|
||||
tool_name_reverse_map: Optional[Dict[str, str]] = None,
|
||||
json_mode: bool | None = None,
|
||||
prefix_prompt: str | None = None,
|
||||
speed: str | None = None,
|
||||
tool_name_reverse_map: dict[str, str] | None = None,
|
||||
):
|
||||
_hidden_params: Dict = {}
|
||||
_hidden_params: dict = {}
|
||||
_hidden_params["additional_headers"] = process_anthropic_headers(
|
||||
dict(raw_response.headers)
|
||||
)
|
||||
|
|
@ -2531,7 +2515,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
model_response._hidden_params = _hidden_params
|
||||
return model_response
|
||||
|
||||
def get_prefix_prompt(self, messages: List[AllMessageValues]) -> Optional[str]:
|
||||
def get_prefix_prompt(self, messages: list[AllMessageValues]) -> str | None:
|
||||
"""
|
||||
Get the prefix prompt from the messages.
|
||||
|
||||
|
|
@ -2560,13 +2544,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: LoggingClass,
|
||||
request_data: Dict,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: Dict,
|
||||
request_data: dict,
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -2591,7 +2575,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
prefix_prompt = self.get_prefix_prompt(messages=messages)
|
||||
speed = optional_params.get("speed")
|
||||
tool_name_reverse_map: Optional[Dict[str, str]] = None
|
||||
tool_name_reverse_map: dict[str, str] | None = None
|
||||
if isinstance(litellm_params, dict):
|
||||
_candidate = litellm_params.get(ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY)
|
||||
if isinstance(_candidate, dict):
|
||||
|
|
@ -2610,14 +2594,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _convert_tool_response_to_message(
|
||||
tool_calls: List[ChatCompletionToolCallChunk],
|
||||
) -> Optional[LitellmMessage]:
|
||||
tool_calls: list[ChatCompletionToolCallChunk],
|
||||
) -> LitellmMessage | None:
|
||||
"""
|
||||
In JSON mode, Anthropic API returns JSON schema as a tool call, we need to convert it to a message to follow the OpenAI format
|
||||
|
||||
"""
|
||||
## HANDLE JSON MODE - anthropic returns single function call
|
||||
json_mode_content_str: Optional[str] = tool_calls[0]["function"].get(
|
||||
json_mode_content_str: str | None = tool_calls[0]["function"].get(
|
||||
"arguments"
|
||||
)
|
||||
try:
|
||||
|
|
@ -2640,7 +2624,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return None
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers]
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
return AnthropicError(
|
||||
status_code=status_code,
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
Translates from OpenAI's `/v1/chat/completions` to ModelScope's `/v1/chat/completions`
|
||||
"""
|
||||
|
||||
from typing import Any, Coroutine, Literal, Optional, Tuple, Union, cast, overload
|
||||
from typing import Any, Coroutine, Literal, Union, cast, overload
|
||||
|
||||
from typing_extensions import override
|
||||
|
||||
|
|
@ -63,8 +63,8 @@ class ModelScopeChatConfig(OpenAIGPTConfig):
|
|||
)
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self, api_base: Optional[str], api_key: Optional[str]
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
self, api_base: str | None, api_key: str | None
|
||||
) -> tuple[str | None, str | None]:
|
||||
api_base = (
|
||||
api_base or get_secret_str("MODELSCOPE_API_BASE") or self.DEFAULT_BASE_URL
|
||||
) # type: ignore
|
||||
|
|
@ -74,12 +74,12 @@ class ModelScopeChatConfig(OpenAIGPTConfig):
|
|||
@override
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
stream: bool | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
If api_base is not provided, use the default ModelScope /chat/completions endpoint.
|
||||
|
|
|
|||
|
|
@ -9,13 +9,9 @@ from typing import (
|
|||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Tuple,
|
||||
Type,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
|
@ -135,7 +131,7 @@ class VertexAIBaseConfig:
|
|||
optional_params[mapped_params[param]] = value
|
||||
return optional_params
|
||||
|
||||
def get_eu_regions(self) -> List[str]:
|
||||
def get_eu_regions(self) -> list[str]:
|
||||
"""
|
||||
Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions
|
||||
"""
|
||||
|
|
@ -152,7 +148,7 @@ class VertexAIBaseConfig:
|
|||
"europe-west9",
|
||||
]
|
||||
|
||||
def get_us_regions(self) -> List[str]:
|
||||
def get_us_regions(self) -> list[str]:
|
||||
"""
|
||||
Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions
|
||||
"""
|
||||
|
|
@ -198,29 +194,29 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
Note: Please make sure to modify the default parameters as required for your use case.
|
||||
"""
|
||||
|
||||
temperature: Optional[float] = None
|
||||
max_output_tokens: Optional[int] = None
|
||||
top_p: Optional[float] = None
|
||||
top_k: Optional[int] = None
|
||||
response_mime_type: Optional[str] = None
|
||||
candidate_count: Optional[int] = None
|
||||
stop_sequences: Optional[list] = None
|
||||
frequency_penalty: Optional[float] = None
|
||||
presence_penalty: Optional[float] = None
|
||||
seed: Optional[int] = None
|
||||
temperature: float | None = None
|
||||
max_output_tokens: int | None = None
|
||||
top_p: float | None = None
|
||||
top_k: int | None = None
|
||||
response_mime_type: str | None = None
|
||||
candidate_count: int | None = None
|
||||
stop_sequences: list | None = None
|
||||
frequency_penalty: float | None = None
|
||||
presence_penalty: float | None = None
|
||||
seed: int | None = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
temperature: Optional[float] = None,
|
||||
max_output_tokens: Optional[int] = None,
|
||||
top_p: Optional[float] = None,
|
||||
top_k: Optional[int] = None,
|
||||
response_mime_type: Optional[str] = None,
|
||||
candidate_count: Optional[int] = None,
|
||||
stop_sequences: Optional[list] = None,
|
||||
frequency_penalty: Optional[float] = None,
|
||||
presence_penalty: Optional[float] = None,
|
||||
seed: Optional[int] = None,
|
||||
temperature: float | None = None,
|
||||
max_output_tokens: int | None = None,
|
||||
top_p: float | None = None,
|
||||
top_k: int | None = None,
|
||||
response_mime_type: str | None = None,
|
||||
candidate_count: int | None = None,
|
||||
stop_sequences: list | None = None,
|
||||
frequency_penalty: float | None = None,
|
||||
presence_penalty: float | None = None,
|
||||
seed: int | None = None,
|
||||
) -> None:
|
||||
locals_ = locals().copy()
|
||||
for key, value in locals_.items():
|
||||
|
|
@ -232,8 +228,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return super().get_config()
|
||||
|
||||
def get_json_schema_from_pydantic_object(
|
||||
self, response_format: Optional[Union[Type["BaseModel"], dict]]
|
||||
) -> Optional[dict]:
|
||||
self, response_format: Union[type["BaseModel"], dict] | None
|
||||
) -> dict | None:
|
||||
"""
|
||||
Override to use Pydantic's model_json_schema() instead of OpenAI's
|
||||
to_strict_json_schema().
|
||||
|
|
@ -292,7 +288,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _forward_gemini_function_call_id(
|
||||
model: str, custom_llm_provider: Optional[str] = None
|
||||
model: str, custom_llm_provider: str | None = None
|
||||
) -> bool:
|
||||
"""
|
||||
Whether to include `id` on function_call / function_response parts.
|
||||
|
|
@ -313,7 +309,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return False
|
||||
return True
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
def get_supported_openai_params(self, model: str) -> list[str]:
|
||||
supported_params = [
|
||||
"temperature",
|
||||
"top_p",
|
||||
|
|
@ -350,7 +346,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
def map_tool_choice_values(
|
||||
self, model: str, tool_choice: Union[str, dict]
|
||||
) -> Optional[ToolConfig]:
|
||||
) -> ToolConfig | None:
|
||||
if tool_choice == "none":
|
||||
return ToolConfig(functionCallingConfig=FunctionCallingConfig(mode="NONE"))
|
||||
elif tool_choice == "required":
|
||||
|
|
@ -494,7 +490,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
def _extract_google_maps_retrieval_config(
|
||||
self, google_maps_config: dict
|
||||
) -> Tuple[dict, Optional[dict]]:
|
||||
) -> tuple[dict, dict | None]:
|
||||
"""
|
||||
Extract location configuration from googleMaps tool for Vertex AI toolConfig.
|
||||
|
||||
|
|
@ -534,7 +530,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
return cleaned_config, retrieval_config
|
||||
|
||||
def get_tool_value(self, tool: dict, tool_name: str) -> Optional[dict]:
|
||||
def get_tool_value(self, tool: dict, tool_name: str) -> dict | None:
|
||||
"""
|
||||
Helper function to get tool value handling both camelCase and underscore_case variants
|
||||
|
||||
|
|
@ -561,10 +557,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _resolve_search_tool_conflict(
|
||||
gtool_func_declarations: list,
|
||||
googleSearch: Optional[dict],
|
||||
googleSearchRetrieval: Optional[dict],
|
||||
enterpriseWebSearch: Optional[dict],
|
||||
urlContext: Optional[dict],
|
||||
googleSearch: dict | None,
|
||||
googleSearchRetrieval: dict | None,
|
||||
enterpriseWebSearch: dict | None,
|
||||
urlContext: dict | None,
|
||||
optional_params: dict,
|
||||
) -> tuple:
|
||||
"""
|
||||
|
|
@ -614,7 +610,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
return googleSearch, googleSearchRetrieval, enterpriseWebSearch, urlContext
|
||||
|
||||
def _map_function(self, value: List[dict], optional_params: dict) -> List[Tools]:
|
||||
def _map_function(self, value: list[dict], optional_params: dict) -> list[Tools]:
|
||||
"""
|
||||
Map OpenAI-style tools/functions to Vertex AI format.
|
||||
|
||||
|
|
@ -630,21 +626,21 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
googleMaps tools contain location data
|
||||
"""
|
||||
gtool_func_declarations = []
|
||||
googleSearch: Optional[dict] = None
|
||||
googleSearchRetrieval: Optional[dict] = None
|
||||
enterpriseWebSearch: Optional[dict] = None
|
||||
urlContext: Optional[dict] = None
|
||||
code_execution: Optional[dict] = None
|
||||
googleMaps: Optional[dict] = None
|
||||
google_maps_retrieval_config: Optional[dict] = None
|
||||
computerUse: Optional[dict] = None
|
||||
googleSearch: dict | None = None
|
||||
googleSearchRetrieval: dict | None = None
|
||||
enterpriseWebSearch: dict | None = None
|
||||
urlContext: dict | None = None
|
||||
code_execution: dict | None = None
|
||||
googleMaps: dict | None = None
|
||||
google_maps_retrieval_config: dict | None = None
|
||||
computerUse: dict | None = None
|
||||
# remove 'additionalProperties' from tools
|
||||
value = _remove_additional_properties(value)
|
||||
# remove 'strict' from tools
|
||||
value = _remove_strict_from_schema(value)
|
||||
|
||||
for tool in value:
|
||||
openai_function_object: Optional[ChatCompletionToolParamFunctionChunk] = (
|
||||
openai_function_object: ChatCompletionToolParamFunctionChunk | None = (
|
||||
None
|
||||
)
|
||||
if "function" in tool: # tools list
|
||||
|
|
@ -763,7 +759,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
# Build list of Tool objects - each Tool should contain exactly one type
|
||||
# per Vertex AI API spec: "A Tool object should contain exactly one type of Tool"
|
||||
_tools_list: List[Tools] = []
|
||||
_tools_list: list[Tools] = []
|
||||
|
||||
(
|
||||
googleSearch,
|
||||
|
|
@ -897,7 +893,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _map_reasoning_effort_to_thinking_budget(
|
||||
reasoning_effort: str,
|
||||
model: Optional[str] = None,
|
||||
model: str | None = None,
|
||||
) -> GeminiThinkingConfig:
|
||||
if reasoning_effort == "minimal":
|
||||
# Use model-specific minimum thinking budget or fallback
|
||||
|
|
@ -948,7 +944,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _map_reasoning_effort_to_thinking_level(
|
||||
reasoning_effort: str,
|
||||
model: Optional[str] = None,
|
||||
model: str | None = None,
|
||||
) -> GeminiThinkingConfig:
|
||||
"""
|
||||
Map reasoning_effort to thinking_level for Gemini 3+ models.
|
||||
|
|
@ -996,12 +992,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
raise ValueError(f"Invalid reasoning effort: {reasoning_effort}")
|
||||
|
||||
@staticmethod
|
||||
def _is_thinking_budget_zero(thinking_budget: Optional[int]) -> bool:
|
||||
def _is_thinking_budget_zero(thinking_budget: int | None) -> bool:
|
||||
return thinking_budget is not None and thinking_budget == 0
|
||||
|
||||
@staticmethod
|
||||
def _validate_thinking_config_conflicts(
|
||||
optional_params: Dict,
|
||||
optional_params: dict,
|
||||
param_name: str,
|
||||
param_description: str = "thinking_budget",
|
||||
) -> None:
|
||||
|
|
@ -1022,7 +1018,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _validate_thinking_level_conflicts(
|
||||
optional_params: Dict,
|
||||
optional_params: dict,
|
||||
) -> None:
|
||||
"""
|
||||
Validate that thinking_level and thinking_budget are not both specified.
|
||||
|
|
@ -1042,7 +1038,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _map_thinking_param(
|
||||
thinking_param: AnthropicThinkingParam,
|
||||
model: Optional[str] = None,
|
||||
model: str | None = None,
|
||||
) -> GeminiThinkingConfig:
|
||||
thinking_enabled = thinking_param.get("type") == "enabled"
|
||||
thinking_budget = thinking_param.get("budget_tokens")
|
||||
|
|
@ -1153,8 +1149,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _apply_include_server_side_tool_invocations(
|
||||
non_default_params: Dict,
|
||||
optional_params: Dict,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
) -> None:
|
||||
"""
|
||||
Set include_server_side_tool_invocations before tools are mapped.
|
||||
|
|
@ -1173,11 +1169,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: Dict,
|
||||
optional_params: Dict,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> Dict:
|
||||
) -> dict:
|
||||
self._apply_include_server_side_tool_invocations(
|
||||
non_default_params, optional_params
|
||||
)
|
||||
|
|
@ -1290,7 +1286,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
elif param == "reasoning_effort":
|
||||
# Extract effort value - handle both string and dict formats
|
||||
# Dict format comes from OpenAI Agents SDK: {"effort": "high", "summary": "auto"}
|
||||
effort_value: Optional[str] = None
|
||||
effort_value: str | None = None
|
||||
if isinstance(value, str):
|
||||
effort_value = value
|
||||
elif isinstance(value, dict):
|
||||
|
|
@ -1374,7 +1370,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
optional_params[mapped_params[param]] = value
|
||||
return optional_params
|
||||
|
||||
def get_eu_regions(self) -> List[str]:
|
||||
def get_eu_regions(self) -> list[str]:
|
||||
"""
|
||||
Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions
|
||||
"""
|
||||
|
|
@ -1411,7 +1407,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return model
|
||||
|
||||
@staticmethod
|
||||
def _is_model_gemini_spec_model(model: Optional[str]) -> bool:
|
||||
def _is_model_gemini_spec_model(model: str | None) -> bool:
|
||||
"""
|
||||
Returns true if user is trying to call custom model in `/gemini` request/response format
|
||||
"""
|
||||
|
|
@ -1434,7 +1430,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return model.split("/")[-1]
|
||||
return model
|
||||
|
||||
def get_flagged_finish_reasons(self) -> Dict[str, str]:
|
||||
def get_flagged_finish_reasons(self) -> dict[str, str]:
|
||||
"""
|
||||
Return Dictionary of finish reasons which indicate response was flagged
|
||||
|
||||
|
|
@ -1471,7 +1467,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def get_finish_reason_mapping() -> Dict[str, OpenAIChatCompletionFinishReason]:
|
||||
def get_finish_reason_mapping() -> dict[str, OpenAIChatCompletionFinishReason]:
|
||||
"""
|
||||
Return Dictionary of Gemini/Vertex AI finish reasons and their
|
||||
OpenAI-compatible mappings.
|
||||
|
|
@ -1495,10 +1491,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return exception_string
|
||||
|
||||
def get_assistant_content_message(
|
||||
self, parts: List[HttpxPartType]
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
content_str: Optional[str] = None
|
||||
reasoning_content_str: Optional[str] = None
|
||||
self, parts: list[HttpxPartType]
|
||||
) -> tuple[str | None, str | None]:
|
||||
content_str: str | None = None
|
||||
reasoning_content_str: str | None = None
|
||||
|
||||
for part in parts:
|
||||
_content_str = ""
|
||||
|
|
@ -1540,8 +1536,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return content_str, reasoning_content_str
|
||||
|
||||
def _extract_thinking_blocks_from_parts(
|
||||
self, parts: List[HttpxPartType]
|
||||
) -> List[ChatCompletionThinkingBlock]:
|
||||
self, parts: list[HttpxPartType]
|
||||
) -> list[ChatCompletionThinkingBlock]:
|
||||
"""Extract thinking blocks from parts if present.
|
||||
|
||||
Per Google's docs (https://ai.google.dev/gemini-api/docs/thinking):
|
||||
|
|
@ -1550,7 +1546,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
it does NOT indicate that the content is thinking (a part can have
|
||||
thoughtSignature without thought: true, e.g., function calls)
|
||||
"""
|
||||
thinking_blocks: List[ChatCompletionThinkingBlock] = []
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock] = []
|
||||
for part in parts:
|
||||
if part.get("thought") is True:
|
||||
thinking_text = part.get("text", "")
|
||||
|
|
@ -1565,8 +1561,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return thinking_blocks
|
||||
|
||||
def _extract_thought_signatures_from_parts(
|
||||
self, parts: List[HttpxPartType]
|
||||
) -> Optional[List[str]]:
|
||||
self, parts: list[HttpxPartType]
|
||||
) -> list[str] | None:
|
||||
"""Extract thoughtSignature values from parts.
|
||||
|
||||
Per Google's docs, thoughtSignature is returned for multi-turn context preservation
|
||||
|
|
@ -1576,7 +1572,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
Returns:
|
||||
List of thoughtSignature strings if any are found, None otherwise
|
||||
"""
|
||||
signatures: List[str] = []
|
||||
signatures: list[str] = []
|
||||
for part in parts:
|
||||
signature = part.get("thoughtSignature")
|
||||
if signature is not None:
|
||||
|
|
@ -1585,8 +1581,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _extract_server_side_tool_invocations(
|
||||
parts: List[HttpxPartType],
|
||||
) -> Optional[List[Dict[str, Any]]]:
|
||||
parts: list[HttpxPartType],
|
||||
) -> list[dict[str, Any]] | None:
|
||||
"""Extract server-side tool invocations (toolCall/toolResponse) from parts.
|
||||
|
||||
These are returned by Gemini when context circulation is enabled
|
||||
|
|
@ -1597,15 +1593,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
Returns:
|
||||
List of server-side invocation dicts if any found, None otherwise.
|
||||
"""
|
||||
invocations: List[Dict[str, Any]] = []
|
||||
invocations: list[dict[str, Any]] = []
|
||||
# Index toolCalls by id so we can pair them with responses
|
||||
tool_calls_by_id: Dict[str, Dict[str, Any]] = {}
|
||||
tool_responses_by_id: Dict[str, Dict[str, Any]] = {}
|
||||
tool_calls_by_id: dict[str, dict[str, Any]] = {}
|
||||
tool_responses_by_id: dict[str, dict[str, Any]] = {}
|
||||
|
||||
for part in parts:
|
||||
if "toolCall" in part:
|
||||
tc = part["toolCall"]
|
||||
entry: Dict[str, Any] = {
|
||||
entry: dict[str, Any] = {
|
||||
"tool_type": tc.get("toolType"),
|
||||
"id": tc.get("id"),
|
||||
"args": tc.get("args"),
|
||||
|
|
@ -1645,10 +1641,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return invocations if invocations else None
|
||||
|
||||
def _extract_image_response_from_parts(
|
||||
self, parts: List[HttpxPartType]
|
||||
) -> Optional[List[ImageURLListItem]]:
|
||||
self, parts: list[HttpxPartType]
|
||||
) -> list[ImageURLListItem] | None:
|
||||
"""Extract image response from parts if present"""
|
||||
images: List[ImageURLListItem] = []
|
||||
images: list[ImageURLListItem] = []
|
||||
for part in parts:
|
||||
if "inlineData" in part:
|
||||
inline_data = part.get("inlineData", {})
|
||||
|
|
@ -1667,8 +1663,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return images
|
||||
|
||||
def _extract_audio_response_from_parts(
|
||||
self, parts: List[HttpxPartType]
|
||||
) -> Optional[ChatCompletionAudioResponse]:
|
||||
self, parts: list[HttpxPartType]
|
||||
) -> ChatCompletionAudioResponse | None:
|
||||
"""Extract audio response from parts if present"""
|
||||
for part in parts:
|
||||
if "text" in part:
|
||||
|
|
@ -1710,16 +1706,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _transform_parts(
|
||||
parts: List[HttpxPartType],
|
||||
parts: list[HttpxPartType],
|
||||
cumulative_tool_call_idx: int,
|
||||
is_function_call: Optional[bool],
|
||||
) -> Tuple[
|
||||
Optional[ChatCompletionToolCallFunctionChunk],
|
||||
Optional[List[ChatCompletionToolCallChunk]],
|
||||
is_function_call: bool | None,
|
||||
) -> tuple[
|
||||
ChatCompletionToolCallFunctionChunk | None,
|
||||
list[ChatCompletionToolCallChunk] | None,
|
||||
int,
|
||||
]:
|
||||
function: Optional[ChatCompletionToolCallFunctionChunk] = None
|
||||
_tools: List[ChatCompletionToolCallChunk] = []
|
||||
function: ChatCompletionToolCallFunctionChunk | None = None
|
||||
_tools: list[ChatCompletionToolCallChunk] = []
|
||||
for part in parts:
|
||||
if "functionCall" in part:
|
||||
_function_chunk: ChatCompletionToolCallFunctionChunk = {
|
||||
|
|
@ -1736,7 +1732,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
gemini_call_id = part["functionCall"].get("id")
|
||||
|
||||
if is_function_call is True:
|
||||
function_dict: Dict[str, Any] = dict(_function_chunk)
|
||||
function_dict: dict[str, Any] = dict(_function_chunk)
|
||||
if thought_signature:
|
||||
if "provider_specific_fields" not in function_dict:
|
||||
function_dict["provider_specific_fields"] = {}
|
||||
|
|
@ -1769,22 +1765,22 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
_tools.append(_tool_response_chunk)
|
||||
cumulative_tool_call_idx += 1
|
||||
if len(_tools) == 0:
|
||||
tools: Optional[List[ChatCompletionToolCallChunk]] = None
|
||||
tools: list[ChatCompletionToolCallChunk] | None = None
|
||||
else:
|
||||
tools = _tools
|
||||
return function, tools, cumulative_tool_call_idx
|
||||
|
||||
@staticmethod
|
||||
def _transform_logprobs(
|
||||
logprobs_result: Optional[LogprobsResult],
|
||||
) -> Optional[ChoiceLogprobs]:
|
||||
logprobs_result: LogprobsResult | None,
|
||||
) -> ChoiceLogprobs | None:
|
||||
if logprobs_result is None:
|
||||
return None
|
||||
if "chosenCandidates" not in logprobs_result:
|
||||
return None
|
||||
logprobs_list: List[ChatCompletionTokenLogprob] = []
|
||||
logprobs_list: list[ChatCompletionTokenLogprob] = []
|
||||
for index, candidate in enumerate(logprobs_result["chosenCandidates"]):
|
||||
top_logprobs: List[TopLogprob] = []
|
||||
top_logprobs: list[TopLogprob] = []
|
||||
if "topCandidates" in logprobs_result and index < len(
|
||||
logprobs_result["topCandidates"]
|
||||
):
|
||||
|
|
@ -1914,16 +1910,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
raise ValueError(
|
||||
f"usageMetadata not found in completion_response. Got={completion_response}"
|
||||
)
|
||||
cached_tokens: Optional[int] = None
|
||||
cached_tokens: int | None = None
|
||||
# Separate variables for prompt tokens by modality
|
||||
prompt_audio_tokens: Optional[int] = None
|
||||
prompt_image_tokens: Optional[int] = None
|
||||
prompt_text_tokens: Optional[int] = None
|
||||
prompt_video_tokens: Optional[int] = None
|
||||
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None
|
||||
reasoning_tokens: Optional[int] = None
|
||||
response_tokens: Optional[int] = None
|
||||
response_tokens_details: Optional[CompletionTokensDetailsWrapper] = None
|
||||
prompt_audio_tokens: int | None = None
|
||||
prompt_image_tokens: int | None = None
|
||||
prompt_text_tokens: int | None = None
|
||||
prompt_video_tokens: int | None = None
|
||||
prompt_tokens_details: PromptTokensDetailsWrapper | None = None
|
||||
reasoning_tokens: int | None = None
|
||||
response_tokens: int | None = None
|
||||
response_tokens_details: CompletionTokensDetailsWrapper | None = None
|
||||
usage_metadata = completion_response["usageMetadata"]
|
||||
|
||||
def _get_token_count(detail: Mapping[str, Any]) -> int:
|
||||
|
|
@ -2021,10 +2017,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
## Parse cacheTokensDetails (breakdown of cached tokens by modality)
|
||||
## When explicit caching is used, Gemini provides this field to show which modalities were cached
|
||||
cached_text_tokens: Optional[int] = None
|
||||
cached_audio_tokens: Optional[int] = None
|
||||
cached_image_tokens: Optional[int] = None
|
||||
cached_video_tokens: Optional[int] = None
|
||||
cached_text_tokens: int | None = None
|
||||
cached_audio_tokens: int | None = None
|
||||
cached_image_tokens: int | None = None
|
||||
cached_video_tokens: int | None = None
|
||||
|
||||
if "cacheTokensDetails" in usage_metadata:
|
||||
for detail in usage_metadata["cacheTokensDetails"]:
|
||||
|
|
@ -2102,8 +2098,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _check_finish_reason(
|
||||
chat_completion_message: Optional[ChatCompletionResponseMessage],
|
||||
finish_reason: Optional[str],
|
||||
chat_completion_message: ChatCompletionResponseMessage | None,
|
||||
finish_reason: str | None,
|
||||
) -> OpenAIChatCompletionFinishReason:
|
||||
from litellm.litellm_core_utils.core_helpers import map_finish_reason
|
||||
|
||||
|
|
@ -2119,7 +2115,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _check_prompt_level_content_filter(
|
||||
processed_chunk: GenerateContentResponseBody,
|
||||
response_id: Optional[str],
|
||||
response_id: str | None,
|
||||
) -> Optional["ModelResponseStream"]:
|
||||
"""
|
||||
Check if prompt is blocked due to content filtering at the prompt level.
|
||||
|
|
@ -2163,8 +2159,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
def _calculate_web_search_requests(grounding_metadata: List[dict]) -> Optional[int]:
|
||||
web_search_requests: Optional[int] = None
|
||||
def _calculate_web_search_requests(grounding_metadata: list[dict]) -> int | None:
|
||||
web_search_requests: int | None = None
|
||||
|
||||
if (
|
||||
grounding_metadata
|
||||
|
|
@ -2184,10 +2180,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
chat_completion_message: ChatCompletionResponseMessage,
|
||||
candidate: Candidates,
|
||||
idx: int,
|
||||
tools: Optional[List[ChatCompletionToolCallChunk]],
|
||||
functions: Optional[ChatCompletionToolCallFunctionChunk],
|
||||
chat_completion_logprobs: Optional[ChoiceLogprobs],
|
||||
image_response: Optional[List[ImageURLListItem]],
|
||||
tools: list[ChatCompletionToolCallChunk] | None,
|
||||
functions: ChatCompletionToolCallFunctionChunk | None,
|
||||
chat_completion_logprobs: ChoiceLogprobs | None,
|
||||
image_response: list[ImageURLListItem] | None,
|
||||
) -> StreamingChoices:
|
||||
"""
|
||||
Helper method to create a streaming choice object for Vertex AI
|
||||
|
|
@ -2219,7 +2215,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _extract_candidate_metadata(
|
||||
candidate: Candidates,
|
||||
) -> Tuple[List[dict], List[dict], List, List]:
|
||||
) -> tuple[list[dict], list[dict], list, list]:
|
||||
"""
|
||||
Extract metadata from a single candidate response.
|
||||
|
||||
|
|
@ -2229,10 +2225,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
safety_ratings: List
|
||||
citation_metadata: List
|
||||
"""
|
||||
grounding_metadata: List[dict] = []
|
||||
url_context_metadata: List[dict] = []
|
||||
safety_ratings: List = []
|
||||
citation_metadata: List = []
|
||||
grounding_metadata: list[dict] = []
|
||||
url_context_metadata: list[dict] = []
|
||||
safety_ratings: list = []
|
||||
citation_metadata: list = []
|
||||
|
||||
if "groundingMetadata" in candidate:
|
||||
if isinstance(candidate["groundingMetadata"], list):
|
||||
|
|
@ -2277,10 +2273,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
@staticmethod
|
||||
def _set_stream_metadata_on_response(
|
||||
model_response: Any,
|
||||
grounding_metadata: List[dict],
|
||||
url_context_metadata: List[dict],
|
||||
safety_ratings: List[dict],
|
||||
citation_metadata: List[dict],
|
||||
grounding_metadata: list[dict],
|
||||
url_context_metadata: list[dict],
|
||||
safety_ratings: list[dict],
|
||||
citation_metadata: list[dict],
|
||||
) -> None:
|
||||
setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore
|
||||
if grounding_metadata:
|
||||
|
|
@ -2306,10 +2302,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
def apply_assembled_streaming_response_metadata(
|
||||
self,
|
||||
response: ModelResponse,
|
||||
chunks: List[Any],
|
||||
chunks: list[Any],
|
||||
) -> None:
|
||||
for field_name in VERTEX_AI_PROVIDER_METADATA_FIELDS:
|
||||
merged: List[Any] = []
|
||||
merged: list[Any] = []
|
||||
for chunk in chunks:
|
||||
value = VertexGeminiConfig._get_stream_chunk_attr(chunk, field_name)
|
||||
if not value:
|
||||
|
|
@ -2324,14 +2320,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _convert_grounding_metadata_to_annotations(
|
||||
grounding_metadata: List[dict],
|
||||
content_text: Optional[str],
|
||||
) -> List[ChatCompletionAnnotation]:
|
||||
grounding_metadata: list[dict],
|
||||
content_text: str | None,
|
||||
) -> list[ChatCompletionAnnotation]:
|
||||
"""
|
||||
Convert Vertex AI grounding metadata to OpenAI-style annotations.
|
||||
"""
|
||||
|
||||
annotations: List[ChatCompletionAnnotation] = []
|
||||
annotations: list[ChatCompletionAnnotation] = []
|
||||
|
||||
for metadata in grounding_metadata:
|
||||
# Extract groundingSupports - these map text segments to sources
|
||||
|
|
@ -2339,7 +2335,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
grounding_chunks = metadata.get("groundingChunks", [])
|
||||
|
||||
# Build a map of chunk indices to web URIs
|
||||
chunk_to_uri_map: Dict[int, Dict[str, str]] = {}
|
||||
chunk_to_uri_map: dict[int, dict[str, str]] = {}
|
||||
for idx, chunk in enumerate(grounding_chunks):
|
||||
if "web" in chunk:
|
||||
web_data = chunk["web"]
|
||||
|
|
@ -2379,11 +2375,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _process_candidates(
|
||||
_candidates: List[Candidates],
|
||||
_candidates: list[Candidates],
|
||||
model_response: Union[ModelResponse, "ModelResponseStream"],
|
||||
standard_optional_params: dict,
|
||||
cumulative_tool_call_index: int = 0,
|
||||
) -> Tuple[List[dict], List[dict], List, List, int]:
|
||||
) -> tuple[list[dict], list[dict], list, list, int]:
|
||||
"""
|
||||
Helper method to process candidates and extract metadata
|
||||
|
||||
|
|
@ -2399,19 +2395,19 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
from litellm.types.utils import ModelResponseStream
|
||||
|
||||
grounding_metadata: List[dict] = []
|
||||
url_context_metadata: List[dict] = []
|
||||
image_response: Optional[List[ImageURLListItem]] = None
|
||||
safety_ratings: List = []
|
||||
citation_metadata: List = []
|
||||
grounding_metadata: list[dict] = []
|
||||
url_context_metadata: list[dict] = []
|
||||
image_response: list[ImageURLListItem] | None = None
|
||||
safety_ratings: list = []
|
||||
citation_metadata: list = []
|
||||
chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"}
|
||||
chat_completion_logprobs: Optional[ChoiceLogprobs] = None
|
||||
tools: Optional[List[ChatCompletionToolCallChunk]] = []
|
||||
functions: Optional[ChatCompletionToolCallFunctionChunk] = None
|
||||
thinking_blocks: Optional[List[ChatCompletionThinkingBlock]] = None
|
||||
reasoning_content: Optional[str] = None
|
||||
thought_signatures: Optional[Any] = None
|
||||
server_side_tool_invocations: Optional[List[Dict[str, Any]]] = None
|
||||
chat_completion_logprobs: ChoiceLogprobs | None = None
|
||||
tools: list[ChatCompletionToolCallChunk] | None = []
|
||||
functions: ChatCompletionToolCallFunctionChunk | None = None
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock] | None = None
|
||||
reasoning_content: str | None = None
|
||||
thought_signatures: Any | None = None
|
||||
server_side_tool_invocations: list[dict[str, Any]] | None = None
|
||||
|
||||
for idx, candidate in enumerate(_candidates):
|
||||
if "content" not in candidate:
|
||||
|
|
@ -2470,13 +2466,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
|
||||
if audio_response is not None:
|
||||
cast(Dict[str, Any], chat_completion_message)[
|
||||
cast(dict[str, Any], chat_completion_message)[
|
||||
"audio"
|
||||
] = audio_response
|
||||
chat_completion_message["content"] = None # OpenAI spec
|
||||
if image_response is not None:
|
||||
# Handle image response - combine with text content into structured format
|
||||
cast(Dict[str, Any], chat_completion_message)[
|
||||
cast(dict[str, Any], chat_completion_message)[
|
||||
"images"
|
||||
] = image_response
|
||||
if content is not None:
|
||||
|
|
@ -2583,13 +2579,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: LoggingClass,
|
||||
request_data: Dict,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: Dict,
|
||||
litellm_params: Dict,
|
||||
request_data: dict,
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -2664,11 +2660,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
response_id = completion_response.get("responseId")
|
||||
if response_id:
|
||||
model_response.id = response_id
|
||||
url_context_metadata: List[dict] = []
|
||||
url_context_metadata: list[dict] = []
|
||||
try:
|
||||
grounding_metadata: List[dict] = []
|
||||
safety_ratings: List[dict] = []
|
||||
citation_metadata: List[dict] = []
|
||||
grounding_metadata: list[dict] = []
|
||||
safety_ratings: list[dict] = []
|
||||
citation_metadata: list[dict] = []
|
||||
if _candidates:
|
||||
(
|
||||
grounding_metadata,
|
||||
|
|
@ -2750,10 +2746,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
def _transform_messages(
|
||||
self,
|
||||
messages: List[AllMessageValues],
|
||||
model: Optional[str] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> List[ContentType]:
|
||||
messages: list[AllMessageValues],
|
||||
model: str | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
) -> list[ContentType]:
|
||||
return _gemini_convert_messages_with_history(
|
||||
messages=messages,
|
||||
model=model,
|
||||
|
|
@ -2762,7 +2758,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers]
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
return VertexAIError(
|
||||
message=error_message, status_code=status_code, headers=headers
|
||||
|
|
@ -2771,25 +2767,25 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: Dict,
|
||||
litellm_params: Dict,
|
||||
headers: Dict,
|
||||
) -> Dict:
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
raise NotImplementedError(
|
||||
"Vertex AI has a custom implementation of transform_request. Needs sync + async."
|
||||
)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Optional[Dict],
|
||||
headers: dict | None,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: Dict,
|
||||
litellm_params: Dict,
|
||||
api_key: Optional[Union[str, Dict]] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> Dict:
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Union[str, dict] | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict:
|
||||
default_headers = {
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
|
@ -2804,8 +2800,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
|
||||
async def make_call(
|
||||
client: Optional[AsyncHTTPHandler], # module-level client
|
||||
gemini_client: Optional[AsyncHTTPHandler], # if passed by user
|
||||
client: AsyncHTTPHandler | None, # module-level client
|
||||
gemini_client: AsyncHTTPHandler | None, # if passed by user
|
||||
api_base: str,
|
||||
headers: dict,
|
||||
data: str,
|
||||
|
|
@ -2857,8 +2853,8 @@ async def make_call(
|
|||
|
||||
|
||||
def make_sync_call(
|
||||
client: Optional[HTTPHandler], # module-level client
|
||||
gemini_client: Optional[HTTPHandler], # if passed by user
|
||||
client: HTTPHandler | None, # module-level client
|
||||
gemini_client: HTTPHandler | None, # if passed by user
|
||||
api_base: str,
|
||||
headers: dict,
|
||||
data: str,
|
||||
|
|
@ -2914,20 +2910,20 @@ class VertexLLM(VertexBase):
|
|||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
data: dict,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
timeout: Union[float, httpx.Timeout] | None,
|
||||
encoding,
|
||||
logging_obj,
|
||||
stream,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
logger_fn=None,
|
||||
api_base: Optional[str] = None,
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
vertex_project: Optional[str] = None,
|
||||
vertex_location: Optional[str] = None,
|
||||
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None,
|
||||
gemini_api_key: Optional[str] = None,
|
||||
extra_headers: Optional[dict] = None,
|
||||
api_base: str | None = None,
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
vertex_project: str | None = None,
|
||||
vertex_location: str | None = None,
|
||||
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None,
|
||||
gemini_api_key: str | None = None,
|
||||
extra_headers: dict | None = None,
|
||||
) -> CustomStreamWrapper:
|
||||
should_use_v1beta1_features = self.is_using_v1beta1_features(
|
||||
optional_params=optional_params
|
||||
|
|
@ -3016,20 +3012,20 @@ class VertexLLM(VertexBase):
|
|||
custom_llm_provider: Literal[
|
||||
"vertex_ai", "vertex_ai_beta", "gemini"
|
||||
], # if it's vertex_ai or gemini (google ai studio)
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
timeout: Union[float, httpx.Timeout] | None,
|
||||
encoding,
|
||||
logging_obj,
|
||||
stream,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
logger_fn=None,
|
||||
api_base: Optional[str] = None,
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
vertex_project: Optional[str] = None,
|
||||
vertex_location: Optional[str] = None,
|
||||
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None,
|
||||
gemini_api_key: Optional[str] = None,
|
||||
extra_headers: Optional[dict] = None,
|
||||
api_base: str | None = None,
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
vertex_project: str | None = None,
|
||||
vertex_location: str | None = None,
|
||||
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None,
|
||||
gemini_api_key: str | None = None,
|
||||
extra_headers: dict | None = None,
|
||||
) -> Union[ModelResponse, CustomStreamWrapper]:
|
||||
should_use_v1beta1_features = self.is_using_v1beta1_features(
|
||||
optional_params=optional_params
|
||||
|
|
@ -3142,18 +3138,18 @@ class VertexLLM(VertexBase):
|
|||
logging_obj,
|
||||
optional_params: dict,
|
||||
acompletion: bool,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
gemini_api_key: Optional[str],
|
||||
timeout: Union[float, httpx.Timeout] | None,
|
||||
vertex_project: str | None,
|
||||
vertex_location: str | None,
|
||||
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None,
|
||||
gemini_api_key: str | None,
|
||||
litellm_params: dict,
|
||||
logger_fn=None,
|
||||
extra_headers: Optional[dict] = None,
|
||||
client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None,
|
||||
api_base: Optional[str] = None,
|
||||
extra_headers: dict | None = None,
|
||||
client: Union[AsyncHTTPHandler, HTTPHandler] | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> Union[ModelResponse, CustomStreamWrapper]:
|
||||
stream: Optional[bool] = optional_params.pop("stream", None) # type: ignore
|
||||
stream: bool | None = optional_params.pop("stream", None) # type: ignore
|
||||
|
||||
transform_request_params = {
|
||||
"gemini_api_key": gemini_api_key,
|
||||
|
|
@ -3347,7 +3343,7 @@ class ModelResponseIterator:
|
|||
streaming_response,
|
||||
sync_stream: bool,
|
||||
logging_obj: LoggingClass,
|
||||
response_headers: Optional[Dict[str, str]] = None,
|
||||
response_headers: dict[str, str] | None = None,
|
||||
):
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
check_is_function_call,
|
||||
|
|
@ -3390,9 +3386,9 @@ class ModelResponseIterator:
|
|||
|
||||
def _apply_stream_candidates(
|
||||
self,
|
||||
_candidates: List[Candidates],
|
||||
_candidates: list[Candidates],
|
||||
model_response: Any,
|
||||
) -> Tuple[List[dict], List[dict], List[dict], List[dict]]:
|
||||
) -> tuple[list[dict], list[dict], list[dict], list[dict]]:
|
||||
(
|
||||
grounding_metadata,
|
||||
url_context_metadata,
|
||||
|
|
@ -3477,8 +3473,8 @@ class ModelResponseIterator:
|
|||
self,
|
||||
processed_chunk: Any,
|
||||
model_response: Any,
|
||||
grounding_metadata: List[dict],
|
||||
) -> Optional[Usage]:
|
||||
grounding_metadata: list[dict],
|
||||
) -> Usage | None:
|
||||
if "usageMetadata" not in processed_chunk:
|
||||
return None
|
||||
|
||||
|
|
@ -3531,12 +3527,12 @@ class ModelResponseIterator:
|
|||
if blocked_response is not None:
|
||||
model_response = blocked_response
|
||||
|
||||
grounding_metadata: List[dict] = []
|
||||
url_context_metadata: List[dict] = []
|
||||
safety_ratings: List[dict] = []
|
||||
citation_metadata: List[dict] = []
|
||||
grounding_metadata: list[dict] = []
|
||||
url_context_metadata: list[dict] = []
|
||||
safety_ratings: list[dict] = []
|
||||
citation_metadata: list[dict] = []
|
||||
|
||||
_candidates: Optional[List[Candidates]] = processed_chunk.get("candidates")
|
||||
_candidates: list[Candidates] | None = processed_chunk.get("candidates")
|
||||
if _candidates:
|
||||
(
|
||||
grounding_metadata,
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@
|
|||
# +-------------------------------------------------------------+
|
||||
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional
|
||||
from typing import TYPE_CHECKING, Any, Literal, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -68,11 +68,11 @@ class DeepKeepGuardrail(CustomGuardrail):
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
firewall_id: Optional[str] = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
firewall_id: str | None = None,
|
||||
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
self.async_handler = get_async_httpx_client(
|
||||
|
|
@ -114,7 +114,7 @@ class DeepKeepGuardrail(CustomGuardrail):
|
|||
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
|
||||
unreachable_fallback
|
||||
)
|
||||
self.extra_headers: Dict[str, str] = extra_headers or {}
|
||||
self.extra_headers: dict[str, str] = extra_headers or {}
|
||||
|
||||
# Set supported event hooks
|
||||
if "supported_event_hooks" not in kwargs:
|
||||
|
|
@ -133,7 +133,7 @@ class DeepKeepGuardrail(CustomGuardrail):
|
|||
self.firewall_id,
|
||||
)
|
||||
|
||||
def _extract_user_api_key_metadata(self, request_data: dict) -> Dict[str, Any]:
|
||||
def _extract_user_api_key_metadata(self, request_data: dict) -> dict[str, Any]:
|
||||
"""
|
||||
Extract user API key metadata from request_data for the DeepKeep API.
|
||||
|
||||
|
|
@ -143,7 +143,7 @@ class DeepKeepGuardrail(CustomGuardrail):
|
|||
Returns:
|
||||
Dictionary with user API key metadata fields.
|
||||
"""
|
||||
result_metadata: Dict[str, Any] = {}
|
||||
result_metadata: dict[str, Any] = {}
|
||||
|
||||
litellm_metadata = request_data.get("litellm_metadata", {})
|
||||
top_level_metadata = request_data.get("metadata", {})
|
||||
|
|
@ -177,9 +177,9 @@ class DeepKeepGuardrail(CustomGuardrail):
|
|||
|
||||
return result_metadata
|
||||
|
||||
def _build_request_headers(self) -> Dict[str, str]:
|
||||
def _build_request_headers(self) -> dict[str, str]:
|
||||
"""Build HTTP headers for the DeepKeep API request."""
|
||||
headers: Dict[str, str] = {
|
||||
headers: dict[str, str] = {
|
||||
"Content-Type": "application/json",
|
||||
"X-API-Key": self.deepkeep_api_key,
|
||||
}
|
||||
|
|
@ -194,7 +194,7 @@ class DeepKeepGuardrail(CustomGuardrail):
|
|||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"],
|
||||
error: Exception,
|
||||
http_status_code: Optional[int] = None,
|
||||
http_status_code: int | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""Allow the request to proceed when the guardrail is unreachable (fail-open mode)."""
|
||||
status_suffix = (
|
||||
|
|
@ -241,12 +241,12 @@ class DeepKeepGuardrail(CustomGuardrail):
|
|||
@staticmethod
|
||||
def _build_return_inputs(
|
||||
*,
|
||||
response_json: Dict[str, Any],
|
||||
response_json: dict[str, Any],
|
||||
texts: list,
|
||||
images: Optional[Any],
|
||||
tools: Optional[Any],
|
||||
tool_calls: Optional[Any],
|
||||
structured_messages: Optional[Any],
|
||||
images: Any | None,
|
||||
tools: Any | None,
|
||||
tool_calls: Any | None,
|
||||
structured_messages: Any | None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""Merge original inputs with any guardrail-modified values from the API response."""
|
||||
return_inputs = GenericGuardrailAPIInputs(texts=texts)
|
||||
|
|
@ -311,7 +311,7 @@ class DeepKeepGuardrail(CustomGuardrail):
|
|||
request_body = request_data.get("body") or {}
|
||||
|
||||
# Merge additional provider-specific params from config and dynamic params
|
||||
additional_params: Dict[str, Any] = {"firewall_id": self.firewall_id}
|
||||
additional_params: dict[str, Any] = {"firewall_id": self.firewall_id}
|
||||
dynamic_params = self.get_guardrail_dynamic_request_body_params(request_body)
|
||||
if dynamic_params:
|
||||
additional_params.update(
|
||||
|
|
@ -322,7 +322,7 @@ class DeepKeepGuardrail(CustomGuardrail):
|
|||
user_metadata = self._extract_user_api_key_metadata(request_data)
|
||||
|
||||
# Build request payload
|
||||
guardrail_request: Dict[str, Any] = {
|
||||
guardrail_request: dict[str, Any] = {
|
||||
"litellm_call_id": (logging_obj.litellm_call_id if logging_obj else None),
|
||||
"litellm_trace_id": (logging_obj.litellm_trace_id if logging_obj else None),
|
||||
"texts": texts,
|
||||
|
|
@ -397,7 +397,7 @@ class DeepKeepGuardrail(CustomGuardrail):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> Optional[type]:
|
||||
def get_config_model() -> type | None:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.deepkeep import (
|
||||
DeepKeepGuardrailConfigModel,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
import importlib
|
||||
import os
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Literal, Optional, Set, Type, cast
|
||||
from typing import Any, Literal, Optional, cast
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
|
|
@ -54,7 +54,7 @@ guardrail_initializer_registry = {
|
|||
SupportedGuardrailIntegrations.LLM_AS_A_JUDGE.value: initialize_llm_as_a_judge,
|
||||
}
|
||||
|
||||
guardrail_class_registry: Dict[str, Type[CustomGuardrail]] = {
|
||||
guardrail_class_registry: dict[str, type[CustomGuardrail]] = {
|
||||
SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail
|
||||
}
|
||||
|
||||
|
|
@ -224,7 +224,7 @@ class GuardrailRegistry:
|
|||
############################################################
|
||||
def get_initialized_guardrail_callback(
|
||||
self, guardrail_name: str
|
||||
) -> Optional[CustomGuardrail]:
|
||||
) -> CustomGuardrail | None:
|
||||
"""
|
||||
Returns the initialized guardrail callback for a given guardrail name
|
||||
"""
|
||||
|
|
@ -334,7 +334,7 @@ class GuardrailRegistry:
|
|||
@staticmethod
|
||||
async def get_all_guardrails_from_db(
|
||||
prisma_client: PrismaClient,
|
||||
) -> List[Guardrail]:
|
||||
) -> list[Guardrail]:
|
||||
"""
|
||||
Get all active guardrails from the database.
|
||||
Only rows with status == "active" are returned (pending_review and rejected are excluded).
|
||||
|
|
@ -347,7 +347,7 @@ class GuardrailRegistry:
|
|||
order={"created_at": "desc"},
|
||||
)
|
||||
|
||||
guardrails: List[Guardrail] = []
|
||||
guardrails: list[Guardrail] = []
|
||||
for guardrail in guardrails_from_db:
|
||||
guardrails.append(Guardrail(**(dict(guardrail)))) # type: ignore
|
||||
|
||||
|
|
@ -357,7 +357,7 @@ class GuardrailRegistry:
|
|||
|
||||
async def get_guardrail_by_id_from_db(
|
||||
self, guardrail_id: str, prisma_client: PrismaClient
|
||||
) -> Optional[Guardrail]:
|
||||
) -> Guardrail | None:
|
||||
"""
|
||||
Get a guardrail by its ID from the database
|
||||
"""
|
||||
|
|
@ -375,7 +375,7 @@ class GuardrailRegistry:
|
|||
|
||||
async def get_guardrail_by_name_from_db(
|
||||
self, guardrail_name: str, prisma_client: PrismaClient
|
||||
) -> Optional[Guardrail]:
|
||||
) -> Guardrail | None:
|
||||
"""
|
||||
Get a guardrail by its name from the database
|
||||
"""
|
||||
|
|
@ -398,17 +398,17 @@ class InMemoryGuardrailHandler:
|
|||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.IN_MEMORY_GUARDRAILS: Dict[str, Guardrail] = {}
|
||||
self.IN_MEMORY_GUARDRAILS: dict[str, Guardrail] = {}
|
||||
"""
|
||||
Guardrail id to Guardrail object mapping
|
||||
"""
|
||||
|
||||
self.guardrail_id_to_custom_guardrail: Dict[str, Optional[CustomGuardrail]] = {}
|
||||
self.guardrail_id_to_custom_guardrail: dict[str, CustomGuardrail | None] = {}
|
||||
"""
|
||||
Guardrail id to CustomGuardrail object mapping
|
||||
"""
|
||||
|
||||
self._sources: Dict[str, Literal["db", "config"]] = {}
|
||||
self._sources: dict[str, Literal["db", "config"]] = {}
|
||||
"""
|
||||
Guardrail id to provenance marker. "db" entries are reconciled against
|
||||
the DB on each polling tick; "config" entries are owned by proxy_config.yaml
|
||||
|
|
@ -418,10 +418,10 @@ class InMemoryGuardrailHandler:
|
|||
def initialize_guardrail(
|
||||
self,
|
||||
guardrail: Guardrail,
|
||||
config_file_path: Optional[str] = None,
|
||||
config_file_path: str | None = None,
|
||||
llm_router: Optional["Router"] = None,
|
||||
source: Literal["db", "config"] = "config",
|
||||
) -> Optional[Guardrail]:
|
||||
) -> Guardrail | None:
|
||||
"""
|
||||
Initialize a guardrail from a dictionary and add it to the litellm callback manager
|
||||
|
||||
|
|
@ -439,7 +439,7 @@ class InMemoryGuardrailHandler:
|
|||
self._sources[guardrail_id] = source
|
||||
return self.IN_MEMORY_GUARDRAILS[guardrail_id]
|
||||
|
||||
custom_guardrail_callback: Optional[CustomGuardrail] = None
|
||||
custom_guardrail_callback: CustomGuardrail | None = None
|
||||
litellm_params_data = guardrail["litellm_params"]
|
||||
verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data)
|
||||
|
||||
|
|
@ -520,11 +520,11 @@ class InMemoryGuardrailHandler:
|
|||
|
||||
def initialize_custom_guardrail(
|
||||
self,
|
||||
guardrail: Dict,
|
||||
guardrail: dict,
|
||||
guardrail_type: str,
|
||||
litellm_params: LitellmParams,
|
||||
config_file_path: Optional[str] = None,
|
||||
) -> Optional[CustomGuardrail]:
|
||||
config_file_path: str | None = None,
|
||||
) -> CustomGuardrail | None:
|
||||
"""
|
||||
Initialize a Custom Guardrail from a python file or module path
|
||||
|
||||
|
|
@ -623,25 +623,25 @@ class InMemoryGuardrailHandler:
|
|||
custom_guardrail_callback
|
||||
)
|
||||
|
||||
def list_in_memory_guardrails(self) -> List[Guardrail]:
|
||||
def list_in_memory_guardrails(self) -> list[Guardrail]:
|
||||
"""
|
||||
List all guardrails in memory
|
||||
"""
|
||||
return list(self.IN_MEMORY_GUARDRAILS.values())
|
||||
|
||||
def get_guardrail_by_id(self, guardrail_id: str) -> Optional[Guardrail]:
|
||||
def get_guardrail_by_id(self, guardrail_id: str) -> Guardrail | None:
|
||||
"""
|
||||
Get a guardrail by its ID from memory
|
||||
"""
|
||||
return self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
|
||||
|
||||
def get_source(self, guardrail_id: str) -> Optional[Literal["db", "config"]]:
|
||||
def get_source(self, guardrail_id: str) -> Literal["db", "config"] | None:
|
||||
"""
|
||||
Return the provenance of an in-memory guardrail.
|
||||
"""
|
||||
return self._sources.get(guardrail_id)
|
||||
|
||||
def reconcile_db_guardrails(self, db_guardrail_ids: Set[str]) -> List[str]:
|
||||
def reconcile_db_guardrails(self, db_guardrail_ids: set[str]) -> list[str]:
|
||||
"""
|
||||
Drop in-memory entries that originated from the DB but are no longer
|
||||
present in db_guardrail_ids. Config-loaded guardrails are never touched.
|
||||
|
|
@ -665,8 +665,8 @@ class InMemoryGuardrailHandler:
|
|||
|
||||
@staticmethod
|
||||
def _normalize_litellm_params_for_comparison(
|
||||
params: Optional[Any],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
params: Any | None,
|
||||
) -> dict[str, Any] | None:
|
||||
"""
|
||||
Render litellm_params to a canonical dict so an in-memory LitellmParams and
|
||||
the raw dict loaded from the DB compare equal when they describe the same
|
||||
|
|
@ -738,9 +738,9 @@ class InMemoryGuardrailHandler:
|
|||
def reinitialize_guardrail(
|
||||
self,
|
||||
guardrail: Guardrail,
|
||||
config_file_path: Optional[str] = None,
|
||||
config_file_path: str | None = None,
|
||||
source: Literal["db", "config"] = "config",
|
||||
) -> Optional[Guardrail]:
|
||||
) -> Guardrail | None:
|
||||
"""
|
||||
Force re-initialization of a guardrail even if it exists in memory.
|
||||
Removes old callback from litellm.callbacks and creates fresh instance.
|
||||
|
|
@ -762,8 +762,8 @@ class InMemoryGuardrailHandler:
|
|||
)
|
||||
|
||||
def sync_guardrail_from_db(
|
||||
self, guardrail: Guardrail, config_file_path: Optional[str] = None
|
||||
) -> Optional[Guardrail]:
|
||||
self, guardrail: Guardrail, config_file_path: str | None = None
|
||||
) -> Guardrail | None:
|
||||
"""
|
||||
Sync a guardrail from DB - initializes if new, re-initializes if changed.
|
||||
This is the method to call during DB polling.
|
||||
|
|
|
|||
|
|
@ -151,8 +151,16 @@ class Violation(NamedTuple):
|
|||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def contains_any(t: Type, _seen: set[int] | None = None) -> bool:
|
||||
_MAX_CONTAINS_ANY_DEPTH = 64
|
||||
|
||||
|
||||
def contains_any(t: Type, _seen: set[int] | None = None, _depth: int = 0) -> bool:
|
||||
"""True if a *value* of type ``t`` carries `Any` anywhere meaningful."""
|
||||
if _depth > _MAX_CONTAINS_ANY_DEPTH:
|
||||
# Bail out on deeply-nested / potentially-circular types rather than
|
||||
# overflowing the Python call stack. A type this deep is unlikely to
|
||||
# carry a *meaningful* Any that the developer could actually fix.
|
||||
return False
|
||||
seen = _seen if _seen is not None else set()
|
||||
p = get_proper_type(t)
|
||||
if id(p) in seen:
|
||||
|
|
@ -166,11 +174,11 @@ def contains_any(t: Type, _seen: set[int] | None = None) -> bool:
|
|||
if isinstance(p, AnyType):
|
||||
return p.type_of_any not in _HARMLESS_ANY
|
||||
if isinstance(p, UnionType):
|
||||
return any(contains_any(item, seen) for item in p.items)
|
||||
return any(contains_any(item, seen, _depth + 1) for item in p.items)
|
||||
if isinstance(p, Instance):
|
||||
return any(contains_any(arg, seen) for arg in p.args)
|
||||
return any(contains_any(arg, seen, _depth + 1) for arg in p.args)
|
||||
if isinstance(p, TupleType):
|
||||
return any(contains_any(item, seen) for item in p.items)
|
||||
return any(contains_any(item, seen, _depth + 1) for item in p.items)
|
||||
return False
|
||||
|
||||
|
||||
|
|
@ -249,17 +257,11 @@ def _reason_ok(reason: str | None) -> bool:
|
|||
return reason is not None and len(reason.strip()) >= MIN_REASON_LEN
|
||||
|
||||
|
||||
def scan_any_ok(
|
||||
path: Path, source: str
|
||||
) -> tuple[frozenset[int], tuple[Violation, ...]]:
|
||||
def scan_any_ok(path: Path, source: str) -> tuple[frozenset[int], tuple[Violation, ...]]:
|
||||
"""Return (lines with a valid any-ok suppression, LIT005 violations)."""
|
||||
try:
|
||||
tokens = tokenize.generate_tokens(
|
||||
iter(source.splitlines(keepends=True)).__next__
|
||||
)
|
||||
comments = tuple(
|
||||
(t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT
|
||||
)
|
||||
tokens = tokenize.generate_tokens(iter(source.splitlines(keepends=True)).__next__)
|
||||
comments = tuple((t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT)
|
||||
except tokenize.TokenError:
|
||||
return frozenset(), ()
|
||||
|
||||
|
|
@ -370,9 +372,7 @@ def check_files(rel_paths: Sequence[str]) -> tuple[Violation, ...]:
|
|||
try:
|
||||
source = abs_path.read_text(encoding="utf-8")
|
||||
except (OSError, UnicodeDecodeError) as exc:
|
||||
out.append(
|
||||
Violation(report_path, 0, 0, "LIT000", f"could not read file: {exc}")
|
||||
)
|
||||
out.append(Violation(report_path, 0, 0, "LIT000", f"could not read file: {exc}"))
|
||||
continue
|
||||
|
||||
ok_lines, ok_violations = scan_any_ok(report_path, source)
|
||||
|
|
@ -495,9 +495,7 @@ def _in_scope(v: Violation, line_map: dict[str, LineScope] | None) -> bool:
|
|||
|
||||
|
||||
def main(argv: Sequence[str]) -> int:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Any-discipline gate (changed-only, changed-lines)."
|
||||
)
|
||||
parser = argparse.ArgumentParser(description="Any-discipline gate (changed-only, changed-lines).")
|
||||
parser.add_argument(
|
||||
"paths",
|
||||
nargs="*",
|
||||
|
|
@ -520,9 +518,7 @@ def main(argv: Sequence[str]) -> int:
|
|||
file=sys.stderr,
|
||||
)
|
||||
return 0
|
||||
rel_paths = _to_litellm_relative(
|
||||
(REPO_ROOT / name).resolve() for name in line_map
|
||||
)
|
||||
rel_paths = _to_litellm_relative((REPO_ROOT / name).resolve() for name in line_map)
|
||||
elif args.paths:
|
||||
rel_paths = _to_litellm_relative((REPO_ROOT / p).resolve() for p in args.paths)
|
||||
else:
|
||||
|
|
@ -546,9 +542,7 @@ def main(argv: Sequence[str]) -> int:
|
|||
file=sys.stderr,
|
||||
)
|
||||
return 1
|
||||
print(
|
||||
f"OK: {len(rel_paths)} changed file(s) under litellm/ have no Any-typed values on changed lines"
|
||||
)
|
||||
print(f"OK: {len(rel_paths)} changed file(s) under litellm/ have no Any-typed values on changed lines")
|
||||
return 0
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue