Coverage for open_webui/utils/asgi_middleware.py: 53%

129 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1""" 

2Pure-ASGI replacements for the project's previous 

3`@app.middleware('http')` / `BaseHTTPMiddleware` middlewares. 

4 

5Why this matters 

6---------------- 

7Starlette's `BaseHTTPMiddleware` (which `@app.middleware('http')` is 

8sugar for) runs the downstream app inside an `anyio` task group. When 

9the wrapper exits — for any reason: response complete, client 

10disconnect, an outer middleware bailing out — the task group cancels 

11the inner task. That `CancelledError` then propagates into whatever 

12the inner task was doing, including in-flight DB queries, embedding 

13calls and disk I/O. 

14 

15In Open WebUI this surfaces as: 

16 

17* SQLAlchemy logging multi-page `NotImplementedError: 

18 terminate_force_close()` tracebacks at ERROR every time a request is 

19 cancelled mid-DB-call (the aiosqlite connector cleanup path). 

20* Spurious cancellations cascading through the four stacked 

21 `@app.middleware('http')` wrappers. 

22 

23Pure ASGI middleware does not introduce a cancel scope around the 

24downstream app, so client disconnects propagate the way ASGI was 

25designed to (via `receive()` returning `http.disconnect`) instead of 

26being injected as `CancelledError` into arbitrary `await` points. 

27 

28Reference: https://www.starlette.io/middleware/#limitations 

29""" 

30 

31from __future__ import annotations 

32 

33import logging 

34import re 

35import time 

36from urllib.parse import parse_qs, urlencode 

37 

38from fastapi.responses import JSONResponse, RedirectResponse 

39from fastapi.security import HTTPAuthorizationCredentials 

40from open_webui.env import CUSTOM_API_KEY_HEADER 

41from open_webui.internal.db import ScopedSession 

42from open_webui.utils.auth import get_http_authorization_cred 

43from open_webui.utils.security_headers import set_security_headers 

44from starlette.datastructures import MutableHeaders 

45from starlette.requests import Request 

46from starlette.types import ASGIApp, Message, Receive, Scope, Send 

47 

48log = logging.getLogger(__name__) 

49 

50 

51class AppHTTPMiddleware: 

52 """Open WebUI's pure-ASGI HTTP middleware. 

53 

54 Keeps the app's request-wide behavior in one middleware layer without 

55 hiding the old concerns behind a stack of wrappers: 

56 

57 * reject malformed `/ws/socket.io` upgrade requests 

58 * stash bearer/cookie/API-key credentials on `request.state.token` 

59 * stamp `X-Process-Time` and configured security headers 

60 * serve the legacy `/watch` and `?shared=` redirects 

61 * commit and release the thread-local sync `ScopedSession` 

62 

63 Most requests now use the async session; the sync ScopedSession is 

64 only touched by startup, healthchecks, and a handful of legacy 

65 helpers (notably the pgvector / opengauss vector-DB clients). The 

66 middleware exists so that PostgreSQL connections do not accumulate 

67 as "idle in transaction" and so that any pending sync work made 

68 inside the request is durably persisted. 

69 

70 Failure semantics 

71 ----------------- 

72 * Downstream raised → roll back any pending sync work, release the 

73 connection, and re-raise so the outer exception middleware can 

74 turn it into an error response. We never commit work on a 

75 request that did not complete successfully. 

76 * Downstream returned → commit pending sync work; on commit 

77 failure, log loudly, roll back, and re-raise. Note that in pure 

78 ASGI the response messages have already been emitted by the 

79 time `await self.app(...)` returns, so a commit failure cannot 

80 retroactively change what the client sees on the wire — but 

81 re-raising still surfaces the error in logs and to ASGI servers 

82 that expose it. We deliberately do not buffer the response to 

83 gate it on commit success, because that would defeat streaming 

84 responses (chat completions, SSE) which are core to the app. 

85 

86 For request paths where commit-before-send is required, manage the 

87 sync session explicitly inside the handler instead of relying on 

88 this middleware. 

89 """ 

