This commit is contained in:
Tufail Akram 2026-09-30 22:14:28 +00:00 • committed by GitHub
commit 168e3ae3ee
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
16 changed files with 809 additions and 441 deletions

View file

@ -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,
)
)
)
)

View file

@ -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:

View file

@ -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),
)
)
)
)

View file

@ -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),
)
)
)
)

View file

@ -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),
)
)
)
)

View file

@ -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,
)
)
)
)

View file

@ -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,

View file

@ -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,
)
)
)

View file

@ -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),
)
)
)
)

View file

@ -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,
)
)
)
)

View file

@ -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

View file

@ -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,
)
)
)

View file

@ -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,
)
)
)

View file

@ -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,
)
)
)

View file

@ -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]}...")

View file

@ -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"]