diff --git a/.env.example b/.env.example index 81cd826ada..5ca052422d 100644 --- a/.env.example +++ b/.env.example @@ -13,6 +13,11 @@ OPENAI_API_KEY='' # CORS_ALLOW_ORIGIN='http://localhost:5173;http://localhost:8080' CORS_ALLOW_ORIGIN='*' +# MFA defaults; saved settings take precedence once persistent config is seeded. +ENABLE_MFA=false +MFA_ALLOW_OAUTH_BYPASS=false +MFA_ALLOW_TRUSTED_HEADER_BYPASS=false + # Set to false to keep memory tools enabled without adding memory context to the system context. ENABLE_MEMORY_SYSTEM_CONTEXT=true @@ -35,7 +40,9 @@ ENABLE_FUNCTIONS=true # Set false to disable external tools and terminals, including personal direct connections. ENABLE_TOOL_SERVERS=true -# For production you should set this to match the proxy configuration (127.0.0.1) +# WARNING: * trusts forwarded headers from every connection. Use only behind a trusted +# proxy that sanitizes these headers and prevents direct access to the backend. +# Otherwise, set this to the IP addresses or networks of your trusted reverse proxies. FORWARDED_ALLOW_IPS='*' # DO NOT TRACK diff --git a/backend/open_webui/__init__.py b/backend/open_webui/__init__.py index c56c875566..0611bda1c6 100644 --- a/backend/open_webui/__init__.py +++ b/backend/open_webui/__init__.py @@ -9,6 +9,8 @@ import typer import uvicorn app = typer.Typer() +mfa_app = typer.Typer(help='Manage multi-factor authentication.') +app.add_typer(mfa_app, name='mfa') KEY_FILE = Path.cwd() / '.webui_secret_key' DEFAULT_SECRET_KEY_LENGTH = 24 @@ -85,7 +87,7 @@ def serve( 'open_webui.main:app', host=host, port=port, - forwarded_allow_ips='*', + forwarded_allow_ips=os.getenv('FORWARDED_ALLOW_IPS', '*'), workers=UVICORN_WORKERS, ws_per_message_deflate=UVICORN_WS_PER_MESSAGE_DEFLATE, loop=loop, @@ -105,10 +107,32 @@ def dev( host=host, port=port, reload=reload, - forwarded_allow_ips='*', + forwarded_allow_ips=os.getenv('FORWARDED_ALLOW_IPS', '*'), ws_per_message_deflate=UVICORN_WS_PER_MESSAGE_DEFLATE, ) +@mfa_app.command() +def reset(email: str, reason: Annotated[str, typer.Option('--reason')]): + """Issue a one-time recovery ticket after the operator verifies the user's identity.""" + import asyncio + + if not os.getenv('WEBUI_SECRET_KEY'): + if not KEY_FILE.exists(): + raise typer.BadParameter( + 'Provide the existing WEBUI_SECRET_KEY or run in the directory containing .webui_secret_key.' + ) + os.environ['WEBUI_SECRET_KEY'] = KEY_FILE.read_text() + from open_webui.utils.mfa import reset_mfa + + try: + ticket = asyncio.run(reset_mfa(email, reason)) + except Exception as error: + typer.echo(f'Reset failed: {error}', err=True) + raise typer.Exit(1) from None + typer.echo('Recovery token (expires in 30 minutes; deliver securely to the verified user):') + typer.echo(ticket) + + if __name__ == '__main__': app() diff --git a/backend/open_webui/config.py b/backend/open_webui/config.py index 071f59695f..25651fa580 100644 --- a/backend/open_webui/config.py +++ b/backend/open_webui/config.py @@ -2455,6 +2455,10 @@ Responses from models: {{responses}}""" ENABLE_API_KEYS = os.getenv('ENABLE_API_KEYS', 'False').lower() == 'true' +ENABLE_MFA = os.getenv('ENABLE_MFA', 'False').lower() == 'true' +MFA_ALLOW_OAUTH_BYPASS = os.getenv('MFA_ALLOW_OAUTH_BYPASS', 'False').lower() == 'true' +MFA_ALLOW_TRUSTED_HEADER_BYPASS = os.getenv('MFA_ALLOW_TRUSTED_HEADER_BYPASS', 'False').lower() == 'true' + ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS = ( os.getenv( 'ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS', @@ -3159,6 +3163,9 @@ DEFAULT_CONFIG = { 'auth.api_key.endpoint_restrictions': ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS, 'auth.api_key.allowed_endpoints': API_KEYS_ALLOWED_ENDPOINTS, 'auth.jwt_expiry': JWT_EXPIRES_IN, + 'auth.mfa.enable': ENABLE_MFA, + 'auth.mfa.allow_oauth_bypass': MFA_ALLOW_OAUTH_BYPASS, + 'auth.mfa.allow_trusted_header_bypass': MFA_ALLOW_TRUSTED_HEADER_BYPASS, 'oauth.enable': ENABLE_OAUTH, 'oauth.enable_signup': ENABLE_OAUTH_SIGNUP, 'oauth.auto_redirect': OAUTH_AUTO_REDIRECT, diff --git a/backend/open_webui/events.py b/backend/open_webui/events.py index aa4a49dccd..e09f3bfa70 100644 --- a/backend/open_webui/events.py +++ b/backend/open_webui/events.py @@ -100,9 +100,41 @@ class EventDefinitions(BaseModel): AUTH_SIGNUP: EventDefinition = EventDefinition( name='auth.signup', description='A user account was created through signup.', message='User signed up' ) + AUTH_MFA_ENROLLED: EventDefinition = EventDefinition( + name='auth.mfa.enrolled', description='MFA enrolled.', message='MFA enrolled' + ) + AUTH_MFA_REPLACED: EventDefinition = EventDefinition( + name='auth.mfa.replaced', description='MFA replaced.', message='MFA replaced' + ) + AUTH_MFA_FAILED: EventDefinition = EventDefinition( + name='auth.mfa.failed', description='MFA failed.', message='MFA failed' + ) + AUTH_MFA_THROTTLED: EventDefinition = EventDefinition( + name='auth.mfa.throttled', description='MFA throttled.', message='MFA throttled' + ) + AUTH_MFA_RECOVERY_USED: EventDefinition = EventDefinition( + name='auth.mfa.recovery_used', description='MFA recovery used.', message='MFA recovery used' + ) + AUTH_MFA_RECOVERY_CODES_REGENERATED: EventDefinition = EventDefinition( + name='auth.mfa.recovery_codes_regenerated', + description='MFA recovery codes regenerated.', + message='MFA recovery codes regenerated', + ) + AUTH_MFA_RESET_REQUESTED: EventDefinition = EventDefinition( + name='auth.mfa.reset_requested', description='MFA reset requested.', message='MFA reset requested' + ) + AUTH_MFA_RESET_COMPLETED: EventDefinition = EventDefinition( + name='auth.mfa.reset_completed', description='MFA reset completed.', message='MFA reset completed' + ) + AUTH_MFA_POLICY_CHANGED: EventDefinition = EventDefinition( + name='auth.mfa.policy_changed', description='MFA policy changed.', message='MFA policy changed' + ) AUTH_LOGIN: EventDefinition = EventDefinition( name='auth.login', description='A user successfully logged in.', message='User logged in' ) + AUTH_SESSIONS_REVOKED: EventDefinition = EventDefinition( + name='auth.sessions_revoked', description='All user sessions were revoked.', message='User sessions revoked' + ) AUTH_LOGOUT: EventDefinition = EventDefinition( name='auth.logout', description='A user logged out.', message='User logged out' ) diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 452f65fb8e..e3d9e0e775 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -165,6 +165,7 @@ from open_webui.routers import ( images, knowledge, memories, + mfa, models, notes, notifications, @@ -193,6 +194,7 @@ from open_webui.socket.main import ( get_models_in_use, get_user_id_from_session_pool, periodic_session_pool_cleanup, + periodic_socket_authentication, periodic_usage_pool_cleanup, redis_event_listener, sio, @@ -371,6 +373,9 @@ async def lifespan(app: FastAPI): await import_legacy_config_json() await seed_registered_defaults() + from open_webui.utils.mfa import validate_mfa_configuration + + await validate_mfa_configuration() await initialize_runtime_config(app) await migrate_legacy_webhook_config() await publish_event(app, EVENTS.SYSTEM_STARTUP_STARTED, source='system') @@ -406,6 +411,7 @@ async def lifespan(app: FastAPI): sio.manager_initialized = True sio.manager.initialize() + app.state.periodic_socket_authentication = asyncio.create_task(periodic_socket_authentication()) app.state.periodic_usage_pool_cleanup = asyncio.create_task(periodic_usage_pool_cleanup()) app.state.periodic_session_pool_cleanup = asyncio.create_task(periodic_session_pool_cleanup()) @@ -500,6 +506,7 @@ async def lifespan(app: FastAPI): if hasattr(app.state, 'redis_event_listener'): app.state.redis_event_listener.cancel() + app.state.periodic_socket_authentication.cancel() app.state.periodic_usage_pool_cleanup.cancel() app.state.periodic_session_pool_cleanup.cancel() app.state.scheduler_worker_loop.cancel() @@ -863,6 +870,7 @@ app.include_router(retrieval.router, prefix='/api/v1/retrieval', tags=['retrieva app.include_router(configs.router, prefix='/api/v1/configs', tags=['configs']) app.include_router(auths.router, prefix='/api/v1/auths', tags=['auths']) +app.include_router(mfa.router, prefix='/api/v1/auths/mfa', tags=['auths']) app.include_router(users.router, prefix='/api/v1/users', tags=['users']) diff --git a/backend/open_webui/migrations/versions/a7d3e9f2b641_add_mfa_and_session_stamp.py b/backend/open_webui/migrations/versions/a7d3e9f2b641_add_mfa_and_session_stamp.py new file mode 100644 index 0000000000..6ae864b039 --- /dev/null +++ b/backend/open_webui/migrations/versions/a7d3e9f2b641_add_mfa_and_session_stamp.py @@ -0,0 +1,24 @@ +"""add MFA state and account session stamps + +Revision ID: a7d3e9f2b641 +Revises: d4c1a8e37b62 +""" + +import sqlalchemy as sa +from alembic import op + +revision = 'a7d3e9f2b641' +down_revision = 'd4c1a8e37b62' +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column('auth', sa.Column('mfa', sa.JSON(), nullable=True)) + op.add_column('auth', sa.Column('session_stamp', sa.Text(), nullable=True)) + + +def downgrade() -> None: + with op.batch_alter_table('auth') as batch: + batch.drop_column('session_stamp') + batch.drop_column('mfa') diff --git a/backend/open_webui/models/auths.py b/backend/open_webui/models/auths.py index b9f8b4b2a0..f5350cc737 100644 --- a/backend/open_webui/models/auths.py +++ b/backend/open_webui/models/auths.py @@ -2,16 +2,17 @@ from __future__ import annotations +import datetime as dt import logging import uuid -from typing import Optional +from typing import Literal import bcrypt -from open_webui.internal.db import Base, JSONField, get_async_db_context -from open_webui.models.users import User, UserModel, UserProfileImageResponse, Users +from open_webui.internal.db import Base, get_async_db, get_async_db_context +from open_webui.models.users import User, UserModel, UserProfileImageResponse, Users, UserStatus from open_webui.utils.validate import validate_image_url -from pydantic import BaseModel, field_validator -from sqlalchemy import Boolean, Column, String, Text, delete, select, update +from pydantic import BaseModel, ConfigDict, Field, SecretStr, field_validator +from sqlalchemy import JSON, Boolean, Column, String, Text, delete, select, update from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession @@ -32,15 +33,58 @@ class Auth(Base): # credential ↔ user linkage email = Column(String) # login address, kept in sync with User.email password = Column(Text) # argon2 / bcrypt hash active = Column(Boolean) # account soft-disable toggle + mfa = Column(JSON, nullable=True) + session_stamp = Column(Text, nullable=True) + + +class MfaLimit(BaseModel): + count: int = Field(default=0, ge=0) + expires_at: int + + +class MfaResetTicket(BaseModel): + token_hash: str + expires_at: int + + +class MfaChallenge(BaseModel): + token_hash: str + type: Literal['enroll', 'verify', 'replace', 'recover'] + expires_at: int + attempts: int = Field(default=0, ge=0) + auth_method: Literal['password', 'ldap', 'oauth', 'trusted_header', 'system', 'api'] + auth_time: int + session_stamp: str | None = None + secret: str | None = None + oauth_session_id: str | None = None + provider: str | None = None + + +class MfaData(BaseModel): + model_config = ConfigDict(extra='forbid') + + revision: str = Field(default_factory=lambda: str(uuid.uuid4())) + secret: str | None = None + last_step: int = -1 + recovery_hashes: list[str] = Field(default_factory=list) + login_challenge: MfaChallenge | None = None + manage_challenge: MfaChallenge | None = None + reset_required: bool = False + reset_ticket: MfaResetTicket | None = None + limits: dict[str, MfaLimit] = Field(default_factory=dict) class AuthModel(BaseModel): """Pydantic mirror of the ``auth`` table row.""" + model_config = ConfigDict(from_attributes=True) + id: str email: str password: str active: bool = True + mfa: MfaData | None = None + session_stamp: str | None = None class Token(BaseModel): @@ -58,6 +102,74 @@ class SigninResponse(Token, UserProfileImageResponse): pass +class SessionUserResponse(Token, UserProfileImageResponse): + expires_at: int | None = None + permissions: dict | None = None + + +class SessionUserInfoResponse(SessionUserResponse, UserStatus): + bio: str | None = None + gender: str | None = None + date_of_birth: dt.date | None = None + + +class AddUserResponse(UserProfileImageResponse): + token: str | None = None + token_type: str | None = None + + +class MfaChallengeResponse(BaseModel): + next_step: Literal['enroll', 'verify', 'recover'] + challenge_token: str + expires_in: int + + +class PendingUserResponse(BaseModel): + next_step: Literal['pending'] = 'pending' + + +class MfaStatusResponse(BaseModel): + enabled: bool + required: bool + recovery_codes_remaining: int + + +class MfaSetupResponse(BaseModel): + manual_key: str + qr_code: str + + +class MfaRecoveryCodesResponse(BaseModel): + recovery_codes: list[str] + + +class MfaEnrollmentResponse(SessionUserResponse): + recovery_codes: list[str] + + +class MfaChallengeForm(BaseModel): + model_config = ConfigDict(extra='forbid') + challenge_token: SecretStr = Field(min_length=20, max_length=160) + + +class MfaVerifyForm(MfaChallengeForm): + code: SecretStr = Field(min_length=1, max_length=128) + recovery: bool = False + + +class MfaFactorForm(BaseModel): + model_config = ConfigDict(extra='forbid') + code: SecretStr = Field(min_length=1, max_length=128) + recovery: bool = False + + +class MfaRecoveryForm(MfaChallengeForm): + reset_token: SecretStr = Field(min_length=20, max_length=160) + + +SigninResult = SessionUserResponse | MfaChallengeResponse | PendingUserResponse + + class SigninForm(BaseModel): email: str password: str @@ -101,6 +213,78 @@ class AddUserForm(SignupForm): class AuthsTable: """Provides CRUD operations for the Auth ↔ User lifecycle.""" + async def get_auth_by_id(self, user_id: str, db: AsyncSession | None = None) -> AuthModel | None: + if db is None: + async with get_async_db() as session: + return await self.get_auth_by_id(user_id, db=session) + row = await db.get(Auth, user_id, populate_existing=True) + return AuthModel.model_validate(row) if row else None + + @staticmethod + def clear_mfa_challenges(value: dict | None) -> dict | None: + if value is None: + return None + mfa = MfaData.model_validate(value) + mfa.login_challenge = None + mfa.manage_challenge = None + mfa.reset_ticket = None + mfa.revision = str(uuid.uuid4()) + return mfa.model_dump(mode='json') + + async def update_mfa_by_id( + self, auth: AuthModel, mfa: MfaData, *, revoke: bool = False, db: AsyncSession + ) -> AuthModel | None: + """Compare-and-swap the complete credential state. The caller owns the transaction.""" + revision = Auth.mfa['revision'].as_string() + expected = auth.mfa.revision if auth.mfa else None + mfa = mfa.model_copy(deep=True, update={'revision': str(uuid.uuid4())}) + stamp = str(uuid.uuid4()) if revoke else auth.session_stamp + result = await db.execute( + update(Auth) + .where( + Auth.id == auth.id, + Auth.active.is_(True), + Auth.password == auth.password, + Auth.session_stamp == auth.session_stamp, + revision == expected if expected is not None else revision.is_(None), + ) + .values(mfa=mfa.model_dump(mode='json'), session_stamp=stamp) + .execution_options(synchronize_session=False) + ) + if result.rowcount != 1: + return None + return auth.model_copy(update={'mfa': mfa, 'session_stamp': stamp}) + + async def revoke_sessions_by_user_id(self, user_id: str, *, db: AsyncSession) -> bool: + row = ( + await db.execute(select(Auth).where(Auth.id == user_id).execution_options(populate_existing=True)) + ).scalar_one_or_none() + if row is None: + return False + # Only clear challenges if the JSON still matches; retry instead of overwriting a factor change. + auth = AuthModel.model_validate(row) + revision = Auth.mfa['revision'].as_string() + expected = auth.mfa.revision if auth.mfa else None + result = await db.execute( + update(Auth) + .where( + Auth.id == user_id, + Auth.session_stamp == auth.session_stamp, + revision == expected if expected is not None else revision.is_(None), + ) + .values(session_stamp=str(uuid.uuid4()), mfa=self.clear_mfa_challenges(row.mfa)) + .execution_options(synchronize_session=False) + ) + if result.rowcount != 1: + raise ValueError('Authentication changed in another request. Please try again.') + return True + + async def revoke_all_sessions(self, *, db: AsyncSession) -> list[str]: + user_ids = list((await db.execute(select(Auth.id))).scalars()) + for user_id in user_ids: + await self.revoke_sessions_by_user_id(user_id, db=db) + return user_ids + async def insert_new_auth( self, email: str, @@ -122,6 +306,7 @@ class AuthsTable: email=email, password=password, active=True, + session_stamp=str(uuid.uuid4()), ) session.add(credential) @@ -146,7 +331,7 @@ class AuthsTable: email: str, verify_password: callable, db: AsyncSession | None = None, - ) -> UserModel | None: + ) -> tuple[UserModel, AuthModel] | None: """Verify email + password credentials and return the matching user.""" log.info('authenticate_user: %s', email) resolved = await Users.get_user_by_email(email, db=db) @@ -161,7 +346,7 @@ class AuthsTable: return if not await verify_password(credential.password): return - return resolved + return resolved, AuthModel.model_validate(credential) async def authenticate_user_by_api_key( self, @@ -214,14 +399,21 @@ class AuthsTable: self, user_id: str, new_password: str, + *, + current_auth: AuthModel | None = None, db: AsyncSession | None = None, ) -> bool: """Set a new password hash for an existing user.""" async with get_async_db_context(db) as session: - auth_row = await session.get(Auth, user_id) - if auth_row is None: + auth = current_auth or await self.get_auth_by_id(user_id, db=session) + if auth is None: return False - auth_row.password = new_password + state = auth.mfa.model_copy(deep=True) if auth.mfa else MfaData() + state.login_challenge = state.manage_challenge = state.reset_ticket = None + updated = await self.update_mfa_by_id(auth, state, revoke=True, db=session) + if updated is None: + raise ValueError('Authentication changed in another request. Please try again.') + await session.execute(update(Auth).where(Auth.id == user_id).values(password=new_password)) await session.commit() return True diff --git a/backend/open_webui/models/config.py b/backend/open_webui/models/config.py index 61b41f37cf..fc6592f409 100644 --- a/backend/open_webui/models/config.py +++ b/backend/open_webui/models/config.py @@ -17,6 +17,7 @@ from typing import Any, ClassVar from fastapi.encoders import jsonable_encoder from open_webui.internal.db import Base, get_async_db from sqlalchemy import JSON, BigInteger, Column, Text, delete, select +from sqlalchemy.ext.asyncio import AsyncSession log = logging.getLogger(__name__) @@ -195,7 +196,7 @@ class Config(Base): return values @staticmethod - async def upsert(updates: dict) -> None: + async def upsert(updates: dict, *, db: AsyncSession | None = None) -> None: """Upsert multiple config key-value pairs. Raises on failure.""" persistent_updates = {} for key, value in updates.items(): @@ -208,16 +209,21 @@ class Config(Base): if not persistent_updates: return - async with get_async_db() as db: - now = int(time.time()) - for key, value in persistent_updates.items(): - existing = await db.get(Config, key) - if existing: - existing.value = value - existing.updated_at = now - else: - db.add(Config(key=key, value=value, updated_at=now)) - await db.commit() + if db is None: + async with get_async_db() as session: + await Config.upsert(persistent_updates, db=session) + await session.commit() + return + + now = int(time.time()) + for key, value in persistent_updates.items(): + existing = await db.get(Config, key) + if existing: + existing.value = value + existing.updated_at = now + else: + db.add(Config(key=key, value=value, updated_at=now)) + await db.flush() @staticmethod async def delete(key: str) -> bool: diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index 0ca2f518f8..891882bab5 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -21,7 +21,6 @@ from open_webui.config import ( OAUTH_PROVIDERS, ) from open_webui.constants import ERROR_MESSAGES -from open_webui.events import EVENTS, publish_event from open_webui.env import ( AIOHTTP_CLIENT_SESSION_SSL, ENABLE_INITIAL_ADMIN_SIGNUP, @@ -38,16 +37,18 @@ from open_webui.env import ( WEBUI_AUTH_TRUSTED_NAME_HEADER, WEBUI_AUTH_TRUSTED_ROLE_HEADER, ) +from open_webui.events import EVENTS, publish_event from open_webui.internal.db import get_async_session from open_webui.models.auths import ( AddUserForm, + AddUserResponse, ApiKey, Auths, LdapForm, + SessionUserInfoResponse, SigninForm, - SigninResponse, + SigninResult, SignupForm, - Token, UpdatePasswordForm, ) from open_webui.models.config import Config @@ -58,32 +59,34 @@ from open_webui.models.users import ( UserModel, UserProfileImageResponse, Users, - UserStatus, ) +from open_webui.routers.mfa import MfaRoute from open_webui.utils.access_control import get_permissions, has_permission from open_webui.utils.auth import ( create_api_key, + create_signin_response, create_token, decode_token, get_admin_user, get_current_user, get_http_authorization_cred, + get_human_user, get_password_hash, get_verified_user, invalidate_token, - revoke_user_tokens, validate_password, verify_password, ) from open_webui.utils.groups import apply_default_group_assignment from open_webui.utils.json_codec import JSONCodec +from open_webui.utils.mfa import MFA_CONFIG_KEYS, MfaConfigForm, get_mfa_config, limit_account, update_mfa_config from open_webui.utils.misc import parse_duration, validate_email_format from open_webui.utils.rate_limit import RateLimiter from pydantic import BaseModel, StrictStr, field_validator from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession -router = APIRouter() +router = APIRouter(route_class=MfaRoute) log = logging.getLogger(__name__) @@ -103,6 +106,7 @@ token_exchange_rate_limiter = ( ADMIN_CONFIG_KEYS = { + **MFA_CONFIG_KEYS, 'SHOW_ADMIN_DETAILS': 'auth.admin.show', 'ADMIN_EMAIL': 'auth.admin.email', 'WEBUI_URL': 'webui.url', @@ -165,93 +169,11 @@ def config_updates(data: dict, key_map: dict[str, str]) -> dict: return {key_map[field]: value for field, value in data.items() if field in key_map} -async def create_session_response( - request: Request, - user, - db, - response: Response = None, - set_cookie: bool = False, - source: str = 'api', -) -> dict: - """ - Create JWT token and build session response for a user. - Shared helper for signin, signup, ldap_auth, add_user, and token_exchange endpoints. - - Args: - request: FastAPI request object - user: User object - db: Database session - response: FastAPI response object (required if set_cookie is True) - set_cookie: Whether to set the auth cookie on the response - """ - expires_delta = parse_duration(await Config.get('auth.jwt_expiry')) - expires_at = None - if expires_delta: - expires_at = int(time.time()) + int(expires_delta.total_seconds()) - - token = create_token( - data={'id': user.id}, - expires_delta=expires_delta, - ) - - if set_cookie and response: - datetime_expires_at = datetime.datetime.fromtimestamp(expires_at, datetime.timezone.utc) if expires_at else None - max_age = int(expires_delta.total_seconds()) if expires_delta else None - response.set_cookie( - key='token', - value=token, - expires=datetime_expires_at, - httponly=True, - samesite=WEBUI_AUTH_COOKIE_SAME_SITE, - secure=WEBUI_AUTH_COOKIE_SECURE, - **({'max_age': max_age} if max_age is not None else {}), - ) - - user_permissions = await get_permissions(user.id, await Config.get('user.permissions'), db=db) - await publish_event( - request, - EVENTS.AUTH_LOGIN, - actor=user, - subject_id=user.id, - subject_type='user', - source=source, - data={'auth_method': source}, - ) - - return { - 'token': token, - 'token_type': 'Bearer', - 'expires_at': expires_at, - 'id': user.id, - 'email': user.email, - 'name': user.name, - 'role': user.role, - 'profile_image_url': f'/api/v1/users/{user.id}/profile/image', - 'permissions': user_permissions, - } - - -############################ -# GetSessionUser -############################ - - -class SessionUserResponse(Token, UserProfileImageResponse): - expires_at: int | None = None - permissions: dict | None = None - - -class SessionUserInfoResponse(SessionUserResponse, UserStatus): - bio: str | None = None - gender: str | None = None - date_of_birth: datetime.date | None = None - - @router.get('/', response_model=SessionUserInfoResponse) async def get_session_user( request: Request, response: Response, - user=Depends(get_current_user), + user=Depends(get_human_user), db: AsyncSession = Depends(get_async_session), ): token = None @@ -395,11 +317,12 @@ async def update_password( if WEBUI_AUTH_TRUSTED_EMAIL_HEADER: raise HTTPException(status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.ACTION_PROHIBITED) if session_user: - user = await Auths.authenticate_user( + authenticated = await Auths.authenticate_user( session_user.email, lambda pw: verify_password(form_data.password, pw), db=db, ) + user, auth = authenticated if authenticated else (None, None) if user: try: @@ -407,9 +330,14 @@ async def update_password( except Exception as e: raise HTTPException(400, detail=str(e)) hashed = await get_password_hash(form_data.new_password) - success = await Auths.update_user_password_by_id(user.id, hashed, db=db) + try: + success = await Auths.update_user_password_by_id(user.id, hashed, current_auth=auth, db=db) + except ValueError: + raise HTTPException(409, 'Authentication changed. Please sign in again.') from None if success: - await revoke_user_tokens(request, user.id) + from open_webui.socket.main import disconnect_user_sessions + + await disconnect_user_sessions(user.id) await publish_event( request, EVENTS.AUTH_PASSWORD_CHANGED, @@ -473,7 +401,7 @@ def extract_group_cn_from_dn(group_dn: str) -> str | None: ############################ # LDAP Authentication ############################ -@router.post('/ldap', response_model=SessionUserResponse) +@router.post('/ldap', response_model=SigninResult) async def ldap_auth( request: Request, response: Response, @@ -635,6 +563,10 @@ async def ldap_auth( ) if username_list and form_data.user.lower() in username_list: + if (await get_mfa_config()).ENABLE_MFA: + existing = await Users.get_user_by_email(email, db=db) + if existing: + await limit_account(existing.id, 'password') connection_user = Connection( server, user_dn, @@ -700,11 +632,13 @@ async def ldap_auth( except Exception as e: log.error(f'Failed to sync groups for user {user.id}: {e}') - return await create_session_response(request, user, db, response, set_cookie=True, source='ldap') + return await create_signin_response(request, user, db, response, set_cookie=True, source='ldap') else: raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) else: raise HTTPException(400, 'User record mismatch.') + except HTTPException: + raise except Exception as e: log.error(f'LDAP authentication error: {str(e)}') raise HTTPException(400, detail='LDAP authentication failed.') @@ -715,7 +649,7 @@ async def ldap_auth( ############################ -@router.post('/signin', response_model=SessionUserResponse) +@router.post('/signin', response_model=SigninResult) async def signin( request: Request, response: Response, @@ -729,6 +663,7 @@ async def signin( ) auth_source = 'password' + auth = None if WEBUI_AUTH_TRUSTED_EMAIL_HEADER: auth_source = 'trusted_header' @@ -792,11 +727,12 @@ async def signin( admin_password = 'admin' if await Users.get_user_by_email(admin_email.lower(), db=db): - user = await Auths.authenticate_user( + authenticated = await Auths.authenticate_user( admin_email.lower(), lambda pw: verify_password(admin_password, pw), db=db, ) + user, auth = authenticated if authenticated else (None, None) else: if await Users.has_users(db=db): raise HTTPException(400, detail=ERROR_MESSAGES.EXISTING_USERS) @@ -810,11 +746,12 @@ async def signin( source='system', ) - user = await Auths.authenticate_user( + authenticated = await Auths.authenticate_user( admin_email.lower(), lambda pw: verify_password(admin_password, pw), db=db, ) + user, auth = authenticated if authenticated else (None, None) else: if await signin_rate_limiter.is_limited(request.app.state.redis, form_data.email.lower()): raise HTTPException( @@ -822,14 +759,20 @@ async def signin( detail=ERROR_MESSAGES.RATE_LIMIT_EXCEEDED, ) - user = await Auths.authenticate_user( + if (await get_mfa_config()).ENABLE_MFA: + existing = await Users.get_user_by_email(form_data.email.lower(), db=db) + if existing: + await limit_account(existing.id, 'password') + + authenticated = await Auths.authenticate_user( form_data.email.lower(), lambda pw: verify_password(form_data.password, pw), db=db, ) + user, auth = authenticated if authenticated else (None, None) if user: - return await create_session_response(request, user, db, response, set_cookie=True, source=auth_source) + return await create_signin_response(request, user, db, response, set_cookie=True, source=auth_source, auth=auth) else: raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) @@ -897,7 +840,7 @@ async def signup_handler( return user -@router.post('/signup', response_model=SessionUserResponse) +@router.post('/signup', response_model=SigninResult) async def signup( request: Request, response: Response, @@ -945,7 +888,7 @@ async def signup( subject_type='user', data={'email': user.email}, ) - return await create_session_response(request, user, db, response, set_cookie=True) + return await create_signin_response(request, user, db, response, set_cookie=True, source='password') except HTTPException: raise except Exception as err: @@ -1099,7 +1042,7 @@ async def delete_oauth_session_by_provider( ############################ -@router.post('/add', response_model=SigninResponse) +@router.post('/add', response_model=AddUserResponse, response_model_exclude_none=True) async def add_user( request: Request, form_data: AddUserForm, @@ -1145,7 +1088,12 @@ async def add_user( ) expires_delta = parse_duration(await Config.get('auth.jwt_expiry')) - token = create_token(data={'id': user.id}, expires_delta=expires_delta) + if (await get_mfa_config()).ENABLE_MFA: + return user.model_dump() + auth = await Auths.get_auth_by_id(user.id, db=db) + token = create_token( + data={'id': user.id, 'typ': 'session', 'session_stamp': auth.session_stamp}, expires_delta=expires_delta + ) return { 'token': token, 'token_type': 'Bearer', @@ -1207,7 +1155,7 @@ async def get_admin_config(request: Request, user=Depends(get_admin_user)): return await get_config_values(ADMIN_CONFIG_KEYS) -class AdminConfig(BaseModel): +class AdminConfig(MfaConfigForm): SHOW_ADMIN_DETAILS: bool ADMIN_EMAIL: str | None = None WEBUI_URL: str @@ -1270,6 +1218,9 @@ class AdminConfig(BaseModel): @router.post('/admin/config') async def update_admin_config(request: Request, form_data: AdminConfig, user=Depends(get_admin_user)): updates = config_updates(form_data.model_dump(), ADMIN_CONFIG_KEYS) + for field, key in MFA_CONFIG_KEYS.items(): + if field not in form_data.model_fields_set: + updates.pop(key, None) if 'ENABLE_LOGIN_FORM' not in form_data.model_fields_set: updates.pop('ui.enable_login_form', None) if 'I18N' not in form_data.model_fields_set: @@ -1293,8 +1244,8 @@ async def update_admin_config(request: Request, form_data: AdminConfig, user=Dep if not re.match(pattern, form_data.JWT_EXPIRES_IN): updates.pop('auth.jwt_expiry', None) - await Config.upsert(updates) - return await get_config_values(ADMIN_CONFIG_KEYS) + changed = await update_mfa_config(request, updates) + return {**(await get_config_values(ADMIN_CONFIG_KEYS)), 'sessions_revoked': changed} class LdapServerConfig(BaseModel): @@ -1622,7 +1573,7 @@ async def get_token_client_id(client, token: str) -> str | None: return None -@router.post('/oauth/{provider}/token/exchange', response_model=SessionUserResponse) +@router.post('/oauth/{provider}/token/exchange', response_model=SigninResult) async def token_exchange( request: Request, response: Response, @@ -1778,4 +1729,4 @@ async def token_exchange( db=db, ) - return await create_session_response(request, user, db, source='oauth') + return await create_signin_response(request, user, db, source='oauth') diff --git a/backend/open_webui/routers/configs.py b/backend/open_webui/routers/configs.py index f48e5cf68f..9aca44f967 100644 --- a/backend/open_webui/routers/configs.py +++ b/backend/open_webui/routers/configs.py @@ -99,7 +99,9 @@ class ImportConfigForm(BaseModel): @router.post('/import', response_model=dict) async def import_config(request: Request, form_data: ImportConfigForm, user=Depends(get_admin_user)): - await Config.upsert(form_data.config) + from open_webui.utils.mfa import update_mfa_config + + await update_mfa_config(request, form_data.config) await publish_event( request, EVENTS.CONFIG_IMPORTED, diff --git a/backend/open_webui/routers/mfa.py b/backend/open_webui/routers/mfa.py new file mode 100644 index 0000000000..3c65951854 --- /dev/null +++ b/backend/open_webui/routers/mfa.py @@ -0,0 +1,187 @@ +"""MFA endpoints; challenges never authorize ordinary application requests.""" + +from urllib.parse import urlsplit + +from fastapi import APIRouter, Depends, HTTPException, Request, Response +from fastapi.exceptions import RequestValidationError +from fastapi.responses import JSONResponse +from fastapi.routing import APIRoute +from open_webui.env import WEBUI_AUTH_COOKIE_SAME_SITE, WEBUI_AUTH_COOKIE_SECURE +from open_webui.models.auths import ( + MfaChallengeForm, + MfaChallengeResponse, + MfaEnrollmentResponse, + MfaFactorForm, + MfaRecoveryCodesResponse, + MfaRecoveryForm, + MfaSetupResponse, + MfaStatusResponse, + MfaVerifyForm, + SessionUserResponse, +) +from open_webui.models.config import Config +from open_webui.models.oauth_sessions import OAuthSessions +from open_webui.utils import mfa +from open_webui.utils.auth import create_session_response, get_human_user +from open_webui.utils.rate_limit import RateLimiter +from pydantic import ValidationError +from sqlalchemy.exc import SQLAlchemyError + +ip_limiter = RateLimiter(limit=100, window=900) +CHALLENGE_COOKIE = 'mfa_challenge' + + +class MfaRoute(APIRoute): + def get_route_handler(self): + handler = super().get_route_handler() + + async def protected(request: Request): + try: + if request.method == 'POST' and (await mfa.get_mfa_config()).ENABLE_MFA: + ip = request.client.host if request.client else 'unknown' + if await ip_limiter.is_limited( + getattr(request.app.state, 'redis', None), 'mfa:ip:' + mfa.token_hash(ip) + ): + raise HTTPException( + 429, 'Too many attempts. Please try again later.', headers={'Retry-After': '900'} + ) + response = await handler(request) + except RequestValidationError: + response = JSONResponse({'detail': 'Invalid authentication request.'}, status_code=422) + except HTTPException as error: + response = JSONResponse({'detail': error.detail}, status_code=error.status_code, headers=error.headers) + except (SQLAlchemyError, ValidationError): + response = JSONResponse({'detail': 'Authentication is temporarily unavailable.'}, status_code=503) + response.headers['Cache-Control'] = 'no-store' + return response + + return protected + + +router = APIRouter(route_class=MfaRoute) + + +def clear_challenge_cookie(response: Response): + response.delete_cookie(CHALLENGE_COOKIE, path='/api/v1/auths/mfa') + + +async def finish_login(request, response, user, auth, challenge): + result = await create_session_response( + request, + user, + response=response, + set_cookie=True, + source=challenge.auth_method, + auth=auth, + mfa_verified=True, + auth_time=challenge.auth_time, + ) + if challenge.oauth_session_id: + session = await OAuthSessions.get_session_by_id(challenge.oauth_session_id) + if session and session.user_id == user.id: + response.set_cookie( + 'oauth_session_id', + session.id, + httponly=True, + secure=WEBUI_AUTH_COOKIE_SECURE, + samesite=WEBUI_AUTH_COOKIE_SAME_SITE, + ) + clear_challenge_cookie(response) + return result + + +@router.get('/status', response_model=MfaStatusResponse) +async def get_mfa_status(request: Request, user=Depends(get_human_user)): + auth = await mfa.get_auth(user.id) + state = auth.mfa + return { + 'enabled': bool(state and state.secret), + 'required': mfa.is_mfa_required(request.state.claims.get('auth_method', ''), await mfa.get_mfa_config()), + 'recovery_codes_remaining': len(state.recovery_hashes) if state else 0, + } + + +@router.post('/challenge', response_model=MfaChallengeResponse) +async def get_mfa_challenge(request: Request, response: Response): + # Cookie-only bootstrap is browser-only and must not be driven cross-origin. + origin = request.headers.get('origin') + configured = await Config.get('webui.url') + expected = urlsplit(configured or str(request.base_url)) + if origin != f'{expected.scheme}://{expected.netloc}': + raise HTTPException(403, 'Invalid request origin.') + token = request.cookies.get(CHALLENGE_COOKIE, '') + try: + _, _, _, challenge, _ = await mfa.load_challenge(token, {'enroll', 'verify', 'recover'}) + except HTTPException: + clear_challenge_cookie(response) + return JSONResponse( + {'detail': 'This authentication step expired. Please start again.'}, + status_code=401, + headers={'Set-Cookie': response.headers.get('set-cookie', '')}, + ) + clear_challenge_cookie(response) + return mfa.challenge_response(token, challenge) + + +@router.post('/enroll/start', response_model=MfaSetupResponse) +async def start_mfa_enrollment(form_data: MfaChallengeForm): + return await mfa.start_mfa_enrollment(form_data.challenge_token.get_secret_value()) + + +@router.post('/enroll/confirm', response_model=MfaEnrollmentResponse | MfaRecoveryCodesResponse) +async def confirm_mfa_enrollment(request: Request, response: Response, form_data: MfaVerifyForm): + user, auth, challenge, codes = await mfa.confirm_mfa_enrollment( + form_data.challenge_token.get_secret_value(), form_data.code.get_secret_value(), request=request + ) + from open_webui.socket.main import disconnect_user_sessions + + await disconnect_user_sessions(user.id) + if challenge.type == 'replace': + return {'recovery_codes': codes} + return {**(await finish_login(request, response, user, auth, challenge)), 'recovery_codes': codes} + + +@router.post('/verify', response_model=SessionUserResponse) +async def verify_mfa_challenge(request: Request, response: Response, form_data: MfaVerifyForm): + user, auth, challenge = await mfa.verify_mfa_challenge( + form_data.challenge_token.get_secret_value(), + form_data.code.get_secret_value(), + form_data.recovery, + request=request, + ) + return await finish_login(request, response, user, auth, challenge) + + +@router.post('/replace', response_model=MfaChallengeResponse) +async def start_mfa_replacement(request: Request, form_data: MfaFactorForm, user=Depends(get_human_user)): + return await mfa.manage_mfa( + user.id, + request.state.claims, + form_data.code.get_secret_value(), + form_data.recovery, + replace=True, + request=request, + ) + + +@router.post('/recovery/codes', response_model=MfaRecoveryCodesResponse) +async def regenerate_mfa_recovery_codes(request: Request, form_data: MfaFactorForm, user=Depends(get_human_user)): + result = await mfa.manage_mfa( + user.id, + request.state.claims, + form_data.code.get_secret_value(), + form_data.recovery, + replace=False, + request=request, + ) + from open_webui.socket.main import disconnect_user_sessions + + await disconnect_user_sessions(user.id) + return result + + +@router.post('/recover', response_model=MfaChallengeResponse) +async def redeem_mfa_reset_token(form_data: MfaRecoveryForm): + return await mfa.redeem_mfa_reset_token( + form_data.challenge_token.get_secret_value(), form_data.reset_token.get_secret_value() + ) diff --git a/backend/open_webui/routers/users.py b/backend/open_webui/routers/users.py index 1391727ded..10805373bd 100644 --- a/backend/open_webui/routers/users.py +++ b/backend/open_webui/routers/users.py @@ -912,6 +912,21 @@ async def get_user_active_status_by_id( ############################ +@router.post('/{user_id}/sessions/revoke', response_model=bool) +async def revoke_user_sessions(request: Request, user_id: str, session_user=Depends(get_admin_user)): + target = await Users.get_user_by_id(user_id) + if target is None: + raise HTTPException(404, 'User not found.') + first_user = await Users.get_first_user() + if first_user and first_user.id == user_id and session_user.id != user_id: + raise HTTPException(403, detail=ERROR_MESSAGES.ACTION_PROHIBITED) + await revoke_user_tokens(request, user_id) + await publish_event( + request, EVENTS.AUTH_SESSIONS_REVOKED, actor=session_user, subject_id=user_id, subject_type='user' + ) + return True + + @router.post('/{user_id}/update', response_model=UserModel | None) async def update_user_by_id( request: Request, @@ -967,7 +982,9 @@ async def update_user_by_id( hashed = await get_password_hash(form_data.password) if await Auths.update_user_password_by_id(user_id, hashed, db=db): - await revoke_user_tokens(request, user_id) + from open_webui.socket.main import disconnect_user_sessions + + await disconnect_user_sessions(user_id) # Build update dict from only the provided fields update_data = {} diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index 5c1039adfe..da28eeca00 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -336,12 +336,27 @@ def get_user_id_from_session_pool(sid): return None +LOCAL_AUTHENTICATED_SIDS: set[str] = set() + + +async def periodic_socket_authentication(): + while True: + await asyncio.sleep(30) + for sid in tuple(LOCAL_AUTHENTICATED_SIDS): + await get_socket_session_user(sid) + + async def get_socket_session_user(sid: str) -> dict | None: """Session user from this worker's local Socket.IO store; only locally connected sids are ever looked up.""" try: - return (await sio.get_session(sid)).get('user') - except KeyError: - return None + session = await sio.get_session(sid) + if session.get('user') and await get_verified_user_by_token(session.get('token', ''), REDIS): + return session['user'] + except Exception: + log.debug('Socket authentication expired for %s', sid) + LOCAL_AUTHENTICATED_SIDS.discard(sid) + await sio.disconnect(sid) + return None def get_session_ids_from_room(room): @@ -475,7 +490,8 @@ async def connect(sid, environ, auth): 'last_seen_at': int(time.time()), } SESSION_POOL[sid] = socket_user - await sio.save_session(sid, {'user': socket_user}) + await sio.save_session(sid, {'user': socket_user, 'token': auth['token']}) + LOCAL_AUTHENTICATED_SIDS.add(sid) await sio.enter_room(sid, f'user:{user.id}') @@ -507,7 +523,8 @@ async def user_join(sid, data): } SESSION_POOL[sid] = socket_user - await sio.save_session(sid, {'user': socket_user}) + await sio.save_session(sid, {'user': socket_user, 'token': auth['token']}) + LOCAL_AUTHENTICATED_SIDS.add(sid) await sio.enter_room(sid, f'user:{user.id}') # Join all the channels only if user has channels permission @@ -986,6 +1003,7 @@ async def yjs_awareness_update(sid, data): @sio.event async def disconnect(sid, reason=None): + LOCAL_AUTHENTICATED_SIDS.discard(sid) if sid in SESSION_POOL: del SESSION_POOL[sid] diff --git a/backend/open_webui/utils/audit.py b/backend/open_webui/utils/audit.py index 9c680a7f43..aace99dd54 100644 --- a/backend/open_webui/utils/audit.py +++ b/backend/open_webui/utils/audit.py @@ -184,10 +184,13 @@ class AuditLoggingMiddleware: if self._should_skip_auditing(request): return await self.app(scope, receive, send) + capture_body = not (request.url.path.startswith('/api/v1/auths') or request.url.path.startswith('/oauth/')) async with self._audit_context(request) as context: async def send_wrapper(message: ASGISendEvent) -> None: - if self.audit_level == AuditLevel.REQUEST_RESPONSE: + if self.audit_level == AuditLevel.REQUEST_RESPONSE and ( + capture_body or message['type'] == 'http.response.start' + ): await self._capture_response(message, context) await send(message) @@ -198,7 +201,7 @@ class AuditLoggingMiddleware: nonlocal original_receive message = await original_receive() - if self.audit_level in ( + if capture_body and self.audit_level in ( AuditLevel.REQUEST, AuditLevel.REQUEST_RESPONSE, ): @@ -241,6 +244,7 @@ class AuditLoggingMiddleware: '/api/v1/auths/signin', '/api/v1/auths/signout', '/api/v1/auths/signup', + '/api/v1/auths/mfa', ) def _should_skip_auditing(self, request: Request) -> bool: diff --git a/backend/open_webui/utils/auth.py b/backend/open_webui/utils/auth.py index 7d6c52fe26..44d0c8c460 100644 --- a/backend/open_webui/utils/auth.py +++ b/backend/open_webui/utils/auth.py @@ -10,14 +10,11 @@ import uuid from datetime import datetime, timedelta from threading import Lock from time import monotonic -from typing import Optional, Union +from typing import Union import bcrypt import jwt -import pytz import requests -from cryptography.hazmat.primitives import serialization -from cryptography.hazmat.primitives.asymmetric import ed25519 from cryptography.hazmat.primitives.ciphers.aead import AESGCM from fastapi import BackgroundTasks, Depends, HTTPException, Request, Response, status from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer @@ -26,7 +23,6 @@ from open_webui.env import ( ENABLE_OTEL, ENABLE_PASSWORD_VALIDATION, LICENSE_BLOB, - OFFLINE_MODE, PASSWORD_HASH_ALGORITHM, PASSWORD_VALIDATION_HINT, PASSWORD_VALIDATION_REGEX_PATTERN, @@ -274,14 +270,28 @@ revocation_log.addFilter(RateLimitFilter()) 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 password changes and OIDC back-channel - logout when individual jti values are unknown; rejects tokens with iat <= revoked_at. + """Check persistent account revocation, then optional Redis per-token revocation. - Fail open on Redis errors to preserve availability; revoked tokens may be accepted. + Database failures fail closed. Redis failures retain the existing availability policy. """ + from open_webui.utils.mfa import get_auth, get_mfa_config, is_mfa_required + + auth = await get_auth(decoded.get('id', '')) + if decoded.get('session_stamp') != auth.session_stamp: + return False + token_type = decoded.get('typ', 'session') + if token_type not in {'session', 'automation', 'subagent'}: + return False + if token_type == 'session': + config = await get_mfa_config() + if is_mfa_required(decoded.get('auth_method', ''), config): + if ( + decoded.get('mfa_verified') is not True + or not auth.mfa + or not auth.mfa.secret + or auth.mfa.reset_required + ): + return False if not redis: return True @@ -347,27 +357,13 @@ async def invalidate_token(request, token): async def revoke_user_tokens(request, user_id: str): - """Reject every token already issued to a user. Requires Redis.""" - redis = request.app.state.redis - - if not redis: - log.warning( - 'Cannot revoke tokens for user %s: Redis is not configured, existing sessions stay valid until expiry.', - user_id, - ) - return - - # The marker has to outlive every token it revokes, so it never expires when tokens do not - expires_delta = parse_duration(await Config.get('auth.jwt_expiry')) - - await redis.set( - f'{REDIS_KEY_PREFIX}:auth:user:{user_id}:revoked_at', - str(int(datetime.now(UTC).timestamp())), - ex=int(expires_delta.total_seconds()) if expires_delta else None, - ) - + """Persist account-wide revocation, whether or not Redis is configured.""" + from open_webui.internal.db import get_async_db from open_webui.socket.main import disconnect_user_sessions + async with get_async_db() as db: + await Auths.revoke_sessions_by_user_id(user_id, db=db) + await db.commit() await disconnect_user_sessions(user_id) @@ -485,6 +481,7 @@ async def get_current_user( # Scope-backed, so outer middleware (audit) can reuse the resolved user request.state.user = user request.state.auth_type = 'jwt' + request.state.claims = data return user else: raise HTTPException( @@ -663,3 +660,107 @@ async def create_admin_user(email: str, password: str, name: str = 'Admin'): except Exception as e: log.error(f'Error creating admin account: {e}') return None + + +async def create_signin_response( + request, + user, + db=None, + response=None, + set_cookie=False, + source='password', + auth=None, + *, + oauth_session_id=None, + provider=None, +): + """Gate every human login before issuing an application credential.""" + from open_webui.utils.mfa import get_auth, get_mfa_config, is_mfa_required, start_mfa_login + + auth = auth or await get_auth(user.id) + current = await get_auth(user.id) + if current.session_stamp != auth.session_stamp or current.password != auth.password: + raise HTTPException(409, 'Authentication changed. Please sign in again.') + config = await get_mfa_config() + if config.ENABLE_MFA and user.role not in {'admin', 'user'}: + return {'next_step': 'pending'} + if is_mfa_required(source, config): + return await start_mfa_login(current, source, oauth_session_id=oauth_session_id, provider=provider) + return await create_session_response( + request, user, db, response, set_cookie=set_cookie, source=source, auth=current + ) + + +async def create_session_response( + request, + user, + db=None, + response=None, + set_cookie=False, + source='password', + *, + auth, + mfa_verified=False, + auth_time=None, +): + """Issue a completed human session using the credential snapshot that authorized it.""" + import time + + from open_webui.env import WEBUI_AUTH_COOKIE_SAME_SITE, WEBUI_AUTH_COOKIE_SECURE + from open_webui.events import EVENTS, publish_event + from open_webui.utils.access_control import get_permissions + from open_webui.utils.mfa import get_auth, get_mfa_config, is_mfa_required + + current = await get_auth(user.id) + if current.session_stamp != auth.session_stamp: + raise HTTPException(409, 'Authentication changed. Please sign in again.') + if is_mfa_required(source, await get_mfa_config()) and not mfa_verified: + raise HTTPException(401, 'Authenticator verification is required.') + expires_delta = parse_duration(await Config.get('auth.jwt_expiry')) + expires_at = int(time.time()) + int(expires_delta.total_seconds()) if expires_delta else None + token = create_token( + { + 'id': user.id, + 'typ': 'session', + 'auth_method': source, + 'auth_time': auth_time or int(time.time()), + 'session_stamp': auth.session_stamp, + 'mfa_verified': mfa_verified, + }, + expires_delta=expires_delta, + ) + if set_cookie and response is not None: + response.set_cookie( + 'token', + token, + httponly=True, + secure=WEBUI_AUTH_COOKIE_SECURE, + samesite=WEBUI_AUTH_COOKIE_SAME_SITE, + max_age=int(expires_delta.total_seconds()) if expires_delta else None, + ) + await publish_event( + request, + EVENTS.AUTH_LOGIN, + actor=user, + subject_id=user.id, + subject_type='user', + source=source, + data={'auth_method': source}, + ) + return { + 'token': token, + 'token_type': 'Bearer', + 'expires_at': expires_at, + 'id': user.id, + 'email': user.email, + 'name': user.name, + 'role': user.role, + 'profile_image_url': f'/api/v1/users/{user.id}/profile/image', + 'permissions': await get_permissions(user.id, await Config.get('user.permissions'), db=db), + } + + +async def get_human_user(request: Request, user=Depends(get_current_user)): + if request.state.auth_type != 'jwt' or request.state.claims.get('typ', 'session') != 'session': + raise HTTPException(403, 'A human session is required.') + return user diff --git a/backend/open_webui/utils/automations.py b/backend/open_webui/utils/automations.py index d6c0f4dfe4..d27b766f93 100644 --- a/backend/open_webui/utils/automations.py +++ b/backend/open_webui/utils/automations.py @@ -393,8 +393,13 @@ async def execute_automation(app, automation: AutomationModel) -> None: expires_delta = parse_duration(str(await Config.get('automations.auth_token_expires_in', '1h'))) except ValueError: expires_delta = None + from open_webui.models.auths import Auths + + auth = await Auths.get_auth_by_id(user.id) + if auth is None or not auth.active: + raise ValueError('Automation owner is no longer active') token = create_token( - data={'id': user.id, 'typ': 'automation'}, + data={'id': user.id, 'typ': 'automation', 'session_stamp': auth.session_stamp}, expires_delta=expires_delta or timedelta(hours=1), ) diff --git a/backend/open_webui/utils/mfa.py b/backend/open_webui/utils/mfa.py new file mode 100644 index 0000000000..dcd74c38b9 --- /dev/null +++ b/backend/open_webui/utils/mfa.py @@ -0,0 +1,448 @@ +"""Authenticator state and short-lived challenges stored on the credential row.""" + +from __future__ import annotations + +import base64 +import hashlib +import hmac +import io +import os +import secrets +import time +import uuid + +import pyotp +import qrcode +import qrcode.image.svg +from cryptography.fernet import Fernet, InvalidToken +from cryptography.hazmat.primitives import hashes +from cryptography.hazmat.primitives.kdf.hkdf import HKDF +from fastapi import HTTPException +from open_webui.env import WEBUI_AUTH, WEBUI_SECRET_KEY +from open_webui.internal.db import get_async_db +from open_webui.models.auths import Auth, AuthModel, Auths, MfaChallenge, MfaData, MfaLimit, MfaResetTicket +from open_webui.models.config import Config +from open_webui.models.users import Users +from pydantic import BaseModel, StrictBool, ValidationError +from sqlalchemy import select, update +from sqlalchemy.exc import SQLAlchemyError + +CHALLENGE_SECONDS = 300 +LIMIT_SECONDS = 900 +MFA_CONFIG_KEYS = { + 'ENABLE_MFA': 'auth.mfa.enable', + 'MFA_ALLOW_OAUTH_BYPASS': 'auth.mfa.allow_oauth_bypass', + 'MFA_ALLOW_TRUSTED_HEADER_BYPASS': 'auth.mfa.allow_trusted_header_bypass', +} + + +class MfaConfigForm(BaseModel): + ENABLE_MFA: StrictBool = False + MFA_ALLOW_OAUTH_BYPASS: StrictBool = False + MFA_ALLOW_TRUSTED_HEADER_BYPASS: StrictBool = False + + +async def get_mfa_config() -> MfaConfigForm: + values = await Config.get_many(*MFA_CONFIG_KEYS.values()) + return MfaConfigForm( + **{field: values[key] for field, key in MFA_CONFIG_KEYS.items() if values.get(key) is not None} + ) + + +def is_mfa_required(auth_method: str, config: MfaConfigForm) -> bool: + if not config.ENABLE_MFA: + return False + if auth_method == 'oauth' and config.MFA_ALLOW_OAUTH_BYPASS: + return False + return not (auth_method == 'trusted_header' and config.MFA_ALLOW_TRUSTED_HEADER_BYPASS) + + +def cipher() -> Fernet: + override = os.getenv('MFA_ENCRYPTION_KEY') + if override: + return Fernet(override.encode()) + if not WEBUI_SECRET_KEY: + raise ValueError('A persistent WEBUI_SECRET_KEY is required for MFA.') + key = HKDF(algorithm=hashes.SHA256(), length=32, salt=None, info=b'open-webui/mfa-encryption/v1').derive( + WEBUI_SECRET_KEY.encode() + ) + return Fernet(base64.urlsafe_b64encode(key)) + + +def decrypt_secret(secret: str) -> str: + try: + return cipher().decrypt(secret.encode()).decode() + except (InvalidToken, ValueError): + raise HTTPException(503, 'Authenticator configuration is unavailable. Contact the operator.') from None + + +def token_hash(token: str) -> str: + return hashlib.sha256(token.encode()).hexdigest() + + +def matching_step(secret: str, code: str, last_step: int = -1) -> int | None: + if len(code) != 6 or not code.isascii() or not code.isdigit(): + return None + current = int(time.time()) // 30 + totp = pyotp.TOTP(secret) + for step in (current, current - 1, current + 1): + if step > last_step and hmac.compare_digest(totp.at(step * 30), code): + return step + return None + + +def recovery_codes() -> tuple[list[str], list[str]]: + codes = [secrets.token_hex(16) for _ in range(10)] + return codes, [token_hash(code) for code in codes] + + +async def record_mfa_event(event: str, user_id: str | None = None, *, request=None, reason: str | None = None): + # Event delivery is optional; always write a structured security record too. + from loguru import logger + from open_webui.utils.audit import AuditLevel, AuditLogEntry, AuditLogger + + data = {'event': event, 'reason': reason} if reason else {'event': event} + logger.bind(mfa_event=event, user_id=user_id).info('MFA security event: {}', data) + AuditLogger(logger).write( + AuditLogEntry( + id=str(uuid.uuid4()), + user={'id': user_id} if user_id else {}, + audit_level=AuditLevel.METADATA.value, + verb='MFA', + request_uri='/api/v1/auths/mfa', + source_ip=request.client.host if request and request.client else None, + ), + extra=data, + ) + if request is not None: + from open_webui.events import publish_event + + await publish_event(request, event, subject_id=user_id, subject_type='user', data=data) + + +async def get_auth(user_id: str) -> AuthModel: + try: + auth = await Auths.get_auth_by_id(user_id) + except (SQLAlchemyError, ValidationError): + raise HTTPException(503, 'Authentication is temporarily unavailable.') from None + if auth is None or not auth.active: + raise HTTPException(401, 'Invalid authentication request.') + return auth + + +async def save_mfa(auth: AuthModel, mfa: MfaData, *, revoke: bool = False) -> AuthModel: + try: + async with get_async_db() as db: + updated = await Auths.update_mfa_by_id(auth, mfa, revoke=revoke, db=db) + if updated is None: + raise HTTPException(409, 'Authentication changed in another request. Please start again.') + await db.commit() + return updated + except SQLAlchemyError: + raise HTTPException(503, 'Authentication is temporarily unavailable.') from None + + +async def limit_account(user_id: str, kind: str, *, expected_stamp=None, check_stamp: bool = False) -> AuthModel: + for _ in range(5): + auth = await get_auth(user_id) + if check_stamp and auth.session_stamp != expected_stamp: + raise HTTPException(401, 'Session expired. Please sign in again.') + mfa = auth.mfa.model_copy(deep=True) if auth.mfa else MfaData() + now = int(time.time()) + limit = mfa.limits.get(kind) + if limit is None or limit.expires_at <= now: + limit = MfaLimit(count=0, expires_at=now + LIMIT_SECONDS) + if limit.count >= 10: + await record_mfa_event('auth.mfa.throttled', user_id) + raise HTTPException( + 429, 'Too many attempts. Please try again later.', headers={'Retry-After': str(limit.expires_at - now)} + ) + limit.count += 1 + mfa.limits[kind] = limit + try: + return await save_mfa(auth, mfa) + except HTTPException as error: + if error.status_code != 409: + raise + raise HTTPException(409, 'Authentication is busy. Please try again.') + + +def new_challenge(auth: AuthModel, kind: str, auth_method: str, auth_time: int, **context) -> tuple[str, MfaChallenge]: + token = f'{auth.id}.{secrets.token_urlsafe(32)}' + return token, MfaChallenge( + token_hash=token_hash(token), + type=kind, + expires_at=int(time.time()) + CHALLENGE_SECONDS, + auth_method=auth_method, + auth_time=auth_time, + session_stamp=auth.session_stamp, + **context, + ) + + +def challenge_response(token: str, challenge: MfaChallenge) -> dict: + return { + 'next_step': 'enroll' if challenge.type == 'replace' else challenge.type, + 'challenge_token': token, + 'expires_in': max(0, challenge.expires_at - int(time.time())), + } + + +async def start_mfa_login(auth: AuthModel, auth_method: str, *, oauth_session_id=None, provider=None) -> dict: + mfa = auth.mfa.model_copy(deep=True) if auth.mfa else MfaData() + kind = 'recover' if mfa.reset_required else 'verify' if mfa.secret else 'enroll' + token, challenge = new_challenge( + auth, kind, auth_method, int(time.time()), oauth_session_id=oauth_session_id, provider=provider + ) + mfa.login_challenge = challenge + await save_mfa(auth, mfa) + return challenge_response(token, challenge) + + +async def load_challenge(token: str, kinds: set[str], *, attempt: bool = False): + user_id, separator, _ = token.partition('.') + if not separator or len(user_id) > 64: + raise HTTPException(401, 'This authentication step expired. Please start again.') + auth = await get_auth(user_id) + for _ in range(5): + mfa = auth.mfa.model_copy(deep=True) if auth.mfa else MfaData() + slot = 'manage_challenge' if kinds == {'replace'} else 'login_challenge' + if ( + 'replace' in kinds + and mfa.manage_challenge + and hmac.compare_digest(mfa.manage_challenge.token_hash, token_hash(token)) + ): + slot = 'manage_challenge' + challenge = getattr(mfa, slot) + if ( + challenge is None + or challenge.type not in kinds + or challenge.expires_at <= int(time.time()) + or challenge.session_stamp != auth.session_stamp + or not hmac.compare_digest(challenge.token_hash, token_hash(token)) + ): + raise HTTPException(401, 'This authentication step expired. Please start again.') + user = await Users.get_user_by_id(user_id) + if user is None or user.role not in {'admin', 'user'}: + raise HTTPException(403, 'Account is awaiting approval.') + if not (await get_mfa_config()).ENABLE_MFA: + raise HTTPException(403, 'MFA is disabled. Please sign in again.') + if not attempt: + return user, auth, mfa, challenge, slot + if challenge.attempts >= 5: + raise HTTPException( + 429, + 'Too many codes. Please sign in again.', + headers={'Retry-After': str(max(1, challenge.expires_at - int(time.time())))}, + ) + limit = mfa.limits.get('factor') + now = int(time.time()) + if limit is None or limit.expires_at <= now: + limit = MfaLimit(count=0, expires_at=now + LIMIT_SECONDS) + if limit.count >= 10: + raise HTTPException( + 429, 'Too many attempts. Please try again later.', headers={'Retry-After': str(limit.expires_at - now)} + ) + limit.count += 1 + mfa.limits['factor'] = limit + challenge.attempts += 1 + try: + updated = await save_mfa(auth, mfa) + return user, updated, updated.mfa.model_copy(deep=True), challenge, slot + except HTTPException as error: + if error.status_code != 409: + raise + auth = await get_auth(user_id) + raise HTTPException(409, 'Authentication is busy. Please try again.') + + +def setup_details(secret: str, email: str) -> dict: + uri = pyotp.TOTP(secret).provisioning_uri(name=email, issuer_name='Open WebUI') + buffer = io.BytesIO() + qrcode.make(uri, image_factory=qrcode.image.svg.SvgPathImage).save(buffer) + return { + 'manual_key': secret, + 'qr_code': 'data:image/svg+xml;base64,' + base64.b64encode(buffer.getvalue()).decode(), + } + + +async def start_mfa_enrollment(token: str) -> dict: + from starlette.concurrency import run_in_threadpool + + user, auth, mfa, challenge, slot = await load_challenge(token, {'enroll', 'replace'}) + if challenge.secret: + secret = decrypt_secret(challenge.secret) + else: + secret = pyotp.random_base32() + challenge.secret = cipher().encrypt(secret.encode()).decode() + setattr(mfa, slot, challenge) + await save_mfa(auth, mfa) + return await run_in_threadpool(setup_details, secret, user.email) + + +def consume_factor(mfa: MfaData, code: str, recovery: bool) -> None: + if recovery: + digest = token_hash(code.strip().lower()) + if digest in mfa.recovery_hashes: + mfa.recovery_hashes.remove(digest) + return + elif mfa.secret: + step = matching_step(decrypt_secret(mfa.secret), code.strip(), mfa.last_step) + if step is not None: + mfa.last_step = step + return + raise HTTPException(401, 'Invalid or already used code.') + + +async def confirm_mfa_enrollment(token: str, code: str, *, request=None): + user, auth, mfa, challenge, _ = await load_challenge(token, {'enroll', 'replace'}, attempt=True) + step = matching_step(decrypt_secret(challenge.secret), code.strip()) if challenge.secret else None + if step is None: + await record_mfa_event('auth.mfa.failed', user.id, request=request) + raise HTTPException(401, 'Invalid code. Check your authenticator and try again.') + codes, hashes = recovery_codes() + mfa.secret, mfa.last_step, mfa.recovery_hashes = challenge.secret, step, hashes + mfa.login_challenge = mfa.manage_challenge = mfa.reset_ticket = None + mfa.reset_required = False + updated = await save_mfa(auth, mfa, revoke=True) + await record_mfa_event( + 'auth.mfa.replaced' if challenge.type == 'replace' else 'auth.mfa.enrolled', user.id, request=request + ) + return user, updated, challenge, codes + + +async def verify_mfa_challenge(token: str, code: str, recovery: bool, *, request=None): + user, auth, mfa, challenge, _ = await load_challenge(token, {'verify'}, attempt=True) + try: + consume_factor(mfa, code, recovery) + except HTTPException: + await record_mfa_event('auth.mfa.failed', user.id, request=request) + raise + mfa.login_challenge = None + updated = await save_mfa(auth, mfa) + if recovery: + await record_mfa_event('auth.mfa.recovery_used', user.id, request=request) + return user, updated, challenge + + +async def manage_mfa(user_id: str, claims: dict, code: str, recovery: bool, *, replace: bool, request=None): + if int(time.time()) - claims.get('auth_time', 0) > CHALLENGE_SECONDS: + raise HTTPException(403, 'reauthentication_required') + if not (await get_mfa_config()).ENABLE_MFA: + raise HTTPException(403, 'MFA is disabled.') + auth = await limit_account(user_id, 'factor', expected_stamp=claims.get('session_stamp'), check_stamp=True) + mfa = auth.mfa.model_copy(deep=True) + try: + consume_factor(mfa, code, recovery) + except HTTPException: + await record_mfa_event('auth.mfa.failed', user_id, request=request) + raise + if replace: + token, challenge = new_challenge(auth, 'replace', claims['auth_method'], claims['auth_time']) + mfa.manage_challenge = challenge + await save_mfa(auth, mfa) + if recovery: + await record_mfa_event('auth.mfa.recovery_used', user_id, request=request) + return challenge_response(token, challenge) + codes, hashes = recovery_codes() + mfa.recovery_hashes = hashes + mfa.login_challenge = mfa.manage_challenge = mfa.reset_ticket = None + await save_mfa(auth, mfa, revoke=True) + await record_mfa_event('auth.mfa.recovery_codes_regenerated', user_id, request=request) + return {'recovery_codes': codes} + + +async def redeem_mfa_reset_token(token: str, reset_token: str): + user, auth, mfa, challenge, _ = await load_challenge(token, {'recover'}, attempt=True) + ticket = mfa.reset_ticket + if ( + not mfa.reset_required + or ticket is None + or ticket.expires_at <= int(time.time()) + or not hmac.compare_digest(ticket.token_hash, token_hash(reset_token)) + ): + raise HTTPException(401, 'Invalid or expired operator recovery token.') + new_token, enroll = new_challenge( + auth, + 'enroll', + challenge.auth_method, + challenge.auth_time, + oauth_session_id=challenge.oauth_session_id, + provider=challenge.provider, + ) + mfa.reset_ticket = None + mfa.login_challenge = enroll + await save_mfa(auth, mfa) + return challenge_response(new_token, enroll) + + +async def reset_mfa(email: str, reason: str) -> str: + user = await Users.get_user_by_email(email.strip().lower()) + if user is None: + raise ValueError('User not found.') + if not reason.strip(): + raise ValueError('A nonempty reason is required.') + await record_mfa_event('auth.mfa.reset_requested', user.id, reason=reason) + auth = await get_auth(user.id) + mfa = auth.mfa.model_copy(deep=True) if auth.mfa else MfaData() + ticket = secrets.token_urlsafe(32) + mfa.secret = None + mfa.last_step = -1 + mfa.recovery_hashes = [] + mfa.login_challenge = mfa.manage_challenge = None + mfa.reset_required = True + mfa.reset_ticket = MfaResetTicket(token_hash=token_hash(ticket), expires_at=int(time.time()) + 1800) + await save_mfa(auth, mfa, revoke=True) + await record_mfa_event('auth.mfa.reset_completed', user.id, reason=reason) + return ticket + + +async def validate_mfa_configuration(config: MfaConfigForm | None = None): + config = config or await get_mfa_config() + if not config.ENABLE_MFA: + return + if not WEBUI_AUTH or not Config.PERSISTENT_ENABLED: + raise ValueError('MFA requires authentication and persistent configuration.') + cipher() + async with get_async_db() as db: + rows = (await db.execute(select(Auth.mfa).where(Auth.mfa.is_not(None)))).scalars() + for value in rows: + if value is not None: + mfa = MfaData.model_validate(value) + if mfa.secret: + decrypt_secret(mfa.secret) + + +async def update_mfa_config(request, updates: dict) -> bool: + """Apply configuration and any account revocations in the same transaction.""" + relevant = {field: updates[key] for field, key in MFA_CONFIG_KEYS.items() if key in updates} + if not relevant: + await Config.upsert(updates) + return False + if ( + getattr(request.state, 'auth_type', None) != 'jwt' + or getattr(request.state, 'claims', {}).get('typ', 'session') != 'session' + ): + raise HTTPException(403, 'A human administrator session is required to change MFA settings.') + try: + async with get_async_db() as db: + # A no-op write locks the policy row on PostgreSQL and SQLite before reading. + await db.execute(update(Config).where(Config.key == 'auth.mfa.enable').values(value=Config.value)) + rows = (await db.execute(select(Config).where(Config.key.in_(MFA_CONFIG_KEYS.values())))).scalars() + values = {row.key: row.value for row in rows} + current = MfaConfigForm(**{field: values[key] for field, key in MFA_CONFIG_KEYS.items() if key in values}) + desired = MfaConfigForm(**(current.model_dump() | relevant)) + await validate_mfa_configuration(desired) + changed = desired != current + await Config.upsert(updates, db=db) + user_ids = await Auths.revoke_all_sessions(db=db) if changed else [] + await db.commit() + except (ValueError, ValidationError) as error: + raise HTTPException(409, 'Unable to update MFA policy. Check configuration and retry.') from error + if changed: + from open_webui.socket.main import disconnect_user_sessions + + for user_id in user_ids: + await disconnect_user_sessions(user_id) + await record_mfa_event('auth.mfa.policy_changed', getattr(request.state.user, 'id', None), request=request) + return changed diff --git a/backend/open_webui/utils/oauth.py b/backend/open_webui/utils/oauth.py index 1691901018..92cc917de3 100644 --- a/backend/open_webui/utils/oauth.py +++ b/backend/open_webui/utils/oauth.py @@ -85,7 +85,6 @@ from open_webui.models.oauth_sessions import OAuthSessions from open_webui.models.users import Users from open_webui.retrieval.web.utils import get_ssrf_safe_session, validate_url from open_webui.utils.auth import ( - create_token, get_optional_verified_user_from_request, get_password_hash, get_verified_user_by_id, @@ -2182,10 +2181,6 @@ class OAuthManager: detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - jwt_token = create_token( - data={'id': user.id}, - expires_delta=parse_duration(auth_config.JWT_EXPIRES_IN), - ) if auth_config.ENABLE_OAUTH_GROUP_MANAGEMENT: await self.update_user_groups( request=request, @@ -2217,37 +2212,7 @@ class OAuthManager: expires_delta = parse_duration(auth_config.JWT_EXPIRES_IN) cookie_max_age = int(expires_delta.total_seconds()) if expires_delta else None - # Set the cookie token - # Redirect back to the frontend with the JWT token - response.set_cookie( - key='token', - value=jwt_token, - httponly=False, # Required for frontend access - samesite=WEBUI_AUTH_COOKIE_SAME_SITE, - secure=WEBUI_AUTH_COOKIE_SECURE, - **({'max_age': cookie_max_age} if cookie_max_age is not None else {}), - ) - - await publish_event( - request, - EVENTS.AUTH_LOGIN, - actor=user, - subject_id=user.id, - subject_type='user', - source='oauth', - data={'auth_method': 'oauth', 'provider': provider}, - ) - - # Legacy cookies for compatibility with older frontend versions - if ENABLE_OAUTH_ID_TOKEN_COOKIE: - response.set_cookie( - key='oauth_id_token', - value=token.get('id_token'), - httponly=True, - samesite=WEBUI_AUTH_COOKIE_SAME_SITE, - secure=WEBUI_AUTH_COOKIE_SECURE, - **({'max_age': cookie_max_age} if cookie_max_age is not None else {}), - ) + session = None try: _normalize_token_expiry(token) @@ -2272,21 +2237,56 @@ class OAuthManager: db=db, ) + except Exception as e: + log.error(f'Failed to store OAuth session server-side: {e}') + + from open_webui.routers.mfa import CHALLENGE_COOKIE + from open_webui.utils.auth import create_signin_response + + result = await create_signin_response( + request, user, db=db, source='oauth', oauth_session_id=session.id if session else None, provider=provider + ) + response.headers['Cache-Control'] = 'no-store' + if result.get('next_step'): + response.delete_cookie('token') + if result['next_step'] == 'pending': + response.headers['location'] = f'{redirect_url}?pending=1' + else: + response.set_cookie( + CHALLENGE_COOKIE, + result['challenge_token'], + max_age=300, + httponly=True, + path='/api/v1/auths/mfa', + secure=WEBUI_AUTH_COOKIE_SECURE, + samesite=WEBUI_AUTH_COOKIE_SAME_SITE, + ) + response.headers['location'] = f'{redirect_url}?mfa=1' + else: + response.set_cookie( + 'token', + result['token'], + httponly=False, + samesite=WEBUI_AUTH_COOKIE_SAME_SITE, + secure=WEBUI_AUTH_COOKIE_SECURE, + **({'max_age': cookie_max_age} if cookie_max_age is not None else {}), + ) if session: response.set_cookie( - key='oauth_session_id', - value=session.id, + 'oauth_session_id', + session.id, + httponly=True, + samesite=WEBUI_AUTH_COOKIE_SAME_SITE, + secure=WEBUI_AUTH_COOKIE_SECURE, + ) + if ENABLE_OAUTH_ID_TOKEN_COOKIE and token.get('id_token'): + response.set_cookie( + 'oauth_id_token', + token['id_token'], httponly=True, samesite=WEBUI_AUTH_COOKIE_SAME_SITE, secure=WEBUI_AUTH_COOKIE_SECURE, - **({'max_age': cookie_max_age} if cookie_max_age is not None else {}), ) - - log.info('Stored OAuth session server-side for user %s, provider %s', user.id, provider) - else: - log.warning(f'Failed to create OAuth session for user {user.id}, provider {provider}') - except Exception as e: - log.error(f'Failed to store OAuth session server-side: {e}') return response diff --git a/backend/open_webui/utils/subagents.py b/backend/open_webui/utils/subagents.py index 61893a4900..c59a068c1c 100644 --- a/backend/open_webui/utils/subagents.py +++ b/backend/open_webui/utils/subagents.py @@ -13,6 +13,7 @@ from open_webui.internal.db import get_async_db from open_webui.models.chat_messages import ChatMessages from open_webui.models.chats import Chat, ChatForm, Chats from open_webui.models.config import Config +from open_webui.models.auths import Auths from open_webui.models.users import UserModel, Users from open_webui.tasks import create_task, has_active_tasks from open_webui.utils.auth import VERIFIED_USER_ROLES, create_token @@ -45,7 +46,7 @@ _foreground_semaphore: asyncio.Semaphore | None = None _parent_locks: weakref.WeakValueDictionary[str, asyncio.Lock] = weakref.WeakValueDictionary() -def _build_request(source: Request, user_id: str, *, internal: bool) -> Request: +async def _build_request(source: Request, user_id: str, *, internal: bool) -> Request: scope = { 'type': 'http', 'asgi': {'version': '3.0', 'spec_version': '2.0'}, @@ -59,8 +60,11 @@ def _build_request(source: Request, user_id: str, *, internal: bool) -> Request: 'app': source.app, } request = Request(scope) + auth = await Auths.get_auth_by_id(user_id) + if auth is None or not auth.active: + raise ValueError('Subagent owner is no longer active') token = create_token( - data={'id': user_id, 'typ': 'subagent'}, + data={'id': user_id, 'typ': 'subagent', 'session_stamp': auth.session_stamp}, expires_delta=timedelta(hours=1), ) request.state.token = HTTPAuthorizationCredentials(scheme='Bearer', credentials=token) @@ -265,7 +269,7 @@ async def process_pending_internal_messages( if run.get('terminal_id'): form_data['terminal_id'] = run['terminal_id'] - request = _build_request(source_request, user.id, internal=False) + request = await _build_request(source_request, user.id, internal=False) await source_request.app.state.CHAT_COMPLETION_HANDLER(request, form_data, user=user) @@ -451,7 +455,7 @@ async def delegate( async def run_reserved() -> dict: try: - child_request = _build_request(request, user.id, internal=True) + child_request = await _build_request(request, user.id, internal=True) child_request.state.max_tool_call_iterations = max_iterations parent_system_prompt = run.get('system_prompt') or '' subagent_system_prompt = ( diff --git a/pyproject.toml b/pyproject.toml index 776dc418ce..4356fc23a7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -15,6 +15,8 @@ dependencies = [ "python-socketio==5.16.2", "orjson==3.11.9", "cryptography==48.0.0", + "pyotp==2.9.0", + "qrcode==8.2", "bcrypt==5.0.0", "argon2-cffi==25.1.0", "PyJWT[crypto]==2.13.0", diff --git a/src/lib/apis/auths/mfa.ts b/src/lib/apis/auths/mfa.ts new file mode 100644 index 0000000000..6fb460c016 --- /dev/null +++ b/src/lib/apis/auths/mfa.ts @@ -0,0 +1,35 @@ +import { WEBUI_API_BASE_URL } from '$lib/constants'; + +export type MfaChallenge = { + next_step: 'enroll' | 'verify' | 'recover'; + challenge_token: string; + expires_in: number; +}; + +export const mfaRequest = async ( + path: + | 'status' + | 'challenge' + | 'enroll/start' + | 'enroll/confirm' + | 'verify' + | 'replace' + | 'recovery/codes' + | 'recover', + body?: Record, + token?: string +) => { + const response = await fetch(`${WEBUI_API_BASE_URL}/auths/mfa/${path}`, { + method: path === 'status' ? 'GET' : 'POST', + credentials: 'same-origin', + cache: 'no-store', + headers: { + 'Content-Type': 'application/json', + ...(token ? { Authorization: `Bearer ${token}` } : {}) + }, + ...(body ? { body: JSON.stringify(body) } : {}) + }); + const result = await response.json(); + if (!response.ok) throw new Error(result.detail || 'Authentication failed. Please try again.'); + return result; +}; diff --git a/src/lib/apis/users/index.ts b/src/lib/apis/users/index.ts index 0ab07bda1f..d72714862d 100644 --- a/src/lib/apis/users/index.ts +++ b/src/lib/apis/users/index.ts @@ -747,3 +747,13 @@ export const getUserUsage = async ( return res; }; + +export const revokeUserSessions = async (token: string, userId: string) => { + const response = await fetch(`${WEBUI_API_BASE_URL}/users/${userId}/sessions/revoke`, { + method: 'POST', + headers: { Authorization: `Bearer ${token}` } + }); + const result = await response.json(); + if (!response.ok) throw new Error(result.detail || 'Failed to revoke sessions'); + return result; +}; diff --git a/src/lib/components/admin/Settings/General.svelte b/src/lib/components/admin/Settings/General.svelte index 65f91be375..570d834e8c 100644 --- a/src/lib/components/admin/Settings/General.svelte +++ b/src/lib/components/admin/Settings/General.svelte @@ -88,6 +88,11 @@ I18N: cleaned }); if (!res) throw new Error($i18n.t('Failed to update settings')); + if (res.sessions_revoked) { + localStorage.removeItem('token'); + window.location.href = '/auth?state=logout&form=signin'; + return; + } await updateI18n(res.I18N ?? cleaned); await updateBanners(); await config.set(await getBackendConfig()); @@ -270,6 +275,31 @@ {/if} + + + + +

