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

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) 

15 

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 

31 

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 

34 

35 

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 

51 

52 

53class AuditLevel(str, Enum): 

54 NONE = 'NONE' 

55 METADATA = 'METADATA' 

56 REQUEST = 'REQUEST' 

57 REQUEST_RESPONSE = 'REQUEST_RESPONSE' 

58 

59 

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. 

63 

64 Parameters: 

65 logger (Logger): An instance of Loguru’s logger. 

66 """ 

67 

68 def __init__(self, logger: 'Logger'): 

69 self.logger = logger.bind(auditable=True) 

70 

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) 

79 

80 if extra: 

81 entry['extra'] = extra 

82 

83 self.logger.log( 

84 log_level, 

85 '', 

86 **entry, 

87 ) 

88 

89 

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. 

93 

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

100 

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] = {} 

106 

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)]) 

110 

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)]) 

114 

115 

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

120 

121 DEFAULT_AUDITED_METHODS = {'PUT', 'PATCH', 'DELETE', 'POST'} 

122 

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) 

135 

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] 

138 

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 

146 

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 ) 

155 

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 ) 

161 

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) 

170 

171 request = Request(scope=cast(MutableMapping, scope)) 

172 

173 if self._should_skip_auditing(request): 

174 return await self.app(scope, receive, send) 

175 

176 async with self._audit_context(request) as context: 

177 

178 async def send_wrapper(message: ASGISendEvent) -> None: 

179 if self.audit_level == AuditLevel.REQUEST_RESPONSE: 

180 await self._capture_response(message, context) 

181 

182 await send(message) 

183 

184 original_receive = receive 

185 

186 async def receive_wrapper() -> ASGIReceiveEvent: 

187 nonlocal original_receive 

188 message = await original_receive() 

189 

190 if self.audit_level in ( 

191 AuditLevel.REQUEST, 

192 AuditLevel.REQUEST_RESPONSE, 

193 ): 

194 await self._capture_request(message, context) 

195 

196 return message 

197 

198 await self.app(scope, receive_wrapper, send_wrapper) 

199 

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) 

210 

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 

218 

219 auth_header = request.headers.get('Authorization') 

220 

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) 

226 

227 return None 

228 

229 ALWAYS_LOG_ENDPOINTS = ( 

230 '/api/v1/auths/signin', 

231 '/api/v1/auths/signout', 

232 '/api/v1/auths/signup', 

233 ) 

234 

235 def _should_skip_auditing(self, request: Request) -> bool: 

236 if AUDIT_LOG_LEVEL == 'NONE': 

237 return True 

238 

239 if request.method not in self.audited_methods: 

240 return True 

241 

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 

246 

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 

251 

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) 

255 

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 

259 

260 return False 

261 

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) 

266 

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'] 

270 

271 elif message['type'] == 'http.response.body': 

272 body = message.get('body', b'') 

273 context.add_response_chunk(body) 

274 

275 async def _log_audit_entry(self, request: Request, context: AuditContext): 

276 try: 

277 user = await self._get_authenticated_user(request) 

278 

279 user = user.model_dump(include={'id', 'name', 'email', 'role'}) if user else {} 

280 

281 request_body = context.request_body.decode('utf-8', errors='replace') 

282 response_body = context.response_body.decode('utf-8', errors='replace') 

283 

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 ) 

291 

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 ) 

304 

305 self.audit_logger.write(entry) 

306 except Exception as e: 

307 logger.error(f'Failed to log audit entry: {str(e)}')