fix(proxy): keep an empty duration out of the legacy update hook on key regenerate

This commit is contained in:
mateo-berri 2026-09-12 18:22:31 -07:00
parent a159c7d98e
commit 0429c64504
2 changed files with 62 additions and 2 deletions

View file

@ -433,12 +433,17 @@ def _effective_key_for_generate(data: GenerateKeyRequest, now: datetime) -> Lite
)
_EMPTY_DURATION_MEANS_UNCHANGED: Final = frozenset({"duration", "budget_duration"})
def _regenerate_request_as_update_request(key: str, data: RegenerateKeyRequest) -> UpdateKeyRequest | None:
changed_fields: Final = MappingProxyType(
{
field: value
for field, value in data.model_dump(exclude_unset=True).items()
if field in UpdateKeyRequest.model_fields and field != "key"
if field in UpdateKeyRequest.model_fields
and field != "key"
and not (field in _EMPTY_DURATION_MEANS_UNCHANGED and value == "")
}
)
if not changed_fields:

View file

@ -12142,7 +12142,10 @@ async def test_execute_virtual_key_regeneration_allows_when_custom_key_update_ho
@pytest.mark.asyncio
@pytest.mark.parametrize("data", [None, RegenerateKeyRequest()])
@pytest.mark.parametrize(
"data",
[None, RegenerateKeyRequest(), RegenerateKeyRequest(duration=""), RegenerateKeyRequest(budget_duration="")],
)
async def test_execute_virtual_key_regeneration_skips_custom_key_update_hook_without_changes(data):
mock_prisma_client = _make_regenerate_mock_prisma()
@ -12184,6 +12187,58 @@ async def test_execute_virtual_key_regeneration_skips_custom_key_update_hook_wit
assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1
@pytest.mark.asyncio
async def test_execute_virtual_key_regeneration_hides_the_untouched_modal_expiry_from_the_custom_key_update_hook():
mock_prisma_client = _make_regenerate_mock_prisma()
untouched_modal_body = RegenerateKeyRequest(
key_alias=None, max_budget=None, tpm_limit=None, rpm_limit=None, duration="", grace_period=""
)
received_data: list[UpdateKeyRequest] = []
async def hook(data: UpdateKeyRequest) -> dict[str, object]:
received_data.append(data)
if data.duration is not None and duration_in_seconds(data.duration) > duration_in_seconds("7d"):
return {"decision": False, "message": "duration must be <= 7d"}
return {"decision": True}
with (
patch( # test-quality-ok: deterministic token setup for the untouched modal body
"litellm.proxy.management_endpoints.key_management_endpoints.get_new_token",
new_callable=AsyncMock,
return_value="sk-newtoken1234ab12",
),
patch( # test-quality-ok: grace-period path is outside the hook input
"litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key",
new_callable=AsyncMock,
),
patch( # test-quality-ok: cache eviction is outside the hook input
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
),
patch( # test-quality-ok: rotation callback is outside the hook input
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook",
new_callable=AsyncMock,
),
patch("litellm.proxy.proxy_server.user_custom_key_update", hook), # test-quality-ok: inject policy hook
):
await _execute_virtual_key_regeneration(
prisma_client=mock_prisma_client,
key_in_db=_make_regenerate_existing_key(),
hashed_api_key="abc123",
key="abc123",
data=untouched_modal_body,
user_api_key_dict=_make_regenerate_user_api_key_dict(),
litellm_changed_by=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
)
assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1
assert len(received_data) == 1
assert "duration" not in received_data[0].model_fields_set
assert received_data[0].model_fields_set >= {"key", "key_alias", "max_budget", "tpm_limit", "rpm_limit"}
_POLICY_DENIAL_MESSAGE = "key duration must be 7d or less"
_POLICY_HASHED_TOKEN = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b"
_POLICY_GENERATED_KEY = {"key": "sk-test-key", "expires": None, "user_id": "test-user", "team_id": None}