refac
This commit is contained in:
@@ -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'])
|
||||
|
||||
@@ -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'])
|
||||
|
||||
@@ -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',
|
||||
|
||||
Reference in New Issue
Block a user