Coverage for open_webui/utils/auth.py: 31%
375 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
1from __future__ import annotations
3import asyncio
4import base64
5import hashlib
6import hmac
7import logging
8import os
9import uuid
10from datetime import datetime, timedelta
11from threading import Lock
12from time import monotonic
13from typing import Optional, Union
15import bcrypt
16import jwt
17import pytz
18import requests
19from cryptography.hazmat.primitives import serialization
20from cryptography.hazmat.primitives.asymmetric import ed25519
21from cryptography.hazmat.primitives.ciphers.aead import AESGCM
22from fastapi import BackgroundTasks, Depends, HTTPException, Request, Response, status
23from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
24from open_webui.constants import ERROR_MESSAGES
25from open_webui.env import (
26 ENABLE_OTEL,
27 ENABLE_PASSWORD_VALIDATION,
28 LICENSE_BLOB,
29 OFFLINE_MODE,
30 PASSWORD_HASH_ALGORITHM,
31 PASSWORD_VALIDATION_HINT,
32 PASSWORD_VALIDATION_REGEX_PATTERN,
33 REDIS_KEY_PREFIX,
34 STATIC_DIR,
35 TRUSTED_SIGNATURE_KEY,
36 WEBUI_AUTH_TRUSTED_EMAIL_HEADER,
37 WEBUI_SECRET_KEY,
38 pk,
39)
40from open_webui.models.auths import Auths
41from open_webui.models.config import Config
42from open_webui.models.users import Users
43from open_webui.utils.access_control import has_permission
44from open_webui.utils.json_codec import JSONCodec
45from open_webui.utils.misc import parse_duration
46from pytz import UTC
47from redis.exceptions import RedisError
49log = logging.getLogger(__name__)
51SESSION_SECRET = WEBUI_SECRET_KEY
52ALGORITHM = 'HS256'
53PASSWORD_BCRYPT_MAX_BYTES = 72
55##############
56# Auth Utils
57##############
60def verify_signature(payload: str, signature: str) -> bool:
61 """
62 Verifies the HMAC signature of the received payload.
63 """
64 try:
65 expected_signature = base64.b64encode(
66 hmac.new(TRUSTED_SIGNATURE_KEY, payload.encode(), hashlib.sha256).digest()
67 ).decode()
69 # Compare securely to prevent timing attacks
70 return hmac.compare_digest(expected_signature, signature)
72 except Exception:
73 return False
76def override_static(path: str, content: str):
77 # Ensure path is safe
78 if '/' in path or '..' in path:
79 log.error(f'Invalid path: {path}')
80 return
82 file_path = os.path.join(STATIC_DIR, path)
83 os.makedirs(os.path.dirname(file_path), exist_ok=True)
85 with open(file_path, 'wb') as f:
86 f.write(base64.b64decode(content)) # Convert Base64 back to raw binary
89def get_license_data(app, key):
90 def data_handler(data):
91 for k, v in data.items():
92 if k == 'resources':
93 # LICENSE covers these Open WebUI branding assets.
94 # Do not alter, remove, obscure, or replace them except as LICENSE permits:
95 # https://docs.openwebui.com/license.
96 for p, c in v.items():
97 globals().get('override_static', lambda a, b: None)(p, c)
98 elif k == 'count':
99 setattr(app.state, 'USER_COUNT', v)
100 elif k == 'name':
101 # LICENSE covers this Open WebUI product name.
102 # Do not alter, remove, obscure, or replace it except as LICENSE permits:
103 # https://docs.openwebui.com/license.
104 setattr(app.state, 'WEBUI_NAME', v)
105 elif k == 'metadata':
106 setattr(app.state, 'LICENSE_METADATA', v)
108 def handler(u):
109 try:
110 res = requests.post(
111 f'{u}/api/v1/license/',
112 json={'key': key, 'version': '1'},
113 timeout=5,
114 )
115 except Exception as ex:
116 log.error(f'License: retrieval issue from {u}: {ex}')
117 return False
119 if getattr(res, 'ok', False):
120 payload = getattr(res, 'json', lambda: {})()
121 data_handler(payload)
122 return True
123 else:
124 log.error(f'License: retrieval issue: {getattr(res, "text", "unknown error")}')
126 if key:
127 us = [
128 'https://api.openwebui.com',
129 'https://licenses.api.openwebui.com',
130 ]
131 try:
132 for u in us:
133 if handler(u):
134 return True
135 except Exception as ex:
136 log.exception(f'License: Uncaught Exception: {ex}')
138 try:
139 if LICENSE_BLOB:
140 nl = 12
141 kb = hashlib.sha256((key.replace('-', '').upper()).encode()).digest()
143 def nt(b):
144 return b[:nl], b[nl:]
146 lb = base64.b64decode(LICENSE_BLOB)
147 ln, lt = nt(lb)
149 aesgcm = AESGCM(kb)
150 p = JSONCodec.loads(aesgcm.decrypt(ln, lt, None))
151 pk.verify(base64.b64decode(p['s']), p['p'].encode())
153 pb = base64.b64decode(p['p'])
154 pn, pt = nt(pb)
156 data = JSONCodec.loads(aesgcm.decrypt(pn, pt, None).decode())
158 exp = data.get('exp')
159 if exp:
160 if isinstance(exp, str):
161 from datetime import date
163 exp = date.fromisoformat(exp)
164 if exp < datetime.now().date():
165 return False
167 data_handler(data)
168 return True
169 except Exception as e:
170 log.error(f'License: {e}')
172 return False
175bearer_security = HTTPBearer(auto_error=False)
178async def get_password_hash(password: str) -> str:
179 """Hash a password using the configured algorithm in a thread pool."""
180 if PASSWORD_HASH_ALGORITHM == 'argon2': 180 ↛ 181line 180 didn't jump to line 181 because the condition on line 180 was never true
181 from argon2 import PasswordHasher
183 return await asyncio.to_thread(PasswordHasher().hash, password)
184 if PASSWORD_HASH_ALGORITHM == 'bcrypt': 184 ↛ 187line 184 didn't jump to line 187 because the condition on line 184 was always true
185 return (await asyncio.to_thread(bcrypt.hashpw, password.encode('utf-8'), bcrypt.gensalt())).decode('utf-8')
187 raise ValueError(f'Unsupported PASSWORD_HASH_ALGORITHM: {PASSWORD_HASH_ALGORITHM}')
190def validate_password(password: str) -> bool:
191 # bcrypt only accepts 72 bytes; reject long new passwords instead of storing an unusable hash.
192 if PASSWORD_HASH_ALGORITHM == 'bcrypt' and len(password.encode('utf-8')) > PASSWORD_BCRYPT_MAX_BYTES: 192 ↛ 193line 192 didn't jump to line 193 because the condition on line 192 was never true
193 raise Exception(
194 ERROR_MESSAGES.PASSWORD_TOO_LONG,
195 )
197 if ENABLE_PASSWORD_VALIDATION: 197 ↛ 198line 197 didn't jump to line 198 because the condition on line 197 was never true
198 if not PASSWORD_VALIDATION_REGEX_PATTERN.match(password):
199 raise Exception(ERROR_MESSAGES.INVALID_PASSWORD(PASSWORD_VALIDATION_HINT))
201 return True
204async def verify_password(plain_password: str, hashed_password: str) -> bool:
205 """Verify a password using the algorithm encoded in its hash."""
206 if not hashed_password: 206 ↛ 207line 206 didn't jump to line 207 because the condition on line 206 was never true
207 return False
209 if hashed_password.startswith('$argon2'): 209 ↛ 210line 209 didn't jump to line 210 because the condition on line 209 was never true
210 from argon2 import PasswordHasher
211 from argon2.exceptions import InvalidHashError, VerificationError
213 try:
214 return await asyncio.to_thread(PasswordHasher().verify, hashed_password, plain_password)
215 except (InvalidHashError, VerificationError):
216 return False
218 password_bytes = plain_password.encode('utf-8')[:PASSWORD_BCRYPT_MAX_BYTES]
219 try:
220 return await asyncio.to_thread(
221 bcrypt.checkpw,
222 password_bytes,
223 hashed_password.encode('utf-8'),
224 )
225 except ValueError:
226 return False
229# Let the one who signed this token be remembered at every gate,
230# and may the claims therein honor the creator long after
231# the session has closed.
232def create_token(data: dict, expires_delta: Union[timedelta, None] = None) -> str:
233 payload = data.copy()
235 if expires_delta: 235 ↛ 239line 235 didn't jump to line 239 because the condition on line 235 was always true
236 expire = datetime.now(UTC) + expires_delta
237 payload.update({'exp': expire})
239 jti = str(uuid.uuid4())
240 payload.update({'jti': jti, 'iat': datetime.now(UTC)})
242 encoded_jwt = jwt.encode(payload, SESSION_SECRET, algorithm=ALGORITHM)
243 return encoded_jwt
246def decode_token(token: str) -> dict | None:
247 try:
248 decoded = jwt.decode(token, SESSION_SECRET, algorithms=[ALGORITHM])
249 return decoded
250 except Exception:
251 return None
254class RateLimitFilter(logging.Filter):
255 """Limit a logger to one record per interval per process."""
257 def __init__(self, interval=60):
258 super().__init__()
259 self.interval = interval
260 self.next_allowed = float('-inf')
261 self.lock = Lock()
263 def filter(self, record):
264 with self.lock:
265 now = monotonic()
266 if now < self.next_allowed:
267 return False
268 self.next_allowed = now + self.interval
269 return True
272revocation_log = logging.getLogger(f'{__name__}.revocation')
273revocation_log.addFilter(RateLimitFilter())
276async def is_valid_token(decoded, redis=None) -> bool:
277 """
278 Check whether a JWT has been revoked. Two mechanisms:
279 1. Per-token (jti) — used by user-initiated sign-out (known jti).
280 2. Per-user (revoked_at) — used by password changes and OIDC back-channel
281 logout when individual jti values are unknown; rejects tokens with iat <= revoked_at.
283 Fail open on Redis errors to preserve availability; revoked tokens may be accepted.
284 """
285 if not redis: 285 ↛ 288line 285 didn't jump to line 288 because the condition on line 285 was always true
286 return True
288 try:
289 # Per-token revocation
290 jti = decoded.get('jti')
291 if jti:
292 revoked = await redis.get(f'{REDIS_KEY_PREFIX}:auth:token:{jti}:revoked')
293 if revoked:
294 return False
296 # Per-user revocation (password change, OIDC back-channel logout)
297 user_id = decoded.get('id')
298 if user_id:
299 revoked_at = await redis.get(f'{REDIS_KEY_PREFIX}:auth:user:{user_id}:revoked_at')
300 if revoked_at:
301 try:
302 revoked_at_ts = int(revoked_at)
303 token_iat = decoded.get('iat')
304 # No iat means legacy token — reject since we can't verify issue time
305 if token_iat is None or token_iat <= revoked_at_ts:
306 return False
307 except (ValueError, TypeError):
308 pass
309 except RedisError as e:
310 revocation_log.warning('Revocation check failed; accepting token: %s', e)
312 return True
315async def invalidate_token(request, token):
316 decoded = decode_token(token)
318 # If token is invalid/expired, nothing to revoke
319 if not decoded:
320 return
322 # Require Redis to store revoked tokens
323 if request.app.state.redis:
324 jti = decoded.get('jti')
325 exp = decoded.get('exp')
327 if jti and exp:
328 ttl = exp - int(datetime.now(UTC).timestamp()) # Calculate time-to-live for the token
330 if ttl > 0:
331 # Revoked tokens must not be able to disconnect newer sessions.
332 if not await is_valid_token(decoded, request.app.state.redis):
333 return
335 # Store the revoked token in Redis with an expiration time
336 await request.app.state.redis.set(
337 f'{REDIS_KEY_PREFIX}:auth:token:{jti}:revoked',
338 '1',
339 ex=ttl,
340 )
342 user_id = decoded.get('id')
343 if user_id:
344 from open_webui.socket.main import disconnect_user_sessions
346 await disconnect_user_sessions(user_id)
349async def revoke_user_tokens(request, user_id: str):
350 """Reject every token already issued to a user. Requires Redis."""
351 redis = request.app.state.redis
353 if not redis:
354 log.warning(
355 'Cannot revoke tokens for user %s: Redis is not configured, existing sessions stay valid until expiry.',
356 user_id,
357 )
358 return
360 # The marker has to outlive every token it revokes, so it never expires when tokens do not
361 expires_delta = parse_duration(await Config.get('auth.jwt_expiry'))
363 await redis.set(
364 f'{REDIS_KEY_PREFIX}:auth:user:{user_id}:revoked_at',
365 str(int(datetime.now(UTC).timestamp())),
366 ex=int(expires_delta.total_seconds()) if expires_delta else None,
367 )
369 from open_webui.socket.main import disconnect_user_sessions
371 await disconnect_user_sessions(user_id)
374def extract_token_from_auth_header(auth_header: str):
375 return auth_header[len('Bearer ') :]
378def create_api_key():
379 key = str(uuid.uuid4()).replace('-', '')
380 return f'sk-{key}'
383def get_http_authorization_cred(auth_header: str | None):
384 if not auth_header:
385 return None
386 try:
387 scheme, credentials = auth_header.split(' ')
388 return HTTPAuthorizationCredentials(scheme=scheme, credentials=credentials)
389 except Exception:
390 return None
393async def get_current_user(
394 request: Request,
395 response: Response,
396 background_tasks: BackgroundTasks,
397 auth_token: HTTPAuthorizationCredentials = Depends(bearer_security),
398 # NOTE: We intentionally do NOT use Depends(get_session) here.
399 # Sessions are managed internally with short-lived context managers.
400 # This ensures connections are released immediately after auth queries,
401 # not held for the entire request duration (e.g., during 30+ second LLM calls).
402):
403 token = None
405 if auth_token is not None:
406 token = auth_token.credentials
408 if token is None and 'token' in request.cookies: 408 ↛ 409line 408 didn't jump to line 409 because the condition on line 408 was never true
409 token = request.cookies.get('token')
411 # Fallback to request.state.token (set by middleware, e.g. for x-api-key)
412 if token is None and hasattr(request.state, 'token') and request.state.token:
413 token = request.state.token.credentials
415 if token is None:
416 raise HTTPException(status_code=401, detail='Not authenticated')
418 # auth by api key
419 if token.startswith('sk-'): 419 ↛ 420line 419 didn't jump to line 420 because the condition on line 419 was never true
420 user = await get_current_user_by_api_key(request, token)
422 # Add user info to current span
423 if ENABLE_OTEL:
424 from opentelemetry import trace
426 current_span = trace.get_current_span()
427 if current_span:
428 current_span.set_attribute('client.user.id', user.id)
429 current_span.set_attribute('client.user.email', user.email)
430 current_span.set_attribute('client.user.role', user.role)
431 current_span.set_attribute('client.auth.type', 'api_key')
433 # Scope-backed, so outer middleware (audit) can reuse the resolved user
434 request.state.user = user
435 request.state.auth_type = 'api_key'
436 return user
438 # auth by jwt token
439 try:
440 try:
441 data = decode_token(token)
442 except Exception as e:
443 raise HTTPException(
444 status_code=status.HTTP_401_UNAUTHORIZED,
445 detail='Invalid token',
446 )
448 if data is not None and 'id' in data:
449 if not await is_valid_token(data, getattr(request.app.state, 'redis', None)): 449 ↛ 450line 449 didn't jump to line 450 because the condition on line 449 was never true
450 raise HTTPException(
451 status_code=status.HTTP_401_UNAUTHORIZED,
452 detail='Invalid token',
453 )
455 user = await Users.get_user_by_id(data['id'])
456 if user is None: 456 ↛ 457line 456 didn't jump to line 457 because the condition on line 456 was never true
457 raise HTTPException(
458 status_code=status.HTTP_401_UNAUTHORIZED,
459 detail=ERROR_MESSAGES.INVALID_TOKEN,
460 )
461 else:
462 if WEBUI_AUTH_TRUSTED_EMAIL_HEADER: 462 ↛ 463line 462 didn't jump to line 463 because the condition on line 462 was never true
463 trusted_email = request.headers.get(WEBUI_AUTH_TRUSTED_EMAIL_HEADER, '').lower()
464 if trusted_email and user.email != trusted_email:
465 raise HTTPException(
466 status_code=status.HTTP_401_UNAUTHORIZED,
467 detail='User mismatch. Please sign in again.',
468 )
470 # Add user info to current span
471 if ENABLE_OTEL: 471 ↛ 472line 471 didn't jump to line 472 because the condition on line 471 was never true
472 from opentelemetry import trace
474 current_span = trace.get_current_span()
475 if current_span:
476 current_span.set_attribute('client.user.id', user.id)
477 current_span.set_attribute('client.user.email', user.email)
478 current_span.set_attribute('client.user.role', user.role)
479 current_span.set_attribute('client.auth.type', 'jwt')
481 # Refresh the user's last active timestamp
482 # Fire-and-forget via asyncio.create_task to avoid blocking
483 asyncio.create_task(Users.update_last_active_by_id(user.id))
485 # Scope-backed, so outer middleware (audit) can reuse the resolved user
486 request.state.user = user
487 request.state.auth_type = 'jwt'
488 return user
489 else:
490 raise HTTPException(
491 status_code=status.HTTP_401_UNAUTHORIZED,
492 detail=ERROR_MESSAGES.UNAUTHORIZED,
493 )
494 except Exception as e:
495 # Delete the token cookie
496 if request.cookies.get('token'): 496 ↛ 497line 496 didn't jump to line 497 because the condition on line 496 was never true
497 response.delete_cookie('token')
499 if request.cookies.get('oauth_id_token'): 499 ↛ 500line 499 didn't jump to line 500 because the condition on line 499 was never true
500 response.delete_cookie('oauth_id_token')
502 # Delete OAuth session if present
503 if request.cookies.get('oauth_session_id'): 503 ↛ 504line 503 didn't jump to line 504 because the condition on line 503 was never true
504 response.delete_cookie('oauth_session_id')
506 raise e
509async def get_current_user_by_api_key(request, api_key: str):
510 # Each function call manages its own short-lived session internally
511 user = await Users.get_user_by_api_key(api_key)
513 if user is None:
514 raise HTTPException(
515 status_code=status.HTTP_401_UNAUTHORIZED,
516 detail=ERROR_MESSAGES.INVALID_TOKEN,
517 )
519 config_values = await Config.get_many(
520 'auth.enable_api_keys',
521 'user.permissions',
522 'auth.api_key.endpoint_restrictions',
523 'auth.api_key.allowed_endpoints',
524 )
526 if not config_values.get('auth.enable_api_keys'):
527 raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.API_KEY_NOT_ALLOWED)
529 if user.role != 'admin':
530 user_permissions = config_values.get('user.permissions')
531 if not await has_permission(
532 user.id,
533 'features.api_keys',
534 user_permissions,
535 ):
536 raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.API_KEY_NOT_ALLOWED)
538 # Enforce endpoint restrictions — checked here (not in middleware)
539 # so it applies regardless of how the API key was transported
540 # (Authorization header, cookie, x-api-key header, etc.).
541 if config_values.get('auth.api_key.endpoint_restrictions'):
542 allowed_endpoints = config_values.get('auth.api_key.allowed_endpoints', '')
543 allowed_paths = [path.strip() for path in str(allowed_endpoints).split(',') if path.strip()]
544 request_path = request.scope['path'] # Use raw ASGI path — not spoofable via Host header (CVE-2026-48710)
545 is_allowed = any(request_path == allowed or request_path.startswith(allowed + '/') for allowed in allowed_paths)
546 if not is_allowed:
547 raise HTTPException(
548 status_code=status.HTTP_403_FORBIDDEN,
549 detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
550 )
552 # Add user info to current span
553 if ENABLE_OTEL:
554 from opentelemetry import trace
556 current_span = trace.get_current_span()
557 if current_span:
558 current_span.set_attribute('client.user.id', user.id)
559 current_span.set_attribute('client.user.email', user.email)
560 current_span.set_attribute('client.user.role', user.role)
561 current_span.set_attribute('client.auth.type', 'api_key')
563 await Users.update_last_active_by_id(user.id)
564 return user
567VERIFIED_USER_ROLES = {'user', 'admin'}
570def get_verified_user(user=Depends(get_current_user)):
571 if user.role not in VERIFIED_USER_ROLES: 571 ↛ 572line 571 didn't jump to line 572 because the condition on line 571 was never true
572 raise HTTPException(
573 status_code=status.HTTP_401_UNAUTHORIZED,
574 detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
575 )
576 return user
579async def get_verified_user_by_token(token: str, redis=None):
580 """Resolve a verified user from a raw token, for WebSocket handshakes that run outside the HTTP dependency chain."""
581 decoded = decode_token(token)
582 if decoded is None or 'id' not in decoded or not await is_valid_token(decoded, redis):
583 return None
585 user = await Users.get_user_by_id(decoded['id'])
586 if user is None or user.role not in VERIFIED_USER_ROLES:
587 return None
589 return user
592async def get_verified_user_by_id(user_id: str | None):
593 if not user_id:
594 return None
596 user = await Users.get_user_by_id(user_id)
597 if user is None or user.role not in VERIFIED_USER_ROLES:
598 return None
600 return user
603async def get_optional_verified_user_from_request(request: Request):
604 token = None
605 auth_token = get_http_authorization_cred(request.headers.get('Authorization'))
606 if auth_token:
607 token = auth_token.credentials
608 if token is None:
609 token = request.cookies.get('token')
610 if token is None and getattr(request.state, 'token', None):
611 token = request.state.token.credentials
612 if not token:
613 return None
615 try:
616 if token.startswith('sk-'):
617 user = await get_current_user_by_api_key(request, token)
618 return user if user.role in VERIFIED_USER_ROLES else None
620 return await get_verified_user_by_token(token, getattr(request.app.state, 'redis', None))
621 except HTTPException:
622 return None
625def get_admin_user(user=Depends(get_current_user)):
626 if user.role != 'admin': 626 ↛ 627line 626 didn't jump to line 627 because the condition on line 626 was never true
627 raise HTTPException(
628 status_code=status.HTTP_401_UNAUTHORIZED,
629 detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
630 )
631 return user
634async def create_admin_user(email: str, password: str, name: str = 'Admin'):
635 """
636 Create an admin user from environment variables.
637 Used for headless/automated deployments.
638 Returns the created user or None if creation failed.
639 """
641 if not email or not password:
642 return None
644 if await Users.has_users():
645 log.debug('Users already exist, skipping admin creation')
646 return None
648 log.info('Creating admin account from environment variables: %s', email)
649 try:
650 hashed = await get_password_hash(password)
651 user = await Auths.insert_new_auth(
652 email=email.lower(),
653 password=hashed,
654 name=name,
655 role='admin',
656 )
657 if user:
658 log.info('Admin account created successfully: %s', email)
659 return user
660 else:
661 log.error('Failed to create admin account from environment variables')
662 return None
663 except Exception as e:
664 log.error(f'Error creating admin account: {e}')
665 return None