mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix: preserve nested thinking extras and honor drop_params
This commit is contained in:
parent
5aefd558d9
commit
ac6fdd4fe8
4 changed files with 140 additions and 23 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue