-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathaudit_logging.py
More file actions
160 lines (137 loc) · 5.49 KB
/
Copy pathaudit_logging.py
File metadata and controls
160 lines (137 loc) · 5.49 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
import logging
import traceback as traceback_module
from collections.abc import Callable
from fastapi import Request
from starlette.middleware.base import BaseHTTPMiddleware
from src.core.database.postgres.session import AsyncSessionLocal
from src.core.security.audit import (
AuditEvent,
AuditService,
ErrorTrace,
ErrorTraceService,
)
from src.core.security.infrastructure.repositories.audit_log_repository import (
SQLAlchemyAuditRepository,
)
from src.core.security.infrastructure.repositories.error_trace_repository import (
SQLAlchemyErrorTraceRepository,
)
logger = logging.getLogger(__name__)
EXCLUDED_AUDIT_PATHS = frozenset(
{
"/health",
"/live",
"/ready",
"/docs",
"/docs/",
"/redoc",
"/redoc/",
"/openapi.json",
}
)
class AuditLoggingMiddleware(BaseHTTPMiddleware):
def __init__(
self,
app,
audit_service_factory: Callable[[], AuditService] | None = None,
error_trace_service_factory: Callable[[], ErrorTraceService] | None = None,
):
super().__init__(app)
self._audit_service_factory = audit_service_factory
self._error_trace_service_factory = error_trace_service_factory
async def dispatch(self, request: Request, call_next):
if self._is_excluded(request):
return await call_next(request)
try:
response = await call_next(request)
except Exception as exc:
await self._record_error_trace(request, exc)
raise
await self._record_audit_event(request, response.status_code)
return response
def _is_excluded(self, request: Request) -> bool:
return request.url.path.rstrip("/") in {
path.rstrip("/") for path in EXCLUDED_AUDIT_PATHS
}
async def _record_audit_event(self, request: Request, status_code: int) -> None:
event = AuditEvent(
action=f"{request.method} {request.url.path}",
actor_id=getattr(request.state, "user_id", None),
resource_type=self._resource_type(request),
resource_id=self._resource_id(request),
request_id=getattr(request.state, "request_id", None),
metadata={
"method": request.method,
"path": request.url.path,
"status_code": status_code,
"client_ip": self._client_ip(request),
"user_agent": request.headers.get("User-Agent"),
},
)
await self._safe_record_audit(event)
async def _record_error_trace(self, request: Request, exc: Exception) -> None:
trace = ErrorTrace(
error_type=type(exc).__name__,
message=str(exc),
traceback=traceback_module.format_exc(),
method=request.method,
path=request.url.path,
actor_id=getattr(request.state, "user_id", None),
request_id=getattr(request.state, "request_id", None),
metadata={
"client_ip": self._client_ip(request),
"user_agent": request.headers.get("User-Agent"),
},
)
await self._safe_record_error_trace(trace)
async def _safe_record_audit(self, event: AuditEvent) -> None:
try:
factory = self._audit_service_factory or self._default_audit_service_factory
service = factory()
await service.record(event)
except Exception:
logger.exception("failed to record audit event")
async def _safe_record_error_trace(self, trace: ErrorTrace) -> None:
try:
factory = (
self._error_trace_service_factory
or self._default_error_trace_service_factory
)
service = factory()
await service.record(trace)
except Exception:
logger.exception("failed to record error trace")
def _default_audit_service_factory(self) -> AuditService:
return _SessionBackedAuditService()
def _default_error_trace_service_factory(self) -> ErrorTraceService:
return _SessionBackedErrorTraceService()
@staticmethod
def _resource_type(request: Request) -> str | None:
parts = [part for part in request.url.path.split("/") if part]
if len(parts) >= 3 and parts[0] == "api" and parts[1].startswith("v"):
return parts[2]
return parts[0] if parts else None
@staticmethod
def _resource_id(request: Request) -> str | None:
for key in ("id", "todo_id", "role_id", "permission_id", "user_id"):
if key in request.path_params:
return str(request.path_params[key])
return None
@staticmethod
def _client_ip(request: Request) -> str | None:
forwarded = request.headers.get("X-Forwarded-For")
if forwarded:
return forwarded.split(",")[0].strip()
return request.client.host if request.client else None
class _SessionBackedAuditService(AuditService):
async def record(self, event: AuditEvent) -> None:
async with AsyncSessionLocal() as session:
repository = SQLAlchemyAuditRepository(session)
await repository.save(event)
await session.commit()
class _SessionBackedErrorTraceService(ErrorTraceService):
async def record(self, trace: ErrorTrace) -> None:
async with AsyncSessionLocal() as session:
repository = SQLAlchemyErrorTraceRepository(session)
await repository.save(trace)
await session.commit()