Files
kite/backend/api/audit.py
T

299 lines
9.3 KiB
Python

"""Admin audit log API endpoints."""
from __future__ import annotations
import asyncio
import csv
import io
import json
import re
import uuid
from datetime import datetime
from typing import Literal, Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from fastapi.responses import StreamingResponse
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import aliased
from db.models import AuditLog, User
from deps.auth import get_current_admin
from deps.db import get_db
from storage import get_storage_backend
from storage.minio_backend import MinIOBackend
router = APIRouter(prefix="/api/admin", tags=["audit"])
_VALID_EVENT_PREFIXES = frozenset({"auth", "document", "folder", "share", "admin", "cloud"})
_CSV_FIELDS = [
"id",
"event_type",
"user_id",
"actor_id",
"user_handle",
"actor_handle",
"user_email",
"resource_id",
"ip_address",
"metadata_",
"created_at",
]
def _audit_base_fields(entry: AuditLog) -> dict:
return {
"id": entry.id,
"event_type": entry.event_type,
"user_id": str(entry.user_id) if entry.user_id else None,
"actor_id": str(entry.actor_id) if entry.actor_id else None,
"resource_id": str(entry.resource_id) if entry.resource_id else None,
"ip_address": str(entry.ip_address) if entry.ip_address else None,
"metadata_": entry.metadata_,
"created_at": entry.created_at.isoformat(),
}
def _audit_to_dict(entry: AuditLog) -> dict:
"""Whitelisted audit serializer shared with the daily export task."""
return _audit_base_fields(entry)
def _audit_to_dict_with_handles(
entry: AuditLog,
user_handle: Optional[str],
actor_handle: Optional[str],
user_email: Optional[str] = None,
) -> dict:
data = _audit_base_fields(entry)
data.update({
"user_handle": user_handle or None,
"actor_handle": actor_handle or None,
"user_email": user_email or None,
})
return data
def _validate_event_type(event_type: Optional[str]) -> None:
if event_type is not None and event_type not in _VALID_EVENT_PREFIXES:
raise HTTPException(status_code=422, detail="Invalid event_type prefix")
def _apply_audit_filters(
query,
start: Optional[datetime],
end: Optional[datetime],
user_uuid: Optional[uuid.UUID],
event_type: Optional[str],
):
_validate_event_type(event_type)
if start is not None:
query = query.where(AuditLog.created_at >= start)
if end is not None:
query = query.where(AuditLog.created_at <= end)
if user_uuid is not None:
query = query.where(AuditLog.user_id == user_uuid)
if event_type is not None:
query = query.where(AuditLog.event_type.like(f"{event_type}.%"))
return query
def _build_filtered_query(
start: Optional[datetime],
end: Optional[datetime],
user_id: Optional[uuid.UUID],
event_type: Optional[str],
):
return _apply_audit_filters(
select(AuditLog).order_by(AuditLog.created_at.desc()),
start,
end,
user_id,
event_type,
)
def _build_filtered_query_with_handles(
start: Optional[datetime],
end: Optional[datetime],
user_uuid: Optional[uuid.UUID],
event_type: Optional[str],
):
UserSubject = aliased(User)
UserActor = aliased(User)
q = (
select(
AuditLog,
UserSubject.handle.label("user_handle"),
UserActor.handle.label("actor_handle"),
UserSubject.email.label("user_email"),
)
.outerjoin(UserSubject, UserSubject.id == AuditLog.user_id)
.outerjoin(UserActor, UserActor.id == AuditLog.actor_id)
.order_by(AuditLog.created_at.desc())
)
return _apply_audit_filters(q, start, end, user_uuid, event_type)
async def _resolve_user_uuid(session: AsyncSession, user_handle: Optional[str]) -> uuid.UUID | None:
if not user_handle:
return None
result = await session.execute(select(User.id).where(User.handle == user_handle))
return result.scalar_one_or_none()
async def _count_audit_log(
session: AsyncSession,
start: Optional[datetime],
end: Optional[datetime],
user_uuid: Optional[uuid.UUID],
event_type: Optional[str],
) -> int:
count_q = _apply_audit_filters(
select(func.count(AuditLog.id)).where(True),
start,
end,
user_uuid,
event_type,
)
result = await session.execute(count_q)
return result.scalar_one()
def _audit_rows_to_dicts(rows) -> list[dict]:
return [_audit_to_dict_with_handles(row[0], row[1], row[2], row[3]) for row in rows]
def _csv_response(csv_text: str, filename: str = "audit-export.csv") -> StreamingResponse:
return StreamingResponse(
iter([csv_text]),
media_type="text/csv",
headers={"Content-Disposition": f"attachment; filename={filename}"},
)
def _empty_csv_response() -> StreamingResponse:
output = io.StringIO()
csv.DictWriter(output, fieldnames=_CSV_FIELDS).writeheader()
return _csv_response(output.getvalue())
def _audit_csv_response(rows) -> StreamingResponse:
output = io.StringIO()
writer = csv.DictWriter(output, fieldnames=_CSV_FIELDS)
writer.writeheader()
for record in _audit_rows_to_dicts(rows):
record["metadata_"] = json.dumps(record["metadata_"]) if record["metadata_"] is not None else ""
writer.writerow(record)
return _csv_response(output.getvalue())
@router.get("/audit-log/daily-exports")
async def list_daily_exports(
_admin: User = Depends(get_current_admin),
) -> dict:
"""List available Celery daily audit export files from MinIO."""
backend = get_storage_backend()
if not isinstance(backend, MinIOBackend):
return {"items": []}
def _list() -> list:
objects = backend._client.list_objects(
"audit-logs", prefix="audit-logs/", recursive=False
)
items = []
for obj in objects:
name = obj.object_name or ""
if name.endswith(".csv"):
date_str = name.removeprefix("audit-logs/").removesuffix(".csv")
items.append({"date": date_str, "key": name})
items.sort(key=lambda x: x["date"], reverse=True)
return items
items = await asyncio.to_thread(_list)
return {"items": items}
@router.get("/audit-log/daily-exports/{date}")
async def download_daily_export(
date: str,
_admin: User = Depends(get_current_admin),
) -> StreamingResponse:
"""Stream a specific Celery daily audit export file from MinIO."""
if not re.fullmatch(r"\d{4}-\d{2}-\d{2}", date):
raise HTTPException(status_code=404, detail="Invalid date format")
backend = get_storage_backend()
if not isinstance(backend, MinIOBackend):
raise HTTPException(status_code=404, detail="Export not found")
key = f"audit-logs/{date}.csv"
def _get() -> bytes:
response = backend._client.get_object("audit-logs", key)
try:
return response.read()
finally:
response.close()
response.release_conn()
try:
csv_bytes = await asyncio.to_thread(_get)
except Exception:
raise HTTPException(status_code=404, detail="Export not found")
return StreamingResponse(
iter([csv_bytes]),
media_type="text/csv",
headers={"Content-Disposition": f'attachment; filename="audit-{date}.csv"'},
)
@router.get("/audit-log")
async def list_audit_log(
start: Optional[datetime] = Query(default=None),
end: Optional[datetime] = Query(default=None),
user_handle: Optional[str] = Query(default=None),
event_type: Optional[str] = Query(default=None),
page: int = Query(default=1, ge=1),
per_page: int = Query(default=50, ge=1, le=500),
session: AsyncSession = Depends(get_db),
_admin: User = Depends(get_current_admin),
) -> dict:
"""Return paginated, filtered audit log entries."""
user_uuid = await _resolve_user_uuid(session, user_handle)
if user_handle and user_uuid is None:
return {"items": [], "total": 0, "page": page, "per_page": per_page}
total = await _count_audit_log(session, start, end, user_uuid, event_type)
data_q = _build_filtered_query_with_handles(start, end, user_uuid, event_type)
data_q = data_q.limit(per_page).offset((page - 1) * per_page)
result = await session.execute(data_q)
return {
"items": _audit_rows_to_dicts(result.all()),
"total": total,
"page": page,
"per_page": per_page,
}
@router.get("/audit-log/export")
async def export_audit_log(
start: Optional[datetime] = Query(default=None),
end: Optional[datetime] = Query(default=None),
user_handle: Optional[str] = Query(default=None),
event_type: Optional[str] = Query(default=None),
format: Literal["csv"] = Query(default="csv"), # noqa: A002
session: AsyncSession = Depends(get_db),
_admin: User = Depends(get_current_admin),
) -> StreamingResponse:
"""Stream a CSV export of filtered audit log entries."""
user_uuid = await _resolve_user_uuid(session, user_handle)
if user_handle and user_uuid is None:
return _empty_csv_response()
q = _build_filtered_query_with_handles(start, end, user_uuid, event_type)
result = await session.execute(q)
return _audit_csv_response(result.all())