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
« 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.
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.
15In Open WebUI this surfaces as:
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.
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.
28Reference: https://www.starlette.io/middleware/#limitations
29"""
31from __future__ import annotations
33import logging
34import re
35import time
36from urllib.parse import parse_qs, urlencode
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
48log = logging.getLogger(__name__)
51class AppHTTPMiddleware:
52 """Open WebUI's pure-ASGI HTTP middleware.
54 Keeps the app's request-wide behavior in one middleware layer without
55 hiding the old concerns behind a stack of wrappers:
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`
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.
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.
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 """
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())
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
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
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)
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
124 self._commit_session()
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
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)
143 return send_with_headers
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
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
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
160 response = JSONResponse(status_code=400, content={'detail': 'Invalid WebSocket upgrade request'})
161 await response(scope, receive, send)
162 return True
164 async def _redirect_legacy_url(self, scope: Scope, receive: Receive, send: Send) -> bool:
165 if scope.get('method', '').upper() != 'GET':
166 return False
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
176 query_params = parse_qs(raw_query.decode('latin-1', errors='replace'))
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]
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
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
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
206 return False
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
212 try:
213 ScopedSession.rollback()
214 except Exception:
215 log.exception(message)
216 finally:
217 ScopedSession.remove()
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
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()
241def _scope_headers(scope: Scope) -> dict[str, str]:
242 """Return ASGI scope headers as a lower-cased str→str dict.
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