fix: preserve reasoning forwarding on partial updates

This commit is contained in:
jibanez-staticduo 2026-09-15 09:44:48 +02:00
parent c84c73b3df
commit a1db7d6b03
No known key found for this signature in database
3 changed files with 28 additions and 17 deletions

View file

@ -357,7 +357,7 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
)
model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True)
merge_reasoning_content_in_choices: bool | None = False
forward_reasoning_content: bool | None = False
forward_reasoning_content: bool | None = None
reasoning_content_field: Literal["reasoning_content", "reasoning"] | None = None
model_info: dict | None = None
mock_response: str | ModelResponse | Exception | Any | None = None

View file

@ -1169,7 +1169,8 @@ class TestUpdateModel:
@pytest.mark.asyncio
@pytest.mark.parametrize("reasoning_field", [None, "reasoning_content", "reasoning"])
async def test_update_model_clears_cache_after_db_write(self, reasoning_field):
@pytest.mark.parametrize("forward", [None, False, True])
async def test_update_model_clears_cache_after_db_write(self, reasoning_field, forward):
"""
Regression test for the stale-router bug: POST /model/update must refresh
the in-memory router after persisting to LiteLLM_ProxyModelTable, otherwise
@ -1192,6 +1193,7 @@ class TestUpdateModel:
"model": "openai/gpt-4o-mini",
"api_key": "sk-existing",
"reasoning_content_field": "reasoning",
"forward_reasoning_content": True,
}
existing_row.model_dump.return_value = {
"model_name": "gpt-4o-mini",
@ -1239,7 +1241,11 @@ class TestUpdateModel:
):
await update_model(
model_params=updateDeployment(
litellm_params=updateLiteLLMParams(guardrails=["g1"], **({} if reasoning_field is None else {"reasoning_content_field": reasoning_field})),
litellm_params=updateLiteLLMParams(
**({"guardrails": ["g1"]} if reasoning_field is None and forward is None else {}),
**({} if reasoning_field is None else {"reasoning_content_field": reasoning_field}),
**({} if forward is None else {"forward_reasoning_content": forward}),
),
model_info=ModelInfo(id=model_id),
),
user_api_key_dict=admin_user,
@ -1249,6 +1255,7 @@ class TestUpdateModel:
mock_clear_cache.assert_awaited_once_with()
stored = json.loads(mock_prisma.db.litellm_proxymodeltable.update.call_args.kwargs["data"]["litellm_params"])
assert stored["reasoning_content_field"] == (reasoning_field or "reasoning")
assert stored["forward_reasoning_content"] is (True if forward is None else forward)
class TestUpdatePublicModelGroups:
@ -6205,24 +6212,34 @@ class TestAccessGroupModelSync:
@pytest.mark.parametrize("field", [None, "reasoning_content", "reasoning"])
def test_model_patch_preserves_reasoning_field_unless_explicit(field, monkeypatch):
@pytest.mark.parametrize("forward", [None, False, True])
def test_model_patch_preserves_reasoning_field_unless_explicit(field, forward, monkeypatch):
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.proxy.management_endpoints.model_management_endpoints import update_db_model
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-reasoning-field")
deployment = Deployment(
model_name="reasoning-test",
litellm_params=LiteLLM_Params(model="openai/reasoning-test", reasoning_content_field="reasoning"),
litellm_params=LiteLLM_Params(model="openai/reasoning-test", reasoning_content_field="reasoning", forward_reasoning_content=True),
model_info=ModelInfo(id="reasoning-row"),
)
params = updateLiteLLMParams(tpm=123, **({} if field is None else {"reasoning_content_field": field}))
params = updateLiteLLMParams(
**({"tpm": 123} if field is None and forward is None else {}),
**({} if field is None else {"reasoning_content_field": field}),
**({} if forward is None else {"forward_reasoning_content": forward}),
)
if forward is None:
assert "forward_reasoning_content" not in params.model_dump(exclude_none=True)
if field is None:
assert "reasoning_content_field" not in params.model_dump(exclude_none=True)
result = update_db_model(db_model=deployment, updated_patch=updateDeployment(litellm_params=params))
stored = json.loads(result["litellm_params"])
assert stored["tpm"] == 123
if field is None and forward is None:
assert stored["tpm"] == 123
assert stored["forward_reasoning_content"] is (True if forward is None else forward)
assert (
stored["reasoning_content_field"] if field is None
else decrypt_value_helper(value=stored["reasoning_content_field"], key="reasoning_content_field")
) == (field or "reasoning")
assert deployment.litellm_params.reasoning_content_field == "reasoning"
assert deployment.litellm_params.forward_reasoning_content is True

View file

@ -29711,11 +29711,8 @@ export interface components {
default_api_key_tpm_limit?: number | null;
/** Drop Params */
drop_params?: boolean | string | null;
/**
* Forward Reasoning Content
* @default false
*/
forward_reasoning_content: boolean | null;
/** Forward Reasoning Content */
forward_reasoning_content?: boolean | null;
/** Gcs Bucket Name */
gcs_bucket_name?: string | null;
/** Google Maps Grounding Cost Per Query */
@ -39932,11 +39929,8 @@ export interface components {
default_api_key_tpm_limit?: number | null;
/** Drop Params */
drop_params?: boolean | string | null;
/**
* Forward Reasoning Content
* @default false
*/
forward_reasoning_content: boolean | null;
/** Forward Reasoning Content */
forward_reasoning_content?: boolean | null;
/** Gcs Bucket Name */
gcs_bucket_name?: string | null;
/** Google Maps Grounding Cost Per Query */