diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 16c4e367fac..5050fad4677 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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, + ) ) ) ######################################################### diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index fa856f5c2ab..71706a34fd3 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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 diff --git a/litellm/proxy/management_helpers/audit_logs.py b/litellm/proxy/management_helpers/audit_logs.py index b4e35a9b243..5d9142d4ebc 100644 --- a/litellm/proxy/management_helpers/audit_logs.py +++ b/litellm/proxy/management_helpers/audit_logs.py @@ -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: diff --git a/tests/unit/proxy/management_helpers/test_audit_logs_proxy.py b/tests/unit/proxy/management_helpers/test_audit_logs_proxy.py index 8b12e870e09..1e1186972b0 100644 --- a/tests/unit/proxy/management_helpers/test_audit_logs_proxy.py +++ b/tests/unit/proxy/management_helpers/test_audit_logs_proxy.py @@ -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