mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(proxy): track hook spawns and defer registration in audit log drain
This commit is contained in:
parent
0b9f2c222b
commit
908d027113
4 changed files with 127 additions and 44 deletions
|
|
@ -656,11 +656,13 @@ async def new_user(
|
|||
#########################################################
|
||||
########## USER CREATED HOOK ################
|
||||
#########################################################
|
||||
asyncio.create_task(
|
||||
UserManagementEventHooks.async_user_created_hook(
|
||||
data=data,
|
||||
response=new_user_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
track_audit_task(
|
||||
asyncio.create_task(
|
||||
UserManagementEventHooks.async_user_created_hook(
|
||||
data=data,
|
||||
response=new_user_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
)
|
||||
)
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -110,6 +110,7 @@ from litellm.proxy.management_helpers.access_group_key_sync import (
|
|||
sync_key_regeneration_access_group_membership,
|
||||
sync_key_update_access_group_membership,
|
||||
)
|
||||
from litellm.proxy.management_helpers.audit_logs import track_audit_task
|
||||
from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
_set_object_permission,
|
||||
|
|
@ -1550,12 +1551,14 @@ async def _common_key_generation_helper(
|
|||
|
||||
response.token = response.token_id # remap token to use the hash, and leave the key in the `key` field [TODO]: clean up generate_key_helper_fn to do this
|
||||
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_generated_hook(
|
||||
data=data,
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
track_audit_task(
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_generated_hook(
|
||||
data=data,
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -2886,13 +2889,15 @@ async def _process_single_key_update(
|
|||
)
|
||||
|
||||
# Trigger async hook
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_updated_hook(
|
||||
data=key_request,
|
||||
existing_key_row=existing_key_row,
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
track_audit_task(
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_updated_hook(
|
||||
data=key_request,
|
||||
existing_key_row=existing_key_row,
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -3583,13 +3588,15 @@ async def update_key_fn(
|
|||
redis_err,
|
||||
)
|
||||
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_updated_hook(
|
||||
data=data,
|
||||
existing_key_row=existing_key_row,
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
track_audit_task(
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_updated_hook(
|
||||
data=data,
|
||||
existing_key_row=existing_key_row,
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -4159,13 +4166,15 @@ async def delete_key_fn(
|
|||
"/keys/delete - cache after delete: %s", user_api_key_cache.key_object_cache.in_memory_cache.cache_dict
|
||||
)
|
||||
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_deleted_hook(
|
||||
data=data,
|
||||
keys_being_deleted=_keys_being_deleted,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
response=number_deleted_keys,
|
||||
track_audit_task(
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_deleted_hook(
|
||||
data=data,
|
||||
keys_being_deleted=_keys_being_deleted,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
response=number_deleted_keys,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -5686,13 +5695,15 @@ async def _execute_virtual_key_regeneration(
|
|||
)
|
||||
|
||||
response: Final = GenerateKeyResponse.model_validate(updated_token_dict)
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_rotated_hook(
|
||||
data=data,
|
||||
existing_key_row=key_in_db,
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
track_audit_task(
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_rotated_hook(
|
||||
data=data,
|
||||
existing_key_row=key_in_db,
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
)
|
||||
)
|
||||
)
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -150,7 +150,7 @@ class AuditTaskRegistry:
|
|||
if exc is not None:
|
||||
verbose_proxy_logger.error("Audit log callback task failed: %s", exc, exc_info=exc)
|
||||
|
||||
async def drain(self, timeout: float = 10.0) -> None:
|
||||
async def drain(self, timeout: float = 5.0) -> None:
|
||||
deadline: Final = asyncio.get_running_loop().time() + timeout
|
||||
while self._pending_tasks:
|
||||
remaining = deadline - asyncio.get_running_loop().time()
|
||||
|
|
@ -187,7 +187,7 @@ def track_audit_task(task: asyncio.Task[_T]) -> asyncio.Task[_T]:
|
|||
return _audit_task_registry.track(task)
|
||||
|
||||
|
||||
async def drain_audit_tasks(timeout: float = 10.0) -> None:
|
||||
async def drain_audit_tasks(timeout: float = 5.0) -> None:
|
||||
await _audit_task_registry.drain(timeout=timeout)
|
||||
|
||||
|
||||
|
|
@ -261,13 +261,13 @@ async def create_object_audit_log(
|
|||
|
||||
|
||||
async def create_audit_log_for_update(request_data: LiteLLM_AuditLogs):
|
||||
if not is_audit_logging_enabled():
|
||||
return
|
||||
|
||||
current_task: Final = asyncio.current_task()
|
||||
if current_task is not None:
|
||||
track_audit_task(current_task)
|
||||
|
||||
if not is_audit_logging_enabled():
|
||||
return
|
||||
|
||||
from litellm.proxy.proxy_server import premium_user, prisma_client
|
||||
|
||||
if premium_user is not True:
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@ verbose_proxy_logger.setLevel(level=logging.DEBUG)
|
|||
from starlette.datastructures import URL
|
||||
|
||||
from litellm.proxy.management_helpers.audit_logs import (
|
||||
_pending_audit_tasks,
|
||||
create_audit_log_for_update,
|
||||
drain_audit_tasks,
|
||||
get_audit_log_changed_by,
|
||||
|
|
@ -460,3 +461,72 @@ async def test_drain_audit_tasks_captures_tasks_queued_during_drain():
|
|||
await drain_audit_tasks(timeout=2.0)
|
||||
|
||||
assert finished_tasks == ["first", "second"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_audit_log_for_update_does_not_track_when_logging_disabled():
|
||||
with (
|
||||
patch("litellm.store_audit_logs", False),
|
||||
patch("litellm.proxy.management_helpers.audit_logs.is_audit_logging_enabled", return_value=False),
|
||||
):
|
||||
request_data: Final = LiteLLM_AuditLogs(
|
||||
id=str(uuid.uuid4()),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
changed_by="test-user",
|
||||
table_name=LitellmTableNames.KEY_TABLE_NAME,
|
||||
object_id="test-obj",
|
||||
action="updated",
|
||||
updated_values="{}",
|
||||
before_value="{}",
|
||||
)
|
||||
current_tracked_before: Final = len(_pending_audit_tasks)
|
||||
await create_audit_log_for_update(request_data=request_data)
|
||||
assert len(_pending_audit_tasks) == current_tracked_before
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hook_spawn_with_io_delay_drained_before_shutdown():
|
||||
mock_prisma: Final = MagicMock()
|
||||
writes: Final[list[dict[str, Any]]] = []
|
||||
drain_completed: Final[list[bool]] = [False]
|
||||
|
||||
async def _mock_create(data: dict[str, Any]) -> None:
|
||||
assert not drain_completed[0]
|
||||
writes.append(data)
|
||||
|
||||
mock_prisma.db.litellm_auditlog.create = AsyncMock(side_effect=_mock_create)
|
||||
|
||||
async def _hook_with_preceding_io() -> None:
|
||||
await asyncio.sleep(0.05)
|
||||
await create_audit_log_for_update(
|
||||
request_data=LiteLLM_AuditLogs(
|
||||
id=str(uuid.uuid4()),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
changed_by="admin",
|
||||
table_name=LitellmTableNames.KEY_TABLE_NAME,
|
||||
object_id="token-hook",
|
||||
action="created",
|
||||
updated_values="{}",
|
||||
before_value=None,
|
||||
)
|
||||
)
|
||||
|
||||
with (
|
||||
patch("litellm.store_audit_logs", True),
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
|
||||
):
|
||||
hook_task: Final = track_audit_task(asyncio.create_task(_hook_with_preceding_io()))
|
||||
await drain_audit_tasks(timeout=2.0)
|
||||
drain_completed[0] = True
|
||||
|
||||
assert hook_task.done()
|
||||
assert len(writes) == 1
|
||||
assert writes[0]["object_id"] == "token-hook"
|
||||
|
||||
|
||||
def test_drain_audit_tasks_default_timeout():
|
||||
import inspect
|
||||
|
||||
sig: Final = inspect.signature(drain_audit_tasks)
|
||||
assert sig.parameters["timeout"].default == 5.0
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue