fix(proxy): track hook spawns and defer registration in audit log drain

This commit is contained in:
Tufail Akram 2026-10-02 14:56:54 +05:30
parent 0b9f2c222b
commit 908d027113
4 changed files with 127 additions and 44 deletions

View file

@ -656,11 +656,13 @@ async def new_user(
#########################################################
########## USER CREATED HOOK ################
#########################################################
asyncio.create_task(
UserManagementEventHooks.async_user_created_hook(
data=data,
response=new_user_response,
user_api_key_dict=user_api_key_dict,
track_audit_task(
asyncio.create_task(
UserManagementEventHooks.async_user_created_hook(
data=data,
response=new_user_response,
user_api_key_dict=user_api_key_dict,
)
)
)
#########################################################

View file

@ -110,6 +110,7 @@ from litellm.proxy.management_helpers.access_group_key_sync import (
sync_key_regeneration_access_group_membership,
sync_key_update_access_group_membership,
)
from litellm.proxy.management_helpers.audit_logs import track_audit_task
from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at
from litellm.proxy.management_helpers.object_permission_utils import (
_set_object_permission,
@ -1550,12 +1551,14 @@ async def _common_key_generation_helper(
response.token = response.token_id # remap token to use the hash, and leave the key in the `key` field [TODO]: clean up generate_key_helper_fn to do this
asyncio.create_task(
KeyManagementEventHooks.async_key_generated_hook(
data=data,
response=response,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
track_audit_task(
asyncio.create_task(
KeyManagementEventHooks.async_key_generated_hook(
data=data,
response=response,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
)
)
@ -2886,13 +2889,15 @@ async def _process_single_key_update(
)
# Trigger async hook
asyncio.create_task(
KeyManagementEventHooks.async_key_updated_hook(
data=key_request,
existing_key_row=existing_key_row,
response=response,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
track_audit_task(
asyncio.create_task(
KeyManagementEventHooks.async_key_updated_hook(
data=key_request,
existing_key_row=existing_key_row,
response=response,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
)
)
@ -3583,13 +3588,15 @@ async def update_key_fn(
redis_err,
)
asyncio.create_task(
KeyManagementEventHooks.async_key_updated_hook(
data=data,
existing_key_row=existing_key_row,
response=response,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
track_audit_task(
asyncio.create_task(
KeyManagementEventHooks.async_key_updated_hook(
data=data,
existing_key_row=existing_key_row,
response=response,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
)
)
@ -4159,13 +4166,15 @@ async def delete_key_fn(
"/keys/delete - cache after delete: %s", user_api_key_cache.key_object_cache.in_memory_cache.cache_dict
)
asyncio.create_task(
KeyManagementEventHooks.async_key_deleted_hook(
data=data,
keys_being_deleted=_keys_being_deleted,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
response=number_deleted_keys,
track_audit_task(
asyncio.create_task(
KeyManagementEventHooks.async_key_deleted_hook(
data=data,
keys_being_deleted=_keys_being_deleted,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
response=number_deleted_keys,
)
)
)
@ -5686,13 +5695,15 @@ async def _execute_virtual_key_regeneration(
)
response: Final = GenerateKeyResponse.model_validate(updated_token_dict)
asyncio.create_task(
KeyManagementEventHooks.async_key_rotated_hook(
data=data,
existing_key_row=key_in_db,
response=response,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
track_audit_task(
asyncio.create_task(
KeyManagementEventHooks.async_key_rotated_hook(
data=data,
existing_key_row=key_in_db,
response=response,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
)
)
return response

View file

@ -150,7 +150,7 @@ class AuditTaskRegistry:
if exc is not None:
verbose_proxy_logger.error("Audit log callback task failed: %s", exc, exc_info=exc)
async def drain(self, timeout: float = 10.0) -> None:
async def drain(self, timeout: float = 5.0) -> None:
deadline: Final = asyncio.get_running_loop().time() + timeout
while self._pending_tasks:
remaining = deadline - asyncio.get_running_loop().time()
@ -187,7 +187,7 @@ def track_audit_task(task: asyncio.Task[_T]) -> asyncio.Task[_T]:
return _audit_task_registry.track(task)
async def drain_audit_tasks(timeout: float = 10.0) -> None:
async def drain_audit_tasks(timeout: float = 5.0) -> None:
await _audit_task_registry.drain(timeout=timeout)
@ -261,13 +261,13 @@ async def create_object_audit_log(
async def create_audit_log_for_update(request_data: LiteLLM_AuditLogs):
if not is_audit_logging_enabled():
return
current_task: Final = asyncio.current_task()
if current_task is not None:
track_audit_task(current_task)
if not is_audit_logging_enabled():
return
from litellm.proxy.proxy_server import premium_user, prisma_client
if premium_user is not True:

View file

@ -41,6 +41,7 @@ verbose_proxy_logger.setLevel(level=logging.DEBUG)
from starlette.datastructures import URL
from litellm.proxy.management_helpers.audit_logs import (
_pending_audit_tasks,
create_audit_log_for_update,
drain_audit_tasks,
get_audit_log_changed_by,
@ -460,3 +461,72 @@ async def test_drain_audit_tasks_captures_tasks_queued_during_drain():
await drain_audit_tasks(timeout=2.0)
assert finished_tasks == ["first", "second"]
@pytest.mark.asyncio
async def test_create_audit_log_for_update_does_not_track_when_logging_disabled():
with (
patch("litellm.store_audit_logs", False),
patch("litellm.proxy.management_helpers.audit_logs.is_audit_logging_enabled", return_value=False),
):
request_data: Final = LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by="test-user",
table_name=LitellmTableNames.KEY_TABLE_NAME,
object_id="test-obj",
action="updated",
updated_values="{}",
before_value="{}",
)
current_tracked_before: Final = len(_pending_audit_tasks)
await create_audit_log_for_update(request_data=request_data)
assert len(_pending_audit_tasks) == current_tracked_before
@pytest.mark.asyncio
async def test_hook_spawn_with_io_delay_drained_before_shutdown():
mock_prisma: Final = MagicMock()
writes: Final[list[dict[str, Any]]] = []
drain_completed: Final[list[bool]] = [False]
async def _mock_create(data: dict[str, Any]) -> None:
assert not drain_completed[0]
writes.append(data)
mock_prisma.db.litellm_auditlog.create = AsyncMock(side_effect=_mock_create)
async def _hook_with_preceding_io() -> None:
await asyncio.sleep(0.05)
await create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by="admin",
table_name=LitellmTableNames.KEY_TABLE_NAME,
object_id="token-hook",
action="created",
updated_values="{}",
before_value=None,
)
)
with (
patch("litellm.store_audit_logs", True),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
):
hook_task: Final = track_audit_task(asyncio.create_task(_hook_with_preceding_io()))
await drain_audit_tasks(timeout=2.0)
drain_completed[0] = True
assert hook_task.done()
assert len(writes) == 1
assert writes[0]["object_id"] == "token-hook"
def test_drain_audit_tasks_default_timeout():
import inspect
sig: Final = inspect.signature(drain_audit_tasks)
assert sig.parameters["timeout"].default == 5.0