90 

91 def __init__(self, app: ASGIApp) -> None: 

92 self.app = app 

93 # Headers derive only from env vars, which are static for the process 

94 # lifetime — compute them once instead of per response. 

95 self._security_headers = list(set_security_headers().items()) 

96 

97 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: 

98 if scope['type'] != 'http': 

99 await self.app(scope, receive, send) 

100 return 

101 

102 if await self._reject_invalid_websocket(scope, receive, send): 102 ↛ 103line 102 didn't jump to line 103 because the condition on line 102 was never true

103 return 

104 

105 start_time = time.monotonic() 

106 request = Request(scope) 

107 self._set_token(request) 

108 send_with_headers = self._send_with_headers(send, start_time) 

109 

110 try: 

111 if await self._redirect_legacy_url(scope, receive, send_with_headers): 111 ↛ 112line 111 didn't jump to line 112 because the condition on line 111 was never true

112 pass 

113 # Keep health probes independent from sync session commit/remove so DB 

114 # pressure cannot delay or fail probe responses. 

115 elif scope.get('path', '') in {'/health', '/ready', '/health/db'}: 

116 await self.app(scope, receive, send_with_headers) 

117 return 

118 else: 

119 await self.app(scope, receive, send_with_headers) 

120 except BaseException: 

121 self._rollback_session('AppHTTPMiddleware: rollback failed after downstream error') 

122 raise 

123 

124 self._commit_session() 

125 

126 def _set_token(self, request: Request) -> None: 

127 token = get_http_authorization_cred(request.headers.get('Authorization')) 

128 if token is None and (cookie_token := request.cookies.get('token')): 128 ↛ 129line 128 didn't jump to line 129 because the condition on line 128 was never true

129 token = HTTPAuthorizationCredentials(scheme='Bearer', credentials=cookie_token) 

130 if token is None and (api_key := request.headers.get(CUSTOM_API_KEY_HEADER)): 130 ↛ 131line 130 didn't jump to line 131 because the condition on line 130 was never true

131 token = HTTPAuthorizationCredentials(scheme='Bearer', credentials=api_key) 

132 request.state.token = token 

133 

134 def _send_with_headers(self, send: Send, start_time: float) -> Send: 

135 async def send_with_headers(message: Message) -> None: 

136 if message['type'] == 'http.response.start': 

137 headers = MutableHeaders(scope=message) 

138 headers['X-Process-Time'] = f'{time.monotonic() - start_time:.6f}' 

139 for key, value in self._security_headers: 

140 headers[key] = value 

141 await send(message) 

142 

143 return send_with_headers 

144 

145 async def _reject_invalid_websocket(self, scope: Scope, receive: Receive, send: Send) -> bool: 

146 path = scope.get('path', '') 

147 if '/ws/socket.io' not in path: 147 ↛ 150line 147 didn't jump to line 150 because the condition on line 147 was always true

148 return False 

149 

150 query_params = parse_qs(scope.get('query_string', b'').decode('latin-1', errors='replace')) 

151 if query_params.get('transport', [''])[0] != 'websocket': 151 ↛ anywhereline 151 didn't jump anywhere: it always raised an exception.

152 return False 

153 

154 headers = _scope_headers(scope) 

155 upgrade = headers.get('upgrade', '').lower() 

156 connection_tokens = [token.strip() for token in headers.get('connection', '').lower().split(',')] 

157 if upgrade == 'websocket' and 'upgrade' in connection_tokens: 

158 return False 

159 

160 response = JSONResponse(status_code=400, content={'detail': 'Invalid WebSocket upgrade request'}) 

161 await response(scope, receive, send) 

162 return True 

163 

164 async def _redirect_legacy_url(self, scope: Scope, receive: Receive, send: Send) -> bool: 

