fix: support encrypted reasoning field storage

This commit is contained in:
jibanez-staticduo 2026-09-15 10:04:36 +02:00
parent a1db7d6b03
commit 624f2f97b6
No known key found for this signature in database
6 changed files with 73 additions and 27 deletions

View file

@ -10,23 +10,24 @@ from litellm.types.llms.openai import AllMessageValues
def normalize_reasoning_content(
messages: Sequence[AllMessageValues], *, forward: bool = True
messages: Sequence[AllMessageValues], *, forward: bool = True, normalize: bool = True
) -> list[AllMessageValues]: # mutable-ok: provider request contract
def normalize_message(message: AllMessageValues) -> AllMessageValues:
if message["role"] != "assistant":
return message
if not normalize and forward:
return message
history: Final[Mapping[str, object]] = message
removed_fields: Final = ("reasoning", "reasoning_content") if normalize else ("reasoning_content",)
reasoning: Final = (
history.get("reasoning") if history.get("reasoning") is not None else history.get("reasoning_content")
)
normalized: Final[Mapping[str, object]] = MappingProxyType(
{
**MappingProxyType(
{key: value for key, value in history.items() if key not in ("reasoning", "reasoning_content")}
),
**MappingProxyType({key: value for key, value in history.items() if key not in removed_fields}),
**(
MappingProxyType({"reasoning": reasoning})
if forward and reasoning is not None
if normalize and forward and reasoning is not None
else MappingProxyType({})
),
}

View file

@ -155,15 +155,11 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
litellm_params: dict, # mutable-ok: provider request contract
headers: dict, # mutable-ok: provider request contract
) -> dict: # mutable-ok: provider request contract
request_messages: Final = (
normalize_reasoning_content(messages, forward=litellm_params.get("forward_reasoning_content") is True)
if litellm_params.get("reasoning_content_field") == "reasoning"
else deepcopy(messages)
request_messages: Final = normalize_reasoning_content(
messages,
forward=litellm_params.get("forward_reasoning_content") is True,
normalize=litellm_params.get("reasoning_content_field") == "reasoning",
)
if litellm_params.get("forward_reasoning_content") is not True:
for message in request_messages:
if message["role"] == "assistant":
message.pop("reasoning_content", None)
return super().transform_request(model, request_messages, optional_params, litellm_params, headers)
async def async_transform_request(

View file

@ -358,7 +358,10 @@ 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 = None
reasoning_content_field: Literal["reasoning_content", "reasoning"] | None = None
reasoning_content_field: str | None = Field(
default=None,
description="Historical assistant reasoning field: reasoning_content (default) or reasoning.",
)
model_info: dict | None = None
mock_response: str | ModelResponse | Exception | Any | None = None

View file

@ -586,6 +586,7 @@ async def test_reasoning_field_sdk_router_final_wire(
("legacy", {}),
("normalized", {"reasoning_content_field": "reasoning"}),
("explicit-default", {"reasoning_content_field": "reasoning_content"}),
("unknown", {"reasoning_content_field": "unknown"}),
)
],
num_retries=0,
@ -602,7 +603,9 @@ async def test_reasoning_field_sdk_router_final_wire(
"usage": {"prompt_tokens": 10, "completion_tokens": 1, "total_tokens": 11},
},
)
for alias, field in (("normalized", "reasoning"), ("legacy", None), ("explicit-default", "reasoning_content")):
for alias, field in (
("normalized", "reasoning"), ("legacy", None), ("explicit-default", "reasoning_content"), ("unknown", "unknown")
):
kwargs: Final = (
{"model": alias, "messages": messages}
if via_router
@ -637,7 +640,7 @@ async def test_reasoning_field_sdk_router_final_wire(
assert "reasoning_content_field" not in payload
assert "forward_reasoning_content" not in payload
assert messages == original
assert route.call_count == 3
assert route.call_count == 4
@pytest.mark.asyncio

View file

@ -1170,7 +1170,7 @@ class TestUpdateModel:
@pytest.mark.asyncio
@pytest.mark.parametrize("reasoning_field", [None, "reasoning_content", "reasoning"])
@pytest.mark.parametrize("forward", [None, False, True])
async def test_update_model_clears_cache_after_db_write(self, reasoning_field, forward):
async def test_update_model_clears_cache_after_db_write(self, reasoning_field, forward, monkeypatch):
"""
Regression test for the stale-router bug: POST /model/update must refresh
the in-memory router after persisting to LiteLLM_ProxyModelTable, otherwise
@ -1186,13 +1186,16 @@ class TestUpdateModel:
updateLiteLLMParams,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-reasoning-field")
model_id = "db-model-under-test"
existing_row = MagicMock()
existing_row.litellm_params = {
"model": "openai/gpt-4o-mini",
"api_key": "sk-existing",
"reasoning_content_field": "reasoning",
"reasoning_content_field": encrypt_value_helper("reasoning"),
"forward_reasoning_content": True,
}
existing_row.model_dump.return_value = {
@ -1228,10 +1231,6 @@ class TestUpdateModel:
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
side_effect=lambda value: value,
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.clear_cache",
new=AsyncMock(
@ -1254,7 +1253,9 @@ class TestUpdateModel:
mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once()
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 decrypt_value_helper(stored["reasoning_content_field"], key="reasoning_content_field") == (reasoning_field or "reasoning")
if reasoning_field is None:
assert stored["reasoning_content_field"] == existing_row.litellm_params["reasoning_content_field"]
assert stored["forward_reasoning_content"] is (True if forward is None else forward)
@ -6243,3 +6244,39 @@ def test_model_patch_preserves_reasoning_field_unless_explicit(field, forward, m
) == (field or "reasoning")
assert deployment.litellm_params.reasoning_content_field == "reasoning"
assert deployment.litellm_params.forward_reasoning_content is True
@pytest.mark.asyncio
@pytest.mark.parametrize("field", ["reasoning_content", "reasoning"])
async def test_reasoning_field_encrypted_db_round_trip(field, monkeypatch):
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.proxy.management_endpoints.model_management_endpoints import get_db_model, update_db_model
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-reasoning-field")
initial = Deployment(
model_name="reasoning-test",
litellm_params=LiteLLM_Params(model="openai/reasoning-test", forward_reasoning_content=True),
model_info=ModelInfo(id="reasoning-row"),
)
first = update_db_model(
db_model=initial,
updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(reasoning_content_field=field)),
)
encrypted = json.loads(first["litellm_params"])
assert encrypted["reasoning_content_field"] != field
raw = {"model_name": first["model_name"], "litellm_params": encrypted, "model_info": {"id": "reasoning-row"}}
deployment = Deployment(**raw)
assert deployment.litellm_params.reasoning_content_field == encrypted["reasoning_content_field"]
row = MagicMock()
row.model_dump.return_value = raw
prisma = MagicMock()
prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=row)
loaded = await get_db_model("reasoning-row", prisma)
assert loaded.litellm_params.reasoning_content_field == encrypted["reasoning_content_field"]
second = update_db_model(
db_model=loaded, updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(tpm=123))
)
stored = json.loads(second["litellm_params"])
assert stored["reasoning_content_field"] == encrypted["reasoning_content_field"]
assert stored["forward_reasoning_content"] is True
assert decrypt_value_helper(stored["reasoning_content_field"], key="reasoning_content_field") == field

View file

@ -29882,8 +29882,11 @@ export interface components {
} | null;
/** Quality Router Default Model */
quality_router_default_model?: string | null;
/** Reasoning Content Field */
reasoning_content_field?: ("reasoning_content" | "reasoning") | null;
/**
* Reasoning Content Field
* @description Historical assistant reasoning field: reasoning_content (default) or reasoning.
*/
reasoning_content_field?: string | null;
/** Region Name */
region_name?: string | null;
/** Regional Endpoint Uplift Multiplier */
@ -40100,8 +40103,11 @@ export interface components {
} | null;
/** Quality Router Default Model */
quality_router_default_model?: string | null;
/** Reasoning Content Field */
reasoning_content_field?: ("reasoning_content" | "reasoning") | null;
/**
* Reasoning Content Field
* @description Historical assistant reasoning field: reasoning_content (default) or reasoning.
*/
reasoning_content_field?: string | null;
/** Region Name */
region_name?: string | null;
/** Regional Endpoint Uplift Multiplier */