diff --git a/litellm/types/router.py b/litellm/types/router.py index f3703f82711..c8a53f3467b 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 86bbd0d3b4b..ba7bfdcdc21 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -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 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index f0df54706b8..a90fce927ca 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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 */