mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
fix(proxy): keep an empty duration out of the legacy update hook on key regenerate
This commit is contained in:
parent
a159c7d98e
commit
0429c64504
2 changed files with 62 additions and 2 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue