diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index 5dc0f659a8d..7c048637a89 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -51,6 +51,7 @@ class KeyManagementEventHooks: create_audit_log_for_update, get_audit_log_changed_by, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import litellm_proxy_admin_name @@ -63,22 +64,24 @@ class KeyManagementEventHooks: if is_audit_logging_enabled(): _updated_values: Final = response.model_dump_json(exclude_none=True) - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=get_audit_log_changed_by( - litellm_changed_by=litellm_changed_by, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ), - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.KEY_TABLE_NAME, - object_id=response.token_id or "", - action="created", - updated_values=_updated_values, - before_value=None, + track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=get_audit_log_changed_by( + litellm_changed_by=litellm_changed_by, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ), + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.KEY_TABLE_NAME, + object_id=response.token_id or "", + action="created", + updated_values=_updated_values, + before_value=None, + ) ) ) ) @@ -112,6 +115,7 @@ class KeyManagementEventHooks: create_audit_log_for_update, get_audit_log_changed_by, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import litellm_proxy_admin_name @@ -151,7 +155,7 @@ class KeyManagementEventHooks: if "project_id" in data.model_fields_set and data.project_id is None else audit_log ) - asyncio.create_task(create_audit_log_for_update(request_data=request_data)) + track_audit_task(asyncio.create_task(create_audit_log_for_update(request_data=request_data))) @staticmethod async def async_key_rotated_hook( @@ -165,6 +169,7 @@ class KeyManagementEventHooks: create_audit_log_for_update, get_audit_log_changed_by, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import litellm_proxy_admin_name @@ -203,22 +208,24 @@ class KeyManagementEventHooks: # store the audit log if is_audit_logging_enabled() and existing_key_row.token is not None: - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=get_audit_log_changed_by( - litellm_changed_by=litellm_changed_by, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ), - changed_by_api_key=user_api_key_dict.token, - table_name=LitellmTableNames.KEY_TABLE_NAME, - object_id=existing_key_row.token, - action="rotated", - updated_values=response.model_dump_json(exclude_none=True), - before_value=existing_key_row.model_dump_json(exclude_none=True), + track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=get_audit_log_changed_by( + litellm_changed_by=litellm_changed_by, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ), + changed_by_api_key=user_api_key_dict.token, + table_name=LitellmTableNames.KEY_TABLE_NAME, + object_id=existing_key_row.token, + action="rotated", + updated_values=response.model_dump_json(exclude_none=True), + before_value=existing_key_row.model_dump_json(exclude_none=True), + ) ) ) ) @@ -254,6 +261,7 @@ class KeyManagementEventHooks: create_audit_log_for_update, get_audit_log_changed_by, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import litellm_proxy_admin_name @@ -265,22 +273,24 @@ class KeyManagementEventHooks: continue _key_row = key_row.model_dump_json(exclude_none=True) - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=get_audit_log_changed_by( - litellm_changed_by=litellm_changed_by, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ), - changed_by_api_key=user_api_key_dict.token, - table_name=LitellmTableNames.KEY_TABLE_NAME, - object_id=key_row.token, - action="deleted", - updated_values="{}", - before_value=_key_row, + track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=get_audit_log_changed_by( + litellm_changed_by=litellm_changed_by, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ), + changed_by_api_key=user_api_key_dict.token, + table_name=LitellmTableNames.KEY_TABLE_NAME, + object_id=key_row.token, + action="deleted", + updated_values="{}", + before_value=_key_row, + ) ) ) ) diff --git a/litellm/proxy/hooks/user_management_event_hooks.py b/litellm/proxy/hooks/user_management_event_hooks.py index 6d978929c05..411f4f256f5 100644 --- a/litellm/proxy/hooks/user_management_event_hooks.py +++ b/litellm/proxy/hooks/user_management_event_hooks.py @@ -23,6 +23,7 @@ from litellm.proxy._types import ( from litellm.proxy.management_helpers.audit_logs import ( create_audit_log_for_update, is_audit_logging_enabled, + track_audit_task, ) from litellm.repositories.user_repository import UserRepository @@ -63,15 +64,17 @@ class UserManagementEventHooks: user_row: Final = await UserRepository(prisma_client).find_by_id(response.user_id) if user_row is None: raise Exception(f"no user row found for user_id={response.user_id}") - asyncio.create_task( - UserManagementEventHooks.create_internal_user_audit_log( - user_id=user_row.user_id, - action="created", - litellm_changed_by=user_api_key_dict.user_id, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - before_value=None, - after_value=user_row.model_dump_json(exclude_none=True), + track_audit_task( + asyncio.create_task( + UserManagementEventHooks.create_internal_user_audit_log( + user_id=user_row.user_id, + action="created", + litellm_changed_by=user_api_key_dict.user_id, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + before_value=None, + after_value=user_row.model_dump_json(exclude_none=True), + ) ) ) except Exception as e: diff --git a/litellm/proxy/management_endpoints/cache_settings_endpoints.py b/litellm/proxy/management_endpoints/cache_settings_endpoints.py index 40124bd19a4..8f67a735216 100644 --- a/litellm/proxy/management_endpoints/cache_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/cache_settings_endpoints.py @@ -304,24 +304,27 @@ async def _emit_cache_settings_audit_log( from litellm.proxy.management_helpers.audit_logs import ( create_audit_log_for_update, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import litellm_proxy_admin_name if not is_audit_logging_enabled(): return - task: Final = asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=litellm_changed_by or user_api_key_dict.user_id or litellm_proxy_admin_name, - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.CACHE_CONFIG_TABLE_NAME, - object_id="cache_config", - action=action, - updated_values=json.dumps({"settings": _redact_settings(after_settings)}, default=str), - before_value=json.dumps({"settings": _redact_settings(before_settings)}, default=str), + task: Final = track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=litellm_changed_by or user_api_key_dict.user_id or litellm_proxy_admin_name, + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.CACHE_CONFIG_TABLE_NAME, + object_id="cache_config", + action=action, + updated_values=json.dumps({"settings": _redact_settings(after_settings)}, default=str), + before_value=json.dumps({"settings": _redact_settings(before_settings)}, default=str), + ) ) ) ) diff --git a/litellm/proxy/management_endpoints/config_override_endpoints.py b/litellm/proxy/management_endpoints/config_override_endpoints.py index 77e5b0e5674..a310724a7a6 100644 --- a/litellm/proxy/management_endpoints/config_override_endpoints.py +++ b/litellm/proxy/management_endpoints/config_override_endpoints.py @@ -108,24 +108,27 @@ async def _emit_config_override_audit_log( from litellm.proxy.management_helpers.audit_logs import ( create_audit_log_for_update, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import litellm_proxy_admin_name if not is_audit_logging_enabled(): return - task: Final = asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=litellm_changed_by or user_api_key_dict.user_id or litellm_proxy_admin_name, - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.CONFIG_OVERRIDES_TABLE_NAME, - object_id=object_id, - action=action, - updated_values=json.dumps({"config": _redact_config(after_config)}, default=str), - before_value=json.dumps({"config": _redact_config(before_config)}, default=str), + task: Final = track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=litellm_changed_by or user_api_key_dict.user_id or litellm_proxy_admin_name, + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.CONFIG_OVERRIDES_TABLE_NAME, + object_id=object_id, + action=action, + updated_values=json.dumps({"config": _redact_config(after_config)}, default=str), + before_value=json.dumps({"config": _redact_config(before_config)}, default=str), + ) ) ) ) diff --git a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py index c59ee92f073..42461ad6ff0 100644 --- a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py +++ b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py @@ -246,24 +246,27 @@ async def _emit_coordination_redis_audit_log( from litellm.proxy.management_helpers.audit_logs import ( create_audit_log_for_update, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import litellm_proxy_admin_name if not is_audit_logging_enabled(): return - task: Final = asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=litellm_changed_by or user_api_key_dict.user_id or litellm_proxy_admin_name, - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.CONFIG_TABLE_NAME, - object_id=_COORDINATION_REDIS_KEY, - action=action, - updated_values=json.dumps({"settings": _redact_all_values(after_settings)}, default=str), - before_value=json.dumps({"settings": _redact_all_values(before_settings)}, default=str), + task: Final = track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=litellm_changed_by or user_api_key_dict.user_id or litellm_proxy_admin_name, + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.CONFIG_TABLE_NAME, + object_id=_COORDINATION_REDIS_KEY, + action=action, + updated_values=json.dumps({"settings": _redact_all_values(after_settings)}, default=str), + before_value=json.dumps({"settings": _redact_all_values(before_settings)}, default=str), + ) ) ) ) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index d0b3a08bc77..5dc08c916df 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -71,6 +71,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( generate_key_helper_fn, prepare_metadata_fields, ) +from litellm.proxy.management_helpers.audit_logs import track_audit_task from litellm.proxy.management_helpers.object_permission_utils import ( _set_object_permission, handle_update_object_permission_common, @@ -657,11 +658,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, + ) ) ) ######################################################### @@ -1358,15 +1361,19 @@ async def _schedule_user_update_audit_log( updated_user_row: Final = await _user_table(prisma_client).find_first(where={"user_id": response["user_id"]}) if updated_user_row: user_row_typed: Final = LiteLLM_UserTable.model_validate(updated_user_row.model_dump(exclude_none=True)) - asyncio.create_task( - UserManagementEventHooks.create_internal_user_audit_log( - user_id=user_row_typed.user_id, - action="updated", - litellm_changed_by=litellm_changed_by or user_api_key_dict.user_id, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - before_value=(existing_user_row.model_dump_json(exclude_none=True) if existing_user_row else None), - after_value=user_row_typed.model_dump_json(exclude_none=True), + track_audit_task( + asyncio.create_task( + UserManagementEventHooks.create_internal_user_audit_log( + user_id=user_row_typed.user_id, + action="updated", + litellm_changed_by=litellm_changed_by or user_api_key_dict.user_id, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + before_value=( + existing_user_row.model_dump_json(exclude_none=True) if existing_user_row else None + ), + after_value=user_row_typed.model_dump_json(exclude_none=True), + ) ) ) except Exception as audit_error: @@ -2004,15 +2011,17 @@ async def bulk_user_update( # Create single audit log entry for bulk operation try: - asyncio.create_task( - UserManagementEventHooks.create_internal_user_audit_log( - user_id=user_api_key_dict.user_id or "", - action="updated", - litellm_changed_by=litellm_changed_by or user_api_key_dict.user_id, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - before_value=f"Updated {len(all_users_in_db)} users", - after_value=json.dumps(non_default_values), + track_audit_task( + asyncio.create_task( + UserManagementEventHooks.create_internal_user_audit_log( + user_id=user_api_key_dict.user_id or "", + action="updated", + litellm_changed_by=litellm_changed_by or user_api_key_dict.user_id, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + before_value=f"Updated {len(all_users_in_db)} users", + after_value=json.dumps(non_default_values), + ) ) ) except Exception as audit_error: @@ -2514,22 +2523,24 @@ async def delete_user( # make an audit log for each team deleted _user_row = user_row.model_dump_json(exclude_none=True) - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=get_audit_log_changed_by( - litellm_changed_by=litellm_changed_by, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ), - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.USER_TABLE_NAME, - object_id=user_id, - action="deleted", - updated_values="{}", - before_value=_user_row, + track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=get_audit_log_changed_by( + litellm_changed_by=litellm_changed_by, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ), + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.USER_TABLE_NAME, + object_id=user_id, + action="deleted", + updated_values="{}", + before_value=_user_row, + ) ) ) ) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 323b9e434a8..856535bc8a0 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -111,6 +111,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, @@ -1553,12 +1554,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, + ) ) ) @@ -2889,13 +2892,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, + ) ) ) @@ -3587,13 +3592,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, + ) ) ) @@ -4163,13 +4170,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, + ) ) ) @@ -5705,13 +5714,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 @@ -7337,6 +7348,7 @@ async def block_key( from litellm.proxy.management_helpers.audit_logs import ( get_audit_log_changed_by, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import ( create_audit_log_for_update, @@ -7383,32 +7395,34 @@ async def block_key( code=status.HTTP_404_NOT_FOUND, ) - if is_audit_logging_enabled(): - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=get_audit_log_changed_by( - litellm_changed_by=litellm_changed_by, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ), - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.KEY_TABLE_NAME, - object_id=hashed_token, - action="blocked", - updated_values="{}", - before_value=existing_record.model_dump_json(), - ) - ) - ) - record: Final = await _prisma_table(VerificationTokenRepository(prisma_client)).update( where={"token": hashed_token}, data=with_settings_updated_at({"blocked": True}), ) + if is_audit_logging_enabled(): + track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=get_audit_log_changed_by( + litellm_changed_by=litellm_changed_by, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ), + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.KEY_TABLE_NAME, + object_id=hashed_token, + action="blocked", + updated_values="{}", + before_value=existing_record.model_dump_json(), + ) + ) + ) + ) + ## UPDATE KEY CACHE - invalidate so next read re-fetches from DB await _delete_cache_key_object( hashed_token=hashed_token, @@ -7451,6 +7465,7 @@ async def unblock_key( from litellm.proxy.management_helpers.audit_logs import ( get_audit_log_changed_by, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import ( create_audit_log_for_update, @@ -7497,32 +7512,34 @@ async def unblock_key( code=status.HTTP_404_NOT_FOUND, ) - if is_audit_logging_enabled(): - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=get_audit_log_changed_by( - litellm_changed_by=litellm_changed_by, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ), - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.KEY_TABLE_NAME, - object_id=hashed_token, - action="unblocked", - updated_values="{}", - before_value=existing_record.model_dump_json(), - ) - ) - ) - record: Final = await _prisma_table(VerificationTokenRepository(prisma_client)).update( where={"token": hashed_token}, data=with_settings_updated_at({"blocked": False}), ) + if is_audit_logging_enabled(): + track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=get_audit_log_changed_by( + litellm_changed_by=litellm_changed_by, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ), + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.KEY_TABLE_NAME, + object_id=hashed_token, + action="unblocked", + updated_values="{}", + before_value=existing_record.model_dump_json(), + ) + ) + ) + ) + ## UPDATE KEY CACHE - invalidate so next read re-fetches from DB await _delete_cache_key_object( hashed_token=hashed_token, diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 11a0075ffd0..57d55812e62 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -90,7 +90,7 @@ from litellm.proxy.management_helpers.access_group_model_sync import ( sync_access_groups_for_deleted_model, sync_access_groups_for_renamed_model, ) -from litellm.proxy.management_helpers.audit_logs import create_object_audit_log +from litellm.proxy.management_helpers.audit_logs import create_object_audit_log, track_audit_task from litellm.proxy.management_helpers.auto_router_permissions import ( MemberAutoRouterWrite, StoredAutoRouterIdentity, @@ -1325,16 +1325,18 @@ async def patch_model( reload_outcome: Final = await clear_cache() ## CREATE AUDIT LOG ## - asyncio.create_task( - create_object_audit_log( - object_id=model_id, - action="updated", - user_api_key_dict=user_api_key_dict, - table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, - before_value=db_model.model_dump_json(exclude_none=True), - after_value=updated_model.model_dump_json(exclude_none=True), - litellm_changed_by=user_api_key_dict.user_id, - litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME, + track_audit_task( + asyncio.create_task( + create_object_audit_log( + object_id=model_id, + action="updated", + user_api_key_dict=user_api_key_dict, + table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, + before_value=db_model.model_dump_json(exclude_none=True), + after_value=updated_model.model_dump_json(exclude_none=True), + litellm_changed_by=user_api_key_dict.user_id, + litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME, + ) ) ) @@ -1431,18 +1433,22 @@ async def _set_model_blocked_status( live_before_reload: Final = live_model_ids_snapshot() reload_outcome: Final = await clear_cache() - asyncio.create_task( - create_object_audit_log( - object_id=data.model_id, - action=action, - user_api_key_dict=user_api_key_dict, - table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, - before_value=db_model.model_dump_json(exclude_none=True), - after_value=( - updated_model.model_dump_json(exclude_none=True) if isinstance(updated_model, BaseModel) else None - ), - litellm_changed_by=litellm_changed_by, - litellm_proxy_admin_name=litellm_proxy_admin_name, + track_audit_task( + asyncio.create_task( + create_object_audit_log( + object_id=data.model_id, + action=action, + user_api_key_dict=user_api_key_dict, + table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, + before_value=db_model.model_dump_json(exclude_none=True), + after_value=( + updated_model.model_dump_json(exclude_none=True) + if isinstance(updated_model, BaseModel) + else None + ), + litellm_changed_by=litellm_changed_by, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ) ) ) @@ -2377,16 +2383,20 @@ async def delete_model( ) ## CREATE AUDIT LOG ## - asyncio.create_task( - create_object_audit_log( - object_id=model_info.id, - action="deleted", - user_api_key_dict=user_api_key_dict, - table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, - before_value=result.model_dump_json(exclude_none=True), - after_value=None, - litellm_changed_by=user_api_key_dict.user_id, - litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME, + track_audit_task( + asyncio.create_task( + create_object_audit_log( + object_id=model_info.id, + action="deleted", + user_api_key_dict=user_api_key_dict, + table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, + before_value=( + result.model_dump_json(exclude_none=True) if isinstance(result, BaseModel) else None + ), + after_value=None, + litellm_changed_by=user_api_key_dict.user_id, + litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME, + ) ) ) return {"message": f"Model: {result.model_id} deleted successfully"} @@ -2620,18 +2630,22 @@ async def add_new_model( ) ## CREATE AUDIT LOG ## - asyncio.create_task( - create_object_audit_log( - object_id=model_response.model_id, - action="created", - user_api_key_dict=user_api_key_dict, - table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, - before_value=None, - after_value=( - model_response.model_dump_json(exclude_none=True) if isinstance(model_response, BaseModel) else None - ), - litellm_changed_by=user_api_key_dict.user_id, - litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME, + track_audit_task( + asyncio.create_task( + create_object_audit_log( + object_id=model_response.model_id, + action="created", + user_api_key_dict=user_api_key_dict, + table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, + before_value=None, + after_value=( + model_response.model_dump_json(exclude_none=True) + if isinstance(model_response, BaseModel) + else None + ), + litellm_changed_by=user_api_key_dict.user_id, + litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME, + ) ) ) @@ -2839,24 +2853,26 @@ async def update_model( live_before_reload: Final = live_model_ids_snapshot() reload_outcome: Final = await clear_cache() ## CREATE AUDIT LOG ## - asyncio.create_task( - create_object_audit_log( - object_id=_model_id, - action="updated", - user_api_key_dict=user_api_key_dict, - table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, - before_value=( - existing_model_row.model_dump_json(exclude_none=True) - if isinstance(existing_model_row, BaseModel) - else None - ), - after_value=( - model_response.model_dump_json(exclude_none=True) - if isinstance(model_response, BaseModel) - else None - ), - litellm_changed_by=user_api_key_dict.user_id, - litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME, + track_audit_task( + asyncio.create_task( + create_object_audit_log( + object_id=_model_id, + action="updated", + user_api_key_dict=user_api_key_dict, + table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, + before_value=( + existing_model_row.model_dump_json(exclude_none=True) + if isinstance(existing_model_row, BaseModel) + else None + ), + after_value=( + model_response.model_dump_json(exclude_none=True) + if isinstance(model_response, BaseModel) + else None + ), + litellm_changed_by=user_api_key_dict.user_id, + litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME, + ) ) ) diff --git a/litellm/proxy/management_endpoints/team_callback_endpoints.py b/litellm/proxy/management_endpoints/team_callback_endpoints.py index ac6169d25bd..ddeb7a109d3 100644 --- a/litellm/proxy/management_endpoints/team_callback_endpoints.py +++ b/litellm/proxy/management_endpoints/team_callback_endpoints.py @@ -198,6 +198,7 @@ async def _emit_team_callback_audit_log( from litellm.proxy.management_helpers.audit_logs import ( create_audit_log_for_update, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import litellm_proxy_admin_name @@ -207,18 +208,20 @@ async def _emit_team_callback_audit_log( redacted_before: Final = _redact_callback_secrets(before_metadata) redacted_after: Final = _redact_callback_secrets(after_metadata) - task: Final = asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=litellm_changed_by or user_api_key_dict.user_id or litellm_proxy_admin_name, - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.TEAM_TABLE_NAME, - object_id=team_id, - action="updated", - updated_values=json.dumps({"metadata": redacted_after}, default=str), - before_value=json.dumps({"metadata": redacted_before}, default=str), + task: Final = track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=litellm_changed_by or user_api_key_dict.user_id or litellm_proxy_admin_name, + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.TEAM_TABLE_NAME, + object_id=team_id, + action="updated", + updated_values=json.dumps({"metadata": redacted_after}, default=str), + before_value=json.dumps({"metadata": redacted_before}, default=str), + ) ) ) ) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index fe976c861e5..9ad2947d044 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -1431,6 +1431,7 @@ async def new_team( from litellm.proxy.management_helpers.audit_logs import ( get_audit_log_changed_by, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import ( _license_check, @@ -1764,22 +1765,24 @@ async def new_team( _updated_values = json.dumps(_updated_values, default=str) - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=get_audit_log_changed_by( - litellm_changed_by=litellm_changed_by, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ), - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.TEAM_TABLE_NAME, - object_id=data.team_id, - action="created", - updated_values=_updated_values, - before_value=None, + track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=get_audit_log_changed_by( + litellm_changed_by=litellm_changed_by, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ), + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.TEAM_TABLE_NAME, + object_id=data.team_id, + action="created", + updated_values=_updated_values, + before_value=None, + ) ) ) ) @@ -1814,28 +1817,31 @@ async def _create_team_update_audit_log( from litellm.proxy.management_helpers.audit_logs import ( create_audit_log_for_update, get_audit_log_changed_by, + track_audit_task, ) _before_value = existing_team_row.json(exclude_none=True) _before_value = json.dumps(_before_value, default=str) _after_value: Final[str] = json.dumps(updated_kv, default=str) - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=get_audit_log_changed_by( - litellm_changed_by=litellm_changed_by, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ), - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.TEAM_TABLE_NAME, - object_id=team_id, - action="updated", - updated_values=_after_value, - before_value=_before_value, + track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=get_audit_log_changed_by( + litellm_changed_by=litellm_changed_by, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ), + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.TEAM_TABLE_NAME, + object_id=team_id, + action="updated", + updated_values=_after_value, + before_value=_before_value, + ) ) ) ) @@ -3179,21 +3185,24 @@ def _schedule_team_membership_audit_log( from litellm.proxy.management_helpers.audit_logs import ( create_object_audit_log, is_audit_logging_enabled, + track_audit_task, ) if not is_audit_logging_enabled() or tuple(before_members) == tuple(after_members): return - asyncio.create_task( - create_object_audit_log( - object_id=team_id, - action="updated", - litellm_changed_by=None, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - table_name=LitellmTableNames.TEAM_TABLE_NAME, - before_value=_members_audit_value(team_alias, before_members), - after_value=_members_audit_value(team_alias, after_members), + track_audit_task( + asyncio.create_task( + create_object_audit_log( + object_id=team_id, + action="updated", + litellm_changed_by=None, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + table_name=LitellmTableNames.TEAM_TABLE_NAME, + before_value=_members_audit_value(team_alias, before_members), + after_value=_members_audit_value(team_alias, after_members), + ) ) ) @@ -3212,6 +3221,7 @@ def _schedule_team_member_add_audit_logs( from litellm.proxy.management_helpers.audit_logs import ( create_object_audit_log, is_audit_logging_enabled, + track_audit_task, ) if not is_audit_logging_enabled(): @@ -3220,16 +3230,18 @@ def _schedule_team_member_add_audit_logs( for user in updated_users: if user.user_id in existing_user_ids: continue - asyncio.create_task( - create_object_audit_log( - object_id=user.user_id, - action="created", - litellm_changed_by=None, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - table_name=LitellmTableNames.USER_TABLE_NAME, - before_value=None, - after_value=safe_dumps(user.model_dump(exclude_none=True)), + track_audit_task( + asyncio.create_task( + create_object_audit_log( + object_id=user.user_id, + action="created", + litellm_changed_by=None, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + table_name=LitellmTableNames.USER_TABLE_NAME, + before_value=None, + after_value=safe_dumps(user.model_dump(exclude_none=True)), + ) ) ) @@ -4338,6 +4350,7 @@ async def delete_team( from litellm.proxy.management_helpers.audit_logs import ( get_audit_log_changed_by, is_audit_logging_enabled, + track_audit_task, ) from litellm.proxy.proxy_server import ( create_audit_log_for_update, @@ -4393,22 +4406,24 @@ async def delete_team( _team_row = team_row.json(exclude_none=True) - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=get_audit_log_changed_by( - litellm_changed_by=litellm_changed_by, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ), - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.TEAM_TABLE_NAME, - object_id=team_id, - action="deleted", - updated_values="{}", - before_value=_team_row, + track_audit_task( + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=get_audit_log_changed_by( + litellm_changed_by=litellm_changed_by, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ), + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.TEAM_TABLE_NAME, + object_id=team_id, + action="deleted", + updated_values="{}", + before_value=_team_row, + ) ) ) ) diff --git a/litellm/proxy/management_helpers/audit_logs.py b/litellm/proxy/management_helpers/audit_logs.py index ecd6abea3c3..dd2f03f9b52 100644 --- a/litellm/proxy/management_helpers/audit_logs.py +++ b/litellm/proxy/management_helpers/audit_logs.py @@ -5,7 +5,7 @@ Functions to create audit logs for LiteLLM Proxy import asyncio import json from datetime import datetime, timezone -from typing import Final +from typing import Final, TypeVar, cast import litellm from litellm._logging import verbose_proxy_logger @@ -124,14 +124,71 @@ def _build_audit_log_payload( ) -def _audit_log_task_done_callback(task: asyncio.Task) -> None: - """Log exceptions from audit log callback tasks so they don't slip through silently.""" - try: - exc: Final = task.exception() - except asyncio.CancelledError: - return - if exc is not None: - verbose_proxy_logger.error("Audit log callback task failed: %s", exc, exc_info=exc) +_T: Final = TypeVar("_T") + + +class AuditTaskRegistry: + __slots__ = ("_pending_tasks",) + + def __init__(self) -> None: + self._pending_tasks: Final[set[asyncio.Task[object]]] = set() # mutable-ok: background task tracking + + def track(self, task: asyncio.Task[_T]) -> asyncio.Task[_T]: + task_object: Final[asyncio.Task[object]] = cast(asyncio.Task[object], task) # cast-ok: task covariance + if task_object.done() or task_object in self._pending_tasks: + return task + self._pending_tasks.add(task_object) + task_object.add_done_callback(self._on_task_done) + return task + + def _on_task_done(self, task: asyncio.Task[object]) -> None: + self._pending_tasks.discard(task) + try: + exc: Final = task.exception() + except asyncio.CancelledError: + return + 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 = 5.0) -> None: + deadline: Final = asyncio.get_running_loop().time() + timeout + while self._pending_tasks: + remaining = deadline - asyncio.get_running_loop().time() + if remaining <= 0: + timed_out_count = len(self._pending_tasks) + self._pending_tasks.clear() + verbose_proxy_logger.warning("Timed out draining %d audit tasks on shutdown", timed_out_count) + break + tasks = tuple(self._pending_tasks) + try: + _, pending = await asyncio.wait(tasks, timeout=remaining) + if pending: + for task in pending: + self._pending_tasks.discard(task) + verbose_proxy_logger.warning("Timed out draining %d audit tasks on shutdown", len(pending)) + break + except Exception as e: + verbose_proxy_logger.warning("Error draining audit tasks on shutdown: %s", e) + break + + def clear(self) -> None: + self._pending_tasks.clear() + + +_audit_task_registry: Final = AuditTaskRegistry() +_pending_audit_tasks: Final = _audit_task_registry._pending_tasks + + +def _audit_log_task_done_callback(task: asyncio.Task[object]) -> None: + _audit_task_registry._on_task_done(task) + + +def track_audit_task(task: asyncio.Task[_T]) -> asyncio.Task[_T]: + return _audit_task_registry.track(task) + + +async def drain_audit_tasks(timeout: float = 5.0) -> None: + await _audit_task_registry.drain(timeout=timeout) async def _dispatch_audit_log_to_callbacks( @@ -153,8 +210,7 @@ async def _dispatch_audit_log_to_callbacks( continue if isinstance(resolved, CustomLogger): - task = asyncio.create_task(resolved.async_log_audit_log_event(payload)) - task.add_done_callback(_audit_log_task_done_callback) + track_audit_task(asyncio.create_task(resolved.async_log_audit_log_event(payload))) except Exception as e: verbose_proxy_logger.error("Failed dispatching audit log to callback: %s", e) @@ -205,12 +261,13 @@ async def create_object_audit_log( async def create_audit_log_for_update(request_data: LiteLLM_AuditLogs): - """ - Create an audit log for an object. - """ if not is_audit_logging_enabled(): return + current_task: Final = asyncio.current_task() + if current_task is not None: + track_audit_task(current_task) + from litellm.proxy.proxy_server import premium_user, prisma_client if premium_user is not True: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ed2a5b89a32..e008b6b6255 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -695,6 +695,7 @@ from litellm.proxy.management_endpoints.workflow_management_endpoints import ( from litellm.proxy.management_helpers.audit_logs import ( create_audit_log_for_update, create_object_audit_log, + track_audit_task, ) from litellm.proxy.management_helpers.team_metadata_validation import ( TEAM_METADATA_SCHEMA_REGISTRY, @@ -1143,6 +1144,12 @@ async def proxy_shutdown_event(worker_heartbeat: ProxyWorkerHeartbeat | None = N verbose_proxy_logger.info("Shutting down LiteLLM Proxy Server") if worker_heartbeat is not None and prisma_client: await worker_heartbeat.deregister() + try: + from litellm.proxy.management_helpers.audit_logs import drain_audit_tasks + + await drain_audit_tasks() + except Exception as e: # noqa: BLE001 # shutdown must continue even if the drain fails + verbose_proxy_logger.exception("Error draining audit tasks on shutdown: %s", e) if prisma_client: # Drain the SGR fold first: it lives in memory, so an un-drained interval # is lost, and a write attempted after disconnect raises @@ -18179,9 +18186,11 @@ async def update_config( existing["alerting"].append("slack") existing[k] = v await _upsert_section("general_settings", existing) - asyncio.create_task( - create_config_audit_log( - "general_settings", "updated", before_general_settings, existing, user_api_key_dict + track_audit_task( + asyncio.create_task( + create_config_audit_log( + "general_settings", "updated", before_general_settings, existing, user_api_key_dict + ) ) ) @@ -18197,9 +18206,11 @@ async def update_config( proxy_config._encrypt_env_variables_for_db(environment_variables=config_info.environment_variables) ) await _upsert_section("environment_variables", existing) - asyncio.create_task( - create_config_audit_log( - "environment_variables", "updated", before_environment_variables, existing, user_api_key_dict + track_audit_task( + asyncio.create_task( + create_config_audit_log( + "environment_variables", "updated", before_environment_variables, existing, user_api_key_dict + ) ) ) @@ -18229,9 +18240,11 @@ async def update_config( merged["success_callback"] = list(set(incoming_cb)) await _upsert_section("litellm_settings", merged) - asyncio.create_task( - create_config_audit_log( - "litellm_settings", "updated", before_litellm_settings, merged, user_api_key_dict + track_audit_task( + asyncio.create_task( + create_config_audit_log( + "litellm_settings", "updated", before_litellm_settings, merged, user_api_key_dict + ) ) ) @@ -18241,9 +18254,11 @@ async def update_config( before_router_settings: Final = copy.deepcopy(existing) new_router_settings: Final = {**existing, **router_settings_updates} await _upsert_section("router_settings", new_router_settings) - asyncio.create_task( - create_config_audit_log( - "router_settings", "updated", before_router_settings, new_router_settings, user_api_key_dict + track_audit_task( + asyncio.create_task( + create_config_audit_log( + "router_settings", "updated", before_router_settings, new_router_settings, user_api_key_dict + ) ) ) @@ -18456,9 +18471,11 @@ async def update_config_general_settings( if is_resource_list("general_settings", data.field_name): stored_endpoints: Final = general_settings.get("pass_through_endpoints") await proxy_config._serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ()) - asyncio.create_task( - create_config_audit_log( - "general_settings", "updated", before_general_settings, general_settings, user_api_key_dict + track_audit_task( + asyncio.create_task( + create_config_audit_log( + "general_settings", "updated", before_general_settings, general_settings, user_api_key_dict + ) ) ) @@ -18807,7 +18824,9 @@ async def _persist_general_settings_ui_litellm_field( config["litellm_settings"] = {} config["litellm_settings"][field_name] = validated await proxy_config.save_config(new_config=config) - asyncio.create_task(create_config_audit_log(field_name, "updated", before_value, validated, user_api_key_dict)) + track_audit_task( + asyncio.create_task(create_config_audit_log(field_name, "updated", before_value, validated, user_api_key_dict)) + ) return {"message": f"Field {field_name} updated", "status": "success"} @@ -18820,7 +18839,11 @@ async def _reset_general_settings_ui_litellm_field(field_name: str, user_api_key if "litellm_settings" in config: config["litellm_settings"].pop(field_name, None) await proxy_config.save_config(new_config=config) - asyncio.create_task(create_config_audit_log(field_name, "deleted", before_value, default_value, user_api_key_dict)) + track_audit_task( + asyncio.create_task( + create_config_audit_log(field_name, "deleted", before_value, default_value, user_api_key_dict) + ) + ) return {"message": f"Field {field_name} reset", "status": "success"} @@ -19078,9 +19101,11 @@ async def delete_config_general_settings( if is_resource_list("general_settings", data.field_name): stored_endpoints: Final = general_settings.get("pass_through_endpoints") await proxy_config._serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ()) - asyncio.create_task( - create_config_audit_log( - "general_settings", "deleted", before_general_settings, general_settings, user_api_key_dict + track_audit_task( + asyncio.create_task( + create_config_audit_log( + "general_settings", "deleted", before_general_settings, general_settings, user_api_key_dict + ) ) ) @@ -19144,13 +19169,15 @@ async def delete_callback( # Save the updated configuration await proxy_config.save_config(new_config=config) - asyncio.create_task( - create_config_audit_log( - "litellm_settings", - "deleted", - {"success_callback": before_success_callbacks}, - {"success_callback": success_callbacks}, - user_api_key_dict, + track_audit_task( + asyncio.create_task( + create_config_audit_log( + "litellm_settings", + "deleted", + {"success_callback": before_success_callbacks}, + {"success_callback": success_callbacks}, + user_api_key_dict, + ) ) ) diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index a112b22b4e4..36bd4db0d29 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -38,6 +38,7 @@ from litellm.proxy.management_endpoints.team_admin_field_permissions import ( SUPPORTED_TEAM_ADMIN_PERMISSIONS, TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING, ) +from litellm.proxy.management_helpers.audit_logs import track_audit_task from litellm.proxy.spend_tracking.ptu_feature_flag import ( PTU_COST_ATTRIBUTION_ENV_VAR, is_ptu_cost_attribution_enabled, @@ -657,13 +658,15 @@ async def add_allowed_ip( await proxy_config.save_config(new_config=config) - asyncio.create_task( - create_config_audit_log( - param_name="general_settings", - action="updated", - before_value={"allowed_ips": before_allowed_ips}, - after_value={"allowed_ips": config["general_settings"]["allowed_ips"]}, - user_api_key_dict=user_api_key_dict, + track_audit_task( + asyncio.create_task( + create_config_audit_log( + param_name="general_settings", + action="updated", + before_value={"allowed_ips": before_allowed_ips}, + after_value={"allowed_ips": config["general_settings"]["allowed_ips"]}, + user_api_key_dict=user_api_key_dict, + ) ) ) @@ -708,13 +711,15 @@ async def delete_allowed_ip( await proxy_config.save_config(new_config=config) - asyncio.create_task( - create_config_audit_log( - param_name="general_settings", - action="deleted", - before_value={"allowed_ips": before_allowed_ips}, - after_value={"allowed_ips": config["general_settings"]["allowed_ips"]}, - user_api_key_dict=user_api_key_dict, + track_audit_task( + asyncio.create_task( + create_config_audit_log( + param_name="general_settings", + action="deleted", + before_value={"allowed_ips": before_allowed_ips}, + after_value={"allowed_ips": config["general_settings"]["allowed_ips"]}, + user_api_key_dict=user_api_key_dict, + ) ) ) @@ -1056,13 +1061,15 @@ async def _update_litellm_setting( # never surfaces as a 500 after save_config has already committed, # matching the create_object_audit_log pattern used elsewhere # (e.g. model_management_endpoints). - asyncio.create_task( - create_config_audit_log( - param_name=settings_key, - action="updated", - before_value=before_value, - after_value=in_memory_var, - user_api_key_dict=user_api_key_dict, + track_audit_task( + asyncio.create_task( + create_config_audit_log( + param_name=settings_key, + action="updated", + before_value=before_value, + after_value=in_memory_var, + user_api_key_dict=user_api_key_dict, + ) ) ) @@ -1271,14 +1278,16 @@ async def update_sso_settings( }, ) - asyncio.create_task( - create_config_audit_log( - param_name="sso_config", - action="updated", - before_value=before_sso_data, - after_value=sso_data, - user_api_key_dict=user_api_key_dict, - table_name=LitellmTableNames.SSO_CONFIG_TABLE_NAME, + track_audit_task( + asyncio.create_task( + create_config_audit_log( + param_name="sso_config", + action="updated", + before_value=before_sso_data, + after_value=sso_data, + user_api_key_dict=user_api_key_dict, + table_name=LitellmTableNames.SSO_CONFIG_TABLE_NAME, + ) ) ) @@ -1447,13 +1456,15 @@ async def update_ui_theme_settings( # Persist only the two owned env vars, merged against the existing DB row. await proxy_config.save_environment_variables(env_updates) - asyncio.create_task( - create_config_audit_log( - param_name="ui_theme_config", - action="updated", - before_value=before_theme, - after_value=theme_data, - user_api_key_dict=user_api_key_dict, + track_audit_task( + asyncio.create_task( + create_config_audit_log( + param_name="ui_theme_config", + action="updated", + before_value=before_theme, + after_value=theme_data, + user_api_key_dict=user_api_key_dict, + ) ) ) @@ -1947,14 +1958,16 @@ async def update_ui_settings( sanitized: Final = {k: v for k, v in ui_settings.items() if k in ALLOWED_UI_SETTINGS_FIELDS} await user_api_key_cache.async_set_cache(key=UI_SETTINGS_CACHE_KEY, value=sanitized, ttl=UI_SETTINGS_CACHE_TTL) - asyncio.create_task( - create_config_audit_log( - param_name="ui_settings", - action="updated", - before_value=existing, - after_value=ui_settings, - user_api_key_dict=user_api_key_dict, - table_name=LitellmTableNames.UI_SETTINGS_TABLE_NAME, + track_audit_task( + asyncio.create_task( + create_config_audit_log( + param_name="ui_settings", + action="updated", + before_value=existing, + after_value=ui_settings, + user_api_key_dict=user_api_key_dict, + table_name=LitellmTableNames.UI_SETTINGS_TABLE_NAME, + ) ) ) diff --git a/litellm/proxy/ui_crud_endpoints/user_banner_endpoints.py b/litellm/proxy/ui_crud_endpoints/user_banner_endpoints.py index 0ff61592d9f..a350611a674 100644 --- a/litellm/proxy/ui_crud_endpoints/user_banner_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/user_banner_endpoints.py @@ -8,6 +8,7 @@ from pydantic import BaseModel, Field, ValidationError, model_validator from litellm._uuid import uuid4 from litellm.proxy._types import LitellmTableNames, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_helpers.audit_logs import track_audit_task from litellm.repositories.user_banner_repository import USER_BANNER_ROW_ID, UserBannerRepository router: Final = APIRouter() @@ -116,14 +117,16 @@ async def update_user_banner( await repository.upsert_settings(json.dumps(banner.model_dump())) - asyncio.create_task( - create_config_audit_log( - param_name=USER_BANNER_ROW_ID, - action="updated", - before_value=before.model_dump(), - after_value=banner.model_dump(), - user_api_key_dict=user_api_key_dict, - table_name=LitellmTableNames.UI_SETTINGS_TABLE_NAME, + track_audit_task( + asyncio.create_task( + create_config_audit_log( + param_name=USER_BANNER_ROW_ID, + action="updated", + before_value=before.model_dump(), + after_value=banner.model_dump(), + user_api_key_dict=user_api_key_dict, + table_name=LitellmTableNames.UI_SETTINGS_TABLE_NAME, + ) ) ) 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 98922801296..e9678bf1e96 100644 --- a/tests/unit/proxy/management_helpers/test_audit_logs_proxy.py +++ b/tests/unit/proxy/management_helpers/test_audit_logs_proxy.py @@ -1,7 +1,8 @@ import os import traceback from litellm._uuid import uuid -from datetime import datetime +from datetime import datetime, timezone +from typing import Any, Final from dotenv import load_dotenv from fastapi import Request @@ -37,12 +38,15 @@ from litellm.proxy.utils import PrismaClient, ProxyLogging, hash_token, update_s 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, + track_audit_task, ) from litellm.proxy._types import LiteLLM_AuditLogs, LitellmTableNames, UserAPIKeyAuth from litellm.caching.caching import DualCache -from unittest.mock import patch, AsyncMock +from unittest.mock import AsyncMock, MagicMock, patch proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) import json @@ -173,7 +177,6 @@ async def test_create_audit_log_for_update_premium_user(): patch("litellm.store_audit_logs", True), patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, ): - mock_prisma.db.litellm_auditlog.create = AsyncMock() request_data = LiteLLM_AuditLogs( @@ -214,9 +217,7 @@ def prisma_client(): os.environ["DATABASE_URL"] = modified_url # Assuming PrismaClient is a class that needs to be instantiated - prisma_client = PrismaClient( - database_url=os.environ["DATABASE_URL"], proxy_logging_obj=proxy_logging_obj - ) + prisma_client = PrismaClient(database_url=os.environ["DATABASE_URL"], proxy_logging_obj=proxy_logging_obj) return prisma_client @@ -251,10 +252,278 @@ async def test_create_audit_log_in_db(prisma_client): await asyncio.sleep(1) # now read the last log from the db - last_log = await prisma_client.db.litellm_auditlog.find_first( - where={"id": audit_log_id} - ) + last_log = await prisma_client.db.litellm_auditlog.find_first(where={"id": audit_log_id}) assert last_log.id == audit_log_id setattr(litellm, "store_audit_logs", False) + + +@pytest.mark.asyncio +async def test_track_audit_task_lifecycle(): + task_completed: Final = asyncio.Event() + + async def _sample_coroutine() -> None: + await asyncio.sleep(0.01) + task_completed.set() + + task: Final = track_audit_task(asyncio.create_task(_sample_coroutine())) + assert not task.done() + + await drain_audit_tasks(timeout=1.0) + assert task_completed.is_set() + assert task.done() + + completed_task: Final = asyncio.create_task(asyncio.sleep(0)) + await completed_task + assert track_audit_task(completed_task) is completed_task + + +@pytest.mark.asyncio +async def test_drain_audit_tasks_waits_for_all_tasks(): + event1: Final = asyncio.Event() + event2: Final = asyncio.Event() + + async def _worker(event: asyncio.Event) -> None: + await asyncio.sleep(0.02) + event.set() + + task1: Final = track_audit_task(asyncio.create_task(_worker(event1))) + task2: Final = track_audit_task(asyncio.create_task(_worker(event2))) + + assert not task1.done() + assert not task2.done() + + await drain_audit_tasks(timeout=1.0) + + assert event1.is_set() + assert event2.is_set() + assert task1.done() + assert task2.done() + + +@pytest.mark.asyncio +async def test_drain_audit_tasks_handles_failing_task_cleanly(): + async def _failing_worker() -> None: + await asyncio.sleep(0.01) + raise RuntimeError("test audit log write failure") + + task: Final = track_audit_task(asyncio.create_task(_failing_worker())) + assert not task.done() + + await drain_audit_tasks(timeout=1.0) + assert task.done() + assert isinstance(task.exception(), RuntimeError) + + +@pytest.mark.asyncio +async def test_drain_audit_tasks_timeout_does_not_cancel_pending_tasks(): + async def _long_running_worker() -> None: + await asyncio.sleep(10.0) + + task: Final = track_audit_task(asyncio.create_task(_long_running_worker())) + try: + assert not task.done() + await drain_audit_tasks(timeout=0.05) + assert not task.cancelled() + assert not task.done() + finally: + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + +@pytest.mark.asyncio +async def test_drain_audit_tasks_persists_management_audit_write_before_db_disconnect(): + mock_prisma: Final = MagicMock() + writes: Final[list[dict[str, Any]]] = [] + is_disconnected: Final = [False] + + async def _mock_create(data: dict[str, Any]) -> None: + if is_disconnected[0]: + raise RuntimeError("Database already disconnected") + await asyncio.sleep(0.02) + writes.append(data) + + mock_prisma.db.litellm_auditlog.create = AsyncMock(side_effect=_mock_create) + + with ( + patch("litellm.store_audit_logs", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + ): + request_data: Final = LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by="admin-user", + changed_by_api_key="sk-admin", + table_name=LitellmTableNames.KEY_TABLE_NAME, + object_id="token-123", + action="blocked", + updated_values="{}", + before_value='{"key": "test"}', + ) + + task: Final = track_audit_task(asyncio.create_task(create_audit_log_for_update(request_data=request_data))) + assert not task.done() + + await drain_audit_tasks(timeout=2.0) + is_disconnected[0] = True + + assert task.done() + assert len(writes) == 1 + assert writes[0]["object_id"] == "token-123" + assert writes[0]["action"] == "blocked" + + +@pytest.mark.asyncio +async def test_create_audit_log_auto_registers_in_drain(): + mock_prisma: Final = MagicMock() + writes: Final[list[dict[str, Any]]] = [] + + async def _mock_create(data: dict[str, Any]) -> None: + await asyncio.sleep(0.02) + writes.append(data) + + mock_prisma.db.litellm_auditlog.create = AsyncMock(side_effect=_mock_create) + + with ( + patch("litellm.store_audit_logs", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + ): + request_data: Final = LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by="admin-user", + changed_by_api_key="sk-admin", + table_name=LitellmTableNames.KEY_TABLE_NAME, + object_id="token-456", + action="deleted", + updated_values="{}", + before_value='{"key": "test-deleted"}', + ) + + task: Final = asyncio.create_task(create_audit_log_for_update(request_data=request_data)) + await asyncio.sleep(0.001) + + await drain_audit_tasks(timeout=2.0) + + assert task.done() + assert len(writes) == 1 + assert writes[0]["object_id"] == "token-456" + assert writes[0]["action"] == "deleted" + + +@pytest.mark.asyncio +async def test_drain_audit_tasks_discards_timed_out_tasks_on_timeout(): + async def _hung_worker() -> None: + await asyncio.Event().wait() + + hung_task: Final = track_audit_task(asyncio.create_task(_hung_worker())) + try: + start_time: Final = time.perf_counter() + await drain_audit_tasks(timeout=0.05) + duration_first_drain: Final = time.perf_counter() - start_time + assert duration_first_drain >= 0.04 + + second_drain_start: Final = time.perf_counter() + await drain_audit_tasks(timeout=0.5) + duration_second_drain: Final = time.perf_counter() - second_drain_start + assert duration_second_drain < 0.05 + finally: + hung_task.cancel() + try: + await hung_task + except asyncio.CancelledError: + pass + + +@pytest.mark.asyncio +async def test_drain_audit_tasks_captures_tasks_queued_during_drain(): + finished_tasks: Final[list[str]] = [] + + async def _first_worker() -> None: + await asyncio.sleep(0.01) + track_audit_task(asyncio.create_task(_second_worker())) + finished_tasks.append("first") + + async def _second_worker() -> None: + await asyncio.sleep(0.01) + finished_tasks.append("second") + + track_audit_task(asyncio.create_task(_first_worker())) + 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