fix(middleware): replace BaseHTTPMiddleware HTTP middlewares with pure ASGI implementations (#23709)
* fix(middleware): replace BaseHTTPMiddleware HTTP middlewares with pure ASGI implementations
Starlette's BaseHTTPMiddleware (and the @app.middleware('http')
decorator that uses it) wraps the downstream app in an anyio task
group whose cancel scope tears down the inner task on every exit —
client disconnect, response complete, or any outer middleware bailing.
That CancelledError gets injected into whatever the inner task was
awaiting, so DB queries, embedding calls, and other long awaits get
killed mid-flight. Under aiosqlite the cleanup path then logs a
multi-page `terminate_force_close() not implemented` traceback at
ERROR for every cancelled DB call.
Open WebUI had four such middlewares stacked
(`commit_session_after_request`, `check_url`, `inspect_websocket`,
`RedirectMiddleware`) so a single cancellation would compound through
all four.
Move the four middlewares to a new `open_webui.utils.asgi_middleware`
module as plain ASGI classes (`__call__(scope, receive, send)`):
* `CommitSessionMiddleware` — was `commit_session_after_request`;
now also rolls back if commit fails
before releasing the connection.
* `AuthTokenMiddleware` — was `check_url`; sets request.state
token + enable_api_keys + stamps
X-Process-Time via a wrapped send.
* `WebsocketUpgradeGuardMiddleware`
— was `inspect_websocket`; rejects
/ws/socket.io HTTP requests that
claim transport=websocket without a
proper Upgrade/Connection header.
* `RedirectMiddleware` — was the BaseHTTPMiddleware subclass;
same /watch + share-target rewrites.
Pure ASGI does not introduce a cancel scope around the downstream app,
so client disconnects propagate via `receive()` (the way ASGI was
designed) instead of being injected as CancelledError. Middleware
ordering is preserved.
https://claude.ai/code/session_01JSr4NZSskEUQvoJnavVXh8
* fix(middleware): CommitSessionMiddleware — rollback on downstream error, never commit failed requests
The first cut put commit() in a finally block, which meant that even
when a downstream handler raised, the middleware would still commit
whatever partial sync writes that handler had made before the
failure. That regressed the previous BaseHTTPMiddleware semantics
where commit only ran on the success path.
Restructure the failure handling:
* Downstream raised → rollback any pending sync work, release the
connection, re-raise so the outer error middleware turns it into
an error response. We never commit a request that did not complete.
* Downstream returned → commit. On commit failure, log loudly,
rollback, and re-raise. ScopedSession.remove() always runs in
finally so the connection cannot leak.
Document the inherent pure-ASGI limitation explicitly: by the time
`await self.app(...)` returns the response messages have already
been emitted, so a commit failure can no longer change what the
client sees on the wire. Buffering the response to gate it on commit
success would break streaming responses (chat completions, SSE) which
are core to Open WebUI; the trade-off is intentional. Routes that
need commit-before-send must manage the sync session explicitly.
Also drop unused `typing` imports flagged by review.
https://claude.ai/code/session_01JSr4NZSskEUQvoJnavVXh8
---------
Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
+17
-96
@@ -46,7 +46,6 @@ from fastapi.staticfiles import StaticFiles
|
||||
from starlette_compress import CompressMiddleware
|
||||
|
||||
from starlette.exceptions import HTTPException as StarletteHTTPException
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.middleware.sessions import SessionMiddleware
|
||||
from starlette.responses import Response, StreamingResponse
|
||||
from starlette.datastructures import Headers
|
||||
@@ -58,6 +57,12 @@ from starsessions import (
|
||||
from starsessions.stores.redis import RedisStore
|
||||
|
||||
from open_webui.utils import logger
|
||||
from open_webui.utils.asgi_middleware import (
|
||||
AuthTokenMiddleware,
|
||||
CommitSessionMiddleware,
|
||||
RedirectMiddleware,
|
||||
WebsocketUpgradeGuardMiddleware,
|
||||
)
|
||||
from open_webui.utils.audit import AuditLevel, AuditLoggingMiddleware
|
||||
from open_webui.utils.logger import start_logger
|
||||
from open_webui.utils.session_pool import get_session
|
||||
@@ -1359,103 +1364,19 @@ if ENABLE_COMPRESSION_MIDDLEWARE:
|
||||
app.add_middleware(CompressMiddleware)
|
||||
|
||||
|
||||
class RedirectMiddleware(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next):
|
||||
# Check if the request is a GET request
|
||||
if request.method == 'GET':
|
||||
path = request.url.path
|
||||
query_params = dict(parse_qs(urlparse(str(request.url)).query))
|
||||
|
||||
redirect_params = {}
|
||||
|
||||
# Check for the specific watch path and the presence of 'v' parameter
|
||||
if path.endswith('/watch') and 'v' in query_params:
|
||||
# Extract the first 'v' parameter
|
||||
youtube_video_id = query_params['v'][0]
|
||||
redirect_params['youtube'] = youtube_video_id
|
||||
|
||||
if 'shared' in query_params and len(query_params['shared']) > 0:
|
||||
# PWA share_target support
|
||||
|
||||
text = query_params['shared'][0]
|
||||
if text:
|
||||
urls = re.match(r'https://\S+', text)
|
||||
if urls:
|
||||
from open_webui.retrieval.loaders.youtube import _parse_video_id
|
||||
|
||||
if youtube_video_id := _parse_video_id(urls[0]):
|
||||
redirect_params['youtube'] = youtube_video_id
|
||||
else:
|
||||
redirect_params['load-url'] = urls[0]
|
||||
else:
|
||||
redirect_params['q'] = text
|
||||
|
||||
if redirect_params:
|
||||
redirect_url = f'/?{urlencode(redirect_params)}'
|
||||
return RedirectResponse(url=redirect_url)
|
||||
|
||||
# Proceed with the normal flow of other requests
|
||||
response = await call_next(request)
|
||||
return response
|
||||
|
||||
|
||||
# All HTTP middlewares below are pure-ASGI implementations. The previous
|
||||
# `BaseHTTPMiddleware` / `@app.middleware('http')` versions wrapped the
|
||||
# downstream app in an anyio task group whose cancel scope cancelled
|
||||
# in-flight DB calls (and any other awaits) on client disconnect /
|
||||
# response completion — which surfaced as noisy SQLAlchemy
|
||||
# `terminate_force_close` tracebacks under aiosqlite and as random
|
||||
# CancelledError storms across the request path. See
|
||||
# `open_webui.utils.asgi_middleware` for the rationale.
|
||||
app.add_middleware(RedirectMiddleware)
|
||||
app.add_middleware(SecurityHeadersMiddleware)
|
||||
|
||||
|
||||
@app.middleware('http')
|
||||
async def commit_session_after_request(request: Request, call_next):
|
||||
response = await call_next(request)
|
||||
# log.debug("Commit session after request")
|
||||
try:
|
||||
ScopedSession.commit()
|
||||
finally:
|
||||
# CRITICAL: remove() returns the connection to the pool.
|
||||
# Without this, connections remain "checked out" and accumulate
|
||||
# as "idle in transaction" in PostgreSQL.
|
||||
ScopedSession.remove()
|
||||
return response
|
||||
|
||||
|
||||
@app.middleware('http')
|
||||
async def check_url(request: Request, call_next):
|
||||
start_time = int(time.time())
|
||||
request.state.token = get_http_authorization_cred(request.headers.get('Authorization'))
|
||||
# Fallback to cookie token for browser sessions
|
||||
if request.state.token is None and request.cookies.get('token'):
|
||||
from fastapi.security import HTTPAuthorizationCredentials
|
||||
|
||||
request.state.token = HTTPAuthorizationCredentials(scheme='Bearer', credentials=request.cookies.get('token'))
|
||||
|
||||
# Fallback to x-api-key header (Anthropic-compatible clients use this
|
||||
# for ALL requests, including GET /v1/models, not just POST /v1/messages).
|
||||
if request.state.token is None and request.headers.get('x-api-key'):
|
||||
from fastapi.security import HTTPAuthorizationCredentials
|
||||
|
||||
request.state.token = HTTPAuthorizationCredentials(
|
||||
scheme='Bearer', credentials=request.headers.get('x-api-key')
|
||||
)
|
||||
|
||||
request.state.enable_api_keys = app.state.config.ENABLE_API_KEYS
|
||||
response = await call_next(request)
|
||||
process_time = int(time.time()) - start_time
|
||||
response.headers['X-Process-Time'] = str(process_time)
|
||||
return response
|
||||
|
||||
|
||||
@app.middleware('http')
|
||||
async def inspect_websocket(request: Request, call_next):
|
||||
if '/ws/socket.io' in request.url.path and request.query_params.get('transport') == 'websocket':
|
||||
upgrade = (request.headers.get('Upgrade') or '').lower()
|
||||
connection = (request.headers.get('Connection') or '').lower().split(',')
|
||||
# Check that there's the correct headers for an upgrade, else reject the connection
|
||||
# This is to work around this upstream issue: https://github.com/miguelgrinberg/python-engineio/issues/367
|
||||
if upgrade != 'websocket' or 'upgrade' not in connection:
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
content={'detail': 'Invalid WebSocket upgrade request'},
|
||||
)
|
||||
return await call_next(request)
|
||||
app.add_middleware(CommitSessionMiddleware)
|
||||
app.add_middleware(AuthTokenMiddleware, fastapi_app=app)
|
||||
app.add_middleware(WebsocketUpgradeGuardMiddleware)
|
||||
|
||||
|
||||
app.add_middleware(
|
||||
|
||||
Reference in New Issue
Block a user