diff --git a/backend/open_webui/utils/auth.py b/backend/open_webui/utils/auth.py index 612a56ef92..7d6c52fe26 100644 --- a/backend/open_webui/utils/auth.py +++ b/backend/open_webui/utils/auth.py @@ -328,6 +328,10 @@ async def invalidate_token(request, token): ttl = exp - int(datetime.now(UTC).timestamp()) # Calculate time-to-live for the token if ttl > 0: + # Revoked tokens must not be able to disconnect newer sessions. + if not await is_valid_token(decoded, request.app.state.redis): + return + # Store the revoked token in Redis with an expiration time await request.app.state.redis.set( f'{REDIS_KEY_PREFIX}:auth:token:{jti}:revoked', @@ -335,6 +339,12 @@ async def invalidate_token(request, token): ex=ttl, ) + user_id = decoded.get('id') + if user_id: + from open_webui.socket.main import disconnect_user_sessions + + await disconnect_user_sessions(user_id) + async def revoke_user_tokens(request, user_id: str): """Reject every token already issued to a user. Requires Redis.""" @@ -356,6 +366,10 @@ async def revoke_user_tokens(request, user_id: str): ex=int(expires_delta.total_seconds()) if expires_delta else None, ) + from open_webui.socket.main import disconnect_user_sessions + + await disconnect_user_sessions(user_id) + def extract_token_from_auth_header(auth_header: str): return auth_header[len('Bearer ') :]