+ {$i18n.t('Changes sign out all devices. Users enroll on their next sign-in.')} +

+ {#if adminConfig.ENABLE_MFA} + + + {/if} +
+ { + try { + await revokeUserSessions(localStorage.token, selectedUser.id); + toast.success($i18n.t('All sessions revoked')); + } catch (error) { + toast.error(error instanceof Error ? error.message : String(error)); + } + }} +/> +
@@ -217,7 +233,14 @@
-
+
+ {/if} + {/if} + + {#if error}{/if} +
+ +
+
+ + {#if challenge.next_step === 'verify'}{/if} +
+ {/if} +
diff --git a/src/lib/components/auth/MfaRecoveryCodes.svelte b/src/lib/components/auth/MfaRecoveryCodes.svelte new file mode 100644 index 0000000000..fc55b3e4af --- /dev/null +++ b/src/lib/components/auth/MfaRecoveryCodes.svelte @@ -0,0 +1,52 @@ + + +
+
+

+ {$i18n.t('Save your recovery codes')} +

+

+ {$i18n.t('Each code works once. Save them somewhere safe; they will not be shown again.')} +

+
+
    + {#each codes as code}
  • {code}
  • {/each} +
+ + +
+ +
+
diff --git a/src/lib/components/chat/Settings/Account.svelte b/src/lib/components/chat/Settings/Account.svelte index ef5da5f5b4..24a9a697cd 100644 --- a/src/lib/components/chat/Settings/Account.svelte +++ b/src/lib/components/chat/Settings/Account.svelte @@ -14,6 +14,7 @@ import { getUserVariables, updateUserVariables } from '$lib/apis/users'; import UpdatePassword from './Account/UpdatePassword.svelte'; + import Mfa from './Account/Mfa.svelte'; import { generateInitialsImage } from '$lib/utils'; import { copyToClipboard } from '$lib/utils'; import Dropdown from '$lib/components/common/Dropdown.svelte'; @@ -384,6 +385,8 @@ {/if} + + {#if canUseApiKeys({ user: $user, config: $config })} diff --git a/src/lib/components/chat/Settings/Account/Mfa.svelte b/src/lib/components/chat/Settings/Account/Mfa.svelte new file mode 100644 index 0000000000..3d0eb846a5 --- /dev/null +++ b/src/lib/components/chat/Settings/Account/Mfa.svelte @@ -0,0 +1,132 @@ + + +
+ {#if codes.length} +
+ {:else if challenge} +
+ { + challenge = null; + }} + /> +
+ {:else if status} +
+ {$i18n.t( + status.enabled ? 'Authenticator configured' : 'Authenticator not configured' + )} + {#if status.required && status.enabled}{/if} +
+

+ {status.enabled + ? $i18n.t('{{count}} recovery codes remaining', { count: status.recovery_codes_remaining }) + : $i18n.t('Managed by your administrator.')} +

+ {#if show && status.required && status.enabled} +
+

+ {$i18n.t('Verify with a fresh code. Changes sign out all devices.')} +

+ + +
+ +
+
+ {/if} + {/if} + {#if error}{/if} + {#if reauthenticate}{/if} +
diff --git a/src/routes/auth/+page.svelte b/src/routes/auth/+page.svelte index f24800ea96..cca615a041 100644 --- a/src/routes/auth/+page.svelte +++ b/src/routes/auth/+page.svelte @@ -1,5 +1,9 @@