diff --git a/backend/open_webui/config.py b/backend/open_webui/config.py index 1ce646c6c8..e0e853e3ab 100644 --- a/backend/open_webui/config.py +++ b/backend/open_webui/config.py @@ -3940,6 +3940,18 @@ AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT = PersistentConfig( os.getenv('AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT', 'audio-24khz-160kbitrate-mono-mp3'), ) +AUDIO_TTS_MISTRAL_API_KEY = PersistentConfig( + 'AUDIO_TTS_MISTRAL_API_KEY', + 'audio.tts.mistral.api_key', + os.getenv('AUDIO_TTS_MISTRAL_API_KEY', ''), +) + +AUDIO_TTS_MISTRAL_API_BASE_URL = PersistentConfig( + 'AUDIO_TTS_MISTRAL_API_BASE_URL', + 'audio.tts.mistral.api_base_url', + os.getenv('AUDIO_TTS_MISTRAL_API_BASE_URL', 'https://api.mistral.ai/v1'), +) + #################################### # LDAP diff --git a/backend/open_webui/env.py b/backend/open_webui/env.py index f0dbf9c114..e0dcaff276 100644 --- a/backend/open_webui/env.py +++ b/backend/open_webui/env.py @@ -544,6 +544,12 @@ OAUTH_MAX_SESSIONS_PER_USER = int(os.environ.get('OAUTH_MAX_SESSIONS_PER_USER', # Allows external apps to exchange OAuth tokens for OpenWebUI tokens ENABLE_OAUTH_TOKEN_EXCHANGE = os.environ.get('ENABLE_OAUTH_TOKEN_EXCHANGE', 'False').lower() == 'true' +# Back-Channel Logout Configuration +# When enabled, exposes POST /oauth/backchannel-logout for IdP-initiated logout +# per OpenID Connect Back-Channel Logout 1.0 spec. +# Requires Redis for JWT revocation. +ENABLE_OAUTH_BACKCHANNEL_LOGOUT = os.environ.get('ENABLE_OAUTH_BACKCHANNEL_LOGOUT', 'False').lower() == 'true' + #################################### # SCIM Configuration #################################### diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 255c9ace5f..03bb651089 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -212,6 +212,8 @@ from open_webui.config import ( AUDIO_TTS_AZURE_SPEECH_REGION, AUDIO_TTS_AZURE_SPEECH_BASE_URL, AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT, + AUDIO_TTS_MISTRAL_API_KEY, + AUDIO_TTS_MISTRAL_API_BASE_URL, PLAYWRIGHT_WS_URL, PLAYWRIGHT_TIMEOUT, FIRECRAWL_API_BASE_URL, @@ -511,6 +513,8 @@ from open_webui.env import ( WEBUI_ADMIN_NAME, ENABLE_EASTER_EGGS, LOG_FORMAT, + # OAuth Back-Channel Logout + ENABLE_OAUTH_BACKCHANNEL_LOGOUT, ) @@ -1282,6 +1286,9 @@ app.state.config.TTS_AZURE_SPEECH_REGION = AUDIO_TTS_AZURE_SPEECH_REGION app.state.config.TTS_AZURE_SPEECH_BASE_URL = AUDIO_TTS_AZURE_SPEECH_BASE_URL app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT = AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT +app.state.config.TTS_MISTRAL_API_KEY = AUDIO_TTS_MISTRAL_API_KEY +app.state.config.TTS_MISTRAL_API_BASE_URL = AUDIO_TTS_MISTRAL_API_BASE_URL + app.state.faster_whisper_model = None app.state.speech_synthesiser = None @@ -2477,6 +2484,21 @@ async def oauth_login_callback( return await oauth_manager.handle_callback(request, provider, response, db=db) +############################ +# OIDC Back-Channel Logout +############################ + + +@app.post('/oauth/backchannel-logout') +async def oauth_backchannel_logout( + request: Request, + db: Session = Depends(get_session), +): + if not ENABLE_OAUTH_BACKCHANNEL_LOGOUT: + raise HTTPException(status_code=404) + return await oauth_manager.handle_backchannel_logout(request, db=db) + + @app.get('/manifest.json') async def get_manifest_json(): if app.state.EXTERNAL_PWA_MANIFEST_URL: diff --git a/backend/open_webui/models/chat_messages.py b/backend/open_webui/models/chat_messages.py index ac75fcf973..b37c04037e 100644 --- a/backend/open_webui/models/chat_messages.py +++ b/backend/open_webui/models/chat_messages.py @@ -5,6 +5,7 @@ from typing import Any, Optional from sqlalchemy.orm import Session from open_webui.internal.db import Base, get_db_context +from open_webui.utils.response import normalize_usage from pydantic import BaseModel, ConfigDict from sqlalchemy import ( @@ -41,6 +42,12 @@ def _normalize_timestamp(timestamp: int) -> float: return timestamp +def get_usage(data: dict) -> Optional[dict]: + """Extract and normalize usage from message data.""" + usage = data.get('usage') or (data.get('info') or {}).get('usage') + return normalize_usage(usage) if usage else None + + #################### # ChatMessage DB Schema #################### @@ -163,11 +170,8 @@ class ChatMessageTable: existing.status_history = data.get('status_history') or data.get('statusHistory') if 'error' in data: existing.error = data.get('error') - # Extract usage - check direct field first, then info.usage - usage = data.get('usage') - if not usage: - info = data.get('info', {}) - usage = info.get('usage') if info else None + # Extract and normalize usage + usage = get_usage(data) if usage: # Deep-merge: preserve existing keys not present in new data # This prevents background tasks (follow-ups, title, tags) @@ -179,11 +183,8 @@ class ChatMessageTable: return ChatMessageModel.model_validate(existing) else: # Insert new - # Extract usage - check direct field first, then info.usage - usage = data.get('usage') - if not usage: - info = data.get('info', {}) - usage = info.get('usage') if info else None + # Extract and normalize usage + usage = get_usage(data) message = ChatMessage( id=composite_id, chat_id=chat_id, diff --git a/backend/open_webui/models/chats.py b/backend/open_webui/models/chats.py index 7183145c87..5f53e741e4 100644 --- a/backend/open_webui/models/chats.py +++ b/backend/open_webui/models/chats.py @@ -9,6 +9,7 @@ from open_webui.internal.db import Base, JSONField, get_db, get_db_context from open_webui.models.tags import TagModel, Tag, Tags from open_webui.models.folders import Folders from open_webui.models.chat_messages import ChatMessage, ChatMessages +from open_webui.models.automations import AutomationRun from open_webui.utils.misc import sanitize_data_for_db, sanitize_text_for_db from pydantic import BaseModel, ConfigDict @@ -826,9 +827,7 @@ class ChatTable: else: query = query.order_by(Chat.updated_at.desc(), Chat.id) - query = query.with_entities( - Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at - ) + query = query.with_entities(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at) if skip: query = query.offset(skip) @@ -1262,9 +1261,7 @@ class ChatTable: query = query.order_by(Chat.updated_at.desc(), Chat.id) - query = query.with_entities( - Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at - ) + query = query.with_entities(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at) if skip: query = query.offset(skip) @@ -1347,9 +1344,7 @@ class ChatTable: query = query.order_by(Chat.updated_at.desc(), Chat.id) - query = query.with_entities( - Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at - ) + query = query.with_entities(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at) if skip: query = query.offset(skip) @@ -1478,6 +1473,9 @@ class ChatTable: def delete_chat_by_id(self, id: str, db: Optional[Session] = None) -> bool: try: with get_db_context(db) as db: + db.query(AutomationRun).filter_by(chat_id=id).update( + {AutomationRun.chat_id: None}, synchronize_session=False + ) db.query(ChatMessage).filter_by(chat_id=id).delete() db.query(Chat).filter_by(id=id).delete() db.commit() @@ -1489,6 +1487,9 @@ class ChatTable: def delete_chat_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool: try: with get_db_context(db) as db: + db.query(AutomationRun).filter_by(chat_id=id).update( + {AutomationRun.chat_id: None}, synchronize_session=False + ) db.query(ChatMessage).filter_by(chat_id=id).delete() db.query(Chat).filter_by(id=id, user_id=user_id).delete() db.commit() @@ -1503,6 +1504,9 @@ class ChatTable: self.delete_shared_chats_by_user_id(user_id, db=db) chat_id_subquery = db.query(Chat.id).filter_by(user_id=user_id).subquery() + db.query(AutomationRun).filter(AutomationRun.chat_id.in_(chat_id_subquery)).update( + {AutomationRun.chat_id: None}, synchronize_session=False + ) db.query(ChatMessage).filter(ChatMessage.chat_id.in_(chat_id_subquery)).delete( synchronize_session=False ) @@ -1517,6 +1521,9 @@ class ChatTable: try: with get_db_context(db) as db: chat_id_subquery = db.query(Chat.id).filter_by(user_id=user_id, folder_id=folder_id).subquery() + db.query(AutomationRun).filter(AutomationRun.chat_id.in_(chat_id_subquery)).update( + {AutomationRun.chat_id: None}, synchronize_session=False + ) db.query(ChatMessage).filter(ChatMessage.chat_id.in_(chat_id_subquery)).delete( synchronize_session=False ) diff --git a/backend/open_webui/models/feedbacks.py b/backend/open_webui/models/feedbacks.py index f930739f60..9172e2ba8e 100644 --- a/backend/open_webui/models/feedbacks.py +++ b/backend/open_webui/models/feedbacks.py @@ -218,9 +218,7 @@ class FeedbackTable: # Apply model_id filter (exact match) model_id = filter.get('model_id') if model_id: - query = query.filter( - Feedback.data['model_id'].as_string() == model_id - ) + query = query.filter(Feedback.data['model_id'].as_string() == model_id) order_by = filter.get('order_by') direction = filter.get('direction') diff --git a/backend/open_webui/models/models.py b/backend/open_webui/models/models.py index 9d2b5819bc..7cab2c830e 100755 --- a/backend/open_webui/models/models.py +++ b/backend/open_webui/models/models.py @@ -414,6 +414,7 @@ class ModelsTable: with get_db_context(db) as db: # update only the fields that are present in the model data = model.model_dump(exclude={'id', 'access_grants'}) + data['updated_at'] = int(time.time()) result = db.query(Model).filter_by(id=id).update(data) db.commit() @@ -425,6 +426,20 @@ class ModelsTable: log.exception(f'Failed to update the model by id {id}: {e}') return None + def update_model_updated_at_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ModelModel]: + try: + with get_db_context(db) as db: + result = db.query(Model).filter_by(id=id).first() + if not result: + return None + result.updated_at = int(time.time()) + db.commit() + db.refresh(result) + return self._to_model_model(result, db=db) + except Exception as e: + log.exception(f'Failed to update the model updated_at by id {id}: {e}') + return None + def delete_model_by_id(self, id: str, db: Optional[Session] = None) -> bool: try: with get_db_context(db) as db: diff --git a/backend/open_webui/routers/audio.py b/backend/open_webui/routers/audio.py index 8e14387a78..9d8938b419 100644 --- a/backend/open_webui/routers/audio.py +++ b/backend/open_webui/routers/audio.py @@ -168,6 +168,8 @@ class TTSConfigForm(BaseModel): AZURE_SPEECH_REGION: str AZURE_SPEECH_BASE_URL: str AZURE_SPEECH_OUTPUT_FORMAT: str + MISTRAL_API_KEY: str + MISTRAL_API_BASE_URL: str class STTConfigForm(BaseModel): @@ -208,6 +210,8 @@ async def get_audio_config(request: Request, user=Depends(get_admin_user)): 'AZURE_SPEECH_REGION': request.app.state.config.TTS_AZURE_SPEECH_REGION, 'AZURE_SPEECH_BASE_URL': request.app.state.config.TTS_AZURE_SPEECH_BASE_URL, 'AZURE_SPEECH_OUTPUT_FORMAT': request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT, + 'MISTRAL_API_KEY': request.app.state.config.TTS_MISTRAL_API_KEY, + 'MISTRAL_API_BASE_URL': request.app.state.config.TTS_MISTRAL_API_BASE_URL, }, 'stt': { 'OPENAI_API_BASE_URL': request.app.state.config.STT_OPENAI_API_BASE_URL, @@ -242,6 +246,8 @@ async def update_audio_config(request: Request, form_data: AudioConfigUpdateForm request.app.state.config.TTS_AZURE_SPEECH_REGION = form_data.tts.AZURE_SPEECH_REGION request.app.state.config.TTS_AZURE_SPEECH_BASE_URL = form_data.tts.AZURE_SPEECH_BASE_URL request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT = form_data.tts.AZURE_SPEECH_OUTPUT_FORMAT + request.app.state.config.TTS_MISTRAL_API_KEY = form_data.tts.MISTRAL_API_KEY + request.app.state.config.TTS_MISTRAL_API_BASE_URL = form_data.tts.MISTRAL_API_BASE_URL request.app.state.config.STT_OPENAI_API_BASE_URL = form_data.stt.OPENAI_API_BASE_URL request.app.state.config.STT_OPENAI_API_KEY = form_data.stt.OPENAI_API_KEY @@ -280,6 +286,8 @@ async def update_audio_config(request: Request, form_data: AudioConfigUpdateForm 'AZURE_SPEECH_REGION': request.app.state.config.TTS_AZURE_SPEECH_REGION, 'AZURE_SPEECH_BASE_URL': request.app.state.config.TTS_AZURE_SPEECH_BASE_URL, 'AZURE_SPEECH_OUTPUT_FORMAT': request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT, + 'MISTRAL_API_KEY': request.app.state.config.TTS_MISTRAL_API_KEY, + 'MISTRAL_API_BASE_URL': request.app.state.config.TTS_MISTRAL_API_BASE_URL, }, 'stt': { 'OPENAI_API_BASE_URL': request.app.state.config.STT_OPENAI_API_BASE_URL, @@ -551,6 +559,76 @@ async def speech(request: Request, user=Depends(get_verified_user)): return FileResponse(file_path) + elif request.app.state.config.TTS_ENGINE == 'mistral': + api_key = request.app.state.config.TTS_MISTRAL_API_KEY + api_base_url = request.app.state.config.TTS_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1' + + if not api_key: + raise HTTPException( + status_code=400, + detail='Mistral API key is required for Mistral TTS', + ) + + try: + timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) + async with aiohttp.ClientSession(timeout=timeout, trust_env=True) as session: + mistral_payload = { + 'input': payload.get('input', ''), + 'model': request.app.state.config.TTS_MODEL or 'mistral-tts-latest', + 'voice_id': payload.get('voice', ''), + 'response_format': 'mp3', + } + + r = await session.post( + url=f'{api_base_url}/audio/speech', + json=mistral_payload, + headers={ + 'Content-Type': 'application/json', + 'Authorization': f'Bearer {api_key}', + }, + ssl=AIOHTTP_CLIENT_SESSION_SSL, + ) + + r.raise_for_status() + + res = await r.json() + audio_data = res.get('audio_data', '') + if not audio_data: + raise ValueError('No audio_data in Mistral TTS response') + + audio_bytes = base64.b64decode(audio_data) + + async with aiofiles.open(file_path, 'wb') as f: + await f.write(audio_bytes) + + async with aiofiles.open(file_body_path, 'w') as f: + await f.write(json.dumps(payload)) + + return FileResponse(file_path) + + except Exception as e: + log.exception(e) + detail = None + + status_code = 500 + detail = 'Open WebUI: Server Connection Error' + + if r is not None: + status_code = r.status + + try: + res = await r.json() + if 'error' in res: + detail = f'External: {res["error"]}' + elif 'message' in res: + detail = f'External: {res["message"]}' + except Exception: + detail = f'External: {e}' + + raise HTTPException( + status_code=status_code, + detail=detail, + ) def transcription_handler(request, file_path, metadata, user=None): filename = os.path.basename(file_path) @@ -1238,6 +1316,8 @@ def get_available_models(request: Request) -> list[dict]: available_models = [{'name': model['name'], 'id': model['model_id']} for model in models] except requests.RequestException as e: log.error(f'Error fetching voices: {str(e)}') + elif request.app.state.config.TTS_ENGINE == 'mistral': + available_models = [{'id': 'mistral-tts-latest'}] return available_models @@ -1301,6 +1381,29 @@ def get_available_voices(request) -> dict: available_voices[voice['ShortName']] = f'{voice["DisplayName"]} ({voice["ShortName"]})' except requests.RequestException as e: log.error(f'Error fetching voices: {str(e)}') + elif request.app.state.config.TTS_ENGINE == 'mistral': + api_key = request.app.state.config.TTS_MISTRAL_API_KEY + api_base_url = request.app.state.config.TTS_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1' + + if api_key: + try: + response = requests.get( + f'{api_base_url}/audio/voices', + headers={ + 'Authorization': f'Bearer {api_key}', + }, + timeout=AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST, + ) + response.raise_for_status() + voices_data = response.json() + + for voice in voices_data: + voice_id = voice.get('voice_id', voice.get('id', '')) + voice_name = voice.get('name', voice_id) + if voice_id: + available_voices[voice_id] = voice_name + except requests.RequestException as e: + log.error(f'Error fetching Mistral voices: {str(e)}') return available_voices diff --git a/backend/open_webui/routers/evaluations.py b/backend/open_webui/routers/evaluations.py index a743f8f9f9..9805f2ece2 100644 --- a/backend/open_webui/routers/evaluations.py +++ b/backend/open_webui/routers/evaluations.py @@ -321,10 +321,7 @@ async def export_all_feedbacks( ): feedbacks = Feedbacks.get_all_feedbacks(db=db) if model_id: - feedbacks = [ - f for f in feedbacks - if f.data and f.data.get('model_id') == model_id - ] + feedbacks = [f for f in feedbacks if f.data and f.data.get('model_id') == model_id] return feedbacks diff --git a/backend/open_webui/routers/files.py b/backend/open_webui/routers/files.py index 6027545190..48172e744e 100644 --- a/backend/open_webui/routers/files.py +++ b/backend/open_webui/routers/files.py @@ -207,8 +207,8 @@ def upload_file_handler( filename = os.path.basename(unsanitized_filename) file_extension = os.path.splitext(filename)[1] - # Remove the leading dot from the file extension - file_extension = file_extension[1:] if file_extension else '' + # Remove the leading dot from the file extension and lowercase it + file_extension = file_extension[1:].lower() if file_extension else '' if process and request.app.state.config.ALLOWED_FILE_EXTENSIONS: request.app.state.config.ALLOWED_FILE_EXTENSIONS = [ diff --git a/backend/open_webui/routers/models.py b/backend/open_webui/routers/models.py index 76ed48d970..6f7b3d48df 100644 --- a/backend/open_webui/routers/models.py +++ b/backend/open_webui/routers/models.py @@ -578,6 +578,8 @@ async def update_model_access_by_id( AccessGrants.set_access_grants('model', form_data.id, form_data.access_grants, db=db) + Models.update_model_updated_at_by_id(form_data.id, db=db) + return Models.get_model_by_id(form_data.id, db=db) diff --git a/backend/open_webui/utils/auth.py b/backend/open_webui/utils/auth.py index 34412d6041..16bc36500b 100644 --- a/backend/open_webui/utils/auth.py +++ b/backend/open_webui/utils/auth.py @@ -206,7 +206,7 @@ def create_token(data: dict, expires_delta: Union[timedelta, None] = None) -> st payload.update({'exp': expire}) jti = str(uuid.uuid4()) - payload.update({'jti': jti}) + payload.update({'jti': jti, 'iat': datetime.now(UTC)}) encoded_jwt = jwt.encode(payload, SESSION_SECRET, algorithm=ALGORITHM) return encoded_jwt @@ -221,15 +221,36 @@ def decode_token(token: str) -> Optional[dict]: async def is_valid_token(request, decoded) -> bool: - # Require Redis to check revoked tokens + """ + Check whether a JWT has been revoked. Two mechanisms: + 1. Per-token (jti) — used by user-initiated sign-out (known jti). + 2. Per-user (revoked_at) — used by OIDC back-channel logout when + individual jti values are unknown; rejects tokens with iat <= revoked_at. + """ if request.app.state.redis: + # Per-token revocation jti = decoded.get('jti') - if jti: revoked = await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:auth:token:{jti}:revoked') if revoked: return False + # Per-user revocation (OIDC back-channel logout) + user_id = decoded.get('id') + if user_id: + revoked_at = await request.app.state.redis.get( + f'{REDIS_KEY_PREFIX}:auth:user:{user_id}:revoked_at' + ) + if revoked_at: + try: + revoked_at_ts = int(revoked_at) + token_iat = decoded.get('iat') + # No iat means legacy token — reject since we can't verify issue time + if token_iat is None or token_iat <= revoked_at_ts: + return False + except (ValueError, TypeError): + pass + return True diff --git a/backend/open_webui/utils/oauth.py b/backend/open_webui/utils/oauth.py index 0cbe753b77..bf6be461e6 100644 --- a/backend/open_webui/utils/oauth.py +++ b/backend/open_webui/utils/oauth.py @@ -75,6 +75,7 @@ from open_webui.env import ( ENABLE_OAUTH_EMAIL_FALLBACK, OAUTH_CLIENT_INFO_ENCRYPTION_KEY, OAUTH_MAX_SESSIONS_PER_USER, + REDIS_KEY_PREFIX, ) from open_webui.utils.misc import parse_duration from open_webui.utils.auth import get_password_hash, create_token @@ -1357,9 +1358,8 @@ class OAuthManager: client = self.get_client(provider) if client is None: raise HTTPException(404) - redirect_uri = ( - (client.server_metadata or {}).get('redirect_uri') - or request.url_for('oauth_login_callback', provider=provider) + redirect_uri = (client.server_metadata or {}).get('redirect_uri') or request.url_for( + 'oauth_login_callback', provider=provider ) kwargs = {} @@ -1694,3 +1694,181 @@ class OAuthManager: log.error(f'Failed to store OAuth session server-side: {e}') return response + + async def handle_backchannel_logout(self, request, db=None): + """ + Handle an OIDC Back-Channel Logout request. + Validates the logout_token, identifies the user, revokes their + sessions via Redis, and deletes their OAuth sessions. + Returns a JSONResponse per the OIDC Back-Channel Logout 1.0 spec. + """ + import jwt as pyjwt + from fastapi.responses import JSONResponse + + # 1. Extract logout_token from form body + try: + form = await request.form() + logout_token = form.get('logout_token') + except Exception: + logout_token = None + + if not logout_token: + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'Missing logout_token parameter'}, + ) + + # 2. Peek at unverified issuer to match against configured providers + try: + unverified_claims = pyjwt.decode(logout_token, options={'verify_signature': False}) + token_issuer = unverified_claims.get('iss') + except Exception as e: + log.warning(f'Back-channel logout: cannot decode logout_token: {e}') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'Malformed logout_token'}, + ) + + if not token_issuer: + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'logout_token missing iss claim'}, + ) + + # 3. Find the configured provider whose issuer matches the token + matched_provider = None + matched_client_id = None + matched_jwks_uri = None + matched_issuer = None + + for provider_name in OAUTH_PROVIDERS: + server_metadata_url = self.get_server_metadata_url(provider_name) + if not server_metadata_url: + continue + + try: + async with aiohttp.ClientSession(trust_env=True) as session: + async with session.get(server_metadata_url, ssl=AIOHTTP_CLIENT_SESSION_SSL) as r: + if r.status != 200: + continue + oidc_config = await r.json() + + provider_issuer = oidc_config.get('issuer') + if provider_issuer and provider_issuer == token_issuer: + client = self.get_client(provider_name) + matched_provider = provider_name + matched_client_id = client.client_id if client else None + matched_jwks_uri = oidc_config.get('jwks_uri') + matched_issuer = provider_issuer + break + except Exception as e: + log.debug(f'Back-channel logout: error checking provider {provider_name}: {e}') + continue + + if not matched_provider or not matched_client_id or not matched_jwks_uri: + log.warning(f'Back-channel logout: no configured provider matches issuer {token_issuer}') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'No configured provider matches token issuer'}, + ) + + # 4. Validate the logout_token signature and claims + try: + jwks_client = pyjwt.PyJWKClient(matched_jwks_uri) + signing_key = jwks_client.get_signing_key_from_jwt(logout_token) + + claims = pyjwt.decode( + logout_token, + signing_key.key, + algorithms=['RS256', 'RS384', 'RS512', 'ES256', 'ES384', 'ES512'], + audience=matched_client_id, + issuer=matched_issuer, + options={ + 'require': ['iss', 'aud', 'iat', 'events'], + }, + ) + except pyjwt.InvalidTokenError as e: + log.warning(f'Back-channel logout: invalid logout_token: {e}') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': f'Invalid logout_token: {e}'}, + ) + except Exception as e: + log.error(f'Back-channel logout: error validating logout_token: {e}') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'Failed to validate logout_token'}, + ) + + # 5. Validate events claim per spec + events = claims.get('events', {}) + if 'http://schemas.openid.net/event/backchannel-logout' not in events: + log.warning('Back-channel logout: missing required backchannel-logout event claim') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'Missing backchannel-logout event claim'}, + ) + + # 6. Per spec, back-channel logout tokens MUST NOT contain a nonce + if 'nonce' in claims: + log.warning('Back-channel logout: logout_token contains nonce (rejected per spec)') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'logout_token must not contain nonce'}, + ) + + # 7. Extract sub and/or sid — at least one must be present + sub = claims.get('sub') + sid = claims.get('sid') + + if not sub and not sid: + log.warning('Back-channel logout: logout_token contains neither sub nor sid') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'logout_token must contain sub or sid'}, + ) + + # 8. Identify users to log out + users_to_logout = [] + if sub: + user = Users.get_user_by_oauth_sub(matched_provider, sub, db=db) + if user: + users_to_logout.append(user) + + if not users_to_logout and sid: + log.info(f'Back-channel logout: no user found by sub, sid-based lookup not yet supported (sid={sid})') + + if not users_to_logout: + log.info(f'Back-channel logout: no matching user for provider={matched_provider}, sub={sub}, sid={sid}') + return JSONResponse(status_code=200, content={}) + + # 9. Revoke tokens and delete sessions + redis = request.app.state.redis + if not redis: + log.warning( + 'Back-channel logout: Redis not configured, cannot revoke JWT tokens. ' + 'OAuth sessions will be deleted but existing JWTs will remain valid until expiry.' + ) + + revoked_count = 0 + for user in users_to_logout: + sessions = OAuthSessions.get_sessions_by_user_id(user.id, db=db) + for oauth_session in sessions: + OAuthSessions.delete_session_by_id(oauth_session.id, db=db) + + if redis: + revocation_key = f'{REDIS_KEY_PREFIX}:auth:user:{user.id}:revoked_at' + await redis.set( + revocation_key, + str(int(time.time())), + ex=60 * 60 * 24 * 30, + ) + revoked_count += 1 + + log.info( + f'Back-channel logout: revoked sessions for user {user.id} ' + f'(email={user.email}, provider={matched_provider}, sessions_deleted={len(sessions)})' + ) + + log.info(f'Back-channel logout: completed for {len(users_to_logout)} user(s), {revoked_count} revocation(s) set') + return JSONResponse(status_code=200, content={}) diff --git a/src/lib/apis/auths/index.ts b/src/lib/apis/auths/index.ts index f8e953f7ca..c501a36ed7 100644 --- a/src/lib/apis/auths/index.ts +++ b/src/lib/apis/auths/index.ts @@ -414,7 +414,7 @@ export const updateUserProfile = async (token: string, profile: object) => { console.error(err); error = err.detail; if (Array.isArray(error)) { - error = error.map((e: { msg?: string }) => e.msg).join("; "); + error = error.map((e: { msg?: string }) => e.msg).join('; '); } return null; }); diff --git a/src/lib/apis/evaluations/index.ts b/src/lib/apis/evaluations/index.ts index 9253295116..bfb6955bcf 100644 --- a/src/lib/apis/evaluations/index.ts +++ b/src/lib/apis/evaluations/index.ts @@ -189,7 +189,13 @@ export const getFeedbackModelIds = async (token: string = '') => { return res; }; -export const getFeedbackItems = async (token: string = '', orderBy, direction, page, modelId: string = '') => { +export const getFeedbackItems = async ( + token: string = '', + orderBy, + direction, + page, + modelId: string = '' +) => { let error = null; const searchParams = new URLSearchParams(); @@ -235,14 +241,17 @@ export const exportAllFeedbacks = async (token: string = '', modelId: string = ' const searchParams = new URLSearchParams(); if (modelId) searchParams.append('model_id', modelId); - const res = await fetch(`${WEBUI_API_BASE_URL}/evaluations/feedbacks/all/export?${searchParams.toString()}`, { - method: 'GET', - headers: { - Accept: 'application/json', - 'Content-Type': 'application/json', - authorization: `Bearer ${token}` + const res = await fetch( + `${WEBUI_API_BASE_URL}/evaluations/feedbacks/all/export?${searchParams.toString()}`, + { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } } - }) + ) .then(async (res) => { if (!res.ok) throw await res.json(); return res.json(); diff --git a/src/lib/components/admin/Evaluations/Feedbacks.svelte b/src/lib/components/admin/Evaluations/Feedbacks.svelte index 6e11291d5f..da7c8b5798 100644 --- a/src/lib/components/admin/Evaluations/Feedbacks.svelte +++ b/src/lib/components/admin/Evaluations/Feedbacks.svelte @@ -176,10 +176,9 @@ : s; }; - return [ - headers.join(','), - ...rows.map((r) => headers.map((h) => escape(r[h])).join(',')) - ].join('\n'); + return [headers.join(','), ...rows.map((r) => headers.map((h) => escape(r[h])).join(','))].join( + '\n' + ); }; const exportHandler = async (format: 'json' | 'csv' = 'json') => { diff --git a/src/lib/components/admin/Settings/Audio.svelte b/src/lib/components/admin/Settings/Audio.svelte index 064cd00c67..427f2e6965 100644 --- a/src/lib/components/admin/Settings/Audio.svelte +++ b/src/lib/components/admin/Settings/Audio.svelte @@ -37,6 +37,8 @@ let TTS_AZURE_SPEECH_REGION = ''; let TTS_AZURE_SPEECH_BASE_URL = ''; let TTS_AZURE_SPEECH_OUTPUT_FORMAT = ''; + let TTS_MISTRAL_API_KEY = ''; + let TTS_MISTRAL_API_BASE_URL = ''; let STT_OPENAI_API_BASE_URL = ''; let STT_OPENAI_API_KEY = ''; @@ -124,6 +126,8 @@ AZURE_SPEECH_REGION: TTS_AZURE_SPEECH_REGION, AZURE_SPEECH_BASE_URL: TTS_AZURE_SPEECH_BASE_URL, AZURE_SPEECH_OUTPUT_FORMAT: TTS_AZURE_SPEECH_OUTPUT_FORMAT, + MISTRAL_API_KEY: TTS_MISTRAL_API_KEY, + MISTRAL_API_BASE_URL: TTS_MISTRAL_API_BASE_URL, SPLIT_ON: TTS_SPLIT_ON }, stt: { @@ -176,6 +180,8 @@ TTS_AZURE_SPEECH_REGION = res.tts.AZURE_SPEECH_REGION; TTS_AZURE_SPEECH_BASE_URL = res.tts.AZURE_SPEECH_BASE_URL; TTS_AZURE_SPEECH_OUTPUT_FORMAT = res.tts.AZURE_SPEECH_OUTPUT_FORMAT; + TTS_MISTRAL_API_KEY = res.tts.MISTRAL_API_KEY; + TTS_MISTRAL_API_BASE_URL = res.tts.MISTRAL_API_BASE_URL; STT_OPENAI_API_BASE_URL = res.stt.OPENAI_API_BASE_URL; STT_OPENAI_API_KEY = res.stt.OPENAI_API_KEY; @@ -517,6 +523,9 @@ if (e.target?.value === 'openai') { TTS_VOICE = 'alloy'; TTS_MODEL = 'tts-1'; + } else if (e.target?.value === 'mistral') { + TTS_VOICE = ''; + TTS_MODEL = 'mistral-tts-latest'; } else { TTS_VOICE = ''; TTS_MODEL = ''; @@ -528,6 +537,7 @@ + @@ -585,6 +595,19 @@ + {:else if TTS_ENGINE === 'mistral'} +
+
+ + + +
+
{/if}
@@ -791,6 +814,47 @@
+ {:else if TTS_ENGINE === 'mistral'} +
+
+
{$i18n.t('TTS Voice')}
+
+
+ + + + {#each voices as voice} + + {/each} + +
+
+
+
+
{$i18n.t('TTS Model')}
+
+
+ + + + {#each models as model} + +
+
+
+
{/if} diff --git a/src/lib/components/automations/AutomationEditor.svelte b/src/lib/components/automations/AutomationEditor.svelte index 0eebdb0fcf..fbdf448179 100644 --- a/src/lib/components/automations/AutomationEditor.svelte +++ b/src/lib/components/automations/AutomationEditor.svelte @@ -276,7 +276,7 @@ {#if isDirty}