165 if scope.get('method', '').upper() != 'GET': 

166 return False 

167 

168 path = scope.get('path', '') 

169 raw_query = scope.get('query_string', b'') 

170 # This middleware only acts on /watch?v= and ?shared= URLs; skip the 

171 # decode + parse_qs work for every other GET. (A false positive on the 

172 # substring check just falls through to the full parse below.) 

173 if not (path.endswith('/watch') or b'shared' in raw_query): 173 ↛ 176line 173 didn't jump to line 176 because the condition on line 173 was always true

174 return False 

175 

176 query_params = parse_qs(raw_query.decode('latin-1', errors='replace')) 

177 

178 redirect_params: dict[str, str] = {} 

179 if path.endswith('/watch') and 'v' in query_params and query_params['v']: 

180 redirect_params['youtube'] = query_params['v'][0] 

181 

182 if 'shared' in query_params and query_params['shared']: 

183 text = query_params['shared'][0] 

184 if text: 

185 url_match = re.match(r'https://\S+', text) 

186 if url_match: 186 ↛ anywhereline 186 didn't jump anywhere: it always raised an exception.

187 # Local import: youtube loader pulls heavy deps and is 

188 # only needed when a share-target actually contains a 

189 # YouTube URL. 

190 from open_webui.retrieval.loaders.youtube import _parse_video_id 

191 

192 youtube_video_id = _parse_video_id(url_match[0]) 

193 if youtube_video_id: 

194 redirect_params['youtube'] = youtube_video_id 

195 else: 

196 redirect_params['load-url'] = url_match[0] 

197 else: 

198 redirect_params['q'] = text 

199 

200 if redirect_params: 200 ↛ anywhereline 200 didn't jump anywhere: it always raised an exception.

201 redirect_url = f'/?{urlencode(redirect_params)}' 

202 response = RedirectResponse(url=redirect_url) 

203 await response(scope, receive, send) 

204 return True 

205 

206 return False 

207 

208 def _rollback_session(self, message: str) -> None: 

209 if not ScopedSession.registry.has(): 209 ↛ 212line 209 didn't jump to line 212 because the condition on line 209 was always true

210 return 

211 

212 try: 

213 ScopedSession.rollback() 

214 except Exception: 

215 log.exception(message) 

216 finally: 

217 ScopedSession.remove() 

218 

219 def _commit_session(self) -> None: 

220 # Nothing in this request touched the sync session: committing would 

221 # only instantiate one to run an empty transaction. 

222 if not ScopedSession.registry.has(): 222 ↛ 225line 222 didn't jump to line 225 because the condition on line 222 was always true

223 return 

224 

225 try: 

226 ScopedSession.commit() 

227 except Exception: 

228 log.exception('AppHTTPMiddleware: post-request commit failed; response was already sent to client') 

229 try: 

230 ScopedSession.rollback() 

231 except Exception: 

232 log.exception('AppHTTPMiddleware: rollback failed after commit failure') 

233 raise 

234 finally: 

235 # CRITICAL: remove() returns the connection to the pool. 

236 # Without this, connections remain "checked out" and 

237 # accumulate as "idle in transaction" in PostgreSQL. 

238 ScopedSession.remove() 

239 

240 

241def _scope_headers(scope: Scope) -> dict[str, str]: 

242 """Return ASGI scope headers as a lower-cased str→str dict. 

243 

244 ASGI delivers headers as a list of (bytes, bytes) pairs. For 

245 convenience, fold duplicate keys with comma-joining (matching 

246 HTTP/1.1 semantics). 

247 """ 

248 decoded: dict[str, str] = {} 

249 for raw_key, raw_value in scope.get('headers', []): 

250 key = raw_key.decode('latin-1').lower() 

251 value = raw_value.decode('latin-1') 

252 if key in decoded: 

253 decoded[key] = f'{decoded[key]}, {value}' 

254 else: 

255 decoded[key] = value 

256 return decoded