mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: support encrypted reasoning field storage
This commit is contained in:
parent
a1db7d6b03
commit
624f2f97b6
6 changed files with 73 additions and 27 deletions
|
|
@ -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({})
|
||||
),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue