diff --git a/litellm/batch_completion/main.py b/litellm/batch_completion/main.py index 702dd194fda..7ec71dcd4bb 100644 --- a/litellm/batch_completion/main.py +++ b/litellm/batch_completion/main.py @@ -1,13 +1,29 @@ +from collections.abc import Mapping from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait from typing import Final import litellm from litellm._logging import print_verbose -from litellm.utils import get_optional_params +from litellm.utils import get_model_info, get_optional_params from ..llms.vllm.completion import handler as vllm_handler +def _model_info_for_batch( + *, + model: str, + custom_llm_provider: str | None, + kwargs: Mapping[str, object], +) -> Mapping[str, object] | None: + from_kwargs: Final = kwargs.get("model_info") + if isinstance(from_kwargs, Mapping): + return from_kwargs + try: + return get_model_info(model=model, custom_llm_provider=custom_llm_provider) + except Exception: # noqa: BLE001 # get_model_info raises Exception for unmapped models + return None + + def batch_completion( model: str, # Optional OpenAI params: see https://platform.openai.com/docs/api-reference/chat/create @@ -79,9 +95,13 @@ def batch_completion( frequency_penalty=frequency_penalty, logit_bias=logit_bias, user=user, - # params to identify the model model=model, custom_llm_provider=custom_llm_provider, + model_info=_model_info_for_batch( + model=model, + custom_llm_provider=custom_llm_provider, + kwargs=kwargs, + ), ) results = vllm_handler.batch_completions( model=model, diff --git a/litellm/litellm_core_utils/thinking_param_translation.py b/litellm/litellm_core_utils/thinking_param_translation.py index c3a183def1b..9e4beeff231 100644 --- a/litellm/litellm_core_utils/thinking_param_translation.py +++ b/litellm/litellm_core_utils/thinking_param_translation.py @@ -88,6 +88,29 @@ 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]) + if left_map is not None and right_map is not None: + return _deep_merge_pair(left_map, right_map) + return right[key] + + +def _deep_merge_pair(left: Mapping[str, object], right: Mapping[str, object]) -> Mapping[str, object]: + keys: Final = frozenset(left) | frozenset(right) + return MappingProxyType({key: _merged_mapping_value(left, right, key) for key in keys}) + + def _map_thinking_to_extra_body( *, thinking_param: str | None, @@ -117,6 +140,20 @@ def _map_thinking_to_extra_body( return MappingProxyType({}) +def _map_effort_to_extra_body( + *, + effort: object, + effort_values: Sequence[str], + send_via: object, +) -> Mapping[str, object]: + clamped: Final = _clamp_effort(effort, effort_values) + if clamped is not None: + 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({}) + + def translate_thinking_params( *, model_info: Mapping[str, object] | None, @@ -143,33 +180,33 @@ def translate_thinking_params( if thinking is None and effort is None: return state - existing_extra: Final = dict(state.extra_body) - patch: dict[str, object] = {} keep_thinking: Final = send_via == _SEND_VIA_PROVIDER_MAPPED - - if thinking is not None and send_via == _SEND_VIA_EXTRA_BODY: - patch.update( - _map_thinking_to_extra_body( - thinking_param=thinking_param, - thinking=thinking, - thinking_values=thinking_values, - ) + thinking_patch: Final = ( + _map_thinking_to_extra_body( + thinking_param=thinking_param, + thinking=thinking, + thinking_values=thinking_values, ) - - if effort is not None: - clamped: Final = _clamp_effort(effort, effort_values) - if clamped is not None: - patch["reasoning_effort"] = clamped - elif not effort_values and send_via == _SEND_VIA_EXTRA_BODY and isinstance(effort, str): - patch["reasoning_effort"] = effort - + if thinking is not None and send_via == _SEND_VIA_EXTRA_BODY + else MappingProxyType({}) + ) + effort_patch: Final = ( + _map_effort_to_extra_body( + effort=effort, + effort_values=effort_values, + send_via=send_via, + ) + if effort is not None + else MappingProxyType({}) + ) + patch: Final = MappingProxyType({**thinking_patch, **effort_patch}) if not patch: return state thinking_mapped: Final = any(key in patch for key in ("thinking", "enable_thinking", "chat_template_kwargs")) next_thinking: Final = thinking if (keep_thinking or not thinking_mapped) else None next_effort: Final = None if "reasoning_effort" in patch else effort - merged_extra: Final = MappingProxyType({**existing_extra, **patch}) + merged_extra: Final = _deep_merge_pair(state.extra_body, patch) return ThinkingParamsState( thinking=next_thinking, reasoning_effort=next_effort, diff --git a/litellm/utils.py b/litellm/utils.py index c814d988ab0..c1e348f7887 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4079,8 +4079,8 @@ def _apply_model_info_thinking_translation( passed_params: dict, non_default_params: dict, ) -> None: - prior_thinking: Final = non_default_params.get("thinking", passed_params.get("thinking")) - prior_effort: Final = non_default_params.get("reasoning_effort", passed_params.get("reasoning_effort")) + prior_thinking: Final = non_default_params.get("thinking") + prior_effort: Final = non_default_params.get("reasoning_effort") existing_extra_raw: Final = passed_params.get("extra_body") existing_extra: Final = existing_extra_raw if isinstance(existing_extra_raw, Mapping) else None translated: Final = apply_thinking_param_translation( diff --git a/tests/test_litellm/litellm_core_utils/test_thinking_param_translation.py b/tests/test_litellm/litellm_core_utils/test_thinking_param_translation.py index 49ebbff9379..c5c91181278 100644 --- a/tests/test_litellm/litellm_core_utils/test_thinking_param_translation.py +++ b/tests/test_litellm/litellm_core_utils/test_thinking_param_translation.py @@ -1,3 +1,4 @@ +import importlib from types import MappingProxyType from litellm.litellm_core_utils.thinking_param_translation import ( @@ -65,6 +66,20 @@ def test_translate_chat_template_kwargs(): } +def test_translate_chat_template_kwargs_preserves_existing_nested_keys(): + result = apply_thinking_param_translation( + model_info=_extra_body_model_info( + thinking_param="chat_template_kwargs", + thinking_values=[], + ), + thinking={"type": "enabled"}, + reasoning_effort=None, + existing_extra_body={"chat_template_kwargs": {"reasoning_budget": 512}}, + ) + assert result.extra_body["chat_template_kwargs"]["reasoning_budget"] == 512 + assert result.extra_body["chat_template_kwargs"]["enable_thinking"] is True + + def test_translate_clamps_effort_aliases(): result = apply_thinking_param_translation( model_info=_extra_body_model_info(reasoning_effort_values=["low", "high", "max"]), @@ -139,3 +154,48 @@ def test_get_optional_params_openai_drop_without_model_info_drops_params(): assert "reasoning_effort" not in extra_body assert optional_params.get("thinking") is None assert optional_params.get("reasoning_effort") is None + + +def test_get_optional_params_does_not_reintroduce_dropped_thinking(): + optional_params = get_optional_params( + model="deepseek-v4-flash", + custom_llm_provider="openai", + drop_params=True, + thinking={"type": "enabled"}, + additional_drop_params=["thinking"], + model_info=_extra_body_model_info( + thinking_param="chat_template_kwargs", + thinking_values=[], + ), + ) + extra_body = optional_params.get("extra_body") or {} + assert optional_params.get("thinking") is None + assert "enable_thinking" not in extra_body + assert "chat_template_kwargs" not in extra_body + + +def test_batch_completion_vllm_passes_model_info(monkeypatch): + batch_completion_mod = importlib.import_module("litellm.batch_completion.main") + + captured: dict[str, object] = {} + looked_up: dict[str, object] = _extra_body_model_info(thinking_param="enable_thinking") + + def fake_get_optional_params(**kwargs: object) -> dict[str, object]: + captured.update(kwargs) + return {} + + def fake_batch_completions(**kwargs: object) -> list[str]: + return ["ok"] + + def fake_get_model_info(**kwargs: object) -> dict[str, object]: + return looked_up + + monkeypatch.setattr(batch_completion_mod, "get_optional_params", fake_get_optional_params) + monkeypatch.setattr(batch_completion_mod.vllm_handler, "batch_completions", fake_batch_completions) + monkeypatch.setattr(batch_completion_mod, "get_model_info", fake_get_model_info) + + batch_completion_mod.batch_completion( + model="vllm/some-model", + messages=[[{"role": "user", "content": "hi"}]], + ) + assert captured.get("model_info") == looked_up