From 33b91bd8ae8a100a5a306c91441a7d0b422c4cde Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Mon, 29 Jun 2026 03:58:00 -0500 Subject: [PATCH] refac --- backend/open_webui/routers/terminals.py | 4 ++-- backend/open_webui/socket/main.py | 25 ++++++++++++++++++++----- backend/open_webui/utils/auth.py | 10 +++++----- 3 files changed, 27 insertions(+), 12 deletions(-) diff --git a/backend/open_webui/routers/terminals.py b/backend/open_webui/routers/terminals.py index e1236b08b0..3d637010fd 100644 --- a/backend/open_webui/routers/terminals.py +++ b/backend/open_webui/routers/terminals.py @@ -211,7 +211,7 @@ async def _resolve_authenticated_connection(ws: WebSocket, server_id: str): import asyncio import json - from open_webui.utils.auth import decode_token + from open_webui.utils.auth import decode_token, is_valid_token # First-message authentication try: @@ -222,7 +222,7 @@ async def _resolve_authenticated_connection(ws: WebSocket, server_id: str): return None token = payload.get('token', '') data = decode_token(token) - if data is None or 'id' not in data: + if data is None or 'id' not in data or not await is_valid_token(data, getattr(ws.app.state, 'redis', None)): await ws.close(code=4001, reason='Invalid token') return None user = await Users.get_user_by_id(data['id']) diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index ce95cbeb6b..1cdd064b3a 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -38,7 +38,7 @@ from open_webui.models.users import UserNameResponse, Users from open_webui.socket.utils import RedisDict, RedisLock, YdocManager from open_webui.tasks import create_task, stop_item_tasks from open_webui.utils.access_control import has_permission -from open_webui.utils.auth import decode_token +from open_webui.utils.auth import decode_token, is_valid_token from open_webui.utils.redis import ( build_sentinel_url, get_redis_connection, @@ -342,9 +342,12 @@ async def usage(sid, data): async def connect(sid, environ, auth): user = None if auth and 'token' in auth: + scope = (environ or {}).get('asgi.scope') or {} + fastapi_app = scope.get('app') + redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS data = decode_token(auth['token']) - if data is not None and 'id' in data: + if data is not None and 'id' in data and await is_valid_token(data, redis): user = await Users.get_user_by_id(data['id']) if user: @@ -369,8 +372,12 @@ async def user_join(sid, data): if not auth or 'token' not in auth: return + environ = sio.get_environ(sid) or {} + scope = environ.get('asgi.scope') or {} + fastapi_app = scope.get('app') + redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS token_data = decode_token(auth['token']) - if token_data is None or 'id' not in token_data: + if token_data is None or 'id' not in token_data or not await is_valid_token(token_data, redis): return user = await Users.get_user_by_id(token_data['id']) @@ -416,8 +423,12 @@ async def join_channel(sid, data): if not auth or 'token' not in auth: return + environ = sio.get_environ(sid) or {} + scope = environ.get('asgi.scope') or {} + fastapi_app = scope.get('app') + redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS data = decode_token(auth['token']) - if data is None or 'id' not in data: + if data is None or 'id' not in data or not await is_valid_token(data, redis): return user = await Users.get_user_by_id(data['id']) @@ -438,8 +449,12 @@ async def join_note(sid, data): if not auth or 'token' not in auth: return + environ = sio.get_environ(sid) or {} + scope = environ.get('asgi.scope') or {} + fastapi_app = scope.get('app') + redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS token_data = decode_token(auth['token']) - if token_data is None or 'id' not in token_data: + if token_data is None or 'id' not in token_data or not await is_valid_token(token_data, redis): return user = await Users.get_user_by_id(token_data['id']) diff --git a/backend/open_webui/utils/auth.py b/backend/open_webui/utils/auth.py index 282e0ccf5b..c95e23bb85 100644 --- a/backend/open_webui/utils/auth.py +++ b/backend/open_webui/utils/auth.py @@ -237,25 +237,25 @@ def decode_token(token: str) -> dict | None: return None -async def is_valid_token(request, decoded) -> bool: +async def is_valid_token(decoded, redis=None) -> bool: """ Check whether a JWT has been revoked. Two mechanisms: 1. Per-token (jti) — used by user-initiated sign-out (known jti). 2. Per-user (revoked_at) — used by OIDC back-channel logout when individual jti values are unknown; rejects tokens with iat <= revoked_at. """ - if request.app.state.redis: + if redis: # Per-token revocation jti = decoded.get('jti') if jti: - revoked = await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:auth:token:{jti}:revoked') + revoked = await redis.get(f'{REDIS_KEY_PREFIX}:auth:token:{jti}:revoked') if revoked: return False # Per-user revocation (OIDC back-channel logout) user_id = decoded.get('id') if user_id: - revoked_at = await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:auth:user:{user_id}:revoked_at') + revoked_at = await redis.get(f'{REDIS_KEY_PREFIX}:auth:user:{user_id}:revoked_at') if revoked_at: try: revoked_at_ts = int(revoked_at) @@ -365,7 +365,7 @@ async def get_current_user( ) if data is not None and 'id' in data: - if data.get('jti') and not await is_valid_token(request, data): + if not await is_valid_token(data, getattr(request.app.state, 'redis', None)): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail='Invalid token',