diff --git a/litellm/litellm_core_utils/thinking_param_translation.py b/litellm/litellm_core_utils/thinking_param_translation.py index 642881b21cf..0ebdb5b3667 100644 --- a/litellm/litellm_core_utils/thinking_param_translation.py +++ b/litellm/litellm_core_utils/thinking_param_translation.py @@ -5,6 +5,12 @@ from dataclasses import dataclass from types import MappingProxyType from typing import Final +from pydantic import TypeAdapter, ValidationError + +_STR_KEYED_MAPPING: Final = TypeAdapter(Mapping[str, object]) +_OBJECT_TUPLE: Final = TypeAdapter(tuple[object, ...]) +_EMPTY: Final[Mapping[str, object]] = MappingProxyType({}) + _SEND_VIA_EXTRA_BODY: Final = "extra_body" _SEND_VIA_PROVIDER_MAPPED: Final = "provider_mapped" _SEND_VIA_VALUES: Final = frozenset((_SEND_VIA_EXTRA_BODY, _SEND_VIA_PROVIDER_MAPPED)) @@ -29,10 +35,19 @@ class ThinkingParamsState: extra_body: Mapping[str, object] +def str_keyed_mapping_or_none(value: object) -> Mapping[str, object] | None: + if not isinstance(value, Mapping): + return None + try: + return _STR_KEYED_MAPPING.validate_python(value) + except ValidationError: + return None + + def _as_str_tuple(value: object) -> tuple[str, ...]: if not isinstance(value, (list, tuple)): return () - return tuple(item for item in value if isinstance(item, str)) + return tuple(item for item in _OBJECT_TUPLE.validate_python(value) if isinstance(item, str)) def _thinking_enabled(thinking: object) -> bool: @@ -40,14 +55,14 @@ def _thinking_enabled(thinking: object) -> bool: return thinking if isinstance(thinking, str): return thinking.lower() in _THINKING_ENABLED_STRINGS - if isinstance(thinking, Mapping): - typ: Final = thinking.get("type") - if isinstance(typ, str): - return typ.lower() in _THINKING_TYPE_ENABLED - enabled: Final = thinking.get("enabled") - if isinstance(enabled, bool): - return enabled - return False + mapping: Final = str_keyed_mapping_or_none(thinking) + if mapping is None: + return False + typ: Final = mapping.get("type") + if isinstance(typ, str): + return typ.lower() in _THINKING_TYPE_ENABLED + enabled: Final = mapping.get("enabled") + return enabled if isinstance(enabled, bool) else False def _thinking_type_candidate(thinking: object) -> str | None: @@ -56,11 +71,10 @@ def _thinking_type_candidate(thinking: object) -> str | None: return "enabled" if thinking else "disabled" case str(): return thinking - case Mapping(): - raw: Final = thinking.get("type") - return raw if isinstance(raw, str) else None case _: - return None + mapping: Final = str_keyed_mapping_or_none(thinking) + raw: Final = mapping.get("type") if mapping is not None else None + return raw if isinstance(raw, str) else None def _thinking_type_value(thinking: object, allowed: Sequence[str]) -> str | None: @@ -75,9 +89,10 @@ def _thinking_type_value(thinking: object, allowed: Sequence[str]) -> str | None def _thinking_payload(thinking: object, typ: str) -> Mapping[str, object]: - if not isinstance(thinking, Mapping): + mapping: Final = str_keyed_mapping_or_none(thinking) + if mapping is None: return MappingProxyType({"type": typ}) - return MappingProxyType({**thinking, "type": typ}) + return MappingProxyType({**mapping, "type": typ}) def _clamp_effort(value: object, allowed: Sequence[str]) -> str | None: @@ -93,19 +108,13 @@ def _clamp_effort(value: object, allowed: Sequence[str]) -> str | None: return None -def _mapping_or_none(value: object) -> Mapping[str, object] | None: - if isinstance(value, Mapping): - return value - return None - - def _merged_mapping_value(left: Mapping[str, object], right: Mapping[str, object], key: str) -> object: if key not in right: return left[key] if key not in left: return right[key] - left_map: Final = _mapping_or_none(left[key]) - right_map: Final = _mapping_or_none(right[key]) + left_map: Final = str_keyed_mapping_or_none(left[key]) + right_map: Final = str_keyed_mapping_or_none(right[key]) if left_map is not None and right_map is not None: return _deep_merge_pair(left_map, right_map) return right[key] @@ -117,7 +126,7 @@ def _deep_merge_pair(left: Mapping[str, object], right: Mapping[str, object]) -> def _thawed(value: object) -> object: - mapping: Final = _mapping_or_none(value) + mapping: Final = str_keyed_mapping_or_none(value) return value if mapping is None else thaw_mapping(mapping) @@ -135,14 +144,15 @@ def _map_thinking_to_extra_body( case "thinking.type": typ: Final = _thinking_type_value(thinking, thinking_values) if typ is None: - return MappingProxyType({}) + return _EMPTY return MappingProxyType({"thinking": _thinking_payload(thinking, typ)}) case "thinking": - if isinstance(thinking, Mapping): - return MappingProxyType({"thinking": MappingProxyType({k: thinking[k] for k in thinking})}) + thinking_mapping: Final = str_keyed_mapping_or_none(thinking) + if thinking_mapping is not None: + return MappingProxyType({"thinking": MappingProxyType(thinking_mapping)}) typ_only: Final = _thinking_type_value(thinking, thinking_values or ("enabled", "disabled")) if typ_only is None: - return MappingProxyType({}) + return _EMPTY return MappingProxyType({"thinking": MappingProxyType({"type": typ_only})}) case "enable_thinking": return MappingProxyType({"enable_thinking": _thinking_enabled(thinking)}) @@ -151,7 +161,7 @@ def _map_thinking_to_extra_body( {"chat_template_kwargs": MappingProxyType({"enable_thinking": _thinking_enabled(thinking)})} ) case _: - return MappingProxyType({}) + return _EMPTY def _map_effort_to_extra_body( @@ -165,7 +175,7 @@ def _map_effort_to_extra_body( return MappingProxyType({"reasoning_effort": clamped}) if not effort_values and send_via == _SEND_VIA_EXTRA_BODY and isinstance(effort, str): return MappingProxyType({"reasoning_effort": effort}) - return MappingProxyType({}) + return _EMPTY def translate_thinking_params( @@ -202,7 +212,7 @@ def translate_thinking_params( thinking_values=thinking_values, ) if thinking is not None and send_via == _SEND_VIA_EXTRA_BODY - else MappingProxyType({}) + else _EMPTY ) effort_patch: Final = ( _map_effort_to_extra_body( @@ -211,7 +221,7 @@ def translate_thinking_params( send_via=send_via, ) if effort is not None - else MappingProxyType({}) + else _EMPTY ) patch: Final = MappingProxyType({**thinking_patch, **effort_patch}) if not patch: diff --git a/litellm/utils.py b/litellm/utils.py index 0d71c65c3c9..6a5819ceeb5 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -296,6 +296,7 @@ from typing_extensions import assert_never from litellm import utils as litellm_utils from litellm.litellm_core_utils.thinking_param_translation import ( ThinkingParamsState, + str_keyed_mapping_or_none, thaw_mapping, translate_thinking_params, ) @@ -4329,14 +4330,13 @@ def remove_sensitive_keys_from_dict(d: dict) -> dict: def _translate_thinking_in_params( *, model_info: Mapping[str, object] | None, - passed_params: dict, # mutable-ok: get_optional_params hands over its legacy mutable params - non_default_params: dict, # mutable-ok: get_optional_params hands over its legacy mutable params -) -> tuple[dict, dict]: # mutable-ok: get_optional_params keeps mutating both copies downstream - existing_extra_body: Final = passed_params.get("extra_body") + passed_params: dict[str, object], # mutable-ok: get_optional_params hands over its legacy mutable params + non_default_params: dict[str, object], # mutable-ok: get_optional_params hands over its legacy mutable params +) -> tuple[dict[str, object], dict[str, object]]: # mutable-ok: get_optional_params keeps mutating both downstream state: Final = ThinkingParamsState( thinking=non_default_params.get("thinking"), reasoning_effort=non_default_params.get("reasoning_effort"), - extra_body=existing_extra_body if isinstance(existing_extra_body, Mapping) else MappingProxyType({}), + extra_body=str_keyed_mapping_or_none(passed_params.get("extra_body")) or MappingProxyType({}), ) translated: Final = translate_thinking_params(model_info=model_info, state=state) if translated is state: diff --git a/tests/unit/litellm_core_utils/test_thinking_param_translation.py b/tests/unit/litellm_core_utils/test_thinking_param_translation.py index a8cc0016bbd..956022ba7f9 100644 --- a/tests/unit/litellm_core_utils/test_thinking_param_translation.py +++ b/tests/unit/litellm_core_utils/test_thinking_param_translation.py @@ -210,6 +210,7 @@ def test_translate_moves_params_into_extra_body( pytest.param(_model_info(), {"type": "adaptive"}, None, id="thinking_type_not_allowed"), pytest.param(_model_info(), {"budget_tokens": 1024}, None, id="thinking_type_unreadable"), pytest.param(_model_info(), 1, None, id="thinking_type_unsupported_value"), + pytest.param(_model_info(), {1: "enabled"}, None, id="thinking_mapping_with_non_string_keys"), pytest.param( _model_info(thinking_param="thinking", thinking_values=[]), "adaptive", None, id="thinking_value_unmapped" ),