Coverage for open_webui/utils/audit.py: 27%
144 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1import re
2import uuid
3from contextlib import asynccontextmanager
4from dataclasses import asdict, dataclass
5from enum import Enum
6from typing import (
7 TYPE_CHECKING,
8 Any,
9 AsyncGenerator,
10 Dict,
11 MutableMapping,
12 Optional,
13 cast,
14)
16from asgiref.typing import (
17 ASGI3Application,
18 ASGIReceiveCallable,
19 ASGIReceiveEvent,
20 ASGISendCallable,
21 ASGISendEvent,
22)
23from asgiref.typing import (
24 Scope as ASGIScope,
25)
26from loguru import logger
27from open_webui.env import AUDIT_INCLUDED_PATHS, AUDIT_LOG_LEVEL, ENABLE_AUDIT_GET_REQUESTS, MAX_BODY_LOG_SIZE
28from open_webui.models.users import UserModel
29from open_webui.utils.auth import get_current_user, get_http_authorization_cred
30from starlette.requests import Request
32if TYPE_CHECKING: 32 ↛ 33line 32 didn't jump to line 33 because the condition on line 32 was never true
33 from loguru import Logger
36@dataclass(frozen=True)
37class AuditLogEntry:
38 # `Metadata` audit level properties
39 id: str
40 user: Optional[dict[str, Any]]
41 audit_level: str
42 verb: str
43 request_uri: str
44 user_agent: Optional[str] = None
45 source_ip: Optional[str] = None
46 # `Request` audit level properties
47 request_object: Any = None
48 # `Request Response` level
49 response_object: Any = None
50 response_status_code: Optional[int] = None
53class AuditLevel(str, Enum):
54 NONE = 'NONE'
55 METADATA = 'METADATA'
56 REQUEST = 'REQUEST'
57 REQUEST_RESPONSE = 'REQUEST_RESPONSE'
60class AuditLogger:
61 """
62 A helper class that encapsulates audit logging functionality. It uses Loguru’s logger with an auditable binding to ensure that audit log entries are filtered correctly.
64 Parameters:
65 logger (Logger): An instance of Loguru’s logger.
66 """
68 def __init__(self, logger: 'Logger'):
69 self.logger = logger.bind(auditable=True)
71 def write(
72 self,
73 audit_entry: AuditLogEntry,
74 *,
75 log_level: str = 'INFO',
76 extra: Optional[dict] = None,
77 ):
78 entry = asdict(audit_entry)
80 if extra:
81 entry['extra'] = extra
83 self.logger.log(
84 log_level,
85 '',
86 **entry,
87 )
90class AuditContext:
91 """
92 Captures and aggregates the HTTP request and response bodies during the processing of a request. It ensures that only a configurable maximum amount of data is stored to prevent excessive memory usage.
94 Attributes:
95 request_body (bytearray): Accumulated request payload.
96 response_body (bytearray): Accumulated response payload.
97 max_body_size (int): Maximum number of bytes to capture.
98 metadata (Dict[str, Any]): A dictionary to store additional audit metadata (user, http verb, user agent, etc.).
99 """
101 def __init__(self, max_body_size: int = MAX_BODY_LOG_SIZE):
102 self.request_body = bytearray()
103 self.response_body = bytearray()
104 self.max_body_size = max_body_size
105 self.metadata: Dict[str, Any] = {}
107 def add_request_chunk(self, chunk: bytes):
108 if len(self.request_body) < self.max_body_size:
109 self.request_body.extend(chunk[: self.max_body_size - len(self.request_body)])
111 def add_response_chunk(self, chunk: bytes):
112 if len(self.response_body) < self.max_body_size:
113 self.response_body.extend(chunk[: self.max_body_size - len(self.response_body)])
116class AuditLoggingMiddleware:
117 """
118 ASGI middleware that intercepts HTTP requests and responses to perform audit logging. It captures request/response bodies (depending on audit level), headers, HTTP methods, and user information, then logs a structured audit entry at the end of the request cycle.
119 """
121 DEFAULT_AUDITED_METHODS = {'PUT', 'PATCH', 'DELETE', 'POST'}
123 def __init__(
124 self,
125 app: ASGI3Application,
126 *,
127 excluded_paths: Optional[list[str]] = None,
128 included_paths: Optional[list[str]] = None,
129 max_body_size: int = MAX_BODY_LOG_SIZE,
130 audit_level: AuditLevel = AuditLevel.NONE,
131 audit_get_requests: bool = False,
132 ) -> None:
133 self.app = app
134 self.audit_logger = AuditLogger(logger)
136 def normalize_paths(paths: Optional[list[str]]) -> list[str]:
137 return [path for path in (path.strip().lstrip('/') for path in paths or []) if path]
139 self.excluded_paths = normalize_paths(excluded_paths)
140 self.included_paths = normalize_paths(included_paths)
141 self.max_body_size = max_body_size
142 self.audited_methods = set(self.DEFAULT_AUDITED_METHODS)
143 if audit_get_requests:
144 self.audited_methods.add('GET')
145 self.audit_level = audit_level
147 # Paths are fixed for the process lifetime; compile once instead of
148 # per request. None means the corresponding mode has nothing to match.
149 self._included_pattern = (
150 re.compile(r'^/api(?:/v1)?/(' + '|'.join(self.included_paths) + r')\b') if self.included_paths else None
151 )
152 self._excluded_pattern = (
153 re.compile(r'^/api(?:/v1)?/(' + '|'.join(self.excluded_paths) + r')\b') if self.excluded_paths else None
154 )
156 if self.included_paths and self.excluded_paths:
157 logger.warning(
158 'Both AUDIT_INCLUDED_PATHS and AUDIT_EXCLUDED_PATHS are set. '
159 'AUDIT_INCLUDED_PATHS (whitelist) takes precedence.'
160 )
162 async def __call__(
163 self,
164 scope: ASGIScope,
165 receive: ASGIReceiveCallable,
166 send: ASGISendCallable,
167 ) -> None:
168 if scope['type'] != 'http':
169 return await self.app(scope, receive, send)
171 request = Request(scope=cast(MutableMapping, scope))
173 if self._should_skip_auditing(request):
174 return await self.app(scope, receive, send)
176 async with self._audit_context(request) as context:
178 async def send_wrapper(message: ASGISendEvent) -> None:
179 if self.audit_level == AuditLevel.REQUEST_RESPONSE:
180 await self._capture_response(message, context)
182 await send(message)
184 original_receive = receive
186 async def receive_wrapper() -> ASGIReceiveEvent:
187 nonlocal original_receive
188 message = await original_receive()
190 if self.audit_level in (
191 AuditLevel.REQUEST,
192 AuditLevel.REQUEST_RESPONSE,
193 ):
194 await self._capture_request(message, context)
196 return message
198 await self.app(scope, receive_wrapper, send_wrapper)
200 @asynccontextmanager
201 async def _audit_context(self, request: Request) -> AsyncGenerator[AuditContext, None]:
202 """
203 async context manager that ensures that an audit log entry is recorded after the request is processed.
204 """
205 context = AuditContext()
206 try:
207 yield context
208 finally:
209 await self._log_audit_entry(request, context)
211 async def _get_authenticated_user(self, request: Request) -> Optional[UserModel]:
212 # get_current_user stashes the resolved user on the scope-backed state;
213 # reuse it instead of running the full auth pipeline (JWT decode, Redis
214 # revocation checks, DB fetch, last-active write) a second time.
215 user = getattr(request.state, 'user', None)
216 if isinstance(user, UserModel):
217 return user
219 auth_header = request.headers.get('Authorization')
221 try:
222 user = await get_current_user(request, None, None, get_http_authorization_cred(auth_header))
223 return user
224 except Exception as e:
225 logger.debug('Failed to get authenticated user: {}', e)
227 return None
229 ALWAYS_LOG_ENDPOINTS = (
230 '/api/v1/auths/signin',
231 '/api/v1/auths/signout',
232 '/api/v1/auths/signup',
233 )
235 def _should_skip_auditing(self, request: Request) -> bool:
236 if AUDIT_LOG_LEVEL == 'NONE':
237 return True
239 if request.method not in self.audited_methods:
240 return True
242 path = request.url.path.lower()
243 for endpoint in self.ALWAYS_LOG_ENDPOINTS:
244 if path.startswith(endpoint):
245 return False # Do NOT skip logging for auth endpoints
247 # Skip logging if the request is not authenticated
248 # Check both Authorization header (API keys) and token cookie (browser sessions)
249 if not request.headers.get('authorization') and not request.cookies.get('token'):
250 return True
252 # Whitelist mode: only log paths that match included_paths
253 if self._included_pattern:
254 return not self._included_pattern.match(request.url.path)
256 # Blacklist mode: skip paths that match excluded_paths
257 if self._excluded_pattern and self._excluded_pattern.match(request.url.path):
258 return True
260 return False
262 async def _capture_request(self, message: ASGIReceiveEvent, context: AuditContext):
263 if message['type'] == 'http.request':
264 body = message.get('body', b'')
265 context.add_request_chunk(body)
267 async def _capture_response(self, message: ASGISendEvent, context: AuditContext):
268 if message['type'] == 'http.response.start':
269 context.metadata['response_status_code'] = message['status']
271 elif message['type'] == 'http.response.body':
272 body = message.get('body', b'')
273 context.add_response_chunk(body)
275 async def _log_audit_entry(self, request: Request, context: AuditContext):
276 try:
277 user = await self._get_authenticated_user(request)
279 user = user.model_dump(include={'id', 'name', 'email', 'role'}) if user else {}
281 request_body = context.request_body.decode('utf-8', errors='replace')
282 response_body = context.response_body.decode('utf-8', errors='replace')
284 # Redact sensitive information
285 if 'password' in request_body:
286 request_body = re.sub(
287 r'"password":\s*"(.*?)"',
288 '"password": "********"',
289 request_body,
290 )
292 entry = AuditLogEntry(
293 id=str(uuid.uuid4()),
294 user=user,
295 audit_level=self.audit_level.value,
296 verb=request.method,
297 request_uri=str(request.url),
298 response_status_code=context.metadata.get('response_status_code', None),
299 source_ip=request.client.host if request.client else None,
300 user_agent=request.headers.get('user-agent'),
301 request_object=request_body,
302 response_object=response_body,
303 )
305 self.audit_logger.write(entry)
306 except Exception as e:
307 logger.error(f'Failed to log audit entry: {str(e)}')