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 9d182d4e259..175149b8704 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 59d8dd821d8..7edb662b404 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -68,6 +68,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, @@ -1352,15 +1353,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: @@ -1991,15 +1996,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: @@ -2501,22 +2508,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 d37dfe87ad5..c3b733ae6f6 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -7332,6 +7332,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, @@ -7378,32 +7379,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, @@ -7446,6 +7449,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, @@ -7492,32 +7496,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 ae294871afc..182b49ab855 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -82,7 +82,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, @@ -1282,16 +1282,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, + ) ) ) @@ -1388,18 +1390,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, + ) ) ) @@ -2274,16 +2280,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"} @@ -2515,18 +2525,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, + ) ) ) @@ -2735,24 +2749,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 4091d69e44e..74d90d5f7d9 100644 --- a/litellm/proxy/management_endpoints/team_callback_endpoints.py +++ b/litellm/proxy/management_endpoints/team_callback_endpoints.py @@ -200,6 +200,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 @@ -209,18 +210,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 f0e59389d48..0d64551087b 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -1469,6 +1469,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, @@ -1802,22 +1803,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, + ) ) ) ) @@ -1852,28 +1855,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, + ) ) ) ) @@ -3225,21 +3231,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), + ) ) ) @@ -3258,6 +3267,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(): @@ -3266,16 +3276,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)), + ) ) ) @@ -4388,6 +4400,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, @@ -4445,22 +4458,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..b4e35a9b243 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 = 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) + 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 = 10.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 = 10.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,9 +261,10 @@ async def create_object_audit_log( async def create_audit_log_for_update(request_data: LiteLLM_AuditLogs): - """ - Create an audit log for an object. - """ + current_task: Final = asyncio.current_task() + if current_task is not None: + track_audit_task(current_task) + if not is_audit_logging_enabled(): return diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f89dde04e06..c576f13993e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -674,6 +674,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, @@ -1112,6 +1113,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 @@ -18026,9 +18033,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 + ) ) ) @@ -18044,9 +18053,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 + ) ) ) @@ -18076,9 +18087,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 + ) ) ) @@ -18088,9 +18101,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 + ) ) ) @@ -18300,9 +18315,11 @@ async def update_config_general_settings( ) await invalidate_config_param("general_settings") proxy_config.settings.apply_db_row("general_settings", general_settings) - 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 + ) ) ) @@ -18640,7 +18657,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"} @@ -18653,7 +18672,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"} @@ -18908,9 +18931,11 @@ async def delete_config_general_settings( ) await invalidate_config_param("general_settings") proxy_config.settings.apply_db_row("general_settings", general_settings) - 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 + ) ) ) @@ -18974,13 +18999,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 227f0e7f795..e9cf92b6fff 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -37,6 +37,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, @@ -656,13 +657,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, + ) ) ) @@ -707,13 +710,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, + ) ) ) @@ -1057,13 +1062,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, + ) ) ) @@ -1272,14 +1279,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, + ) ) ) @@ -1448,13 +1457,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, + ) ) ) @@ -1945,14 +1956,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 893fda797cd..8d806eb0695 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/interactions/test_openapi_compliance.py b/tests/unit/interactions/test_openapi_compliance.py index d3f1183cea6..98574134ef3 100644 --- a/tests/unit/interactions/test_openapi_compliance.py +++ b/tests/unit/interactions/test_openapi_compliance.py @@ -32,8 +32,7 @@ def _load_openapi_spec_dict() -> Dict[str, Any]: return response.json() except Exception as e: # pragma: no cover - defensive, env-dependent pytest.skip( - f"Skipping Google Interactions OpenAPI compliance tests - " - f"unable to load spec from {OPENAPI_SPEC_URL}: {e}" + f"Skipping Google Interactions OpenAPI compliance tests - unable to load spec from {OPENAPI_SPEC_URL}: {e}" ) @@ -61,11 +60,16 @@ class TestRequestCompliance: def test_create_model_interaction_request_schema(self, spec_dict): """Verify CreateModelInteractionParams schema fields.""" - schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"] + schemas = spec_dict["components"]["schemas"] + schema = schemas.get("ModelInteraction") or schemas.get("CreateModelInteractionParams") + assert schema is not None, "ModelInteraction schema not found" # Required fields per spec assert "model" in schema["required"] - assert "input" in schema["required"] + assert "model" in schema["properties"] + assert "input" in schema["properties"] + if "CreateModelInteractionParams" in spec_dict["components"]["schemas"]: + assert "input" in schema["required"] # Check our supported optional fields exist in spec our_optional_fields = [ @@ -88,7 +92,9 @@ class TestRequestCompliance: def test_input_types_match_spec(self, spec_dict): """Verify input field supports string, Content, Content[], Turn[].""" - schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"] + schemas = spec_dict["components"]["schemas"] + schema = schemas.get("ModelInteraction") or schemas.get("CreateModelInteractionParams") + assert schema is not None, "ModelInteraction schema not found" input_schema = schema["properties"]["input"] # The input property may be inline oneOf or a $ref to InteractionsInput @@ -125,22 +131,18 @@ class TestRequestCompliance: discriminator = content_schema.get("discriminator") if discriminator is not None: - assert ( - discriminator.get("propertyName") == "type" - ), f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'" + assert discriminator.get("propertyName") == "type", ( + f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'" + ) variant_names = [ - option["$ref"].split("/")[-1] - for option in content_schema.get("oneOf", []) - if "$ref" in option + option["$ref"].split("/")[-1] for option in content_schema.get("oneOf", []) if "$ref" in option ] assert variant_names, f"Content is not a union of named variants: {content_schema}" mapping = (discriminator or {}).get("mapping") or {} type_values = { - variant: mapping_value - for mapping_value, ref in mapping.items() - for variant in [ref.split("/")[-1]] + variant: mapping_value for mapping_value, ref in mapping.items() for variant in [ref.split("/")[-1]] } or { variant: _declared_type_value(spec_dict["components"]["schemas"].get(variant, {})) for variant in variant_names @@ -191,7 +193,9 @@ class TestRequestCompliance: for option in spec_dict["components"]["schemas"]["Step"]["oneOf"] if "$ref" in option } - assert {"UserInputStep", "ModelOutputStep"} <= step_variants, f"Step union is missing role steps: {step_variants}" + assert {"UserInputStep", "ModelOutputStep"} <= step_variants, ( + f"Step union is missing role steps: {step_variants}" + ) for step_name, type_value in [("UserInputStep", "user_input"), ("ModelOutputStep", "model_output")]: step_schema = spec_dict["components"]["schemas"][step_name] @@ -261,9 +265,7 @@ class TestResponseCompliance: expected_fields = ["total_input_tokens", "total_output_tokens", "total_tokens"] for field in expected_fields: - assert ( - field in usage_schema["properties"] - ), f"Usage field '{field}' not in spec" + assert field in usage_schema["properties"], f"Usage field '{field}' not in spec" print(f"✓ Usage field '{field}' exists") @@ -282,9 +284,7 @@ class TestToolsCompliance: """Verify FunctionDeclaration schema for function tools.""" if "FunctionDeclaration" in spec_dict["components"]["schemas"]: func_schema = spec_dict["components"]["schemas"]["FunctionDeclaration"] - assert "name" in func_schema.get( - "properties", {} - ) or "name" in func_schema.get("required", []) + assert "name" in func_schema.get("properties", {}) or "name" in func_schema.get("required", []) print("✓ FunctionDeclaration schema found") else: print("⚠ FunctionDeclaration schema not found (may be nested)") @@ -313,7 +313,7 @@ class TestEndpointCompliance: get_path = None for path, methods in paths.items(): - if "{id}" in path and "interactions" in path and "get" in methods: + if ("{id}" in path or "{interactionsId}" in path) and "interactions" in path and "get" in methods: get_path = path break @@ -326,7 +326,7 @@ class TestEndpointCompliance: delete_path = None for path, methods in paths.items(): - if "{id}" in path and "interactions" in path and "delete" in methods: + if ("{id}" in path or "{interactionsId}" in path) and "interactions" in path and "delete" in methods: delete_path = path break @@ -350,6 +350,4 @@ if __name__ == "__main__": if method in ["get", "post", "delete", "put", "patch"]: print(f" {method.upper()} {path}") - print( - f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}..." - ) + print(f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}...") 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 878e19f5b6f..8b12e870e09 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 @@ -41,11 +42,13 @@ from starlette.datastructures import URL from litellm.proxy.management_helpers.audit_logs import ( 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 @@ -176,7 +179,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( @@ -217,9 +219,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 @@ -254,10 +254,209 @@ 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"]