Merge branch 'dev' into fix/analytics-message-count-role-filter

This commit is contained in:
Classic298 2026-04-14 15:45:20 +02:00 committed by GitHub
commit 51d7cc8f18
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
99 changed files with 2916 additions and 1612 deletions

View file

@ -1,3 +1,4 @@
import asyncio
import json
import logging
import os
@ -35,7 +36,7 @@ from open_webui.env import (
WEBUI_NAME,
log,
)
from open_webui.internal.db import Base, get_db
from open_webui.internal.db import Base, get_db, get_async_db
from open_webui.utils.redis import get_redis_connection
@ -90,6 +91,7 @@ def load_json_config():
def save_to_db(data):
"""Sync save — used ONLY at startup/import time."""
with get_db() as db:
existing_config = db.query(Config).first()
if not existing_config:
@ -102,12 +104,39 @@ def save_to_db(data):
db.commit()
async def async_save_to_db(data):
"""Async save — used for ALL runtime config persistence."""
from sqlalchemy import select
async with get_async_db() as db:
result = await db.execute(select(Config).limit(1))
existing_config = result.scalars().first()
if not existing_config:
new_config = Config(data=data, version=0)
db.add(new_config)
else:
existing_config.data = data
existing_config.updated_at = datetime.now()
db.add(existing_config)
await db.commit()
def reset_config():
"""Sync reset — used ONLY at startup."""
with get_db() as db:
db.query(Config).delete()
db.commit()
async def async_reset_config():
"""Async reset — used at runtime."""
from sqlalchemy import delete as sa_delete
async with get_async_db() as db:
await db.execute(sa_delete(Config))
await db.commit()
# When initializing, check if config.json exists and migrate it to the database
if os.path.exists(f'{DATA_DIR}/config.json'):
data = load_json_config()
@ -144,6 +173,7 @@ PERSISTENT_CONFIG_REGISTRY = []
def save_config(config):
"""Sync save — used ONLY at startup/import time."""
global CONFIG_DATA
global PERSISTENT_CONFIG_REGISTRY
try:
@ -159,6 +189,23 @@ def save_config(config):
return True
async def async_save_config(config):
"""Async save — used for ALL runtime config persistence."""
global CONFIG_DATA
global PERSISTENT_CONFIG_REGISTRY
try:
await async_save_to_db(config)
CONFIG_DATA = config
# Trigger updates on all registered PersistentConfig entries
for config_item in PERSISTENT_CONFIG_REGISTRY:
config_item.update()
except Exception as e:
log.exception(e)
return False
return True
T = TypeVar('T')
ENABLE_PERSISTENT_CONFIG = os.environ.get('ENABLE_PERSISTENT_CONFIG', 'True').lower() == 'true'
@ -202,6 +249,7 @@ class PersistentConfig(Generic[T]):
log.info(f'Updated {self.env_name} to new value {self.value}')
def save(self):
"""Sync save — used ONLY at startup/import time."""
log.info(f"Saving '{self.env_name}' to the database")
path_parts = self.config_path.split('.')
sub_config = CONFIG_DATA
@ -213,6 +261,19 @@ class PersistentConfig(Generic[T]):
save_to_db(CONFIG_DATA)
self.config_value = self.value
async def async_save(self):
"""Async save — used for ALL runtime config persistence."""
log.info(f"Saving '{self.env_name}' to the database")
path_parts = self.config_path.split('.')
sub_config = CONFIG_DATA
for key in path_parts[:-1]:
if key not in sub_config:
sub_config[key] = {}
sub_config = sub_config[key]
sub_config[path_parts[-1]] = self.value
await async_save_to_db(CONFIG_DATA)
self.config_value = self.value
class AppConfig:
_redis: Union[redis.Redis, redis.cluster.RedisCluster] = None
@ -246,12 +307,27 @@ class AppConfig:
self._state[key] = value
else:
self._state[key].value = value
self._state[key].save()
# At runtime (inside the event loop) persist via the async engine
# to avoid blocking the loop and contending with the async DB pool.
# At startup/import time, fall back to sync.
try:
loop = asyncio.get_running_loop()
loop.create_task(self._async_persist(key))
except RuntimeError:
self._state[key].save()
if self._redis and ENABLE_PERSISTENT_CONFIG:
redis_key = f'{self._redis_key_prefix}:config:{key}'
self._redis.set(redis_key, json.dumps(self._state[key].value))
async def _async_persist(self, key):
"""Persist a single config key via the async engine."""
try:
await self._state[key].async_save()
except Exception as e:
log.error(f'Failed to async-persist config key {key}: {e}')
def __getattr__(self, key):
if key not in self._state:
raise AttributeError(f"Config key '{key}' not found")
@ -915,6 +991,7 @@ if CUSTOM_NAME:
####################################
STORAGE_PROVIDER = os.environ.get('STORAGE_PROVIDER', 'local') # defaults to local, s3
STORAGE_LOCAL_CACHE = os.environ.get('STORAGE_LOCAL_CACHE', 'true').lower() == 'true'
S3_ACCESS_KEY_ID = os.environ.get('S3_ACCESS_KEY_ID', None)
S3_SECRET_ACCESS_KEY = os.environ.get('S3_SECRET_ACCESS_KEY', None)
@ -1148,10 +1225,16 @@ ENABLE_SIGNUP = PersistentConfig(
ENABLE_LOGIN_FORM = PersistentConfig(
'ENABLE_LOGIN_FORM',
'ui.ENABLE_LOGIN_FORM',
'ui.enable_login_form',
os.environ.get('ENABLE_LOGIN_FORM', 'True').lower() == 'true',
)
ENABLE_PASSWORD_CHANGE_FORM = PersistentConfig(
'ENABLE_PASSWORD_CHANGE_FORM',
'ui.enable_password_change_form',
os.environ.get('ENABLE_PASSWORD_CHANGE_FORM', 'True').lower() == 'true',
)
ENABLE_PASSWORD_AUTH = os.environ.get('ENABLE_PASSWORD_AUTH', 'True').lower() == 'true'
DEFAULT_LOCALE = PersistentConfig(

View file

@ -91,6 +91,19 @@ class ERROR_MESSAGES(str, Enum):
INVALID_PASSWORD = lambda err='': err if err else 'The password does not meet the required validation criteria.'
AUTOMATION_LIMIT_EXCEEDED = lambda size='': f'Automation limit reached ({size})'
AUTOMATION_TOO_FREQUENT = (
lambda interval='': f'Schedule too frequent. Minimum interval is {interval} seconds.'
)
AUTOMATION_INVALID_RRULE = lambda err='': f'Invalid RRULE: {err}'
AUTOMATION_NO_FUTURE_RUNS = 'RRULE has no future occurrences'
FEATURE_DISABLED = lambda name='': f'{name} is disabled'
INPUT_TOO_LONG = lambda size='': f'Input prompt exceeds maximum length of {size}'
SERVER_CONNECTION_ERROR = 'Open WebUI: Server Connection Error'
REQUIRED_FIELD_EMPTY = lambda name='': f'Required field {name} is empty'
OAUTH_NOT_CONFIGURED = lambda name='': f"Provider '{name}' is not configured"
class TASKS(str, Enum):
def __str__(self) -> str:

View file

@ -431,9 +431,7 @@ except ValueError:
# enabled, the kernel sends TCP keepalive probes on idle connections so
# half-closed sockets (e.g. after a silent firewall/LB reset or a NIC
# flap) are detected before the next command lands on them.
REDIS_SOCKET_KEEPALIVE = (
os.environ.get('REDIS_SOCKET_KEEPALIVE', 'False').lower() == 'true'
)
REDIS_SOCKET_KEEPALIVE = os.environ.get('REDIS_SOCKET_KEEPALIVE', 'False').lower() == 'true'
# How often (in seconds) redis-py should PING an idle pooled connection
# before reusing it. Opt-in: defaults to unset (empty string) so behavior
@ -519,6 +517,11 @@ PASSWORD_VALIDATION_HINT = os.environ.get('PASSWORD_VALIDATION_HINT', '')
BYPASS_MODEL_ACCESS_CONTROL = os.environ.get('BYPASS_MODEL_ACCESS_CONTROL', 'False').lower() == 'true'
# When enabled, skips pydub-based preprocessing (format conversion, compression,
# and chunked splitting) before sending files to processing engines. Useful when
# the upstream provider handles these steps or when ffmpeg is unavailable.
BYPASS_PYDUB_PREPROCESSING = os.environ.get('BYPASS_PYDUB_PREPROCESSING', 'False').lower() == 'true'
# When disabled (default), the OpenAI catch-all proxy endpoint (/{path:path})
# is blocked. Enable only if you need direct passthrough to upstream OpenAI-
# compatible APIs for endpoints not natively handled by Open WebUI.
@ -806,6 +809,36 @@ else:
AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER = AIOHTTP_CLIENT_TIMEOUT
####################################
# AIOHTTP Connection Pool
####################################
AIOHTTP_POOL_CONNECTIONS = os.environ.get('AIOHTTP_POOL_CONNECTIONS', '')
if AIOHTTP_POOL_CONNECTIONS == '':
AIOHTTP_POOL_CONNECTIONS = None
else:
try:
AIOHTTP_POOL_CONNECTIONS = int(AIOHTTP_POOL_CONNECTIONS)
except ValueError:
AIOHTTP_POOL_CONNECTIONS = None
AIOHTTP_POOL_CONNECTIONS_PER_HOST = os.environ.get('AIOHTTP_POOL_CONNECTIONS_PER_HOST', '')
if AIOHTTP_POOL_CONNECTIONS_PER_HOST == '':
AIOHTTP_POOL_CONNECTIONS_PER_HOST = None
else:
try:
AIOHTTP_POOL_CONNECTIONS_PER_HOST = int(AIOHTTP_POOL_CONNECTIONS_PER_HOST)
except ValueError:
AIOHTTP_POOL_CONNECTIONS_PER_HOST = None
AIOHTTP_POOL_DNS_TTL = os.environ.get('AIOHTTP_POOL_DNS_TTL', '300')
try:
AIOHTTP_POOL_DNS_TTL = int(AIOHTTP_POOL_DNS_TTL)
if AIOHTTP_POOL_DNS_TTL < 0:
AIOHTTP_POOL_DNS_TTL = 300
except ValueError:
AIOHTTP_POOL_DNS_TTL = 300
RAG_EMBEDDING_TIMEOUT = os.environ.get('RAG_EMBEDDING_TIMEOUT', '')
if RAG_EMBEDDING_TIMEOUT == '':
@ -909,6 +942,9 @@ AUDIT_INCLUDED_PATHS = os.getenv('AUDIT_INCLUDED_PATHS', '').split(',')
AUDIT_INCLUDED_PATHS = [path.strip() for path in AUDIT_INCLUDED_PATHS]
AUDIT_INCLUDED_PATHS = [path.lstrip('/') for path in AUDIT_INCLUDED_PATHS if path]
# When enabled, GET requests are also audited (disabled by default to avoid log noise)
ENABLE_AUDIT_GET_REQUESTS = os.getenv('ENABLE_AUDIT_GET_REQUESTS', 'False').lower() == 'true'
####################################
# OPENTELEMETRY

View file

@ -214,6 +214,7 @@ if 'sqlite' in ASYNC_SQLALCHEMY_DATABASE_URL:
)
if DATABASE_ENABLE_SQLITE_WAL:
@event.listens_for(async_engine.sync_engine, 'connect')
def _set_sqlite_wal(dbapi_connection, connection_record):
cursor = dbapi_connection.cursor()
@ -229,7 +230,6 @@ else:
pool_timeout=DATABASE_POOL_TIMEOUT,
pool_recycle=DATABASE_POOL_RECYCLE,
pool_pre_ping=True,
poolclass=QueuePool,
)
else:
async_engine = create_async_engine(

View file

@ -21,7 +21,7 @@ from typing import Optional
from aiocache import cached
import aiohttp
import anyio.to_thread
import requests
from redis import Redis
@ -60,6 +60,7 @@ from starsessions.stores.redis import RedisStore
from open_webui.utils import logger
from open_webui.utils.audit import AuditLevel, AuditLoggingMiddleware
from open_webui.utils.logger import start_logger
from open_webui.utils.session_pool import get_session
from open_webui.socket.main import (
MODELS,
app as socket_app,
@ -67,6 +68,7 @@ from open_webui.socket.main import (
periodic_session_pool_cleanup,
get_event_emitter,
get_models_in_use,
get_user_id_from_session_pool,
)
from open_webui.routers import (
analytics,
@ -114,7 +116,7 @@ from open_webui.internal.db import ScopedSession, engine, get_async_session
from open_webui.models.functions import Functions
from open_webui.models.models import Models
from open_webui.models.users import UserModel, Users
from open_webui.models.chats import Chats
from open_webui.models.chats import Chats, ChatForm
from open_webui.config import (
# Ollama
@ -378,6 +380,7 @@ from open_webui.config import (
JWT_EXPIRES_IN,
ENABLE_SIGNUP,
ENABLE_LOGIN_FORM,
ENABLE_PASSWORD_CHANGE_FORM,
ENABLE_API_KEYS,
ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS,
API_KEYS_ALLOWED_ENDPOINTS,
@ -469,12 +472,14 @@ from open_webui.config import (
AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH,
AppConfig,
reset_config,
async_reset_config,
)
from open_webui.env import (
ENABLE_CUSTOM_MODEL_FALLBACK,
LICENSE_KEY,
AUDIT_EXCLUDED_PATHS,
AUDIT_INCLUDED_PATHS,
ENABLE_AUDIT_GET_REQUESTS,
AUDIT_LOG_LEVEL,
CHANGELOG,
REDIS_URL,
@ -553,6 +558,7 @@ from open_webui.utils.oauth import (
get_oauth_client_info_with_static_credentials,
encrypt_data,
decrypt_data,
resolve_oauth_client_info,
OAuthManager,
OAuthClientManager,
OAuthClientInformationFull,
@ -565,6 +571,7 @@ from open_webui.tasks import (
list_task_ids_by_item_id,
create_task,
stop_task,
stop_item_tasks,
list_tasks,
) # Import from tasks.py
@ -622,7 +629,7 @@ async def lifespan(app: FastAPI):
start_logger()
if RESET_CONFIG_ON_START:
reset_config()
await async_reset_config()
if LICENSE_KEY:
get_license_data(app, LICENSE_KEY)
@ -718,6 +725,11 @@ async def lifespan(app: FastAPI):
yield
# Shutdown: clean up shared resources
from open_webui.utils.session_pool import close_session
await close_session()
if hasattr(app.state, 'redis_task_command_listener'):
app.state.redis_task_command_listener.cancel()
@ -845,6 +857,7 @@ app.state.BASE_MODELS = []
app.state.config.WEBUI_URL = WEBUI_URL
app.state.config.ENABLE_SIGNUP = ENABLE_SIGNUP
app.state.config.ENABLE_LOGIN_FORM = ENABLE_LOGIN_FORM
app.state.config.ENABLE_PASSWORD_CHANGE_FORM = ENABLE_PASSWORD_CHANGE_FORM
app.state.config.ENABLE_API_KEYS = ENABLE_API_KEYS
app.state.config.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS = ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS
@ -1390,51 +1403,6 @@ app.add_middleware(RedirectMiddleware)
app.add_middleware(SecurityHeadersMiddleware)
class APIKeyRestrictionMiddleware:
def __init__(self, app):
self.app = app
async def __call__(self, scope, receive, send):
if scope['type'] == 'http':
request = Request(scope)
auth_header = request.headers.get('Authorization')
token = None
if auth_header:
parts = auth_header.split(' ', 1)
if len(parts) == 2:
token = parts[1]
# Only apply restrictions if an sk- API key is used
if token and token.startswith('sk-'):
# Check if restrictions are enabled
if app.state.config.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS:
allowed_paths = [
path.strip()
for path in str(app.state.config.API_KEYS_ALLOWED_ENDPOINTS).split(',')
if path.strip()
]
request_path = request.url.path
# Match exact path or prefix path
is_allowed = any(
request_path == allowed or request_path.startswith(allowed + '/') for allowed in allowed_paths
)
if not is_allowed:
await JSONResponse(
status_code=status.HTTP_403_FORBIDDEN,
content={'detail': 'API key not allowed to access this endpoint.'},
)(scope, receive, send)
return
await self.app(scope, receive, send)
app.add_middleware(APIKeyRestrictionMiddleware)
@app.middleware('http')
async def commit_session_after_request(request: Request, call_next):
response = await call_next(request)
@ -1459,15 +1427,14 @@ async def check_url(request: Request, call_next):
request.state.token = HTTPAuthorizationCredentials(scheme='Bearer', credentials=request.cookies.get('token'))
# Fallback to x-api-key header for Anthropic Messages API routes
# Fallback to x-api-key header (Anthropic-compatible clients use this
# for ALL requests, including GET /v1/models, not just POST /v1/messages).
if request.state.token is None and request.headers.get('x-api-key'):
request_path = request.url.path
if request_path in ('/api/message', '/api/v1/messages') or request_path.startswith('/ollama/v1/messages'):
from fastapi.security import HTTPAuthorizationCredentials
from fastapi.security import HTTPAuthorizationCredentials
request.state.token = HTTPAuthorizationCredentials(
scheme='Bearer', credentials=request.headers.get('x-api-key')
)
request.state.token = HTTPAuthorizationCredentials(
scheme='Bearer', credentials=request.headers.get('x-api-key')
)
request.state.enable_api_keys = app.state.config.ENABLE_API_KEYS
response = await call_next(request)
@ -1560,6 +1527,7 @@ if audit_level != AuditLevel.NONE:
audit_level=audit_level,
excluded_paths=AUDIT_EXCLUDED_PATHS,
included_paths=AUDIT_INCLUDED_PATHS,
audit_get_requests=ENABLE_AUDIT_GET_REQUESTS,
max_body_size=MAX_BODY_LOG_SIZE,
)
##################################
@ -1727,13 +1695,30 @@ async def chat_completion(
if model_info_params.get('reasoning_tags') is not None:
reasoning_tags = model_info_params.get('reasoning_tags')
# parent_id signals intent:
# null → new chat (root message, no parent)
# value → follow-up (user message's parentId = prev assistant)
# absent → legacy caller, no chat management
is_new_chat = 'parent_id' in form_data and form_data['parent_id'] is None and not form_data.get('chat_id')
parent_id = form_data.pop('parent_id', None)
form_data.pop('new_chat', None) # Legacy field
# Multi-model: {model_id: assistant_message_id}
# Single-model fallback: built from 'model' + 'id'
message_ids = form_data.pop('message_ids', None)
if not message_ids:
message_ids = {model_id: form_data.pop('id', None)}
else:
form_data.pop('id', None)
user_message = form_data.pop('user_message', None) or form_data.pop('parent_message', None)
metadata = {
'user_id': user.id,
'chat_id': form_data.pop('chat_id', None),
'message_id': form_data.pop('id', None),
'parent_message': form_data.pop('parent_message', None),
'parent_message_id': form_data.pop('parent_id', None),
'user_message': user_message,
'user_message_id': user_message.get('id') if user_message else None,
'session_id': form_data.pop('session_id', None),
'folder_id': form_data.pop('folder_id', None),
'filter_ids': form_data.pop('filter_ids', []),
'tool_ids': form_data.get('tool_ids', None),
'tool_servers': form_data.pop('tool_servers', None),
@ -1756,36 +1741,160 @@ async def chat_completion(
},
}
if is_new_chat:
metadata['chat_id'] = str(uuid4())
if metadata.get('chat_id') and user:
if not metadata['chat_id'].startswith('local:'): # temporary chats are not stored
# Verify chat ownership — lightweight EXISTS check avoids
# deserializing the full chat JSON blob just to confirm the row exists
if (
not await Chats.is_chat_owner(metadata['chat_id'], user.id) and user.role != 'admin'
): # admins can access any chat
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.DEFAULT(),
chat_id = metadata['chat_id']
if not chat_id.startswith('local:'): # temporary chats are not stored
if is_new_chat:
# Build the full history upfront with ALL assistant placeholders
user_message = metadata.get('user_message') or {}
user_message_id = user_message.get('id') if user_message else None
history_messages = {}
all_assistant_ids = [assistant_id for assistant_id in message_ids.values() if assistant_id]
if user_message_id and user_message:
user_message['childrenIds'] = all_assistant_ids
history_messages[user_message_id] = user_message
for target_model_id, assistant_message_id in message_ids.items():
if assistant_message_id:
history_messages[assistant_message_id] = {
'id': assistant_message_id,
'parentId': user_message_id,
'childrenIds': [],
'role': 'assistant',
'content': '',
'done': False,
'model': target_model_id,
'timestamp': int(time.time()),
}
await Chats.insert_new_chat(
chat_id,
user.id,
ChatForm(
chat={
'id': chat_id,
'title': 'New Chat',
'models': list(message_ids.keys()),
'history': {
'currentId': all_assistant_ids[0] if all_assistant_ids else user_message_id,
'messages': history_messages,
},
'messages': [
{'role': 'user', 'content': user_message.get('content', '')},
] if user_message_id else [],
'tags': [],
'timestamp': int(time.time() * 1000),
},
folder_id=metadata.get('folder_id'),
),
)
# Insert chat files from parent message if any
parent_message = metadata.get('parent_message') or {}
parent_message_files = parent_message.get('files', [])
if parent_message_files:
try:
await Chats.insert_chat_files(
metadata['chat_id'],
parent_message.get('id'),
[
file_item.get('id')
for file_item in parent_message_files
if file_item.get('type') == 'file'
],
user.id,
# Insert chat files from user message if any
user_message_files = user_message.get('files', [])
if user_message_files:
try:
await Chats.insert_chat_files(
chat_id,
user_message_id,
[
file_item.get('id')
for file_item in user_message_files
if file_item.get('type') == 'file'
],
user.id,
)
except Exception as e:
log.debug(f'Error inserting chat files: {e}')
pass
else:
# Existing chat — verify ownership
if (
not await Chats.is_chat_owner(chat_id, user.id) and user.role != 'admin'
):
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.DEFAULT(),
)
except Exception as e:
log.debug(f'Error inserting chat files: {e}')
pass
# Save user message to DB
user_message = metadata.get('user_message') or {}
if user_message and user_message.get('id'):
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id,
user_message['id'],
user_message,
)
# Link grandparent → user message (childrenIds)
grandparent_id = user_message.get('parentId')
if grandparent_id:
grandparent = await Chats.get_message_by_id_and_message_id(chat_id, grandparent_id)
if grandparent:
child_ids = grandparent.get('childrenIds', [])
if user_message['id'] not in child_ids:
child_ids.append(user_message['id'])
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id, grandparent_id, {'childrenIds': child_ids}
)
# Insert chat files from user message if any
user_message_files = user_message.get('files', [])
if user_message_files:
try:
await Chats.insert_chat_files(
chat_id,
user_message.get('id'),
[
file_item.get('id')
for file_item in user_message_files
if file_item.get('type') == 'file'
],
user.id,
)
except Exception as e:
log.debug(f'Error inserting chat files: {e}')
pass
# Save ALL assistant placeholders
user_message_id = metadata.get('user_message_id')
all_assistant_ids = [assistant_id for assistant_id in message_ids.values() if assistant_id]
# Link user message → all assistant messages (childrenIds)
if user_message_id and all_assistant_ids:
existing_user_message = await Chats.get_message_by_id_and_message_id(
chat_id, user_message_id
)
if existing_user_message:
child_ids = existing_user_message.get('childrenIds', [])
for assistant_id in all_assistant_ids:
if assistant_id not in child_ids:
child_ids.append(assistant_id)
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id, user_message_id, {'childrenIds': child_ids},
)
# Save each assistant placeholder
for target_model_id, assistant_message_id in message_ids.items():
if assistant_message_id:
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id,
assistant_message_id,
{
'id': assistant_message_id,
'parentId': user_message_id,
'childrenIds': [],
'role': 'assistant',
'content': '',
'done': False,
'model': target_model_id,
'timestamp': int(time.time()),
},
)
request.state.metadata = metadata
form_data['metadata'] = metadata
@ -1802,19 +1911,21 @@ async def chat_completion(
form_data, metadata, events = await process_chat_payload(request, form_data, user, metadata, model)
response = await chat_completion_handler(request, form_data, user)
if metadata.get('chat_id') and metadata.get('message_id'):
# When the upstream provider returns an error (e.g. HTTP 400
# content-filter, quota exceeded), generate_chat_completion
# returns a JSONResponse instead of raising. Detect this and
# raise so the except-block below emits chat:message:error +
# chat:tasks:cancel, unblocking the frontend.
if isinstance(response, JSONResponse) and response.status_code >= 400:
try:
if not metadata['chat_id'].startswith('local:'):
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
'parentId': metadata.get('parent_message_id', None),
'model': model_id,
},
)
error_body = json.loads(response.body.decode('utf-8', 'replace'))
detail = error_body.get('error', error_body) if isinstance(error_body, dict) else error_body
if isinstance(detail, dict):
detail = detail.get('message', detail.get('detail', str(detail)))
except Exception:
pass
detail = f'Provider returned HTTP {response.status_code}'
raise Exception(detail)
ctx = await build_chat_response_context(request, form_data, user, model, metadata, tasks, events)
@ -1823,17 +1934,18 @@ async def chat_completion(
log.info('Chat processing was cancelled')
try:
event_emitter = await get_event_emitter(metadata)
await asyncio.shield(
event_emitter(
{'type': 'chat:tasks:cancel'},
if event_emitter:
await asyncio.shield(
event_emitter(
{'type': 'chat:tasks:cancel'},
)
)
)
except Exception as e:
pass
finally:
raise # re-raise to ensure proper task cancellation handling
except Exception as e:
log.debug(f'Error processing chat payload: {e}')
log.error('Error processing chat payload: %s', e)
if metadata.get('chat_id') and metadata.get('message_id'):
# Update the chat message with the error
try:
@ -1842,21 +1954,22 @@ async def chat_completion(
metadata['chat_id'],
metadata['message_id'],
{
'parentId': metadata.get('parent_message_id', None),
'parentId': metadata.get('user_message_id', None),
'error': {'content': str(e)},
},
)
event_emitter = await get_event_emitter(metadata)
await event_emitter(
{
'type': 'chat:message:error',
'data': {'error': {'content': str(e)}},
}
)
await event_emitter(
{'type': 'chat:tasks:cancel'},
)
if event_emitter:
await event_emitter(
{
'type': 'chat:message:error',
'data': {'error': {'content': str(e)}},
}
)
await event_emitter(
{'type': 'chat:tasks:cancel'},
)
except Exception:
pass
@ -1892,19 +2005,55 @@ async def chat_completion(
except Exception as e:
log.debug(f'Error emitting chat:active: {e}')
if metadata.get('session_id') and metadata.get('chat_id') and metadata.get('message_id'):
# Asynchronous Chat Processing
task_id, _ = await create_task(
request.app.state.redis,
process_chat(request, form_data, user, metadata, model),
id=metadata['chat_id'],
)
# Emit chat:active=true when task starts
event_emitter = await get_event_emitter(metadata, update_db=False)
if event_emitter:
await event_emitter({'type': 'chat:active', 'data': {'active': True}})
return {'status': True, 'task_id': task_id}
# Fan out: one task per model
if metadata.get('session_id') and metadata.get('chat_id'):
task_ids = []
chat_id = metadata['chat_id']
for target_model_id, assistant_message_id in message_ids.items():
if not assistant_message_id:
continue
# Per-model metadata: own message_id + model
per_model_metadata = {
**metadata,
'message_id': assistant_message_id,
}
# Per-model form_data: own model
model_form_data = {
**form_data,
'model': target_model_id,
'metadata': per_model_metadata,
}
# Resolve the model object for this specific model
resolved_model = request.app.state.MODELS.get(target_model_id, model)
task_id, _ = await create_task(
request.app.state.redis,
process_chat(request, model_form_data, user, per_model_metadata, resolved_model),
id=chat_id,
)
task_ids.append(task_id)
# Emit chat:active=true
if task_ids:
event_emitter = await get_event_emitter(
{**metadata, 'message_id': list(message_ids.values())[0]},
update_db=False,
)
if event_emitter:
await event_emitter({'type': 'chat:active', 'data': {'active': True}})
return {
'status': True,
'task_ids': task_ids,
'chat_id': chat_id,
}
else:
# Legacy/direct: single model, synchronous
metadata['message_id'] = list(message_ids.values())[0]
return await process_chat(request, form_data, user, metadata, model)
@ -1979,6 +2128,8 @@ async def generate_messages(
@app.post('/api/chat/completed')
async def chat_completed(request: Request, form_data: dict, user=Depends(get_verified_user)):
"""Deprecated: outlet filters now run inline during chat completion.
Kept for backward compatibility with external integrations."""
try:
model_item = form_data.pop('model_item', {})
@ -2012,7 +2163,7 @@ async def chat_action(request: Request, action_id: str, form_data: dict, user=De
@app.post('/api/tasks/stop/{task_id}')
async def stop_task_endpoint(request: Request, task_id: str, user=Depends(get_verified_user)):
async def stop_task_endpoint(request: Request, task_id: str, user=Depends(get_admin_user)):
try:
result = await stop_task(request.app.state.redis, task_id)
return result
@ -2021,15 +2172,21 @@ async def stop_task_endpoint(request: Request, task_id: str, user=Depends(get_ve
@app.get('/api/tasks')
async def list_tasks_endpoint(request: Request, user=Depends(get_verified_user)):
async def list_tasks_endpoint(request: Request, user=Depends(get_admin_user)):
return {'tasks': await list_tasks(request.app.state.redis)}
@app.get('/api/tasks/chat/{chat_id}')
@app.get('/api/tasks/chat/{chat_id:path}')
async def list_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=Depends(get_verified_user)):
chat = await Chats.get_chat_by_id(chat_id)
if chat is None or chat.user_id != user.id:
return {'task_ids': []}
if chat_id.startswith('local:'):
socket_id = chat_id[len('local:') :]
owner_id = get_user_id_from_session_pool(socket_id)
if owner_id != user.id and user.role != 'admin':
return {'task_ids': []}
else:
chat = await Chats.get_chat_by_id(chat_id)
if chat is None or (chat.user_id != user.id and user.role != 'admin'):
return {'task_ids': []}
task_ids = await list_task_ids_by_item_id(request.app.state.redis, chat_id)
@ -2037,6 +2194,21 @@ async def list_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=De
return {'task_ids': task_ids}
@app.post('/api/tasks/chat/{chat_id:path}/stop')
async def stop_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=Depends(get_verified_user)):
if chat_id.startswith('local:'):
socket_id = chat_id[len('local:') :]
owner_id = get_user_id_from_session_pool(socket_id)
if owner_id != user.id and user.role != 'admin':
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
else:
chat = await Chats.get_chat_by_id(chat_id)
if chat is None or (chat.user_id != user.id and user.role != 'admin'):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
result = await stop_item_tasks(request.app.state.redis, chat_id)
return result
##################################
#
# Config Endpoints
@ -2091,6 +2263,7 @@ async def get_app_config(request: Request):
'enable_api_keys': app.state.config.ENABLE_API_KEYS,
'enable_signup': app.state.config.ENABLE_SIGNUP,
'enable_login_form': app.state.config.ENABLE_LOGIN_FORM,
'enable_password_change_form': app.state.config.ENABLE_PASSWORD_CHANGE_FORM,
'enable_websocket': ENABLE_WEBSOCKET_SUPPORT,
'enable_version_update_check': ENABLE_VERSION_UPDATE_CHECK,
'enable_public_active_users_count': ENABLE_PUBLIC_ACTIVE_USERS_COUNT,
@ -2301,10 +2474,8 @@ if len(app.state.config.TOOL_SERVER_CONNECTIONS) > 0:
auth_type = tool_server_connection.get('auth_type', 'none')
if server_id and auth_type in ('oauth_2.1', 'oauth_2.1_static'):
oauth_client_info = tool_server_connection.get('info', {}).get('oauth_client_info', '')
try:
oauth_client_info = decrypt_data(oauth_client_info)
oauth_client_info = resolve_oauth_client_info(tool_server_connection)
app.state.oauth_client_manager.add_client(
f'mcp:{server_id}',
OAuthClientInformationFull(**oauth_client_info),
@ -2365,18 +2536,25 @@ async def register_client(request, client_id: str) -> bool:
try:
if auth_type == 'oauth_2.1_static':
# Static credentials: rebuild from stored credentials + fresh metadata
existing_client_info = connection.get('info', {}).get('oauth_client_info', '')
if not existing_client_info:
log.error(f'No stored OAuth client info for static client {client_id}')
return False
existing_data = decrypt_data(existing_client_info)
# Static credentials: rebuild from admin-provided credentials + fresh metadata
info = connection.get('info', {})
oauth_client_id = info.get('oauth_client_id') or ''
oauth_client_secret = info.get('oauth_client_secret') or ''
if not oauth_client_id or not oauth_client_secret:
# Fall back to blob for backward compatibility
existing_client_info = info.get('oauth_client_info', '')
if not existing_client_info:
log.error(f'No stored OAuth client info for static client {client_id}')
return False
existing_data = decrypt_data(existing_client_info)
oauth_client_id = oauth_client_id or existing_data.get('client_id', '')
oauth_client_secret = oauth_client_secret or existing_data.get('client_secret', '')
oauth_client_info = await get_oauth_client_info_with_static_credentials(
request,
client_id,
server_url,
oauth_client_id=existing_data.get('client_id', ''),
oauth_client_secret=existing_data.get('client_secret', ''),
oauth_client_id=oauth_client_id,
oauth_client_secret=oauth_client_secret,
)
else:
oauth_client_info = await get_oauth_client_info_with_dynamic_client_registration(
@ -2509,7 +2687,13 @@ async def oauth_backchannel_logout(
@app.get('/manifest.json')
async def get_manifest_json():
if app.state.EXTERNAL_PWA_MANIFEST_URL:
return requests.get(app.state.EXTERNAL_PWA_MANIFEST_URL).json()
session = await get_session()
async with session.get(
app.state.EXTERNAL_PWA_MANIFEST_URL,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
return await r.json()
else:
return {
'name': app.state.WEBUI_NAME,

View file

@ -295,8 +295,7 @@ class AccessGrantsTable:
async with get_async_db_context(db) as db:
# Check for existing grant
result = await db.execute(
select(AccessGrant)
.filter_by(
select(AccessGrant).filter_by(
resource_type=resource_type,
resource_id=resource_id,
principal_type=principal_type,
@ -334,8 +333,7 @@ class AccessGrantsTable:
"""Remove a single access grant."""
async with get_async_db_context(db) as db:
result = await db.execute(
delete(AccessGrant)
.filter_by(
delete(AccessGrant).filter_by(
resource_type=resource_type,
resource_id=resource_id,
principal_type=principal_type,
@ -355,8 +353,7 @@ class AccessGrantsTable:
"""Remove all access grants for a resource."""
async with get_async_db_context(db) as db:
result = await db.execute(
delete(AccessGrant)
.filter_by(
delete(AccessGrant).filter_by(
resource_type=resource_type,
resource_id=resource_id,
)
@ -451,8 +448,7 @@ class AccessGrantsTable:
"""
async with get_async_db_context(db) as db:
result = await db.execute(
select(AccessGrant)
.filter_by(
select(AccessGrant).filter_by(
resource_type=resource_type,
resource_id=resource_id,
)
@ -470,8 +466,7 @@ class AccessGrantsTable:
"""Get all grants for a specific resource."""
async with get_async_db_context(db) as db:
result = await db.execute(
select(AccessGrant)
.filter_by(
select(AccessGrant).filter_by(
resource_type=resource_type,
resource_id=resource_id,
)
@ -490,8 +485,7 @@ class AccessGrantsTable:
return {}
async with get_async_db_context(db) as db:
result = await db.execute(
select(AccessGrant)
.filter(
select(AccessGrant).filter(
AccessGrant.resource_type == resource_type,
AccessGrant.resource_id.in_(resource_ids),
)
@ -634,8 +628,7 @@ class AccessGrantsTable:
async with get_async_db_context(db) as db:
result = await db.execute(
select(AccessGrant)
.filter_by(
select(AccessGrant).filter_by(
resource_type=resource_type,
resource_id=resource_id,
permission=permission,

View file

@ -141,7 +141,9 @@ class AuthsTable:
except Exception:
return None
async def authenticate_user_by_api_key(self, api_key: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
async def authenticate_user_by_api_key(
self, api_key: str, db: Optional[AsyncSession] = None
) -> Optional[UserModel]:
log.info(f'authenticate_user_by_api_key')
# if no api_key, return None
if not api_key:
@ -159,9 +161,7 @@ class AuthsTable:
async with get_async_db_context(db) as db:
# Single JOIN query instead of two separate queries
result = await db.execute(
select(Auth, User)
.join(User, Auth.id == User.id)
.filter(Auth.email == email, Auth.active == True)
select(Auth, User).join(User, Auth.id == User.id).filter(Auth.email == email, Auth.active == True)
)
row = result.first()
if row:

View file

@ -145,9 +145,7 @@ class AutomationTable:
async def count_by_user(self, user_id: str, db: Optional[AsyncSession] = None) -> int:
async with get_async_db_context(db) as db:
result = await db.execute(
select(func.count()).select_from(Automation).filter_by(user_id=user_id)
)
result = await db.execute(select(func.count()).select_from(Automation).filter_by(user_id=user_id))
return result.scalar()
async def get_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationModel]:
@ -185,9 +183,7 @@ class AutomationTable:
stmt = stmt.order_by(Automation.created_at.desc())
# Get total count
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip:
@ -343,18 +339,14 @@ class AutomationRunTable:
.subquery()
)
result = await db.execute(
select(AutomationRun)
.join(
select(AutomationRun).join(
subq,
(AutomationRun.automation_id == subq.c.automation_id)
& (AutomationRun.created_at == subq.c.max_created),
)
)
rows = result.scalars().all()
return {
row.automation_id: AutomationRunModel.model_validate(row)
for row in rows
}
return {row.automation_id: AutomationRunModel.model_validate(row) for row in rows}
async def get_by_automation(
self,

View file

@ -414,9 +414,13 @@ class ChannelTable:
all_channels = list(membership_channels) + list(standard_channels)
channel_ids = [c.id for c in all_channels]
grants_map = await AccessGrants.get_grants_by_resources('channel', channel_ids, db=db)
return [await self._to_channel_model(c, access_grants=grants_map.get(c.id, []), db=db) for c in all_channels]
return [
await self._to_channel_model(c, access_grants=grants_map.get(c.id, []), db=db) for c in all_channels
]
async def get_dm_channel_by_user_ids(self, user_ids: list[str], db: Optional[AsyncSession] = None) -> Optional[ChannelModel]:
async def get_dm_channel_by_user_ids(
self, user_ids: list[str], db: Optional[AsyncSession] = None
) -> Optional[ChannelModel]:
async with get_async_db_context(db) as db:
# Ensure uniqueness in case a list with duplicates is passed
unique_user_ids = list(set(user_ids))
@ -462,9 +466,7 @@ class ChannelTable:
# 1. Collect all user_ids including groups + inviter
requested_users = await self._collect_unique_user_ids(invited_by, user_ids, group_ids)
result = await db.execute(
select(ChannelMember.user_id).filter(ChannelMember.channel_id == channel_id)
)
result = await db.execute(select(ChannelMember.user_id).filter(ChannelMember.channel_id == channel_id))
existing_users = {row[0] for row in result.all()}
new_user_ids = requested_users - existing_users
@ -512,7 +514,9 @@ class ChannelTable:
membership = result.scalars().first()
return membership is not None
async def join_channel(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[ChannelMemberModel]:
async def join_channel(
self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[ChannelMemberModel]:
async with get_async_db_context(db) as db:
# Check if the membership already exists
result = await db.execute(
@ -581,11 +585,11 @@ class ChannelTable:
membership = result.scalars().first()
return ChannelMemberModel.model_validate(membership) if membership else None
async def get_members_by_channel_id(self, channel_id: str, db: Optional[AsyncSession] = None) -> list[ChannelMemberModel]:
async def get_members_by_channel_id(
self, channel_id: str, db: Optional[AsyncSession] = None
) -> list[ChannelMemberModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(ChannelMember.channel_id == channel_id)
)
result = await db.execute(select(ChannelMember).filter(ChannelMember.channel_id == channel_id))
memberships = result.scalars().all()
return [ChannelMemberModel.model_validate(membership) for membership in memberships]
@ -613,7 +617,9 @@ class ChannelTable:
await db.commit()
return True
async def update_member_last_read_at(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async def update_member_last_read_at(
self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None
) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(
@ -658,11 +664,13 @@ class ChannelTable:
async def is_user_channel_member(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(
select(ChannelMember)
.filter(
ChannelMember.channel_id == channel_id,
ChannelMember.user_id == user_id,
ChannelMember.is_active.is_(True),
).limit(1)
)
.limit(1)
)
membership = result.scalars().first()
return membership is not None
@ -726,11 +734,13 @@ class ChannelTable:
# --- Case A: group or dm => user must be an active member ---
if channel.type in ['group', 'dm']:
result = await db.execute(
select(ChannelMember).filter(
select(ChannelMember)
.filter(
ChannelMember.channel_id == channel.id,
ChannelMember.user_id == user_id,
ChannelMember.is_active.is_(True),
).limit(1)
)
.limit(1)
)
membership = result.scalars().first()
if membership:
@ -774,11 +784,13 @@ class ChannelTable:
# If the channel is a group or dm, read access requires membership (active)
if channel.type in ['group', 'dm']:
result = await db.execute(
select(ChannelMember).filter(
select(ChannelMember)
.filter(
ChannelMember.channel_id == id,
ChannelMember.user_id == user_id,
ChannelMember.is_active.is_(True),
).limit(1)
)
.limit(1)
)
membership = result.scalars().first()
if membership:
@ -863,9 +875,7 @@ class ChannelTable:
) -> bool:
try:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id)
)
result = await db.execute(select(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id))
channel_file = result.scalars().first()
if not channel_file:
return False
@ -878,7 +888,9 @@ class ChannelTable:
except Exception:
return False
async def remove_file_from_channel_by_id(self, channel_id: str, file_id: str, db: Optional[AsyncSession] = None) -> bool:
async def remove_file_from_channel_by_id(
self, channel_id: str, file_id: str, db: Optional[AsyncSession] = None
) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(delete(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id))
@ -921,13 +933,17 @@ class ChannelTable:
await db.commit()
return webhook
async def get_webhooks_by_channel_id(self, channel_id: str, db: Optional[AsyncSession] = None) -> list[ChannelWebhookModel]:
async def get_webhooks_by_channel_id(
self, channel_id: str, db: Optional[AsyncSession] = None
) -> list[ChannelWebhookModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.channel_id == channel_id))
webhooks = result.scalars().all()
return [ChannelWebhookModel.model_validate(w) for w in webhooks]
async def get_webhook_by_id(self, webhook_id: str, db: Optional[AsyncSession] = None) -> Optional[ChannelWebhookModel]:
async def get_webhook_by_id(
self, webhook_id: str, db: Optional[AsyncSession] = None
) -> Optional[ChannelWebhookModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.id == webhook_id))
webhook = result.scalars().first()

View file

@ -272,13 +272,10 @@ class ChatMessageTable:
"""Get distinct chat_ids that used a specific model."""
async with get_async_db_context(db) as db:
stmt = (
select(
ChatMessage.chat_id,
func.max(ChatMessage.created_at).label('last_message_at'),
)
.filter(ChatMessage.model_id == model_id)
)
stmt = select(
ChatMessage.chat_id,
func.max(ChatMessage.created_at).label('last_message_at'),
).filter(ChatMessage.model_id == model_id)
if start_date:
stmt = stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
@ -313,13 +310,10 @@ class ChatMessageTable:
async with get_async_db_context(db) as db:
from open_webui.models.groups import GroupMember
stmt = (
select(ChatMessage.model_id, func.count(ChatMessage.id).label('count'))
.filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
stmt = select(ChatMessage.model_id, func.count(ChatMessage.id).label('count')).filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
if start_date:
@ -365,19 +359,16 @@ class ChatMessageTable:
else:
raise NotImplementedError(f'Unsupported dialect: {dialect}')
stmt = (
select(
ChatMessage.model_id,
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
func.count(ChatMessage.id).label('message_count'),
)
.filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
ChatMessage.usage.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
stmt = select(
ChatMessage.model_id,
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
func.count(ChatMessage.id).label('message_count'),
).filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
ChatMessage.usage.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
if start_date:
@ -430,19 +421,16 @@ class ChatMessageTable:
else:
raise NotImplementedError(f'Unsupported dialect: {dialect}')
stmt = (
select(
ChatMessage.user_id,
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
func.count(ChatMessage.id).label('message_count'),
)
.filter(
ChatMessage.role == 'assistant',
ChatMessage.user_id.isnot(None),
ChatMessage.usage.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
stmt = select(
ChatMessage.user_id,
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
func.count(ChatMessage.id).label('message_count'),
).filter(
ChatMessage.role == 'assistant',
ChatMessage.user_id.isnot(None),
ChatMessage.usage.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
if start_date:
@ -506,9 +494,8 @@ class ChatMessageTable:
async with get_async_db_context(db) as db:
from open_webui.models.groups import GroupMember
stmt = (
select(ChatMessage.chat_id, func.count(ChatMessage.id).label('count'))
.filter(~ChatMessage.user_id.like('shared-%'))
stmt = select(ChatMessage.chat_id, func.count(ChatMessage.id).label('count')).filter(
~ChatMessage.user_id.like('shared-%')
)
if start_date:
@ -535,13 +522,10 @@ class ChatMessageTable:
from datetime import datetime, timedelta
from open_webui.models.groups import GroupMember
stmt = (
select(ChatMessage.created_at, ChatMessage.model_id)
.filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
stmt = select(ChatMessage.created_at, ChatMessage.model_id).filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
if start_date:
@ -585,13 +569,10 @@ class ChatMessageTable:
async with get_async_db_context(db) as db:
from datetime import datetime, timedelta
stmt = (
select(ChatMessage.created_at, ChatMessage.model_id)
.filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
stmt = select(ChatMessage.created_at, ChatMessage.model_id).filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
if start_date:

View file

@ -292,9 +292,10 @@ class ChatTable:
return changed
async def insert_new_chat(self, user_id: str, form_data: ChatForm, db: Optional[AsyncSession] = None) -> Optional[ChatModel]:
async def insert_new_chat(
self, id: str, user_id: str, form_data: ChatForm, db: Optional[AsyncSession] = None
) -> Optional[ChatModel]:
async with get_async_db_context(db) as db:
id = str(uuid.uuid4())
chat = ChatModel(
**{
'id': id,
@ -551,7 +552,9 @@ class ChatTable:
await self.update_chat_by_id(id, chat, db=db)
return message_files
async def insert_shared_chat_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> Optional[ChatModel]:
async def insert_shared_chat_by_chat_id(
self, chat_id: str, db: Optional[AsyncSession] = None
) -> Optional[ChatModel]:
async with get_async_db_context(db) as db:
# Get the existing chat to share
chat = await db.get(Chat, chat_id)
@ -585,7 +588,9 @@ class ChatTable:
await db.commit()
return shared_chat if shared_result else None
async def update_shared_chat_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> Optional[ChatModel]:
async def update_shared_chat_by_chat_id(
self, chat_id: str, db: Optional[AsyncSession] = None
) -> Optional[ChatModel]:
try:
async with get_async_db_context(db) as db:
chat = await db.get(Chat, chat_id)
@ -689,7 +694,9 @@ class ChatTable:
db: Optional[AsyncSession] = None,
) -> list[ChatTitleIdResponse]:
async with get_async_db_context(db) as db:
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at).filter_by(user_id=user_id, archived=True)
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at).filter_by(
user_id=user_id, archived=True
)
if filter:
query_key = filter.get('query')
@ -740,7 +747,11 @@ class ChatTable:
db: Optional[AsyncSession] = None,
) -> list[SharedChatResponse]:
async with get_async_db_context(db) as db:
stmt = select(Chat.id, Chat.title, Chat.share_id, Chat.updated_at, Chat.created_at).filter_by(user_id=user_id).filter(Chat.share_id.isnot(None))
stmt = (
select(Chat.id, Chat.title, Chat.share_id, Chat.updated_at, Chat.created_at)
.filter_by(user_id=user_id)
.filter(Chat.share_id.isnot(None))
)
if filter:
query_key = filter.get('query')
@ -793,7 +804,9 @@ class ChatTable:
db: Optional[AsyncSession] = None,
) -> list[ChatTitleIdResponse]:
async with get_async_db_context(db) as db:
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(user_id=user_id)
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(
user_id=user_id
)
if not include_archived:
stmt = stmt.filter_by(archived=False)
@ -846,7 +859,9 @@ class ChatTable:
db: Optional[AsyncSession] = None,
) -> list[ChatTitleIdResponse]:
async with get_async_db_context(db) as db:
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(user_id=user_id)
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(
user_id=user_id
)
if not include_folders:
stmt = stmt.filter_by(folder_id=None)
@ -889,10 +904,7 @@ class ChatTable:
) -> list[ChatModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Chat)
.filter(Chat.id.in_(chat_ids))
.filter_by(archived=False)
.order_by(Chat.updated_at.desc())
select(Chat).filter(Chat.id.in_(chat_ids)).filter_by(archived=False).order_by(Chat.updated_at.desc())
)
all_chats = result.scalars().all()
return [ChatModel.model_validate(chat) for chat in all_chats]
@ -925,7 +937,9 @@ class ChatTable:
except Exception:
return None
async def get_chat_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[ChatModel]:
async def get_chat_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[ChatModel]:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Chat).filter_by(id=id, user_id=user_id))
@ -941,9 +955,7 @@ class ChatTable:
"""
try:
async with get_async_db_context(db) as db:
result = await db.execute(
select(exists().where(and_(Chat.id == id, Chat.user_id == user_id)))
)
result = await db.execute(select(exists().where(and_(Chat.id == id, Chat.user_id == user_id))))
return result.scalar()
except Exception:
return False
@ -997,9 +1009,7 @@ class ChatTable:
else:
stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id)
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip is not None:
@ -1017,7 +1027,9 @@ class ChatTable:
}
)
async def get_pinned_chats_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[ChatTitleIdResponse]:
async def get_pinned_chats_by_user_id(
self, user_id: str, db: Optional[AsyncSession] = None
) -> list[ChatTitleIdResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at)
@ -1060,7 +1072,9 @@ class ChatTable:
search_text = sanitize_text_for_db(search_text).lower().strip()
if not search_text:
return await self.get_chat_list_by_user_id(user_id, include_archived, filter={}, skip=skip, limit=limit, db=db)
return await self.get_chat_list_by_user_id(
user_id, include_archived, filter={}, skip=skip, limit=limit, db=db
)
search_text_words = search_text.split(' ')
@ -1305,7 +1319,9 @@ class ChatTable:
except Exception:
return None
async def get_chat_tags_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> list[TagModel]:
async def get_chat_tags_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> list[TagModel]:
async with get_async_db_context(db) as db:
chat = await db.get(Chat, id)
tag_ids = chat.meta.get('tags', [])
@ -1320,7 +1336,9 @@ class ChatTable:
db: Optional[AsyncSession] = None,
) -> list[ChatTitleIdResponse]:
async with get_async_db_context(db) as db:
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(user_id=user_id)
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(
user_id=user_id
)
tag_id = tag_name.replace(' ', '_').lower()
bind = await db.connection()
@ -1378,7 +1396,9 @@ class ChatTable:
except Exception:
return None
async def count_chats_by_tag_name_and_user_id(self, tag_name: str, user_id: str, db: Optional[AsyncSession] = None) -> int:
async def count_chats_by_tag_name_and_user_id(
self, tag_name: str, user_id: str, db: Optional[AsyncSession] = None
) -> int:
async with get_async_db_context(db) as db:
stmt = select(func.count(Chat.id)).filter_by(user_id=user_id, archived=False)
tag_id = tag_name.replace(' ', '_').lower()
@ -1424,11 +1444,11 @@ class ChatTable:
orphans.append(tag_id)
await Tags.delete_tags_by_ids_and_user_id(orphans, user_id, db=db)
async def count_chats_by_folder_id_and_user_id(self, folder_id: str, user_id: str, db: Optional[AsyncSession] = None) -> int:
async def count_chats_by_folder_id_and_user_id(
self, folder_id: str, user_id: str, db: Optional[AsyncSession] = None
) -> int:
async with get_async_db_context(db) as db:
result = await db.execute(
select(func.count(Chat.id)).filter_by(user_id=user_id, folder_id=folder_id)
)
result = await db.execute(select(func.count(Chat.id)).filter_by(user_id=user_id, folder_id=folder_id))
count = result.scalar()
log.info(f"Count of chats for folder '{folder_id}': {count}")
@ -1470,9 +1490,7 @@ class ChatTable:
async def delete_chat_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(
update(AutomationRun).filter_by(chat_id=id).values(chat_id=None)
)
await db.execute(update(AutomationRun).filter_by(chat_id=id).values(chat_id=None))
await db.execute(delete(ChatMessage).filter_by(chat_id=id))
await db.execute(delete(Chat).filter_by(id=id))
await db.commit()
@ -1484,9 +1502,7 @@ class ChatTable:
async def delete_chat_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(
update(AutomationRun).filter_by(chat_id=id).values(chat_id=None)
)
await db.execute(update(AutomationRun).filter_by(chat_id=id).values(chat_id=None))
await db.execute(delete(ChatMessage).filter_by(chat_id=id))
await db.execute(delete(Chat).filter_by(id=id, user_id=user_id))
await db.commit()
@ -1502,7 +1518,9 @@ class ChatTable:
chat_id_subquery = select(Chat.id).filter_by(user_id=user_id).scalar_subquery()
await db.execute(
update(AutomationRun).filter(AutomationRun.chat_id.in_(select(Chat.id).filter_by(user_id=user_id))).values(chat_id=None)
update(AutomationRun)
.filter(AutomationRun.chat_id.in_(select(Chat.id).filter_by(user_id=user_id)))
.values(chat_id=None)
)
await db.execute(
delete(ChatMessage).filter(ChatMessage.chat_id.in_(select(Chat.id).filter_by(user_id=user_id)))
@ -1514,16 +1532,16 @@ class ChatTable:
except Exception:
return False
async def delete_chats_by_user_id_and_folder_id(self, user_id: str, folder_id: str, db: Optional[AsyncSession] = None) -> bool:
async def delete_chats_by_user_id_and_folder_id(
self, user_id: str, folder_id: str, db: Optional[AsyncSession] = None
) -> bool:
try:
async with get_async_db_context(db) as db:
chat_ids_stmt = select(Chat.id).filter_by(user_id=user_id, folder_id=folder_id)
await db.execute(
update(AutomationRun).filter(AutomationRun.chat_id.in_(chat_ids_stmt)).values(chat_id=None)
)
await db.execute(
delete(ChatMessage).filter(ChatMessage.chat_id.in_(chat_ids_stmt))
)
await db.execute(delete(ChatMessage).filter(ChatMessage.chat_id.in_(chat_ids_stmt)))
await db.execute(delete(Chat).filter_by(user_id=user_id, folder_id=folder_id))
await db.commit()
@ -1619,9 +1637,7 @@ class ChatTable:
) -> list[ChatFileModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChatFile)
.filter_by(chat_id=chat_id, message_id=message_id)
.order_by(ChatFile.created_at.asc())
select(ChatFile).filter_by(chat_id=chat_id, message_id=message_id).order_by(ChatFile.created_at.asc())
)
all_chat_files = result.scalars().all()
return [ChatFileModel.model_validate(chat_file) for chat_file in all_chat_files]

View file

@ -251,9 +251,7 @@ class FeedbackTable:
stmt = stmt.order_by(Feedback.created_at.desc())
# Count BEFORE pagination
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip:
@ -280,8 +278,9 @@ class FeedbackTable:
async def get_all_feedback_ids(self, db: Optional[AsyncSession] = None) -> list[FeedbackIdResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Feedback.id, Feedback.user_id, Feedback.created_at, Feedback.updated_at)
.order_by(Feedback.updated_at.desc())
select(Feedback.id, Feedback.user_id, Feedback.created_at, Feedback.updated_at).order_by(
Feedback.updated_at.desc()
)
)
return [
FeedbackIdResponse(
@ -378,16 +377,12 @@ class FeedbackTable:
async def get_feedbacks_by_type(self, type: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Feedback).filter_by(type=type).order_by(Feedback.updated_at.desc())
)
result = await db.execute(select(Feedback).filter_by(type=type).order_by(Feedback.updated_at.desc()))
return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()]
async def get_feedbacks_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Feedback).filter_by(user_id=user_id).order_by(Feedback.updated_at.desc())
)
result = await db.execute(select(Feedback).filter_by(user_id=user_id).order_by(Feedback.updated_at.desc()))
return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()]
async def update_feedback_by_id(

View file

@ -125,7 +125,9 @@ class FileUpdateForm(BaseModel):
class FilesTable:
async def insert_new_file(self, user_id: str, form_data: FileForm, db: Optional[AsyncSession] = None) -> Optional[FileModel]:
async def insert_new_file(
self, user_id: str, form_data: FileForm, db: Optional[AsyncSession] = None
) -> Optional[FileModel]:
async with get_async_db_context(db) as db:
file_data = form_data.model_dump()
@ -167,7 +169,9 @@ class FilesTable:
except Exception:
return None
async def get_file_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[FileModel]:
async def get_file_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[FileModel]:
async with get_async_db_context(db) as db:
try:
result = await db.execute(select(File).filter_by(id=id, user_id=user_id))
@ -179,7 +183,9 @@ class FilesTable:
except Exception:
return None
async def get_file_metadata_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FileMetadataResponse]:
async def get_file_metadata_by_id(
self, id: str, db: Optional[AsyncSession] = None
) -> Optional[FileMetadataResponse]:
async with get_async_db_context(db) as db:
try:
file = await db.get(File, id)
@ -211,12 +217,12 @@ class FilesTable:
async def get_files_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[FileModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(File).filter(File.id.in_(ids)).order_by(File.updated_at.desc())
)
result = await db.execute(select(File).filter(File.id.in_(ids)).order_by(File.updated_at.desc()))
return [FileModel.model_validate(file) for file in result.scalars().all()]
async def get_file_metadatas_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[FileMetadataResponse]:
async def get_file_metadatas_by_ids(
self, ids: list[str], db: Optional[AsyncSession] = None
) -> list[FileMetadataResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(File.id, File.hash, File.meta, File.created_at, File.updated_at)
@ -251,18 +257,11 @@ class FilesTable:
if user_id:
stmt = stmt.filter_by(user_id=user_id)
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
result = await db.execute(
stmt.order_by(File.updated_at.desc(), File.id.desc()).offset(skip).limit(limit)
)
items = [
FileModelResponse.model_validate(file, from_attributes=True)
for file in result.scalars().all()
]
result = await db.execute(stmt.order_by(File.updated_at.desc(), File.id.desc()).offset(skip).limit(limit))
items = [FileModelResponse.model_validate(file, from_attributes=True) for file in result.scalars().all()]
return FileListResponse(items=items, total=total)
@ -320,9 +319,7 @@ class FilesTable:
if pattern != '%':
stmt = stmt.filter(File.filename.ilike(pattern, escape='\\'))
result = await db.execute(
stmt.order_by(File.created_at.desc(), File.id.desc()).offset(skip).limit(limit)
)
result = await db.execute(stmt.order_by(File.created_at.desc(), File.id.desc()).offset(skip).limit(limit))
return [FileModel.model_validate(file) for file in result.scalars().all()]
async def update_file_by_id(
@ -349,7 +346,9 @@ class FilesTable:
log.exception(f'Error updating file completely by id: {e}')
return None
async def update_file_hash_by_id(self, id: str, hash: Optional[str], db: Optional[AsyncSession] = None) -> Optional[FileModel]:
async def update_file_hash_by_id(
self, id: str, hash: Optional[str], db: Optional[AsyncSession] = None
) -> Optional[FileModel]:
async with get_async_db_context(db) as db:
try:
result = await db.execute(select(File).filter_by(id=id))
@ -362,7 +361,9 @@ class FilesTable:
except Exception:
return None
async def update_file_data_by_id(self, id: str, data: dict, db: Optional[AsyncSession] = None) -> Optional[FileModel]:
async def update_file_data_by_id(
self, id: str, data: dict, db: Optional[AsyncSession] = None
) -> Optional[FileModel]:
async with get_async_db_context(db) as db:
try:
result = await db.execute(select(File).filter_by(id=id))
@ -374,7 +375,9 @@ class FilesTable:
except Exception as e:
return None
async def update_file_metadata_by_id(self, id: str, meta: dict, db: Optional[AsyncSession] = None) -> Optional[FileModel]:
async def update_file_metadata_by_id(
self, id: str, meta: dict, db: Optional[AsyncSession] = None
) -> Optional[FileModel]:
async with get_async_db_context(db) as db:
try:
result = await db.execute(select(File).filter_by(id=id))

View file

@ -74,14 +74,14 @@ class FolderForm(BaseModel):
data: Optional[dict] = None
meta: Optional[dict] = None
parent_id: Optional[str] = None
model_config = ConfigDict(extra='allow')
model_config = ConfigDict(extra='forbid')
class FolderUpdateForm(BaseModel):
name: Optional[str] = None
data: Optional[dict] = None
meta: Optional[dict] = None
model_config = ConfigDict(extra='allow')
model_config = ConfigDict(extra='forbid')
class FolderTable:
@ -171,9 +171,7 @@ class FolderTable:
async with get_async_db_context(db) as db:
# Check if folder exists
result = await db.execute(
select(Folder)
.filter_by(parent_id=parent_id, user_id=user_id)
.filter(Folder.name.ilike(name))
select(Folder).filter_by(parent_id=parent_id, user_id=user_id).filter(Folder.name.ilike(name))
)
folder = result.scalars().first()
@ -235,8 +233,7 @@ class FolderTable:
form_data = form_data.model_dump(exclude_unset=True)
existing_result = await db.execute(
select(Folder)
.filter_by(
select(Folder).filter_by(
name=form_data.get('name'),
parent_id=folder.parent_id,
user_id=user_id,
@ -289,7 +286,9 @@ class FolderTable:
log.error(f'update_folder: {e}')
return
async def delete_folder_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> list[str]:
async def delete_folder_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> list[str]:
try:
folder_ids = []
async with get_async_db_context(db) as db:

View file

@ -160,7 +160,9 @@ class FunctionsTable:
for func in functions:
if func.id in existing_ids:
await db.execute(
update(Function).filter_by(id=func.id).values(
update(Function)
.filter_by(id=func.id)
.values(
**func.model_dump(),
user_id=user_id,
updated_at=int(time.time()),
@ -233,9 +235,7 @@ class FunctionsTable:
async def get_function_list(self, db: Optional[AsyncSession] = None) -> list[FunctionUserResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Function).order_by(Function.updated_at.desc())
)
result = await db.execute(select(Function).order_by(Function.updated_at.desc()))
functions = result.scalars().all()
user_ids = list(set(func.user_id for func in functions))
@ -261,7 +261,9 @@ class FunctionsTable:
for func in functions
]
async def get_functions_by_type(self, type: str, active_only=False, db: Optional[AsyncSession] = None) -> list[FunctionModel]:
async def get_functions_by_type(
self, type: str, active_only=False, db: Optional[AsyncSession] = None
) -> list[FunctionModel]:
async with get_async_db_context(db) as db:
if active_only:
result = await db.execute(select(Function).filter_by(type=type, is_active=True))
@ -342,7 +344,9 @@ class FunctionsTable:
log.exception(f'Error updating function metadata by id {id}: {e}')
return None
async def get_user_valves_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[dict]:
async def get_user_valves_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[dict]:
try:
user = await Users.get_user_by_id(user_id, db=db)
user_settings = user.settings.model_dump() if user.settings else {}
@ -381,11 +385,15 @@ class FunctionsTable:
log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}')
return None
async def update_function_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[FunctionModel]:
async def update_function_by_id(
self, id: str, updated: dict, db: Optional[AsyncSession] = None
) -> Optional[FunctionModel]:
async with get_async_db_context(db) as db:
try:
await db.execute(
update(Function).filter_by(id=id).values(
update(Function)
.filter_by(id=id)
.values(
**updated,
updated_at=int(time.time()),
)

View file

@ -261,12 +261,10 @@ class GroupTable:
if 'share' in filter:
share_value = filter['share']
stmt = stmt.filter(Group.data.op('->>') ('share') == str(share_value))
stmt = stmt.filter(Group.data.op('->>')('share') == str(share_value))
# Get total count
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
member_count = (
@ -348,7 +346,9 @@ class GroupTable:
return [m[0] for m in members]
async def get_group_user_ids_by_ids(self, group_ids: list[str], db: Optional[AsyncSession] = None) -> dict[str, list[str]]:
async def get_group_user_ids_by_ids(
self, group_ids: list[str], db: Optional[AsyncSession] = None
) -> dict[str, list[str]]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(GroupMember.group_id, GroupMember.user_id).filter(GroupMember.group_id.in_(group_ids))
@ -362,7 +362,9 @@ class GroupTable:
return group_user_ids
async def set_group_user_ids_by_id(self, group_id: str, user_ids: list[str], db: Optional[AsyncSession] = None) -> None:
async def set_group_user_ids_by_id(
self, group_id: str, user_ids: list[str], db: Optional[AsyncSession] = None
) -> None:
async with get_async_db_context(db) as db:
# Delete existing members
await db.execute(delete(GroupMember).filter(GroupMember.group_id == group_id))
@ -411,7 +413,9 @@ class GroupTable:
try:
async with get_async_db_context(db) as db:
await db.execute(
update(Group).filter_by(id=id).values(
update(Group)
.filter_by(id=id)
.values(
**form_data.model_dump(exclude_none=True),
updated_at=int(time.time()),
)
@ -455,14 +459,10 @@ class GroupTable:
# Remove the user from each group
for group in groups:
await db.execute(
delete(GroupMember).filter(
GroupMember.group_id == group.id, GroupMember.user_id == user_id
)
delete(GroupMember).filter(GroupMember.group_id == group.id, GroupMember.user_id == user_id)
)
await db.execute(
update(Group).filter_by(id=group.id).values(updated_at=int(time.time()))
)
await db.execute(update(Group).filter_by(id=group.id).values(updated_at=int(time.time())))
await db.commit()
return True
@ -507,7 +507,9 @@ class GroupTable:
continue
return new_groups
async def sync_groups_by_group_names(self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None) -> bool:
async def sync_groups_by_group_names(
self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None
) -> bool:
async with get_async_db_context(db) as db:
try:
now = int(time.time())
@ -538,9 +540,7 @@ class GroupTable:
)
)
await db.execute(
update(Group).filter(Group.id.in_(groups_to_remove)).values(updated_at=now)
)
await db.execute(update(Group).filter(Group.id.in_(groups_to_remove)).values(updated_at=now))
# 5. Bulk insert missing memberships
for group_id in groups_to_add:
@ -555,9 +555,7 @@ class GroupTable:
)
if groups_to_add:
await db.execute(
update(Group).filter(Group.id.in_(groups_to_add)).values(updated_at=now)
)
await db.execute(update(Group).filter(Group.id.in_(groups_to_add)).values(updated_at=now))
await db.commit()
return True

View file

@ -196,11 +196,13 @@ class KnowledgeTable:
knowledge_bases.append(
KnowledgeUserModel.model_validate(
{
**(await self._to_knowledge_model(
knowledge,
access_grants=grants_map.get(knowledge.id, []),
db=db,
)).model_dump(),
**(
await self._to_knowledge_model(
knowledge,
access_grants=grants_map.get(knowledge.id, []),
db=db,
)
).model_dump(),
'user': user.model_dump() if user else None,
}
)
@ -249,9 +251,7 @@ class KnowledgeTable:
stmt = stmt.order_by(Knowledge.updated_at.desc(), Knowledge.id.asc())
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip:
stmt = stmt.offset(skip)
@ -269,11 +269,13 @@ class KnowledgeTable:
knowledge_bases.append(
KnowledgeUserModel.model_validate(
{
**(await self._to_knowledge_model(
knowledge_base,
access_grants=grants_map.get(knowledge_base.id, []),
db=db,
)).model_dump(),
**(
await self._to_knowledge_model(
knowledge_base,
access_grants=grants_map.get(knowledge_base.id, []),
db=db,
)
).model_dump(),
'user': (UserModel.model_validate(user).model_dump() if user else None),
}
)
@ -321,9 +323,7 @@ class KnowledgeTable:
stmt = stmt.order_by(File.updated_at.desc(), File.id.asc())
# Count before pagination
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip:
@ -490,9 +490,7 @@ class KnowledgeTable:
stmt = stmt.order_by(primary_sort, File.id.asc())
# Count BEFORE pagination
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip:
@ -530,7 +528,9 @@ class KnowledgeTable:
except Exception:
return []
async def get_file_metadatas_by_id(self, knowledge_id: str, db: Optional[AsyncSession] = None) -> list[FileMetadataResponse]:
async def get_file_metadatas_by_id(
self, knowledge_id: str, db: Optional[AsyncSession] = None
) -> list[FileMetadataResponse]:
try:
files = await self.get_files_by_id(knowledge_id, db=db)
return [FileMetadataResponse(**file.model_dump()) for file in files]
@ -579,7 +579,9 @@ class KnowledgeTable:
except Exception:
return False
async def remove_file_from_knowledge_by_id(self, knowledge_id: str, file_id: str, db: Optional[AsyncSession] = None) -> bool:
async def remove_file_from_knowledge_by_id(
self, knowledge_id: str, file_id: str, db: Optional[AsyncSession] = None
) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(delete(KnowledgeFile).filter_by(knowledge_id=knowledge_id, file_id=file_id))
@ -596,9 +598,7 @@ class KnowledgeTable:
await db.commit()
# Update the knowledge entry's updated_at timestamp
await db.execute(
update(Knowledge).filter_by(id=id).values(updated_at=int(time.time()))
)
await db.execute(update(Knowledge).filter_by(id=id).values(updated_at=int(time.time())))
await db.commit()
return await self.get_knowledge_by_id(id=id, db=db)
@ -616,7 +616,9 @@ class KnowledgeTable:
try:
async with get_async_db_context(db) as db:
await db.execute(
update(Knowledge).filter_by(id=id).values(
update(Knowledge)
.filter_by(id=id)
.values(
**form_data.model_dump(exclude={'access_grants'}),
updated_at=int(time.time()),
)
@ -635,7 +637,9 @@ class KnowledgeTable:
try:
async with get_async_db_context(db) as db:
await db.execute(
update(Knowledge).filter_by(id=id).values(
update(Knowledge)
.filter_by(id=id)
.values(
data=data,
updated_at=int(time.time()),
)

View file

@ -250,11 +250,11 @@ class MessageTable:
}
return None
async def get_thread_replies_by_message_id(self, id: str, db: Optional[AsyncSession] = None) -> list[MessageReplyToResponse]:
async def get_thread_replies_by_message_id(
self, id: str, db: Optional[AsyncSession] = None
) -> list[MessageReplyToResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Message).filter_by(parent_id=id).order_by(Message.created_at.desc())
)
result = await db.execute(select(Message).filter_by(parent_id=id).order_by(Message.created_at.desc()))
all_messages = result.scalars().all()
messages = []
@ -369,7 +369,9 @@ class MessageTable:
)
return messages
async def get_last_message_by_channel_id(self, channel_id: str, db: Optional[AsyncSession] = None) -> Optional[MessageModel]:
async def get_last_message_by_channel_id(
self, channel_id: str, db: Optional[AsyncSession] = None
) -> Optional[MessageModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Message).filter_by(channel_id=channel_id).order_by(Message.created_at.desc()).limit(1)
@ -453,9 +455,7 @@ class MessageTable:
) -> Optional[MessageReactionModel]:
async with get_async_db_context(db) as db:
# check for existing reaction
result = await db.execute(
select(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name)
)
result = await db.execute(select(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name))
existing_reaction = result.scalars().first()
if existing_reaction:
return MessageReactionModel.model_validate(existing_reaction)

View file

@ -1,3 +1,4 @@
import json
import logging
import time
from typing import Optional
@ -200,7 +201,8 @@ class ModelsTable:
model_ids = [model.id for model in all_models]
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
return [
await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db) for model in all_models
await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db)
for model in all_models
]
async def get_models(self, db: Optional[AsyncSession] = None) -> list[ModelUserResponse]:
@ -221,11 +223,13 @@ class ModelsTable:
models.append(
ModelUserResponse.model_validate(
{
**(await self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,
)).model_dump(),
**(
await self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,
)
).model_dump(),
'user': user.model_dump() if user else None,
}
)
@ -239,7 +243,8 @@ class ModelsTable:
model_ids = [model.id for model in all_models]
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
return [
await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db) for model in all_models
await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db)
for model in all_models
]
async def get_models_by_user_id(
@ -315,9 +320,20 @@ class ModelsTable:
tag = filter.get('tag')
if tag:
like_pattern = f'%"{tag.lower()}"%'
meta_text = func.lower(cast(Model.meta, String))
stmt = stmt.filter(meta_text.like(like_pattern))
# SQLite stores JSON text via json.dumps(ensure_ascii=True),
# so non-ASCII chars are \uXXXX-escaped. PostgreSQL native JSONB
# stores literal Unicode. Use the right pattern for each.
if db.bind.dialect.name == 'sqlite':
if tag.isascii():
meta_text = func.lower(cast(Model.meta, String))
pattern = f'%{json.dumps(tag.lower())}%'
else:
meta_text = cast(Model.meta, String)
pattern = f'%{json.dumps(tag)}%'
else:
meta_text = func.lower(cast(Model.meta, String))
pattern = f'%{json.dumps(tag.lower(), ensure_ascii=False)}%'
stmt = stmt.filter(meta_text.like(pattern))
order_by = filter.get('order_by')
direction = filter.get('direction')
@ -342,9 +358,7 @@ class ModelsTable:
stmt = stmt.order_by(Model.created_at.desc())
# Count BEFORE pagination
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip:
@ -362,11 +376,13 @@ class ModelsTable:
for model, user in items:
models.append(
ModelUserResponse(
**(await self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,
)).model_dump(),
**(
await self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,
)
).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
)
@ -416,7 +432,9 @@ class ModelsTable:
except Exception:
return None
async def update_model_by_id(self, id: str, model: ModelForm, db: Optional[AsyncSession] = None) -> Optional[ModelModel]:
async def update_model_by_id(
self, id: str, model: ModelForm, db: Optional[AsyncSession] = None
) -> Optional[ModelModel]:
try:
async with get_async_db_context(db) as db:
# update only the fields that are present in the model
@ -473,7 +491,9 @@ class ModelsTable:
except Exception:
return False
async def sync_models(self, user_id: str, models: list[ModelModel], db: Optional[AsyncSession] = None) -> list[ModelModel]:
async def sync_models(
self, user_id: str, models: list[ModelModel], db: Optional[AsyncSession] = None
) -> list[ModelModel]:
try:
async with get_async_db_context(db) as db:
# Get existing models
@ -488,7 +508,9 @@ class ModelsTable:
for model in models:
if model.id in existing_ids:
await db.execute(
update(Model).filter_by(id=model.id).values(
update(Model)
.filter_by(id=model.id)
.values(
**model.model_dump(exclude={'access_grants'}),
user_id=user_id,
updated_at=int(time.time()),

View file

@ -113,7 +113,9 @@ class NoteTable:
permission=permission,
)
async def insert_new_note(self, user_id: str, form_data: NoteForm, db: Optional[AsyncSession] = None) -> Optional[NoteModel]:
async def insert_new_note(
self, user_id: str, form_data: NoteForm, db: Optional[AsyncSession] = None
) -> Optional[NoteModel]:
async with get_async_db_context(db) as db:
note = NoteModel(
**{
@ -159,18 +161,22 @@ class NoteTable:
if filter:
query_key = filter.get('query')
if query_key:
# Normalize search by removing hyphens and spaces (e.g., "todo" matches "to-do" and "to do")
normalized_query = query_key.replace('-', '').replace(' ', '')
stmt = stmt.filter(
or_(
func.replace(func.replace(Note.title, '-', ''), ' ', '').ilike(f'%{normalized_query}%'),
func.replace(
func.replace(cast(Note.data['content']['md'], Text), '-', ''),
' ',
'',
).ilike(f'%{normalized_query}%'),
# Split query into individual words and normalize each
# (strip hyphens so "todo" matches "to-do").
# All words must match somewhere in title OR content (AND semantics).
search_words = query_key.split()
normalized_words = [w.replace('-', '') for w in search_words if w.replace('-', '')]
for word in normalized_words:
stmt = stmt.filter(
or_(
func.replace(func.replace(Note.title, '-', ''), ' ', '').ilike(f'%{word}%'),
func.replace(
func.replace(cast(Note.data['content']['md'], Text), '-', ''),
' ',
'',
).ilike(f'%{word}%'),
)
)
)
view_option = filter.get('view_option')
if view_option == 'created':
@ -216,9 +222,7 @@ class NoteTable:
stmt = stmt.order_by(Note.updated_at.desc())
# Count BEFORE pagination
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip:
@ -236,11 +240,13 @@ class NoteTable:
for note, user in items:
notes.append(
NoteUserResponse(
**(await self._to_note_model(
note,
access_grants=grants_map.get(note.id, []),
db=db,
)).model_dump(),
**(
await self._to_note_model(
note,
access_grants=grants_map.get(note.id, []),
db=db,
)
).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
)

View file

@ -123,7 +123,7 @@ class OAuthSessionTable:
'user_id': user_id,
'provider': provider,
'token': self._encrypt_token(token),
'expires_at': token.get('expires_at'),
'expires_at': token.get('expires_at') or int(time.time() + 3600),
'created_at': current_time,
'updated_at': current_time,
}
@ -151,7 +151,9 @@ class OAuthSessionTable:
log.error(f'Error creating OAuth session: {e}')
return None
async def get_session_by_id(self, session_id: str, db: Optional[AsyncSession] = None) -> Optional[OAuthSessionModel]:
async def get_session_by_id(
self, session_id: str, db: Optional[AsyncSession] = None
) -> Optional[OAuthSessionModel]:
"""Get OAuth session by ID"""
try:
async with get_async_db_context(db) as db:
@ -235,15 +237,17 @@ class OAuthSessionTable:
results = []
for session in sessions:
try:
results.append(OAuthSessionModel(
id=session.id,
user_id=session.user_id,
provider=session.provider,
token=self._decrypt_token(session.token),
expires_at=session.expires_at,
created_at=session.created_at,
updated_at=session.updated_at,
))
results.append(
OAuthSessionModel(
id=session.id,
user_id=session.user_id,
provider=session.provider,
token=self._decrypt_token(session.token),
expires_at=session.expires_at,
created_at=session.created_at,
updated_at=session.updated_at,
)
)
except Exception as e:
log.warning(
f'Skipping OAuth session {session.id} due to decryption failure, deleting corrupted session: {type(e).__name__}: {e}'
@ -266,9 +270,11 @@ class OAuthSessionTable:
current_time = int(time.time())
await db.execute(
update(OAuthSession).filter_by(id=session_id).values(
update(OAuthSession)
.filter_by(id=session_id)
.values(
token=self._encrypt_token(token),
expires_at=token.get('expires_at'),
expires_at=token.get('expires_at') or int(time.time() + 3600),
updated_at=current_time,
)
)

View file

@ -1,3 +1,4 @@
import json
import time
import uuid
from typing import Optional
@ -213,11 +214,13 @@ class PromptsTable:
prompts.append(
PromptUserResponse.model_validate(
{
**(await self._to_prompt_model(
prompt,
access_grants=grants_map.get(prompt.id, []),
db=db,
)).model_dump(),
**(
await self._to_prompt_model(
prompt,
access_grants=grants_map.get(prompt.id, []),
db=db,
)
).model_dump(),
'user': user.model_dump() if user else None,
}
)
@ -290,10 +293,22 @@ class PromptsTable:
tag = filter.get('tag')
if tag:
# Search for tag in JSON array field
like_pattern = f'%"{tag.lower()}"%'
tags_text = func.lower(cast(Prompt.tags, String))
stmt = stmt.filter(tags_text.like(like_pattern))
# SQLite stores JSON text via json.dumps(ensure_ascii=True),
# so non-ASCII chars are \uXXXX-escaped. PostgreSQL native JSONB
# stores literal Unicode. Use the right pattern for each.
if db.bind.dialect.name == 'sqlite':
if tag.isascii():
tags_text = func.lower(cast(Prompt.tags, String))
pattern = f'%{json.dumps(tag.lower())}%'
else:
# LOWER() is ASCII-only; non-ASCII codepoints would
# produce different \uXXXX escapes when lowered.
tags_text = cast(Prompt.tags, String)
pattern = f'%{json.dumps(tag)}%'
else:
tags_text = func.lower(cast(Prompt.tags, String))
pattern = f'%{json.dumps(tag.lower(), ensure_ascii=False)}%'
stmt = stmt.filter(tags_text.like(pattern))
order_by = filter.get('order_by')
direction = filter.get('direction')
@ -319,9 +334,7 @@ class PromptsTable:
stmt = stmt.order_by(Prompt.updated_at.desc())
# Count BEFORE pagination
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip:
@ -339,11 +352,13 @@ class PromptsTable:
for prompt, user in items:
prompts.append(
PromptUserResponse(
**(await self._to_prompt_model(
prompt,
access_grants=grants_map.get(prompt.id, []),
db=db,
)).model_dump(),
**(
await self._to_prompt_model(
prompt,
access_grants=grants_map.get(prompt.id, []),
db=db,
)
).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
)

View file

@ -184,11 +184,13 @@ class SkillsTable:
skills.append(
SkillUserModel.model_validate(
{
**(await self._to_skill_model(
skill,
access_grants=grants_map.get(skill.id, []),
db=db,
)).model_dump(),
**(
await self._to_skill_model(
skill,
access_grants=grants_map.get(skill.id, []),
db=db,
)
).model_dump(),
'user': user.model_dump() if user else None,
}
)
@ -262,9 +264,7 @@ class SkillsTable:
stmt = stmt.order_by(Skill.updated_at.desc())
# Count BEFORE pagination
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip:
@ -282,11 +282,13 @@ class SkillsTable:
for skill, user in items:
skills.append(
SkillUserResponse(
**(await self._to_skill_model(
skill,
access_grants=grants_map.get(skill.id, []),
db=db,
)).model_dump(),
**(
await self._to_skill_model(
skill,
access_grants=grants_map.get(skill.id, []),
db=db,
)
).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
)
@ -296,7 +298,9 @@ class SkillsTable:
log.exception(f'Error searching skills: {e}')
return SkillListResponse(items=[], total=0)
async def update_skill_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[SkillModel]:
async def update_skill_by_id(
self, id: str, updated: dict, db: Optional[AsyncSession] = None
) -> Optional[SkillModel]:
try:
async with get_async_db_context(db) as db:
access_grants = updated.pop('access_grants', None)

View file

@ -71,7 +71,9 @@ class TagTable:
log.exception(f'Error inserting a new tag: {e}')
return None
async def get_tag_by_name_and_user_id(self, name: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[TagModel]:
async def get_tag_by_name_and_user_id(
self, name: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[TagModel]:
try:
id = name.replace(' ', '_').lower()
async with get_async_db_context(db) as db:
@ -86,7 +88,9 @@ class TagTable:
result = await db.execute(select(Tag).filter_by(user_id=user_id))
return [TagModel.model_validate(tag) for tag in result.scalars().all()]
async def get_tags_by_ids_and_user_id(self, ids: list[str], user_id: str, db: Optional[AsyncSession] = None) -> list[TagModel]:
async def get_tags_by_ids_and_user_id(
self, ids: list[str], user_id: str, db: Optional[AsyncSession] = None
) -> list[TagModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Tag).filter(Tag.id.in_(ids), Tag.user_id == user_id))
return [TagModel.model_validate(tag) for tag in result.scalars().all()]
@ -103,7 +107,9 @@ class TagTable:
log.error(f'delete_tag: {e}')
return False
async def delete_tags_by_ids_and_user_id(self, ids: list[str], user_id: str, db: Optional[AsyncSession] = None) -> bool:
async def delete_tags_by_ids_and_user_id(
self, ids: list[str], user_id: str, db: Optional[AsyncSession] = None
) -> bool:
"""Delete all tags whose id is in *ids* for the given user, in one query."""
if not ids:
return True

View file

@ -172,11 +172,13 @@ class ToolsTable:
tools.append(
ToolUserModel.model_validate(
{
**(await self._to_tool_model(
tool,
access_grants=grants_map.get(tool.id, []),
db=db,
)).model_dump(),
**(
await self._to_tool_model(
tool,
access_grants=grants_map.get(tool.id, []),
db=db,
)
).model_dump(),
'user': user.model_dump() if user else None,
}
)
@ -218,18 +220,20 @@ class ToolsTable:
log.exception(f'Error getting tool valves by id {id}')
return None
async def update_tool_valves_by_id(self, id: str, valves: dict, db: Optional[AsyncSession] = None) -> Optional[ToolValves]:
async def update_tool_valves_by_id(
self, id: str, valves: dict, db: Optional[AsyncSession] = None
) -> Optional[ToolValves]:
try:
async with get_async_db_context(db) as db:
await db.execute(
update(Tool).filter_by(id=id).values(valves=valves, updated_at=int(time.time()))
)
await db.execute(update(Tool).filter_by(id=id).values(valves=valves, updated_at=int(time.time())))
await db.commit()
return await self.get_tool_by_id(id, db=db)
except Exception:
return None
async def get_user_valves_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[dict]:
async def get_user_valves_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[dict]:
try:
user = await Users.get_user_by_id(user_id, db=db)
user_settings = user.settings.model_dump() if user.settings else {}
@ -272,9 +276,7 @@ class ToolsTable:
try:
async with get_async_db_context(db) as db:
access_grants = updated.pop('access_grants', None)
await db.execute(
update(Tool).filter_by(id=id).values(**updated, updated_at=int(time.time()))
)
await db.execute(update(Tool).filter_by(id=id).values(**updated, updated_at=int(time.time())))
await db.commit()
if access_grants is not None:
await AccessGrants.set_access_grants('tool', id, access_grants, db=db)

View file

@ -31,11 +31,13 @@ import datetime
# daily bread of every session. Let none go hungry.
####################
class UserSettings(BaseModel):
ui: Optional[dict] = {}
model_config = ConfigDict(extra='allow')
pass
class User(Base):
__tablename__ = 'user'
@ -69,6 +71,7 @@ class User(Base):
updated_at = Column(BigInteger)
created_at = Column(BigInteger)
class UserModel(BaseModel):
id: str
@ -109,11 +112,13 @@ class UserModel(BaseModel):
self.profile_image_url = f'/api/v1/users/{self.id}/profile/image'
return self
class UserStatusModel(UserModel):
is_active: bool = False
model_config = ConfigDict(from_attributes=True)
class ApiKey(Base):
__tablename__ = 'api_key'
@ -126,6 +131,7 @@ class ApiKey(Base):
created_at = Column(BigInteger, nullable=False)
updated_at = Column(BigInteger, nullable=False)
class ApiKeyModel(BaseModel):
id: str
user_id: str
@ -138,10 +144,12 @@ class ApiKeyModel(BaseModel):
model_config = ConfigDict(from_attributes=True)
####################
# Forms
####################
class UpdateProfileForm(BaseModel):
profile_image_url: str
name: str
@ -154,25 +162,31 @@ class UpdateProfileForm(BaseModel):
def check_profile_image_url(cls, v: str) -> str:
return validate_profile_image_url(v)
class UserGroupIdsModel(UserModel):
group_ids: list[str] = []
class UserModelResponse(UserModel):
model_config = ConfigDict(extra='allow')
class UserListResponse(BaseModel):
users: list[UserModelResponse]
total: int
class UserGroupIdsListResponse(BaseModel):
users: list[UserGroupIdsModel]
total: int
class UserStatus(BaseModel):
status_emoji: Optional[str] = None
status_message: Optional[str] = None
status_expires_at: Optional[int] = None
class UserInfoResponse(UserStatus):
id: str
name: str
@ -182,51 +196,63 @@ class UserInfoResponse(UserStatus):
groups: Optional[list] = []
is_active: bool = False
class UserIdNameResponse(BaseModel):
id: str
name: str
class UserIdNameStatusResponse(UserStatus):
id: str
name: str
is_active: Optional[bool] = None
class UserInfoListResponse(BaseModel):
users: list[UserInfoResponse]
total: int
class UserIdNameListResponse(BaseModel):
users: list[UserIdNameResponse]
total: int
class UserNameResponse(BaseModel):
id: str
name: str
role: str
class UserResponse(UserNameResponse):
email: str
class UserProfileImageResponse(UserNameResponse):
email: str
profile_image_url: str
class UserRoleUpdateForm(BaseModel):
id: str
role: str
class UserUpdateForm(BaseModel):
role: str
name: str
email: str
profile_image_url: str
role: Optional[str] = None
name: Optional[str] = None
email: Optional[str] = None
profile_image_url: Optional[str] = None
password: Optional[str] = None
@field_validator('profile_image_url')
@field_validator('profile_image_url', mode='before')
@classmethod
def check_profile_image_url(cls, v: str) -> str:
def check_profile_image_url(cls, v: Optional[str]) -> Optional[str]:
if v is None:
return v
return validate_profile_image_url(v)
class UsersTable:
async def insert_new_user(
self,
@ -292,7 +318,9 @@ class UsersTable:
except Exception:
return None
async def get_user_by_oauth_sub(self, provider: str, sub: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
async def get_user_by_oauth_sub(
self, provider: str, sub: str, db: Optional[AsyncSession] = None
) -> Optional[UserModel]:
try:
async with get_async_db_context(db) as db:
dialect_name = db.bind.dialect.name
@ -457,9 +485,7 @@ class UsersTable:
stmt = stmt.order_by(User.created_at.desc())
# Count BEFORE pagination
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
# correct pagination logic
@ -478,20 +504,18 @@ class UsersTable:
async def get_users_by_group_id(self, group_id: str, db: Optional[AsyncSession] = None) -> list[UserModel]:
async with get_async_db_context(db) as db:
from open_webui.models.groups import GroupMember
result = await db.execute(
select(User)
.join(GroupMember, User.id == GroupMember.user_id)
.filter(GroupMember.group_id == group_id)
select(User).join(GroupMember, User.id == GroupMember.user_id).filter(GroupMember.group_id == group_id)
)
users = result.scalars().all()
return [UserModel.model_validate(user) for user in users]
async def get_users_by_user_ids(self, user_ids: list[str], db: Optional[AsyncSession] = None) -> list[UserStatusModel]:
async def get_users_by_user_ids(
self, user_ids: list[str], db: Optional[AsyncSession] = None
) -> list[UserStatusModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(User).filter(User.id.in_(user_ids))
)
result = await db.execute(select(User).filter(User.id.in_(user_ids)))
users = result.scalars().all()
return [UserModel.model_validate(user) for user in users]
@ -536,7 +560,9 @@ class UsersTable:
)
return result.scalar()
async def update_user_role_by_id(self, id: str, role: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
async def update_user_role_by_id(
self, id: str, role: str, db: Optional[AsyncSession] = None
) -> Optional[UserModel]:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
@ -674,7 +700,9 @@ class UsersTable:
print(e)
return None
async def update_user_settings_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
async def update_user_settings_by_id(
self, id: str, updated: dict, db: Optional[AsyncSession] = None
) -> Optional[UserModel]:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
@ -802,4 +830,5 @@ class UsersTable:
return user.last_active_at >= three_minutes_ago
return False
Users = UsersTable()

View file

@ -34,6 +34,7 @@ from open_webui.retrieval.vector.main import (
SearchResult,
GetResult,
)
from open_webui.utils.misc import sanitize_text_for_db
from open_webui.config import (
PGVECTOR_DB_URL,
PGVECTOR_INITIALIZE_MAX_VECTOR_LENGTH,
@ -289,7 +290,9 @@ class PgvectorClient(VectorDBBase):
vector = self.adjust_vector_length(item['vector'])
# Use raw SQL for BYTEA/pgcrypto
# Ensure metadata is converted to its JSON text representation
json_metadata = json.dumps(item['metadata'])
# Sanitize to strip null bytes / surrogates that PostgreSQL cannot store
json_metadata = sanitize_text_for_db(json.dumps(item['metadata']))
item_text = sanitize_text_for_db(item['text'])
self.session.execute(
text("""
INSERT INTO document_chunk
@ -305,7 +308,7 @@ class PgvectorClient(VectorDBBase):
'id': item['id'],
'vector': vector,
'collection_name': collection_name,
'text': item['text'],
'text': item_text,
'metadata_text': json_metadata,
'key': PGVECTOR_PGCRYPTO_KEY,
},
@ -338,7 +341,9 @@ class PgvectorClient(VectorDBBase):
if PGVECTOR_PGCRYPTO:
for item in items:
vector = self.adjust_vector_length(item['vector'])
json_metadata = json.dumps(item['metadata'])
# Sanitize to strip null bytes / surrogates that PostgreSQL cannot store
json_metadata = sanitize_text_for_db(json.dumps(item['metadata']))
item_text = sanitize_text_for_db(item['text'])
self.session.execute(
text("""
INSERT INTO document_chunk
@ -358,7 +363,7 @@ class PgvectorClient(VectorDBBase):
'id': item['id'],
'vector': vector,
'collection_name': collection_name,
'text': item['text'],
'text': item_text,
'metadata_text': json_metadata,
'key': PGVECTOR_PGCRYPTO_KEY,
},

View file

@ -1,5 +1,7 @@
from datetime import datetime
from open_webui.utils.misc import sanitize_text_for_db
KEYS_TO_EXCLUDE = ['content', 'pages', 'tables', 'paragraphs', 'sections', 'figures']
@ -12,7 +14,8 @@ def filter_metadata(metadata: dict[str, any]) -> dict[str, any]:
def process_metadata(
metadata: dict[str, any],
) -> dict[str, any]:
# Removes large fields and converts non-serializable types (datetime, list, dict) to strings.
# Removes large fields, converts non-serializable types (datetime, list, dict) to strings,
# and sanitizes strings for database storage (strips null bytes and invalid surrogates).
result = {}
for key, value in metadata.items():
# Skip large fields
@ -20,7 +23,8 @@ def process_metadata(
continue
# Convert non-serializable fields to strings
if isinstance(value, (datetime, list, dict)):
result[key] = str(value)
result[key] = sanitize_text_for_db(str(value))
else:
result[key] = value
result[key] = sanitize_text_for_db(value)
return result

View file

@ -1,6 +1,7 @@
import logging
from typing import Optional, List
import requests
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
log = logging.getLogger(__name__)
@ -14,23 +15,38 @@ def search_firecrawl(
filter_list: Optional[List[str]] = None,
) -> List[SearchResult]:
try:
from firecrawl import FirecrawlApp
url = firecrawl_url.rstrip('/')
response = requests.post(
f'{url}/v1/search',
headers={
'Content-Type': 'application/json',
'Authorization': f'Bearer {firecrawl_api_key}',
},
json={
'query': query,
'limit': count,
'timeout': count * 3000,
},
timeout=count * 3 + 10,
)
response.raise_for_status()
data = response.json().get('data', {})
firecrawl = FirecrawlApp(api_key=firecrawl_api_key, api_url=firecrawl_url)
response = firecrawl.search(query=query, limit=count, ignore_invalid_urls=True, timeout=count * 3)
results = response.web
if filter_list:
results = get_filtered_results(results, filter_list)
results = [
SearchResult(
link=result.url,
title=result.title,
snippet=result.description,
link=r.get('url', ''),
title=r.get('title', ''),
snippet=r.get('description', ''),
)
for result in results[:count]
for r in data.get('web', [])
]
log.info(f'External search results: {results}')
if filter_list:
results = get_filtered_results(results, filter_list)
results = results[:count]
log.info(f'FireCrawl search results: {results}')
return results
except Exception as e:
log.error(f'Error in External search: {e}')
log.error(f'Error in FireCrawl search: {e}')
return []

View file

@ -1,4 +1,5 @@
import asyncio
import ipaddress
import logging
import socket
import ssl
@ -84,11 +85,9 @@ def validate_url(url: Union[str, Sequence[str]]):
ipv4_addresses, ipv6_addresses = resolve_hostname(parsed_url.hostname)
# Check if any of the resolved addresses are private
# This is technically still vulnerable to DNS rebinding attacks, as we don't control WebBaseLoader
for ip in ipv4_addresses:
if validators.ipv4(ip, private=True):
raise ValueError(ERROR_MESSAGES.INVALID_URL)
for ip in ipv6_addresses:
if validators.ipv6(ip, private=True):
for ip in ipv4_addresses + ipv6_addresses:
addr = ipaddress.ip_address(ip)
if not addr.is_global:
raise ValueError(ERROR_MESSAGES.INVALID_URL)
return True
elif isinstance(url, Sequence):
@ -193,27 +192,6 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
proxy: Optional[Dict[str, str]] = None,
params: Optional[Dict] = None,
):
"""Concurrent document loader for FireCrawl operations.
Executes multiple FireCrawlLoader instances concurrently using thread pooling
to improve bulk processing efficiency.
Args:
web_paths: List of URLs/paths to process.
verify_ssl: If True, verify SSL certificates.
trust_env: If True, use proxy settings from environment variables.
requests_per_second: Number of requests per second to limit to.
continue_on_failure (bool): If True, continue loading other URLs on failure.
api_key: API key for FireCrawl service. Defaults to None
(uses FIRE_CRAWL_API_KEY environment variable if not provided).
api_url: Base URL for FireCrawl API. Defaults to official API endpoint.
mode: Operation mode selection:
- 'crawl': Website crawling mode
- 'scrape': Direct page scraping (default)
- 'map': Site map generation
proxy: Proxy override settings for the FireCrawl API.
params: The parameters to pass to the Firecrawl API.
For more details, visit: https://docs.firecrawl.dev/sdks/python#batch-scrape
"""
proxy_server = proxy.get('server') if proxy else None
if trust_env and not proxy_server:
env_proxies = urllib.request.getproxies()
@ -230,44 +208,43 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
self.trust_env = trust_env
self.continue_on_failure = continue_on_failure
self.api_key = api_key
self.api_url = api_url
self.api_url = (api_url or 'https://api.firecrawl.dev').rstrip('/')
self.timeout = timeout
self.mode = mode
self.params = params or {}
def lazy_load(self) -> Iterator[Document]:
"""Load documents using FireCrawl batch_scrape."""
log.debug(
'Starting FireCrawl batch scrape for %d URLs, mode: %s, params: %s',
len(self.web_paths),
self.mode,
self.params,
)
try:
from firecrawl import FirecrawlApp
headers = {
'Content-Type': 'application/json',
'Authorization': f'Bearer {self.api_key}',
}
firecrawl = FirecrawlApp(api_key=self.api_key, api_url=self.api_url)
result = firecrawl.batch_scrape(
self.web_paths,
formats=['markdown'],
skip_tls_verification=not self.verify_ssl,
ignore_invalid_urls=True,
remove_base64_images=True,
max_age=300000, # 5 minutes https://docs.firecrawl.dev/features/fast-scraping#common-maxage-values
wait_timeout=self.timeout if self.timeout else len(self.web_paths) * 3,
**self.params,
)
for url in self.web_paths:
payload = {
'url': url,
'formats': ['markdown'],
**self.params,
}
if self.timeout:
payload['timeout'] = self.timeout * 1000
if result.status != 'completed':
raise RuntimeError(f'FireCrawl batch scrape did not complete successfully. result: {result}')
for data in result.data:
metadata = data.metadata or {}
yield Document(
page_content=data.markdown or '',
metadata={'source': metadata.url or metadata.source_url or ''},
response = requests.post(
f'{self.api_url}/v1/scrape',
headers=headers,
json=payload,
timeout=self.timeout or 60,
verify=self.verify_ssl,
)
response.raise_for_status()
data = response.json().get('data', {})
metadata = data.get('metadata', {})
source = metadata.get('url') or metadata.get('sourceURL') or url
yield Document(
page_content=data.get('markdown', ''),
metadata={'source': source},
)
except Exception as e:
if self.continue_on_failure:
log.exception(f'Error extracting content from URLs: {e}')
@ -275,38 +252,10 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
raise e
async def alazy_load(self):
"""Async version of lazy_load."""
log.debug(
'Starting FireCrawl batch scrape for %d URLs, mode: %s, params: %s',
len(self.web_paths),
self.mode,
self.params,
)
try:
from firecrawl import FirecrawlApp
firecrawl = FirecrawlApp(api_key=self.api_key, api_url=self.api_url)
result = firecrawl.batch_scrape(
self.web_paths,
formats=['markdown'],
skip_tls_verification=not self.verify_ssl,
ignore_invalid_urls=True,
remove_base64_images=True,
max_age=300000, # 5 minutes https://docs.firecrawl.dev/features/fast-scraping#common-maxage-values
wait_timeout=self.timeout if self.timeout else len(self.web_paths) * 3,
**self.params,
)
if result.status != 'completed':
raise RuntimeError(f'FireCrawl batch scrape did not complete successfully. result: {result}')
for data in result.data:
metadata = data.metadata or {}
yield Document(
page_content=data.markdown or '',
metadata={'source': metadata.url or metadata.source_url or ''},
)
docs = await run_in_threadpool(lambda: list(self.lazy_load()))
for doc in docs:
yield doc
except Exception as e:
if self.continue_on_failure:
log.exception(f'Error extracting content from URLs: {e}')

View file

@ -62,7 +62,9 @@ async def get_model_analytics(
db: AsyncSession = Depends(get_async_session),
):
"""Get message counts per model."""
counts = await ChatMessages.get_message_count_by_model(start_date=start_date, end_date=end_date, group_id=group_id, db=db)
counts = await ChatMessages.get_message_count_by_model(
start_date=start_date, end_date=end_date, group_id=group_id, db=db
)
models = [
ModelAnalyticsEntry(model_id=model_id, count=count)
for model_id, count in sorted(counts.items(), key=lambda x: -x[1])
@ -80,7 +82,9 @@ async def get_user_analytics(
db: AsyncSession = Depends(get_async_session),
):
"""Get message counts and token usage per user with user info."""
counts = await ChatMessages.get_message_count_by_user(start_date=start_date, end_date=end_date, group_id=group_id, db=db)
counts = await ChatMessages.get_message_count_by_user(
start_date=start_date, end_date=end_date, group_id=group_id, db=db
)
token_usage = await ChatMessages.get_token_usage_by_user(
start_date=start_date, end_date=end_date, group_id=group_id, db=db
)
@ -227,7 +231,9 @@ async def get_token_usage(
db: AsyncSession = Depends(get_async_session),
):
"""Get token usage aggregated by model."""
usage = await ChatMessages.get_token_usage_by_model(start_date=start_date, end_date=end_date, group_id=group_id, db=db)
usage = await ChatMessages.get_token_usage_by_model(
start_date=start_date, end_date=end_date, group_id=group_id, db=db
)
models = [
TokenUsageEntry(model_id=model_id, **data)

View file

@ -54,6 +54,7 @@ from open_webui.env import (
AIOHTTP_CLIENT_SESSION_SSL,
AIOHTTP_CLIENT_TIMEOUT,
AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST,
BYPASS_PYDUB_PREPROCESSING,
DEVICE_TYPE,
ENABLE_FORWARD_USER_INFO_HEADERS,
)
@ -330,7 +331,9 @@ async def speech(request: Request, user=Depends(get_verified_user)):
detail=ERROR_MESSAGES.NOT_FOUND,
)
if user.role != 'admin' and not await has_permission(user.id, 'chat.tts', request.app.state.config.USER_PERMISSIONS):
if user.role != 'admin' and not await has_permission(
user.id, 'chat.tts', request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
@ -630,6 +633,7 @@ async def speech(request: Request, user=Depends(get_verified_user)):
detail=detail,
)
def transcription_handler(request, file_path, metadata, user=None):
filename = os.path.basename(file_path)
file_dir = os.path.dirname(file_path)
@ -1095,24 +1099,28 @@ def transcription_handler(request, file_path, metadata, user=None):
def transcribe(request: Request, file_path: str, metadata: Optional[dict] = None, user=None):
log.info(f'transcribe: {file_path} {metadata}')
if is_audio_conversion_required(file_path):
file_path = convert_audio_to_mp3(file_path)
if BYPASS_PYDUB_PREPROCESSING:
log.info('Bypassing pydub preprocessing (BYPASS_PYDUB_PREPROCESSING=true)')
chunk_paths = [file_path]
else:
if is_audio_conversion_required(file_path):
file_path = convert_audio_to_mp3(file_path)
try:
file_path = compress_audio(file_path)
except Exception as e:
log.exception(e)
try:
file_path = compress_audio(file_path)
except Exception as e:
log.exception(e)
# Always produce a list of chunk paths (could be one entry if small)
try:
chunk_paths = split_audio(file_path, MAX_FILE_SIZE)
print(f'Chunk paths: {chunk_paths}')
except Exception as e:
log.exception(e)
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.DEFAULT(e),
)
# Always produce a list of chunk paths (could be one entry if small)
try:
chunk_paths = split_audio(file_path, MAX_FILE_SIZE)
print(f'Chunk paths: {chunk_paths}')
except Exception as e:
log.exception(e)
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.DEFAULT(e),
)
results = []
try:
@ -1214,7 +1222,9 @@ async def transcription(
language: Optional[str] = Form(None),
user=Depends(get_verified_user),
):
if user.role != 'admin' and not await has_permission(user.id, 'chat.stt', request.app.state.config.USER_PERMISSIONS):
if user.role != 'admin' and not await has_permission(
user.id, 'chat.stt', request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,

View file

@ -54,6 +54,7 @@ from open_webui.config import (
OAUTH_PROVIDERS,
OAUTH_MERGE_ACCOUNTS_BY_EMAIL,
)
from open_webui.utils.oauth import auth_manager_config
from pydantic import BaseModel
from open_webui.utils.misc import parse_duration, validate_email_format
@ -96,7 +97,9 @@ log = logging.getLogger(__name__)
signin_rate_limiter = RateLimiter(redis_client=get_redis_client(), limit=5 * 3, window=60 * 3)
async def create_session_response(request: Request, user, db, response: Response = None, set_cookie: bool = False) -> dict:
async def create_session_response(
request: Request, user, db, response: Response = None, set_cookie: bool = False
) -> dict:
"""
Create JWT token and build session response for a user.
Shared helper for signin, signup, ldap_auth, add_user, and token_exchange endpoints.
@ -322,6 +325,14 @@ async def ldap_auth(
detail=ERROR_MESSAGES.ACTION_PROHIBITED,
)
# Reject empty passwords before attempting the LDAP bind.
# Per RFC 4513 §5.1.2, a Simple Bind with a non-empty DN but empty
# password is "unauthenticated simple authentication" — many LDAP
# servers (OpenLDAP default, some AD configs) return success for these,
# which would grant access without valid credentials.
if not form_data.password or not form_data.password.strip():
raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
# NOW load LDAP config variables
LDAP_SERVER_LABEL = request.app.state.config.LDAP_SERVER_LABEL
LDAP_SERVER_HOST = request.app.state.config.LDAP_SERVER_HOST
@ -909,7 +920,9 @@ async def add_user(
@router.get('/admin/details')
async def get_admin_details(request: Request, user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)):
async def get_admin_details(
request: Request, user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)
):
if request.app.state.config.SHOW_ADMIN_DETAILS:
admin_email = request.app.state.config.ADMIN_EMAIL
admin_name = None
@ -1118,7 +1131,7 @@ async def update_ldap_server(request: Request, form_data: LdapServerConfig, user
for key in required_fields:
value = getattr(form_data, key)
if not value:
raise HTTPException(400, detail=f'Required field {key} is empty')
raise HTTPException(400, detail=ERROR_MESSAGES.REQUIRED_FIELD_EMPTY(key))
request.app.state.config.LDAP_SERVER_LABEL = form_data.label
request.app.state.config.LDAP_SERVER_HOST = form_data.host
@ -1173,7 +1186,9 @@ async def update_ldap_config(request: Request, form_data: LdapConfigForm, user=D
# create api key
@router.post('/api_key', response_model=ApiKey)
async def generate_api_key(request: Request, user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)):
async def generate_api_key(
request: Request, user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)
):
if not request.app.state.config.ENABLE_API_KEYS or (
user.role != 'admin'
and not await has_permission(user.id, 'features.api_keys', request.app.state.config.USER_PERMISSIONS)
@ -1245,7 +1260,7 @@ async def token_exchange(
if provider not in OAUTH_PROVIDERS:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Provider '{provider}' is not configured",
detail=ERROR_MESSAGES.OAUTH_NOT_CONFIGURED(provider),
)
# Get the OAuth client for this provider
oauth_manager = request.app.state.oauth_manager
@ -1253,7 +1268,7 @@ async def token_exchange(
if not client:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"OAuth client for '{provider}' not found",
detail=ERROR_MESSAGES.OAUTH_NOT_CONFIGURED(provider),
)
# Validate the token by calling the userinfo endpoint
@ -1295,6 +1310,17 @@ async def token_exchange(
)
email = email.lower()
# Enforce domain allowlist — same check as the normal OAuth callback
if (
'*' not in auth_manager_config.OAUTH_ALLOWED_DOMAINS
and email.split('@')[-1] not in auth_manager_config.OAUTH_ALLOWED_DOMAINS
):
log.warning(f'Token exchange denied: email domain not in allowed domains list')
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
# Try to find the user by OAuth sub
user = await Users.get_user_by_oauth_sub(provider, sub, db=db)

View file

@ -74,7 +74,7 @@ async def check_automation_limits(request, user, rrule_str: str, db, is_create:
if max_count > 0 and await Automations.count_by_user(user.id, db=db) >= max_count:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f'Automation limit reached ({max_count})',
detail=ERROR_MESSAGES.AUTOMATION_LIMIT_EXCEEDED(max_count),
)
# Min interval (create + update)
@ -86,7 +86,7 @@ async def check_automation_limits(request, user, rrule_str: str, db, is_create:
if interval is not None and interval < min_interval:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f'Schedule too frequent. Minimum interval is {min_interval} seconds.',
detail=ERROR_MESSAGES.AUTOMATION_TOO_FREQUENT(min_interval),
)
@ -167,15 +167,6 @@ async def create_new_automation(
await check_automation_limits(request, user, form_data.data.rrule, db, is_create=True)
# Validate terminal server exists if linked
if form_data.data.terminal and form_data.data.terminal.server_id:
connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
if not any(c.get('id') == form_data.data.terminal.server_id for c in connections):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail='Terminal server not found',
)
tz = user.timezone
automation = await Automations.insert(user.id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db)
return await enrich_automation(automation, db, tz=tz)
@ -226,15 +217,6 @@ async def update_automation_by_id(
await check_automation_limits(request, user, form_data.data.rrule, db, is_create=False)
# Validate terminal server exists if linked
if form_data.data.terminal and form_data.data.terminal.server_id:
connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
if not any(c.get('id') == form_data.data.terminal.server_id for c in connections):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail='Terminal server not found',
)
tz = user.timezone
updated = await Automations.update_by_id(id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db)
return await enrich_automation(updated, db, tz=tz)

View file

@ -61,7 +61,7 @@ from open_webui.utils.chat import generate_chat_completion
from open_webui.utils.auth import get_admin_user, get_verified_user
from open_webui.utils.access_control import has_permission
from open_webui.utils.access_control import has_permission, filter_allowed_access_grants
from open_webui.utils.webhook import post_webhook
from open_webui.utils.channels import extract_mentions, replace_mentions
from open_webui.internal.db import get_async_session
@ -94,7 +94,9 @@ async def channel_has_access(
return False
async def get_channel_users_with_access(channel: ChannelModel, permission: str = 'read', db: Optional[AsyncSession] = None):
async def get_channel_users_with_access(
channel: ChannelModel, permission: str = 'read', db: Optional[AsyncSession] = None
):
return await AccessGrants.get_users_with_access(
resource_type='channel',
resource_id=channel.id,
@ -138,7 +140,7 @@ async def check_channels_access(request: Request, user: Optional[UserModel] = No
if not request.app.state.config.ENABLE_CHANNELS:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail='Channels are not enabled',
detail=ERROR_MESSAGES.FEATURE_DISABLED('Channels'),
)
if user:
@ -303,6 +305,14 @@ async def create_new_channel(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
form_data.access_grants = filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
form_data.access_grants,
'sharing.public_channels',
)
try:
if form_data.type == 'dm':
existing_channel = await Channels.get_dm_channel_by_user_ids([user.id, *form_data.user_ids], db=db)
@ -633,6 +643,14 @@ async def update_channel_by_id(
if channel.user_id != user.id and user.role != 'admin':
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
form_data.access_grants = filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
form_data.access_grants,
'sharing.public_channels',
)
try:
channel = await Channels.update_channel_by_id(id, form_data, db=db)
return ChannelModel(**channel.model_dump())
@ -877,11 +895,13 @@ async def model_response_handler(request, channel, message, user, db=None):
if model:
try:
# reverse to get in chronological order
thread_messages = (await Messages.get_messages_by_parent_id(
channel.id,
message.parent_id if message.parent_id else message.id,
db=db,
))[::-1]
thread_messages = (
await Messages.get_messages_by_parent_id(
channel.id,
message.parent_id if message.parent_id else message.id,
db=db,
)
)[::-1]
response_message, channel = await new_message_handler(
request,
@ -1104,7 +1124,9 @@ async def post_new_message(
try:
if files := message.data.get('files', []):
for file in files:
await Channels.set_file_message_id_in_channel_by_id(channel.id, file.get('id', ''), message.id, db=db)
await Channels.set_file_message_id_in_channel_by_id(
channel.id, file.get('id', ''), message.id, db=db
)
except Exception as e:
log.debug(e)
@ -1769,7 +1791,7 @@ async def post_webhook_message(
if not webhook:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail='Invalid webhook URL',
detail=ERROR_MESSAGES.INVALID_URL,
)
channel = await Channels.get_channel_by_id(webhook.channel_id, db=db)
@ -1787,7 +1809,7 @@ async def post_webhook_message(
if not message:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail='Failed to create message',
detail=ERROR_MESSAGES.DEFAULT('Failed to create message'),
)
# Update last_used_at

View file

@ -1,6 +1,7 @@
import json
import logging
from typing import Optional
from uuid import uuid4
from sqlalchemy.ext.asyncio import AsyncSession
import asyncio
from fastapi.responses import StreamingResponse
@ -493,7 +494,9 @@ async def delete_all_user_chats(
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
if user.role == 'user' and not await has_permission(user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS):
if user.role == 'user' and not await has_permission(
user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
@ -538,7 +541,9 @@ async def get_user_chat_list_by_user_id(
if direction:
filter['direction'] = direction
return await Chats.get_chat_list_by_user_id(user_id, include_archived=True, filter=filter, skip=skip, limit=limit, db=db)
return await Chats.get_chat_list_by_user_id(
user_id, include_archived=True, filter=filter, skip=skip, limit=limit, db=db
)
############################
@ -553,7 +558,7 @@ async def create_new_chat(
db: AsyncSession = Depends(get_async_session),
):
try:
chat = await Chats.insert_new_chat(user.id, form_data, db=db)
chat = await Chats.insert_new_chat(str(uuid4()), user.id, form_data, db=db)
return ChatResponse(**chat.model_dump())
except Exception as e:
log.exception(e)
@ -620,7 +625,9 @@ async def search_user_chats(
@router.get('/folder/{folder_id}', response_model=list[ChatResponse])
async def get_chats_by_folder_id(folder_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_chats_by_folder_id(
folder_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
folder_ids = [folder_id]
children_folders = await Folders.get_children_folders_by_id_and_user_id(folder_id, user.id, db=db)
if children_folders:
@ -815,7 +822,9 @@ async def get_shared_session_user_chat_list(
@router.get('/share/{share_id}', response_model=Optional[ChatResponse])
async def get_shared_chat_by_id(share_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_shared_chat_by_id(
share_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
if user.role == 'pending':
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND)
@ -851,7 +860,9 @@ async def get_user_chat_list_by_tag_name(
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
chats = await Chats.get_chat_list_by_user_id_and_tag_name(user.id, form_data.name, form_data.skip, form_data.limit, db=db)
chats = await Chats.get_chat_list_by_user_id_and_tag_name(
user.id, form_data.name, form_data.skip, form_data.limit, db=db
)
if len(chats) == 0:
await Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db)
@ -1056,7 +1067,9 @@ async def delete_chat_by_id(
@router.get('/{id}/pinned', response_model=Optional[bool])
async def get_pinned_status_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_pinned_status_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
return chat.pinned
@ -1137,7 +1150,9 @@ async def clone_chat_by_id(
@router.post('/{id}/clone/shared', response_model=Optional[ChatResponse])
async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def clone_shared_chat_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
if user.role == 'admin':
chat = await Chats.get_chat_by_id(id, db=db)
else:
@ -1250,7 +1265,9 @@ async def share_chat_by_id(
@router.delete('/{id}/share', response_model=Optional[bool])
async def delete_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def delete_shared_chat_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
if not chat.share_id:
@ -1371,7 +1388,9 @@ async def delete_tag_by_id_and_tag_name(
@router.delete('/{id}/tags/all', response_model=Optional[bool])
async def delete_all_tags_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def delete_all_tags_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
old_tags = chat.meta.get('tags', [])

View file

@ -8,7 +8,7 @@ from typing import Optional
from open_webui.env import AIOHTTP_CLIENT_TIMEOUT
from open_webui.utils.auth import get_admin_user, get_verified_user
from open_webui.config import get_config, save_config
from open_webui.config import get_config, save_config, async_save_config
from open_webui.config import BannerModel
from open_webui.utils.tools import (
@ -27,6 +27,7 @@ from open_webui.utils.oauth import (
get_oauth_client_info_with_static_credentials,
encrypt_data,
decrypt_data,
resolve_oauth_client_info,
OAuthClientInformationFull,
)
from mcp.shared.auth import OAuthMetadata
@ -49,7 +50,7 @@ class ImportConfigForm(BaseModel):
@router.post('/import', response_model=dict)
async def import_config(form_data: ImportConfigForm, user=Depends(get_admin_user)):
save_config(form_data.config)
await async_save_config(form_data.config)
return get_config()
@ -203,9 +204,7 @@ async def set_tool_servers_config(
if auth_type in ('oauth_2.1', 'oauth_2.1_static') and server_id:
try:
oauth_client_info = connection.get('info', {}).get('oauth_client_info', '')
oauth_client_info = decrypt_data(oauth_client_info)
oauth_client_info = resolve_oauth_client_info(connection)
request.app.state.oauth_client_manager.add_client(
f'{server_type}:{server_id}',
OAuthClientInformationFull(**oauth_client_info),

View file

@ -415,7 +415,9 @@ async def update_feedback_by_id(
@router.delete('/feedback/{id}')
async def delete_feedback_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def delete_feedback_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
if user.role == 'admin':
success = await Feedbacks.delete_feedback_by_id(id=id, db=db)
else:

View file

@ -22,7 +22,7 @@ from fastapi import (
from fastapi.responses import FileResponse, StreamingResponse
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import get_async_session, SessionLocal
from open_webui.internal.db import get_async_session, get_async_db_context
from open_webui.constants import ERROR_MESSAGES
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
@ -48,7 +48,7 @@ from open_webui.routers.audio import transcribe
from open_webui.storage.provider import Storage
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL, STORAGE_LOCAL_CACHE, STORAGE_PROVIDER, UPLOAD_DIR
from open_webui.utils.auth import get_admin_user, get_verified_user
from open_webui.utils.misc import strict_match_mime_type
from pydantic import BaseModel
@ -88,6 +88,20 @@ def _is_text_file(file_path: str, chunk_size: int = 8192) -> bool:
return False
def _cleanup_local_cache(file_path: str) -> None:
"""Remove the local cached copy of a cloud-stored file after processing."""
if STORAGE_LOCAL_CACHE or STORAGE_PROVIDER == 'local':
return
try:
local_filename = os.path.basename(file_path)
local_path = os.path.join(UPLOAD_DIR, local_filename)
if os.path.isfile(local_path):
os.remove(local_path)
log.debug(f'Cleaned up local cache: {local_path}')
except OSError as e:
log.warning(f'Failed to clean up local cache for {file_path}: {e}')
async def process_uploaded_file(
request,
file,
@ -113,7 +127,7 @@ async def process_uploaded_file(
file_path_processed = Storage.get_file(file_path)
result = transcribe(request, file_path_processed, file_metadata, user)
process_file(
await process_file(
request,
ProcessFileForm(file_id=file_item.id, content=result.get('text', '')),
user=user,
@ -122,7 +136,7 @@ async def process_uploaded_file(
elif (not content_type.startswith(('image/', 'video/'))) or (
request.app.state.config.CONTENT_EXTRACTION_ENGINE == 'external'
):
process_file(
await process_file(
request,
ProcessFileForm(file_id=file_item.id),
user=user,
@ -132,7 +146,7 @@ async def process_uploaded_file(
raise Exception(f'File type {content_type} is not supported for processing')
else:
log.info(f'File type {file.content_type} is not provided, but trying to process anyway')
process_file(
await process_file(
request,
ProcessFileForm(file_id=file_item.id),
user=user,
@ -150,11 +164,14 @@ async def process_uploaded_file(
db=db_session,
)
if db:
_process_handler(db)
else:
with SessionLocal() as db_session:
_process_handler(db_session)
try:
if db:
await _process_handler(db)
else:
async with get_async_db_context() as db_session:
await _process_handler(db_session)
finally:
_cleanup_local_cache(file_path)
@router.post('/', response_model=FileModelResponse)
@ -495,7 +512,9 @@ async def get_file_process_status(
@router.get('/{id}/data/content')
async def get_file_data_content_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_file_data_content_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
file = await Files.get_file_by_id(id, db=db)
if not file:
@ -540,7 +559,7 @@ async def update_file_data_content_by_id(
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'write', user, db=db):
try:
process_file(
await process_file(
request,
ProcessFileForm(file_id=id, content=form_data.content),
user=user,
@ -560,7 +579,7 @@ async def update_file_data_content_by_id(
# Remove old embeddings for this file from the KB collection
VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'file_id': id})
# Re-add from the now-updated file-{file_id} collection
process_file(
await process_file(
request,
ProcessFileForm(file_id=id, collection_name=knowledge.id),
user=user,
@ -646,7 +665,9 @@ async def get_file_content_by_id(
@router.get('/{id}/content/html')
async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_html_file_content_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
file = await Files.get_file_by_id(id, db=db)
if not file:
@ -693,7 +714,9 @@ async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user),
@router.get('/{id}/content/{file_name}')
async def get_file_content_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_file_content_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
file = await Files.get_file_by_id(id, db=db)
if not file:

View file

@ -89,7 +89,9 @@ async def get_folders(
valid_files.append(file)
folder.data['files'] = valid_files
await Folders.update_folder_by_id_and_user_id(folder.id, user.id, FolderUpdateForm(data=folder.data), db=db)
await Folders.update_folder_by_id_and_user_id(
folder.id, user.id, FolderUpdateForm(data=folder.data), db=db
)
folder_list.append(FolderNameIdResponse(**folder.model_dump()))
@ -107,7 +109,9 @@ async def create_folder(
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
folder = await Folders.get_folder_by_parent_id_and_user_id_and_name(form_data.parent_id, user.id, form_data.name, db=db)
folder = await Folders.get_folder_by_parent_id_and_user_id_and_name(
form_data.parent_id, user.id, form_data.name, db=db
)
if folder:
raise HTTPException(
@ -250,7 +254,9 @@ async def update_folder_is_expanded_by_id(
folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
if folder:
try:
folder = await Folders.update_folder_is_expanded_by_id_and_user_id(id, user.id, form_data.is_expanded, db=db)
folder = await Folders.update_folder_is_expanded_by_id_and_user_id(
id, user.id, form_data.is_expanded, db=db
)
return folder
except Exception as e:
log.exception(e)

View file

@ -128,7 +128,7 @@ async def load_function_from_url(request: Request, form_data: LoadUrlForm, user=
'content': data,
}
except Exception as e:
raise HTTPException(status_code=500, detail=f'Error importing function: {e}')
raise HTTPException(status_code=500, detail=ERROR_MESSAGES.DEFAULT(e))
############################
@ -373,7 +373,9 @@ async def delete_function_by_id(
@router.get('/id/{id}/valves', response_model=Optional[dict])
async def get_function_valves_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
async def get_function_valves_by_id(
id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)
):
function = await Functions.get_function_by_id(id, db=db)
if function:
try:
@ -473,7 +475,9 @@ async def update_function_valves_by_id(
@router.get('/id/{id}/valves/user', response_model=Optional[dict])
async def get_function_user_valves_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_function_user_valves_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
function = await Functions.get_function_by_id(id, db=db)
if function:
try:

View file

@ -10,6 +10,7 @@ from pathlib import Path
from typing import Optional
from urllib.parse import quote
import aiohttp
import requests
from fastapi import APIRouter, Depends, HTTPException, Request, UploadFile
from fastapi.responses import FileResponse
@ -21,7 +22,8 @@ from open_webui.config import (
)
from open_webui.constants import ERROR_MESSAGES
from open_webui.retrieval.web.utils import validate_url
from open_webui.env import ENABLE_FORWARD_USER_INFO_HEADERS
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, ENABLE_FORWARD_USER_INFO_HEADERS
from open_webui.utils.session_pool import get_session
from open_webui.models.chats import Chats
from open_webui.routers.files import upload_file_handler, get_file_content_by_id
@ -313,12 +315,14 @@ def get_automatic1111_api_auth(request: Request):
async def verify_url(request: Request, user=Depends(get_admin_user)):
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'automatic1111':
try:
r = requests.get(
session = await get_session()
async with session.get(
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
headers={'authorization': get_automatic1111_api_auth(request)},
)
r.raise_for_status()
return True
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
return True
except Exception:
request.app.state.config.ENABLE_IMAGE_GENERATION = False
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.INVALID_URL)
@ -327,12 +331,14 @@ async def verify_url(request: Request, user=Depends(get_admin_user)):
if request.app.state.config.COMFYUI_API_KEY:
headers = {'Authorization': f'Bearer {request.app.state.config.COMFYUI_API_KEY}'}
try:
r = requests.get(
session = await get_session()
async with session.get(
url=f'{request.app.state.config.COMFYUI_BASE_URL}/object_info',
headers=headers,
)
r.raise_for_status()
return True
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
return True
except Exception:
request.app.state.config.ENABLE_IMAGE_GENERATION = False
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.INVALID_URL)
@ -357,11 +363,13 @@ async def get_models(request: Request, user=Depends(get_verified_user)):
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'comfyui':
# TODO - get models from comfyui
headers = {'Authorization': f'Bearer {request.app.state.config.COMFYUI_API_KEY}'}
r = requests.get(
session = await get_session()
async with session.get(
url=f'{request.app.state.config.COMFYUI_BASE_URL}/object_info',
headers=headers,
)
info = r.json()
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
info = await r.json()
workflow = json.loads(request.app.state.config.COMFYUI_WORKFLOW)
model_node_id = None
@ -399,11 +407,13 @@ async def get_models(request: Request, user=Depends(get_verified_user)):
request.app.state.config.IMAGE_GENERATION_ENGINE == 'automatic1111'
or request.app.state.config.IMAGE_GENERATION_ENGINE == ''
):
r = requests.get(
session = await get_session()
async with session.get(
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/sd-models',
headers={'authorization': get_automatic1111_api_auth(request)},
)
models = r.json()
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
models = await r.json()
return list(
map(
lambda model: {'id': model['title'], 'name': model['model_name']},
@ -533,7 +543,7 @@ async def image_generations(
model = get_image_model(request)
r = None
try:
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'openai':
headers = {
@ -552,7 +562,11 @@ async def image_generations(
'model': model,
'prompt': form_data.prompt,
'n': form_data.n,
'size': (form_data.size if form_data.size else request.app.state.config.IMAGE_SIZE),
**(
{'size': form_data.size or request.app.state.config.IMAGE_SIZE}
if (form_data.size or request.app.state.config.IMAGE_SIZE)
else {}
),
**(
{}
if re.match(
@ -568,16 +582,15 @@ async def image_generations(
),
}
# Use asyncio.to_thread for the requests.post call
r = await asyncio.to_thread(
requests.post,
session = await get_session()
async with session.post(
url=url,
json=data,
headers=headers,
)
r.raise_for_status()
res = r.json()
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
res = await r.json()
images = []
@ -619,16 +632,15 @@ async def image_generations(
model = f'{model}:generateContent'
data = {'contents': [{'parts': [{'text': form_data.prompt}]}]}
# Use asyncio.to_thread for the requests.post call
r = await asyncio.to_thread(
requests.post,
session = await get_session()
async with session.post(
url=f'{request.app.state.config.IMAGES_GEMINI_API_BASE_URL}/models/{model}',
json=data,
headers=headers,
)
r.raise_for_status()
res = r.json()
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
res = await r.json()
images = []
@ -727,15 +739,14 @@ async def image_generations(
if request.app.state.config.AUTOMATIC1111_PARAMS:
data = {**data, **request.app.state.config.AUTOMATIC1111_PARAMS}
# Use asyncio.to_thread for the requests.post call
r = await asyncio.to_thread(
requests.post,
session = await get_session()
async with session.post(
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/txt2img',
json=data,
headers={'authorization': get_automatic1111_api_auth(request)},
)
res = r.json()
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
res = await r.json()
log.debug(f'res: {res}')
images = []
@ -753,10 +764,8 @@ async def image_generations(
return images
except Exception as e:
error = e
if r != None:
data = r.json()
if 'error' in data:
error = data['error']['message']
if isinstance(e, aiohttp.ClientResponseError):
error = e.message
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(error))
@ -798,11 +807,12 @@ async def image_edits(
if data.startswith('http://') or data.startswith('https://'):
# Validate URL to prevent SSRF attacks against local/private networks
validate_url(data)
r = await asyncio.to_thread(requests.get, data)
r.raise_for_status()
session = await get_session()
async with session.get(data, ssl=AIOHTTP_CLIENT_SESSION_SSL) as r:
r.raise_for_status()
image_data = base64.b64encode(r.content).decode('utf-8')
return f'data:{r.headers["content-type"]};base64,{image_data}'
image_data = base64.b64encode(await r.read()).decode('utf-8')
return f'data:{r.headers["content-type"]};base64,{image_data}'
else:
file_id = None
@ -832,7 +842,7 @@ async def image_edits(
except Exception as e:
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(e))
async def get_image_file_item(base64_string, param_name='image'):
def get_image_file_item(base64_string, param_name='image'):
data = base64_string
header, encoded = data.split(',', 1)
mime_type = header.split(';')[0].lstrip('data:')
@ -846,7 +856,7 @@ async def image_edits(
),
)
r = None
try:
if request.app.state.config.IMAGE_EDIT_ENGINE == 'openai':
headers = {
@ -883,17 +893,30 @@ async def image_edits(
if request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION:
url_search_params += f'?api-version={request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION}'
# Use asyncio.to_thread for the requests.post call
r = await asyncio.to_thread(
requests.post,
# Build multipart form data for aiohttp
form = aiohttp.FormData()
for key, value in data.items():
if isinstance(value, dict):
form.add_field(key, json.dumps(value))
else:
form.add_field(key, str(value))
for param_name, (filename, file_obj, content_type_val) in files:
form.add_field(
param_name,
file_obj,
filename=filename,
content_type=content_type_val,
)
session = await get_session()
async with session.post(
url=f'{request.app.state.config.IMAGES_EDIT_OPENAI_API_BASE_URL}/images/edits{url_search_params}',
headers=headers,
files=files,
data=data,
)
r.raise_for_status()
res = r.json()
data=form,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
res = await r.json()
images = []
for image in res['data']:
@ -940,16 +963,15 @@ async def image_edits(
]
)
# Use asyncio.to_thread for the requests.post call
r = await asyncio.to_thread(
requests.post,
session = await get_session()
async with session.post(
url=f'{request.app.state.config.IMAGES_EDIT_GEMINI_API_BASE_URL}/models/{model}',
json=data,
headers=headers,
)
r.raise_for_status()
res = r.json()
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
res = await r.json()
images = []
for image in res['candidates']:
@ -1048,13 +1070,7 @@ async def image_edits(
return images
except Exception as e:
error = e
if r != None:
data = r.text
try:
data = json.loads(data)
if 'error' in data:
error = data['error']['message']
except Exception:
error = data
if isinstance(e, aiohttp.ClientResponseError):
error = e.message
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(error))

View file

@ -2,7 +2,7 @@ from typing import List, Optional
from pydantic import BaseModel
from fastapi import APIRouter, Depends, HTTPException, status, Request, Query
from fastapi.responses import StreamingResponse
from fastapi.concurrency import run_in_threadpool
import logging
import io
import zipfile
@ -319,8 +319,7 @@ async def reindex_knowledge_files(
failed_files = []
for file in files:
try:
await run_in_threadpool(
process_file,
await process_file(
request,
ProcessFileForm(file_id=file.id, collection_name=knowledge_base.id),
user=user,
@ -543,7 +542,7 @@ async def update_knowledge_access_by_id(
await AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db)
return KnowledgeFilesResponse(
**await Knowledges.get_knowledge_by_id(id=id, db=db).model_dump(),
**(await Knowledges.get_knowledge_by_id(id=id, db=db)).model_dump(),
files=await Knowledges.get_file_metadatas_by_id(id, db=db),
)
@ -659,7 +658,7 @@ async def add_file_to_knowledge_by_id(
# Add content to the vector database
try:
process_file(
await process_file(
request,
ProcessFileForm(file_id=form_data.file_id, collection_name=id),
user=user,
@ -737,7 +736,7 @@ async def update_file_from_knowledge_by_id(
# Add content to the vector database
try:
process_file(
await process_file(
request,
ProcessFileForm(file_id=form_data.file_id, collection_name=id),
user=user,
@ -859,7 +858,9 @@ async def remove_file_from_knowledge_by_id(
@router.delete('/{id}/delete', response_model=bool)
async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def delete_knowledge_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
@ -932,7 +933,9 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: A
@router.post('/{id}/reset', response_model=Optional[KnowledgeResponse])
async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def reset_knowledge_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
@ -962,7 +965,7 @@ async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: As
log.debug(e)
pass
knowledge = Knowledges.reset_knowledge_by_id(id=id, db=db)
knowledge = await Knowledges.reset_knowledge_by_id(id=id, db=db)
return knowledge

View file

@ -268,7 +268,7 @@ async def update_memory_by_id(
memory = await Memories.update_memory_by_id_and_user_id(memory_id, user.id, form_data.content)
if memory is None:
raise HTTPException(status_code=404, detail='Memory not found')
raise HTTPException(status_code=404, detail=ERROR_MESSAGES.NOT_FOUND)
if form_data.content is not None:
vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)

View file

@ -50,7 +50,10 @@ from open_webui.models.groups import Groups
from open_webui.utils.access_control import check_model_access
from open_webui.utils.misc import (
calculate_sha256,
)
from open_webui.utils.session_pool import (
cleanup_response,
get_session,
stream_wrapper,
)
from open_webui.utils.payload import (
@ -122,10 +125,7 @@ async def send_request(
r = None
streaming = False
try:
session = aiohttp.ClientSession(
trust_env=True,
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
session = await get_session()
headers = {
'Content-Type': 'application/json',
@ -138,8 +138,12 @@ async def send_request(
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id')
r = await session.request(
method, url, data=payload, headers=headers,
method,
url,
data=payload,
headers=headers,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
if not r.ok:
@ -153,7 +157,7 @@ async def send_request(
log.error(f'Failed to parse error response: {e}')
raise HTTPException(
status_code=r.status,
detail='Open WebUI: Server Connection Error',
detail=ERROR_MESSAGES.SERVER_CONNECTION_ERROR,
)
r.raise_for_status()
@ -165,7 +169,7 @@ async def send_request(
streaming = True
return StreamingResponse(
stream_wrapper(r, session),
stream_wrapper(r),
status_code=r.status,
headers=response_headers,
)
@ -180,11 +184,11 @@ async def send_request(
except Exception as e:
raise HTTPException(
status_code=r.status if r else 500,
detail=f'Ollama: {e}' if str(e) else 'Open WebUI: Server Connection Error',
detail=f'Ollama: {e}' if str(e) else ERROR_MESSAGES.SERVER_CONNECTION_ERROR,
)
finally:
if not streaming:
await cleanup_response(r, session)
await cleanup_response(r)
def get_api_key(idx, url, configs):
@ -247,7 +251,7 @@ async def verify_connection(form_data: ConnectionVerificationForm, user=Depends(
return data
except aiohttp.ClientError as e:
log.exception(f'Client error: {str(e)}')
raise HTTPException(status_code=500, detail='Open WebUI: Server Connection Error')
raise HTTPException(status_code=500, detail=ERROR_MESSAGES.SERVER_CONNECTION_ERROR)
except Exception as e:
log.exception(f'Unexpected error: {e}')
error_detail = f'Unexpected error: {str(e)}'
@ -424,7 +428,7 @@ async def get_filtered_models(models, user, db=None):
@router.get('/api/tags/{url_idx}')
async def get_ollama_tags(request: Request, url_idx: Optional[int] = None, user=Depends(get_verified_user)):
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
models = []
@ -617,7 +621,7 @@ async def pull_model(
user=Depends(get_admin_user),
):
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
form_data = form_data.model_dump(exclude_none=True)
form_data['model'] = form_data.get('model', form_data.get('name'))
@ -652,7 +656,7 @@ async def push_model(
user=Depends(get_admin_user),
):
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
if url_idx is None:
await get_all_models(request, user=user)
@ -695,7 +699,7 @@ async def create_model(
user=Depends(get_admin_user),
):
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
log.debug(f'form_data: {form_data}')
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
@ -723,7 +727,7 @@ async def copy_model(
user=Depends(get_admin_user),
):
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
if url_idx is None:
await get_all_models(request, user=user)
@ -758,7 +762,7 @@ async def delete_model(
user=Depends(get_admin_user),
):
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
form_data = form_data.model_dump(exclude_none=True)
form_data['model'] = form_data.get('model', form_data.get('name'))
@ -781,7 +785,8 @@ async def delete_model(
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
await send_request(
f'{url}/api/delete', 'DELETE',
f'{url}/api/delete',
'DELETE',
payload=json.dumps(form_data),
key=key,
user=user,
@ -792,7 +797,7 @@ async def delete_model(
@router.post('/api/show')
async def show_model_info(request: Request, form_data: ModelNameForm, user=Depends(get_verified_user)):
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
form_data = form_data.model_dump(exclude_none=True)
form_data['model'] = form_data.get('model', form_data.get('name'))
@ -845,7 +850,7 @@ async def embed(
user=Depends(get_verified_user),
):
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
log.info(f'generate_ollama_batch_embeddings {form_data}')
@ -904,7 +909,7 @@ async def embeddings(
user=Depends(get_verified_user),
):
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
log.info(f'generate_ollama_embeddings {form_data}')
@ -971,7 +976,7 @@ async def generate_completion(
user=Depends(get_verified_user),
):
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
# Enforce per-model access control
await check_model_access(user, await Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL)
@ -1062,7 +1067,7 @@ async def generate_chat_completion(
bypass_system_prompt: bool = False,
):
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
# NOTE: We intentionally do NOT use Depends(get_async_session) here.
# Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions.
@ -1308,7 +1313,7 @@ async def generate_anthropic_messages(
See https://docs.ollama.com/api/anthropic-compatibility
"""
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
payload = {**form_data}
model_id = payload.get('model', '')
@ -1366,7 +1371,7 @@ async def generate_responses(
See https://ollama.com/blog/responses-api
"""
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
payload = form_data.model_dump()
model_id = form_data.model
@ -1391,13 +1396,13 @@ async def generate_responses(
):
raise HTTPException(
status_code=403,
detail='Model not found',
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
else:
if user.role != 'admin':
raise HTTPException(
status_code=403,
detail='Model not found',
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
url, url_idx = await get_ollama_url(request, payload['model'], url_idx)
@ -1649,9 +1654,7 @@ async def upload_model(
url = f'{ollama_url}/api/blobs/sha256:{file_hash}'
upload_timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT)
async with aiohttp.ClientSession(timeout=upload_timeout, trust_env=True) as upload_session:
async with upload_session.post(
url, data=blob_data, ssl=AIOHTTP_CLIENT_SESSION_SSL
) as response:
async with upload_session.post(url, data=blob_data, ssl=AIOHTTP_CLIENT_SESSION_SSL) as response:
if not response.ok:
raise Exception('Ollama: Could not create blob, Please try again.')

View file

@ -8,7 +8,7 @@ from urllib.parse import quote, urlparse
import aiohttp
from aiocache import cached
import requests
from azure.identity import DefaultAzureCredential, get_bearer_token_provider
@ -52,9 +52,12 @@ from open_webui.utils.payload import (
apply_system_prompt_to_body,
)
from open_webui.utils.misc import (
cleanup_response,
convert_logit_bias_input_to_json,
stream_chunks_handler,
)
from open_webui.utils.session_pool import (
cleanup_response,
get_session,
stream_wrapper,
)
@ -309,19 +312,20 @@ async def speech(request: Request, user=Depends(get_verified_user)):
r = None
try:
r = requests.post(
session = await get_session()
r = await session.post(
url=f'{url}/audio/speech',
data=body,
headers=headers,
cookies=cookies,
stream=True,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
)
r.raise_for_status()
# Save the streaming content to a file
with open(file_path, 'wb') as f:
for chunk in r.iter_content(chunk_size=8192):
async for chunk in r.content.iter_chunked(8192):
f.write(chunk)
with open(file_body_path, 'w') as f:
@ -336,14 +340,14 @@ async def speech(request: Request, user=Depends(get_verified_user)):
detail = None
if r is not None:
try:
res = r.json()
res = await r.json()
if 'error' in res:
detail = f'External: {res["error"]}'
except Exception:
detail = f'External: {e}'
raise HTTPException(
status_code=r.status_code if r else 500,
status_code=r.status if r else 500,
detail=detail if detail else 'Open WebUI: Server Connection Error',
)
@ -688,7 +692,7 @@ async def verify_connection(
elif is_anthropic_url(url):
result = await get_anthropic_models(url, key)
if result is None:
raise HTTPException(status_code=500, detail='Failed to connect to Anthropic API')
raise HTTPException(status_code=500, detail=ERROR_MESSAGES.SERVER_CONNECTION_ERROR)
if 'error' in result:
raise HTTPException(status_code=500, detail=result['error'])
return result
@ -715,10 +719,10 @@ async def verify_connection(
except aiohttp.ClientError as e:
# ClientError covers all aiohttp requests issues
log.exception(f'Client error: {str(e)}')
raise HTTPException(status_code=500, detail='Open WebUI: Server Connection Error')
raise HTTPException(status_code=500, detail=ERROR_MESSAGES.SERVER_CONNECTION_ERROR)
except Exception as e:
log.exception(f'Unexpected error: {e}')
raise HTTPException(status_code=500, detail='Open WebUI: Server Connection Error')
raise HTTPException(status_code=500, detail=ERROR_MESSAGES.SERVER_CONNECTION_ERROR)
def get_azure_allowed_params(api_version: str) -> set[str]:
@ -1080,7 +1084,7 @@ async def generate_chat_completion(
else:
raise HTTPException(
status_code=404,
detail='Model not found',
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Get the API config for the model
@ -1174,12 +1178,11 @@ async def generate_chat_completion(
payload = json.dumps(payload)
r = None
session = None
streaming = False
response = None
try:
session = aiohttp.ClientSession(trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT))
session = await get_session()
r = await session.request(
method='POST',
@ -1188,13 +1191,33 @@ async def generate_chat_completion(
headers=headers,
cookies=cookies,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
# Check if response is SSE
if 'text/event-stream' in r.headers.get('Content-Type', ''):
# If the provider returned an error status with SSE content-type,
# read the body and return a proper error response instead of
# streaming the error back (which hides the error from logs).
if r.status >= 400:
error_body = await r.text()
log.error(
'Provider returned HTTP %d with SSE content-type: %s',
r.status,
error_body[:1000],
)
try:
error_json = json.loads(error_body)
return JSONResponse(status_code=r.status, content=error_json)
except json.JSONDecodeError:
return JSONResponse(
status_code=r.status,
content={'error': {'message': error_body, 'code': r.status}},
)
streaming = True
return StreamingResponse(
stream_wrapper(r, session, stream_chunks_handler),
stream_wrapper(r, content_handler=stream_chunks_handler),
status_code=r.status,
headers=dict(r.headers),
)
@ -1221,11 +1244,11 @@ async def generate_chat_completion(
raise HTTPException(
status_code=r.status if r else 500,
detail='Open WebUI: Server Connection Error',
detail=ERROR_MESSAGES.SERVER_CONNECTION_ERROR,
)
finally:
if not streaming:
await cleanup_response(r, session)
await cleanup_response(r)
async def embeddings(request: Request, form_data: dict, user):
@ -1261,27 +1284,24 @@ async def embeddings(request: Request, form_data: dict, user):
)
r = None
session = None
streaming = False
headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user)
try:
session = aiohttp.ClientSession(
trust_env=True,
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
session = await get_session()
r = await session.request(
method='POST',
url=f'{url}/embeddings',
data=body,
headers=headers,
cookies=cookies,
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
if 'text/event-stream' in r.headers.get('Content-Type', ''):
streaming = True
return StreamingResponse(
stream_wrapper(r, session),
stream_wrapper(r),
status_code=r.status,
headers=dict(r.headers),
)
@ -1302,11 +1322,11 @@ async def embeddings(request: Request, form_data: dict, user):
log.exception(e)
raise HTTPException(
status_code=r.status if r else 500,
detail='Open WebUI: Server Connection Error',
detail=ERROR_MESSAGES.SERVER_CONNECTION_ERROR,
)
finally:
if not streaming:
await cleanup_response(r, session)
await cleanup_response(r)
class ResponsesForm(BaseModel):
@ -1365,7 +1385,6 @@ async def responses(
)
r = None
session = None
streaming = False
try:
@ -1388,10 +1407,7 @@ async def responses(
else:
request_url = f'{url}/responses'
session = aiohttp.ClientSession(
trust_env=True,
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
session = await get_session()
r = await session.request(
method='POST',
url=request_url,
@ -1399,13 +1415,14 @@ async def responses(
headers=headers,
cookies=cookies,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
# Check if response is SSE
if 'text/event-stream' in r.headers.get('Content-Type', ''):
streaming = True
return StreamingResponse(
stream_wrapper(r, session),
stream_wrapper(r),
status_code=r.status,
headers=dict(r.headers),
)
@ -1429,11 +1446,11 @@ async def responses(
log.exception(e)
raise HTTPException(
status_code=r.status if r else 500,
detail='Open WebUI: Server Connection Error',
detail=ERROR_MESSAGES.SERVER_CONNECTION_ERROR,
)
finally:
if not streaming:
await cleanup_response(r, session)
await cleanup_response(r)
@router.api_route('/{path:path}', methods=['GET', 'POST', 'PUT', 'DELETE'])
@ -1479,7 +1496,6 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
)
r = None
session = None
streaming = False
try:
@ -1508,10 +1524,7 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
else:
request_url = f'{url}/{path}'
session = aiohttp.ClientSession(
trust_env=True,
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
session = await get_session()
r = await session.request(
method=request.method,
url=request_url,
@ -1519,13 +1532,14 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
headers=headers,
cookies=cookies,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
# Check if response is SSE
if 'text/event-stream' in r.headers.get('Content-Type', ''):
streaming = True
return StreamingResponse(
stream_wrapper(r, session),
stream_wrapper(r),
status_code=r.status,
headers=dict(r.headers),
)
@ -1553,4 +1567,4 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
)
finally:
if not streaming:
await cleanup_response(r, session)
await cleanup_response(r)

View file

@ -198,7 +198,9 @@ async def create_new_prompt(
@router.get('/command/{command}', response_model=Optional[PromptAccessResponse])
async def get_prompt_by_command(command: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_prompt_by_command(
command: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
prompt = await Prompts.get_prompt_by_command(command, db=db)
if prompt:
@ -240,7 +242,9 @@ async def get_prompt_by_command(command: str, user=Depends(get_verified_user), d
@router.get('/id/{prompt_id}', response_model=Optional[PromptAccessResponse])
async def get_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_prompt_by_id(
prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if prompt:
@ -320,7 +324,7 @@ async def update_prompt_by_id(
if existing_prompt and existing_prompt.id != prompt.id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Command '/{form_data.command}' is already in use by another prompt",
detail=ERROR_MESSAGES.COMMAND_TAKEN,
)
form_data.access_grants = await filter_allowed_access_grants(
@ -385,10 +389,12 @@ async def update_prompt_metadata(
if existing_prompt and existing_prompt.id != prompt.id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Command '/{form_data.command}' is already in use",
detail=ERROR_MESSAGES.COMMAND_TAKEN,
)
updated_prompt = await Prompts.update_prompt_metadata(prompt.id, form_data.name, form_data.command, form_data.tags, db=db)
updated_prompt = await Prompts.update_prompt_metadata(
prompt.id, form_data.name, form_data.command, form_data.tags, db=db
)
if updated_prompt:
return updated_prompt
else:
@ -497,7 +503,9 @@ async def update_prompt_access_by_id(
@router.post('/id/{prompt_id}/toggle', response_model=Optional[PromptModel])
async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def toggle_prompt_active(
prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
@ -537,7 +545,9 @@ async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user),
@router.delete('/id/{prompt_id}/delete', response_model=bool)
async def delete_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def delete_prompt_by_id(
prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
@ -741,7 +751,7 @@ async def get_prompt_diff(
if not diff:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail='One or both history entries not found',
detail=ERROR_MESSAGES.NOT_FOUND,
)
return diff

View file

@ -40,7 +40,7 @@ from open_webui.models.files import FileModel, FileUpdateForm, Files
from open_webui.utils.access_control.files import has_access_to_file
from open_webui.models.knowledge import Knowledges
from open_webui.storage.provider import Storage
from open_webui.internal.db import get_async_session, get_db
from open_webui.internal.db import get_async_db, get_async_session
from sqlalchemy.ext.asyncio import AsyncSession
@ -151,7 +151,7 @@ def get_ef(
model_kwargs=SENTENCE_TRANSFORMERS_MODEL_KWARGS,
)
except Exception as e:
log.debug(f'Error loading SentenceTransformer: {e}')
log.error(f'Error loading SentenceTransformer: {e}')
return ef
@ -1694,11 +1694,16 @@ async def process_file(
try:
# Commit any pending changes before the slow embedding step.
# Note: file is already a Pydantic model (not ORM), so no expunge needed.
db.commit()
await db.commit()
# External embedding API takes time (5-60s+).
# Subsequent updates use fresh sessions via get_db().
result = save_docs_to_vector_db(
# Subsequent updates use fresh async sessions.
# NOTE: save_docs_to_vector_db is a sync function that
# calls asyncio.run_coroutine_threadsafe(..., main_loop).result()
# which blocks the calling thread. We MUST run it in a
# worker thread to avoid deadlocking the event loop.
result = await run_in_threadpool(
save_docs_to_vector_db,
request,
docs=docs,
collection_name=collection_name,
@ -1714,7 +1719,7 @@ async def process_file(
if result:
# Fresh session for the final update.
with get_db() as session:
async with get_async_db() as session:
await Files.update_file_metadata_by_id(
file.id,
{
@ -1744,7 +1749,7 @@ async def process_file(
except Exception as e:
log.exception(e)
# Fresh session for error status update.
with get_db() as session:
async with get_async_db() as session:
await Files.update_file_data_by_id(
file.id,
{'status': 'failed'},

View file

@ -18,7 +18,7 @@ from open_webui.utils.task import (
moa_response_generation_template,
)
from open_webui.utils.auth import get_admin_user, get_verified_user
from open_webui.constants import TASKS
from open_webui.constants import ERROR_MESSAGES, TASKS
from open_webui.routers.pipelines import process_pipeline_inlet_filter
@ -168,7 +168,7 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail='Model not found',
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
@ -245,7 +245,7 @@ async def generate_follow_ups(request: Request, form_data: dict, user=Depends(ge
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail='Model not found',
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
@ -313,7 +313,7 @@ async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail='Model not found',
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
@ -375,7 +375,7 @@ async def generate_image_prompt(request: Request, form_data: dict, user=Depends(
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail='Model not found',
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
@ -431,13 +431,13 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v
if not request.app.state.config.ENABLE_SEARCH_QUERY_GENERATION:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f'Search query generation is disabled',
detail=ERROR_MESSAGES.FEATURE_DISABLED('Search query generation'),
)
elif type == 'retrieval':
if not request.app.state.config.ENABLE_RETRIEVAL_QUERY_GENERATION:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f'Query generation is disabled',
detail=ERROR_MESSAGES.FEATURE_DISABLED('Query generation'),
)
if getattr(request.state, 'cached_queries', None):
@ -455,7 +455,7 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail='Model not found',
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
@ -508,7 +508,7 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend
if not request.app.state.config.ENABLE_AUTOCOMPLETE_GENERATION:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f'Autocompletion generation is disabled',
detail=ERROR_MESSAGES.FEATURE_DISABLED('Autocompletion generation'),
)
type = form_data.get('type')
@ -519,7 +519,7 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend
if len(prompt) > request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f'Input prompt exceeds maximum length of {request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH}',
detail=ERROR_MESSAGES.INPUT_TOO_LONG(request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH),
)
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
@ -533,7 +533,7 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail='Model not found',
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
@ -595,7 +595,7 @@ async def generate_emoji(request: Request, form_data: dict, user=Depends(get_ver
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail='Model not found',
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
@ -661,7 +661,7 @@ async def generate_moa_response(request: Request, form_data: dict, user=Depends(
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail='Model not found',
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
template = DEFAULT_MOA_GENERATION_PROMPT_TEMPLATE

View file

@ -285,7 +285,7 @@ async def load_tool_from_url(request: Request, form_data: LoadUrlForm, user=Depe
'content': data,
}
except Exception as e:
raise HTTPException(status_code=500, detail=f'Error importing tool: {e}')
raise HTTPException(status_code=500, detail=ERROR_MESSAGES.DEFAULT(e))
############################
@ -622,7 +622,9 @@ async def delete_tools_by_id(
@router.get('/id/{id}/valves', response_model=Optional[dict])
async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_tools_valves_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(
@ -775,7 +777,9 @@ async def update_tools_valves_by_id(
@router.get('/id/{id}/valves/user', response_model=Optional[dict])
async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_tools_user_valves_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(

View file

@ -40,6 +40,7 @@ from open_webui.utils.auth import (
validate_password,
)
from open_webui.utils.access_control import get_permissions, has_permission
from open_webui.socket.main import disconnect_user_sessions
log = logging.getLogger(__name__)
@ -272,7 +273,9 @@ async def update_default_user_permissions(request: Request, form_data: UserPermi
@router.get('/user/settings', response_model=Optional[UserSettings])
async def get_user_settings_by_session_user(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_user_settings_by_session_user(
user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
user = await Users.get_user_by_id(user.id, db=db)
if user:
return user.settings
@ -467,7 +470,9 @@ async def get_user_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSe
@router.get('/{user_id}/info', response_model=UserInfoResponse)
async def get_user_info_by_id(user_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def get_user_info_by_id(
user_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
user = await Users.get_user_by_id(user_id, db=db)
if user:
groups = await Groups.get_groups_by_member_id(user_id, db=db)
@ -486,7 +491,9 @@ async def get_user_info_by_id(user_id: str, user=Depends(get_verified_user), db:
@router.get('/{user_id}/oauth/sessions')
async def get_user_oauth_sessions_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
async def get_user_oauth_sessions_by_id(
user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)
):
sessions = await OAuthSessions.get_sessions_by_user_id(user_id, db=db)
if sessions and len(sessions) > 0:
return sessions
@ -573,7 +580,7 @@ async def update_user_by_id(
detail=ERROR_MESSAGES.ACTION_PROHIBITED,
)
if form_data.role != 'admin':
if form_data.role is not None and form_data.role != 'admin':
# If the primary admin is trying to change their own role, prevent it
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
@ -590,7 +597,7 @@ async def update_user_by_id(
user = await Users.get_user_by_id(user_id, db=db)
if user:
if form_data.email.lower() != user.email:
if form_data.email is not None and form_data.email.lower() != user.email:
email_user = await Users.get_user_by_email(form_data.email.lower(), db=db)
if email_user:
raise HTTPException(
@ -607,19 +614,32 @@ async def update_user_by_id(
hashed = get_password_hash(form_data.password)
await Auths.update_user_password_by_id(user_id, hashed, db=db)
await Auths.update_email_by_id(user_id, form_data.email.lower(), db=db)
updated_user = await Users.update_user_by_id(
user_id,
{
'role': form_data.role,
'name': form_data.name,
'email': form_data.email.lower(),
'profile_image_url': form_data.profile_image_url,
},
db=db,
)
# Build update dict from only the provided fields
update_data = {}
if form_data.role is not None:
update_data['role'] = form_data.role
if form_data.name is not None:
update_data['name'] = form_data.name
if form_data.email is not None:
update_data['email'] = form_data.email.lower()
await Auths.update_email_by_id(user_id, form_data.email.lower(), db=db)
if form_data.profile_image_url is not None:
update_data['profile_image_url'] = form_data.profile_image_url
if update_data:
updated_user = await Users.update_user_by_id(
user_id,
update_data,
db=db,
)
else:
updated_user = user
if updated_user:
# If the role changed, disconnect all socket sessions so stale
# privileges cached in SESSION_POOL are invalidated.
if updated_user.role != user.role:
await disconnect_user_sessions(user_id)
return updated_user
raise HTTPException(
@ -659,6 +679,7 @@ async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: Asyn
result = await Auths.delete_auth_by_id(user_id, db=db)
if result:
await disconnect_user_sessions(user_id)
return True
raise HTTPException(
@ -679,5 +700,7 @@ async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: Asyn
@router.get('/{user_id}/groups')
async def get_user_groups_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
async def get_user_groups_by_id(
user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)
):
return await Groups.get_groups_by_member_id(user_id, db=db)

View file

@ -45,7 +45,7 @@ async def execute_code(request: Request, form_data: CodeForm, user=Depends(get_v
if not request.app.state.config.ENABLE_CODE_EXECUTION:
raise HTTPException(
status_code=403,
detail='Code execution is disabled',
detail=ERROR_MESSAGES.FEATURE_DISABLED('Code execution'),
)
if request.app.state.config.CODE_EXECUTION_ENGINE == 'jupyter':
@ -69,7 +69,7 @@ async def execute_code(request: Request, form_data: CodeForm, user=Depends(get_v
else:
raise HTTPException(
status_code=400,
detail='Code execution engine not supported',
detail=ERROR_MESSAGES.DEFAULT('Code execution engine not supported'),
)

View file

@ -312,6 +312,24 @@ async def enter_room_for_users(room: str, user_ids: list[str]):
log.debug(f'Failed to make users {user_ids} join room {room}: {e}')
async def disconnect_user_sessions(user_id: str):
"""Disconnect all Socket.IO sessions belonging to a user.
Call this when a user's role is changed or the user is deleted so that
stale role/permission data cached in SESSION_POOL is invalidated.
The client will automatically reconnect and re-authenticate with
fresh data from the database.
"""
try:
session_ids = get_session_ids_from_room(f'user:{user_id}')
for sid in session_ids:
await sio.disconnect(sid)
if session_ids:
log.info(f'Disconnected {len(session_ids)} session(s) for user {user_id}')
except Exception as e:
log.warning(f'Failed to disconnect sessions for user {user_id}: {e}')
@sio.on('usage')
async def usage(sid, data):
if sid in SESSION_POOL:

View file

@ -140,7 +140,7 @@ class S3StorageProvider(StorageProvider):
def upload_file(self, file: BinaryIO, filename: str, tags: Dict[str, str]) -> Tuple[bytes, str]:
"""Handles uploading of the file to S3 storage."""
_, file_path = LocalStorageProvider.upload_file(file, filename, tags)
contents, file_path = LocalStorageProvider.upload_file(file, filename, tags)
s3_key = os.path.join(self.key_prefix, filename)
try:
self.s3_client.upload_file(file_path, self.bucket_name, s3_key)
@ -153,7 +153,7 @@ class S3StorageProvider(StorageProvider):
Tagging=tagging,
)
return (
open(file_path, 'rb').read(),
contents,
f's3://{self.bucket_name}/{s3_key}',
)
except ClientError as e:

View file

@ -760,14 +760,26 @@ async def search_notes(
content_snippet = ''
if note.data and note.data.get('content', {}).get('md'):
md_content = note.data['content']['md']
lower_content = md_content.lower()
lower_query = query.lower()
idx = lower_content.find(lower_query)
if idx != -1:
start = max(0, idx - 50)
end = min(len(md_content), idx + len(query) + 100)
content_lower = md_content.lower()
# Find the first matching word to center the snippet around.
search_words = query.lower().split()
match_pos = -1
match_len = len(query)
for word in search_words:
found_pos = content_lower.find(word)
if found_pos != -1:
match_pos = found_pos
match_len = len(word)
break
if match_pos != -1:
snippet_start = max(0, match_pos - 50)
snippet_end = min(len(md_content), match_pos + match_len + 100)
content_snippet = (
('...' if start > 0 else '') + md_content[start:end] + ('...' if end < len(md_content) else '')
('...' if snippet_start > 0 else '')
+ md_content[snippet_start:snippet_end]
+ ('...' if snippet_end < len(md_content) else '')
)
else:
content_snippet = md_content[:150] + ('...' if len(md_content) > 150 else '')
@ -1213,7 +1225,7 @@ async def search_channel_messages(
end_ts = end_timestamp * 1_000_000_000 if end_timestamp else None
# Search messages using the model method
matching_messages = Messages.search_messages_by_channel_ids(
matching_messages = await Messages.search_messages_by_channel_ids(
channel_ids=channel_ids,
query=query,
start_timestamp=start_ts,
@ -1274,7 +1286,7 @@ async def view_channel_message(
try:
user_id = __user__.get('id')
message = Messages.get_message_by_id(message_id)
message = await Messages.get_message_by_id(message_id)
if not message:
return json.dumps({'error': 'Message not found'})
@ -1336,7 +1348,7 @@ async def view_channel_thread(
user_id = __user__.get('id')
# Get the parent message
parent_message = Messages.get_message_by_id(parent_message_id)
parent_message = await Messages.get_message_by_id(parent_message_id)
if not parent_message:
return json.dumps({'error': 'Message not found'})
@ -1353,7 +1365,7 @@ async def view_channel_thread(
return json.dumps({'error': 'Access denied'})
# Get all thread replies
thread_replies = Messages.get_thread_replies_by_message_id(parent_message_id)
thread_replies = await Messages.get_thread_replies_by_message_id(parent_message_id)
# Build the response
messages = []
@ -2337,13 +2349,40 @@ VALID_TASK_STATUSES = {'pending', 'in_progress', 'completed', 'cancelled'}
class TaskItem(BaseModel):
id: Optional[str] = Field(None, description='Unique identifier for the task. Auto-generated if omitted.')
content: Optional[str] = Field(None, description='Task description. Aliases: title, name, description.')
content: str = Field(..., description='Task description.')
status: Literal['pending', 'in_progress', 'completed', 'cancelled'] = Field('pending', description='Task status.')
async def tasks(
tasks: Optional[list[TaskItem]] = None,
overwrite: bool = True,
def _task_summary(all_tasks: list[dict]) -> dict:
"""Build summary counts for a task list."""
pending = sum(1 for t in all_tasks if t['status'] == 'pending')
in_progress = sum(1 for t in all_tasks if t['status'] == 'in_progress')
completed = sum(1 for t in all_tasks if t['status'] == 'completed')
cancelled = sum(1 for t in all_tasks if t['status'] == 'cancelled')
return {
'total': len(all_tasks),
'pending': pending,
'in_progress': in_progress,
'completed': completed,
'cancelled': cancelled,
}
async def _emit_tasks(event_emitter, all_tasks: list[dict]):
"""Persist task state to the UI."""
if event_emitter:
await event_emitter(
{
'type': 'chat:message:tasks',
'data': {
'tasks': all_tasks,
},
}
)
async def create_tasks(
tasks: list[TaskItem],
__chat_id__: str = None,
__message_id__: str = None,
__event_emitter__: callable = None,
@ -2351,144 +2390,42 @@ async def tasks(
__user__: dict = None,
) -> str:
"""
Track progress on multi-step work by maintaining a task checklist.
Use this whenever a request involves multiple steps or could take
significant effort. Call to set the full list, then call again
with overwrite=false after completing each task to mark it
completed. Do not leave tasks in_progress when the work is done.
Each task has an id, content, and status (pending, in_progress,
completed, cancelled).
Create a task checklist to track progress on multi-step work.
Call this once at the start to define all steps, then use
update_task to mark each task as you complete it.
:param tasks: Optional list of task items. Each item: id (string), content (string, required for new tasks), status (pending|in_progress|completed|cancelled). Leave empty to fetch without modifying.
:param overwrite: If true (default), replaces the entire task list. If false, updates/adds tasks by id while keeping existing ones.
:param tasks: List of task items. Each item: content (string, required), status (pending|in_progress|completed|cancelled, default pending), id (optional, auto-generated).
:return: JSON with the full task list and summary counts
"""
if __chat_id__ is None:
return json.dumps({'error': 'Chat context not available'})
try:
def _to_dict(task) -> dict:
"""Convert TaskItem or dict to plain dict."""
all_tasks = []
for idx, task in enumerate(tasks):
if hasattr(task, 'model_dump'):
d = task.model_dump(exclude_none=True)
# Include any extra fields the model sent
if hasattr(task, 'model_extra') and task.model_extra:
d.update(task.model_extra)
return d
return dict(task) if not isinstance(task, dict) else task
elif isinstance(task, dict):
d = task
else:
d = dict(task)
def _resolve_content(d: dict) -> str:
"""Accept content, title, name, or description as the task text."""
for key in ('content', 'title', 'name', 'description'):
val = str(d.get(key, '')).strip()
if val:
return val
return ''
content = str(d.get('content', '')).strip()
if not content:
continue
def _resolve_id(d: dict, idx: int) -> str:
"""Use provided id, or auto-generate from index."""
item_id = str(d.get('id', '') or '').strip()
return item_id if item_id else str(idx + 1)
item_id = str(d.get('id', '') or '').strip() or str(idx + 1)
status = str(d.get('status', 'pending')).strip().lower()
if status not in VALID_TASK_STATUSES:
status = 'pending'
if tasks is None:
# Read-only - return current list
all_tasks = await Chats.get_chat_tasks_by_id(__chat_id__)
elif overwrite:
# Full replacement - validate and write
all_tasks = []
for idx, task in enumerate(tasks):
d = _to_dict(task)
item_id = _resolve_id(d, idx)
content = _resolve_content(d)
if not content:
continue
all_tasks.append({'id': item_id, 'content': content, 'status': status})
status = str(d.get('status', 'pending')).strip().lower()
if status not in VALID_TASK_STATUSES:
status = 'pending'
all_tasks.append(
{
'id': item_id,
'content': content,
'status': status,
}
)
else:
# Partial update - merge by id
existing_tasks = await Chats.get_chat_tasks_by_id(__chat_id__)
existing_by_id = {t['id']: t for t in existing_tasks}
seen_ids = set()
for idx, task in enumerate(tasks):
d = _to_dict(task)
item_id = _resolve_id(d, len(existing_tasks) + idx)
seen_ids.add(item_id)
if item_id in existing_by_id:
resolved = _resolve_content(d)
if resolved:
existing_by_id[item_id]['content'] = resolved
status = str(d.get('status', '')).strip().lower()
if status and status in VALID_TASK_STATUSES:
existing_by_id[item_id]['status'] = status
else:
content = _resolve_content(d)
if not content:
continue
status = str(d.get('status', 'pending')).strip().lower()
if status not in VALID_TASK_STATUSES:
status = 'pending'
existing_by_id[item_id] = {
'id': item_id,
'content': content,
'status': status,
}
# Preserve order of existing, append new
all_tasks = []
for t in existing_tasks:
if t['id'] in existing_by_id:
all_tasks.append(existing_by_id[t['id']])
for item_id in seen_ids:
if not any(t['id'] == item_id for t in existing_tasks):
all_tasks.append(existing_by_id[item_id])
# Persist to DB and emit (skip for read-only)
if tasks is not None:
await Chats.update_chat_tasks_by_id(__chat_id__, all_tasks)
if __event_emitter__:
await __event_emitter__(
{
'type': 'chat:message:tasks',
'data': {
'tasks': all_tasks,
},
}
)
# Build summary counts
pending = sum(1 for t in all_tasks if t['status'] == 'pending')
in_progress = sum(1 for t in all_tasks if t['status'] == 'in_progress')
completed = sum(1 for t in all_tasks if t['status'] == 'completed')
cancelled = sum(1 for t in all_tasks if t['status'] == 'cancelled')
await Chats.update_chat_tasks_by_id(__chat_id__, all_tasks)
await _emit_tasks(__event_emitter__, all_tasks)
return json.dumps(
{
'tasks': all_tasks,
'summary': {
'total': len(all_tasks),
'pending': pending,
'in_progress': in_progress,
'completed': completed,
'cancelled': cancelled,
},
},
{'tasks': all_tasks, 'summary': _task_summary(all_tasks)},
ensure_ascii=False,
)
except Exception as e:
@ -2496,6 +2433,58 @@ async def tasks(
return json.dumps({'error': str(e)})
async def update_task(
id: str,
status: str = 'completed',
__chat_id__: str = None,
__message_id__: str = None,
__event_emitter__: callable = None,
__request__: Request = None,
__user__: dict = None,
) -> str:
"""
Mark a single task as completed, in_progress, pending, or cancelled.
Call this after finishing each step. You MUST call this for every
task, including the very last one.
:param id: The task ID to update
:param status: New status: completed, in_progress, pending, or cancelled (default: completed)
:return: JSON with the updated task list and summary counts
"""
if __chat_id__ is None:
return json.dumps({'error': 'Chat context not available'})
try:
status = status.strip().lower()
if status not in VALID_TASK_STATUSES:
return json.dumps(
{'error': f'Invalid status: {status}. Must be one of: {", ".join(sorted(VALID_TASK_STATUSES))}'}
)
all_tasks = await Chats.get_chat_tasks_by_id(__chat_id__)
found = False
for task in all_tasks:
if task['id'] == id:
task['status'] = status
found = True
break
if not found:
return json.dumps({'error': f'Task with id "{id}" not found'})
await Chats.update_chat_tasks_by_id(__chat_id__, all_tasks)
await _emit_tasks(__event_emitter__, all_tasks)
return json.dumps(
{'tasks': all_tasks, 'summary': _task_summary(all_tasks)},
ensure_ascii=False,
)
except Exception as e:
log.exception(f'update_task_status error: {e}')
return json.dumps({'error': str(e)})
# =============================================================================
# AUTOMATION TOOLS
# =============================================================================

View file

@ -181,17 +181,133 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
)
elif block_type == 'tool_result':
# Tool results become separate tool messages in OpenAI format
tool_content = block.get('content', '')
if isinstance(tool_content, list):
tool_text_parts = []
for tc in tool_content:
if isinstance(tc, dict) and tc.get('type') == 'text':
tool_text_parts.append(tc.get('text', ''))
tool_content = '\n'.join(tool_text_parts)
tool_result_content = block.get('content', '')
tool_content: str | list = ''
if isinstance(tool_result_content, str):
tool_content = tool_result_content
elif isinstance(tool_result_content, list):
# Build a multimodal content array to preserve
# images and other non-text content types.
converted_parts = []
for content_block in tool_result_content:
if not isinstance(content_block, dict):
continue
content_type = content_block.get('type', 'text')
if content_type == 'text':
converted_parts.append(
{
'type': 'text',
'text': content_block.get('text', ''),
}
)
elif content_type == 'image':
source = content_block.get('source', {})
if source.get('type') == 'base64':
media_type = source.get(
'media_type', 'image/png'
)
data = source.get('data', '')
converted_parts.append(
{
'type': 'image_url',
'image_url': {
'url': f'data:{media_type};base64,{data}',
},
}
)
elif source.get('type') == 'url':
converted_parts.append(
{
'type': 'image_url',
'image_url': {
'url': source.get('url', ''),
},
}
)
elif content_type == 'document':
# Documents have no direct OpenAI equivalent;
# convert to a text representation.
document_source = content_block.get(
'source', {}
)
document_title = content_block.get(
'title', 'Document'
)
document_context = content_block.get(
'context', ''
)
document_text = (
f'[Document: {document_title}]'
)
if document_context:
document_text += f'\n{document_context}'
if (
document_source.get('type') == 'text'
and document_source.get('data')
):
document_text += (
f'\n{document_source["data"]}'
)
converted_parts.append(
{'type': 'text', 'text': document_text}
)
elif content_type == 'search_result':
# Convert search results to a text
# representation with source attribution.
search_title = content_block.get('title', '')
search_url = content_block.get('source', '')
search_content_blocks = content_block.get(
'content', []
)
search_texts = []
for search_block in search_content_blocks:
if (
isinstance(search_block, dict)
and search_block.get('type') == 'text'
):
search_texts.append(
search_block.get('text', '')
)
search_body = '\n'.join(search_texts)
search_text = (
f'[Search Result: {search_title}]'
)
if search_url:
search_text += f'\nSource: {search_url}'
if search_body:
search_text += f'\n{search_body}'
converted_parts.append(
{'type': 'text', 'text': search_text}
)
# Flatten to string when only text parts are present
if all(
part.get('type') == 'text'
for part in converted_parts
):
tool_content = '\n'.join(
part.get('text', '')
for part in converted_parts
)
elif converted_parts:
tool_content = converted_parts
else:
tool_content = ''
# Propagate error status if present
if block.get('is_error'):
tool_content = f'Error: {tool_content}'
if isinstance(tool_content, str):
tool_content = f'Error: {tool_content}'
elif isinstance(tool_content, list):
tool_content.insert(
0,
{
'type': 'text',
'text': 'Error: ',
},
)
messages.append(
{

View file

@ -24,7 +24,7 @@ from asgiref.typing import (
from loguru import logger
from starlette.requests import Request
from open_webui.env import AUDIT_LOG_LEVEL, AUDIT_INCLUDED_PATHS, MAX_BODY_LOG_SIZE
from open_webui.env import AUDIT_LOG_LEVEL, ENABLE_AUDIT_GET_REQUESTS, AUDIT_INCLUDED_PATHS, MAX_BODY_LOG_SIZE
from open_webui.utils.auth import get_current_user, get_http_authorization_cred
from open_webui.models.users import UserModel
@ -117,7 +117,7 @@ class AuditLoggingMiddleware:
ASGI middleware that intercepts HTTP requests and responses to perform audit logging. It captures request/response bodies (depending on audit level), headers, HTTP methods, and user information, then logs a structured audit entry at the end of the request cycle.
"""
AUDITED_METHODS = {'PUT', 'PATCH', 'DELETE', 'POST'}
DEFAULT_AUDITED_METHODS = {'PUT', 'PATCH', 'DELETE', 'POST'}
def __init__(
self,
@ -127,12 +127,16 @@ class AuditLoggingMiddleware:
included_paths: Optional[list[str]] = None,
max_body_size: int = MAX_BODY_LOG_SIZE,
audit_level: AuditLevel = AuditLevel.NONE,
audit_get_requests: bool = False,
) -> None:
self.app = app
self.audit_logger = AuditLogger(logger)
self.excluded_paths = excluded_paths or []
self.included_paths = included_paths or []
self.max_body_size = max_body_size
self.audited_methods = set(self.DEFAULT_AUDITED_METHODS)
if audit_get_requests:
self.audited_methods.add('GET')
self.audit_level = audit_level
if self.included_paths and self.excluded_paths:
@ -202,7 +206,10 @@ class AuditLoggingMiddleware:
return None
def _should_skip_auditing(self, request: Request) -> bool:
if request.method not in {'POST', 'PUT', 'PATCH', 'DELETE'} or AUDIT_LOG_LEVEL == 'NONE':
if AUDIT_LOG_LEVEL == 'NONE':
return True
if request.method not in self.audited_methods:
return True
ALWAYS_LOG_ENDPOINTS = {

View file

@ -20,7 +20,6 @@ from pytz import UTC
from typing import Optional, Union, List, Dict
from open_webui.utils.access_control import has_permission
from open_webui.models.users import Users
from open_webui.models.auths import Auths
@ -238,9 +237,7 @@ async def is_valid_token(request, decoded) -> bool:
# Per-user revocation (OIDC back-channel logout)
user_id = decoded.get('id')
if user_id:
revoked_at = await request.app.state.redis.get(
f'{REDIS_KEY_PREFIX}:auth:user:{user_id}:revoked_at'
)
revoked_at = await 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)
@ -385,6 +382,7 @@ async def get_current_user(
# Refresh the user's last active timestamp
# Fire-and-forget via asyncio.create_task to avoid blocking
import asyncio
asyncio.create_task(Users.update_last_active_by_id(user.id))
return user
else:
@ -427,6 +425,21 @@ async def get_current_user_by_api_key(request, api_key: str):
):
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.API_KEY_NOT_ALLOWED)
# Enforce endpoint restrictions — checked here (not in middleware)
# so it applies regardless of how the API key was transported
# (Authorization header, cookie, x-api-key header, etc.).
if request.app.state.config.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS:
allowed_paths = [
path.strip() for path in str(request.app.state.config.API_KEYS_ALLOWED_ENDPOINTS).split(',') if path.strip()
]
request_path = request.url.path
is_allowed = any(request_path == allowed or request_path.startswith(allowed + '/') for allowed in allowed_paths)
if not is_allowed:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
# Add user info to current span
if ENABLE_OTEL:
from opentelemetry import trace

View file

@ -22,11 +22,12 @@ from dateutil.rrule import rrulestr
from fastapi import Request
from starlette.datastructures import Headers
from open_webui.constants import ERROR_MESSAGES
from open_webui.models.automations import Automations, AutomationRuns, AutomationModel
from open_webui.models.chats import ChatForm, Chats
from open_webui.models.users import Users
from open_webui.utils.task import prompt_template
from open_webui.internal.db import get_db
from open_webui.internal.db import get_async_db
log = logging.getLogger(__name__)
@ -59,9 +60,9 @@ def validate_rrule(s: str) -> None:
try:
rule = _parse_rule(s)
except Exception as e:
raise ValueError(f'Invalid RRULE: {e}')
raise ValueError(ERROR_MESSAGES.AUTOMATION_INVALID_RRULE(e))
if rule.after(datetime.now()) is None:
raise ValueError('RRULE has no future occurrences')
raise ValueError(ERROR_MESSAGES.AUTOMATION_NO_FUTURE_RUNS)
def next_run_ns(s: str, tz: str = None) -> Optional[int]:
@ -125,7 +126,7 @@ async def automation_worker_loop(app) -> None:
log.info(f'Automation worker started (poll interval: {AUTOMATION_POLL_INTERVAL}s)')
while True:
try:
with get_db() as db:
async with get_async_db() as db:
batch = await Automations.claim_due(int(time.time_ns()), limit=10, db=db)
if batch:
log.info(f'Claimed {len(batch)} due automation(s)')
@ -224,6 +225,16 @@ def _resolve_model_filter_ids(app, model_id: str) -> list[str]:
return list(filter_ids) if filter_ids else []
def _resolve_model_terminal_id(app, model_id: str) -> Optional[str]:
"""Read model default terminal_id from model config.
The frontend does this in Chat.svelte (model.info.meta.terminalId).
"""
models = getattr(app.state, 'MODELS', {})
model = models.get(model_id, {})
return model.get('info', {}).get('meta', {}).get('terminalId') or None
async def _set_terminal_cwd(app, server_id: str, user, cwd: str, chat_id: str) -> None:
"""Set the working directory on a terminal server via the proxy.
@ -296,8 +307,9 @@ async def execute_automation(app, automation: AutomationModel) -> None:
user_msg_id = str(uuid4())
assistant_msg_id = str(uuid4())
# Create the chat with user message (same structure as frontend)
chat_id = str(uuid4())
chat = await Chats.insert_new_chat(
chat_id,
automation.user_id,
ChatForm(
chat={
@ -357,13 +369,8 @@ async def execute_automation(app, automation: AutomationModel) -> None:
features = _resolve_model_features(app, model_id)
filter_ids = _resolve_model_filter_ids(app, model_id)
# If a terminal is linked, set the CWD before building the payload
terminal_id = None
if terminal_config and terminal_config.get('server_id'):
terminal_id = terminal_config['server_id']
cwd = terminal_config.get('cwd')
if cwd:
await _set_terminal_cwd(app, terminal_id, user, cwd, chat.id)
# Resolve terminal from model config
terminal_id = _resolve_model_terminal_id(app, model_id)
# Build the same payload the frontend sends to /api/chat/completions
form_data = {
@ -372,7 +379,13 @@ async def execute_automation(app, automation: AutomationModel) -> None:
'stream': True,
'chat_id': chat.id,
'id': assistant_msg_id,
'parent_id': user_msg_id,
'parent_id': None, # Root message (chat already created above)
'user_message': {
'id': user_msg_id,
'parentId': None,
'role': 'user',
'content': prompt,
},
'session_id': f'automation:{automation.id}',
'background_tasks': {},
}
@ -423,5 +436,5 @@ async def _record_run(
error: str = None,
):
"""Insert a run record into automation_run."""
with get_db() as db:
async with get_async_db() as db:
await AutomationRuns.insert(automation_id, status, chat_id=chat_id, error=error, db=db)

View file

@ -25,7 +25,8 @@ import base64
import io
import re
import requests
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL
from open_webui.utils.session_pool import get_session
BASE64_IMAGE_URL_PREFIX = re.compile(r'data:image/\w+;base64,', re.IGNORECASE)
MARKDOWN_IMAGE_URL_PATTERN = re.compile(r'!\[(.*?)\]\((.+?)\)', re.IGNORECASE)
@ -37,12 +38,13 @@ async def get_image_base64_from_url(url: str) -> Optional[str]:
# Validate URL to prevent SSRF attacks against local/private networks
validate_url(url)
# Download the image from the URL
response = requests.get(url)
response.raise_for_status()
image_data = response.content
encoded_string = base64.b64encode(image_data).decode('utf-8')
content_type = response.headers.get('Content-Type', 'image/png')
return f'data:{content_type};base64,{encoded_string}'
session = await get_session()
async with session.get(url, ssl=AIOHTTP_CLIENT_SESSION_SSL) as response:
response.raise_for_status()
image_data = await response.read()
encoded_string = base64.b64encode(image_data).decode('utf-8')
content_type = response.headers.get('Content-Type', 'image/png')
return f'data:{content_type};base64,{encoded_string}'
else:
file = await Files.get_file_by_id(url)
@ -88,7 +90,7 @@ async def convert_markdown_base64_images(request, content: str, metadata, user):
last_end = 0
for match in MARKDOWN_IMAGE_URL_PATTERN.finditer(content):
result_parts.append(content[last_end:match.start()])
result_parts.append(content[last_end : match.start()])
base64_string = match.group(2)
if len(base64_string) > MIN_REPLACEMENT_URL_LENGTH:
url = await get_image_url_from_base64(request, base64_string, metadata, user)

View file

@ -444,7 +444,7 @@ def serialize_output(output: list) -> str:
files = result_item.get('files')
embeds = result_item.get('embeds', '')
content += f'<details type="tool_calls" done="true" id="{call_id}" name="{name}" arguments="{html.escape(json.dumps(arguments))}" result="{html.escape(json.dumps(result_text, ensure_ascii=False))}" files="{html.escape(json.dumps(files)) if files else ""}" embeds="{html.escape(json.dumps(embeds))}">\n<summary>Tool Executed</summary>\n</details>\n'
content += f'<details type="tool_calls" done="true" id="{call_id}" name="{name}" arguments="{html.escape(json.dumps(arguments))}" files="{html.escape(json.dumps(files)) if files else ""}" embeds="{html.escape(json.dumps(embeds))}">\n<summary>Tool Executed</summary>\n{html.escape(json.dumps(result_text, ensure_ascii=False))}\n</details>\n'
else:
content += f'<details type="tool_calls" done="false" id="{call_id}" name="{name}" arguments="{html.escape(json.dumps(arguments))}">\n<summary>Executing...</summary>\n</details>\n'
@ -889,10 +889,14 @@ def get_source_context(sources: list, source_ids: dict = None, include_content:
if source_id not in source_ids:
source_ids[source_id] = len(source_ids) + 1
src_name = source.get('source', {}).get('name')
src_type = source.get('source', {}).get('type')
src_rid = source.get('source', {}).get('id')
body = doc if include_content else ''
context_string += (
f'<source id="{source_ids[source_id]}"'
+ (f' name="{src_name}"' if src_name else '')
+ (f' resource-type="{src_type}"' if src_type else '')
+ (f' resource-id="{src_rid}"' if src_rid else '')
+ f'>{body}</source>\n'
)
return context_string
@ -2050,13 +2054,16 @@ async def convert_url_images_to_base64(form_data):
continue
try:
base64_data = await asyncio.to_thread(get_image_base64_from_url, image_url)
new_content.append(
{
'type': 'image_url',
'image_url': {'url': base64_data},
}
)
base64_data = await get_image_base64_from_url(image_url)
if base64_data:
new_content.append(
{
'type': 'image_url',
'image_url': {'url': base64_data},
}
)
else:
new_content.append(item)
except Exception as e:
log.debug(f'Error converting image URL to base64: {e}')
new_content.append(item)
@ -2146,10 +2153,10 @@ async def process_chat_payload(request, form_data, user, metadata, model):
# Load messages from DB when available — DB preserves structured 'output' items
# which the frontend strips, causing tool calls to be merged into content.
chat_id = metadata.get('chat_id')
parent_message_id = metadata.get('parent_message_id')
user_message_id = metadata.get('user_message_id')
if chat_id and parent_message_id and not chat_id.startswith('local:'):
db_messages = await load_messages_from_db(chat_id, parent_message_id)
if chat_id and user_message_id and not chat_id.startswith('local:'):
db_messages = await load_messages_from_db(chat_id, user_message_id)
if db_messages:
system_message = get_system_message(form_data.get('messages', []))
form_data['messages'] = [system_message, *db_messages] if system_message else db_messages
@ -3054,6 +3061,110 @@ async def background_tasks_handler(ctx):
pass
async def outlet_filter_handler(ctx):
"""Run outlet filters inline after chat completion.
Replaces the separate POST /api/chat/completed round-trip.
Persists outlet-modified content to DB and emits a chat:outlet event
so the frontend can sync its in-memory state.
"""
request = ctx['request']
user = ctx['user']
model = ctx['model']
metadata = ctx['metadata']
event_emitter = ctx.get('event_emitter')
event_caller = ctx.get('event_caller')
chat_id = metadata.get('chat_id', '')
message_id = metadata.get('message_id')
if not chat_id or chat_id.startswith('local:') or not message_id:
return
try:
messages_map = await Chats.get_messages_map_by_chat_id(chat_id)
if not messages_map:
return
message_list = get_message_list(messages_map, message_id)
if not message_list:
return
model_id = model.get('id') if isinstance(model, dict) else model
outlet_data = {
'model': model_id,
'messages': [
{
'id': m.get('id'),
'role': m.get('role'),
'content': m.get('content', ''),
'info': m.get('info'),
'timestamp': m.get('timestamp'),
**(({'usage': m['usage']} if m.get('usage') else {})),
**(({'sources': m['sources']} if m.get('sources') else {})),
}
for m in message_list
],
'filter_ids': metadata.get('filter_ids', []),
'chat_id': chat_id,
'session_id': metadata.get('session_id'),
'id': message_id,
}
# Pipeline outlet filters
models = request.app.state.MODELS
try:
outlet_data = await process_pipeline_outlet_filter(request, outlet_data, user, models)
except Exception as e:
log.debug(f'Pipeline outlet filter error: {e}')
# Function outlet filters
extra_params = {
'__event_emitter__': event_emitter,
'__event_call__': event_caller,
'__user__': user.model_dump() if isinstance(user, UserModel) else {},
'__metadata__': metadata,
'__request__': request,
'__model__': model,
}
filter_ids = await get_sorted_filter_ids(request, model, metadata.get('filter_ids', []))
filter_functions = await Functions.get_functions_by_ids(filter_ids)
outlet_result, _ = await process_filter_functions(
request=request,
filter_functions=filter_functions,
filter_type='outlet',
form_data=outlet_data,
extra_params=extra_params,
)
# Persist outlet-modified content and notify frontend
if outlet_result and outlet_result.get('messages'):
for msg in outlet_result['messages']:
msg_id = msg.get('id')
if msg_id and msg_id in messages_map:
original = messages_map[msg_id]
if original.get('content') != msg.get('content'):
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id,
msg_id,
{
'content': msg['content'],
'originalContent': original.get('content'),
},
)
if event_emitter:
await event_emitter({
'type': 'chat:outlet',
'data': {'messages': outlet_result['messages']},
})
except Exception as e:
log.debug(f'Error running outlet filters: {e}')
async def non_streaming_chat_response_handler(response, ctx):
request = ctx['request']
@ -3077,6 +3188,8 @@ async def non_streaming_chat_response_handler(response, ctx):
else:
error = str(error)
log.error('Provider returned error (non-streaming): %s', error)
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
@ -3173,6 +3286,7 @@ async def non_streaming_chat_response_handler(response, ctx):
)
await background_tasks_handler(ctx)
await outlet_filter_handler(ctx)
response = build_response_object(response, merge_events_into_response(response_data, events))
except Exception as e:
@ -3645,6 +3759,7 @@ async def streaming_chat_response_handler(response, ctx):
if not choices:
error = data.get('error', {})
if error:
log.error('Provider returned error (streaming): %s', error)
try:
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
@ -4060,11 +4175,12 @@ async def streaming_chat_response_handler(response, ctx):
if responses_api_tool_calls:
tool_calls.append(_split_tool_calls(responses_api_tool_calls))
try:
await stream_body_handler(response, form_data)
finally:
if response.background:
await response.background()
await stream_body_handler(response, form_data)
tool_call_retries = 0
tool_call_sources = [] # Track citation sources from tool results
all_tool_call_sources = [] # Accumulated sources across all iterations
@ -4682,27 +4798,36 @@ async def streaming_chat_response_handler(response, ctx):
)
await background_tasks_handler(ctx)
await outlet_filter_handler(ctx)
except asyncio.CancelledError:
log.warning('Task was cancelled!')
await event_emitter({'type': 'chat:tasks:cancel'})
try:
await asyncio.shield(event_emitter({'type': 'chat:tasks:cancel'}))
if not ENABLE_REALTIME_CHAT_SAVE:
# Save message in the database
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
'done': True,
'content': serialize_output(output),
'output': output,
},
)
else:
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{'done': True},
)
if not ENABLE_REALTIME_CHAT_SAVE:
# Save message in the database
await asyncio.shield(
Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
'done': True,
'content': serialize_output(output),
'output': output,
},
)
)
else:
await asyncio.shield(
Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{'done': True},
)
)
except Exception:
pass
raise # re-raise CancelledError for proper propagation
if response.background is not None:
await response.background()

View file

@ -905,9 +905,17 @@ async def cleanup_response(
session: Optional[aiohttp.ClientSession],
):
if response:
response.close()
if not response.closed:
# aiohttp 3.9+ made ClientResponse.close() synchronous (returns None).
# Older versions returned a coroutine. Handle both gracefully.
result = response.close()
if result is not None:
await result
if session:
await session.close()
if not session.closed:
result = session.close()
if result is not None:
await result
async def stream_wrapper(response, session, content_handler=None):

View file

@ -18,6 +18,7 @@ from typing import Literal
import aiohttp
from authlib.integrations.starlette_client import OAuth
from authlib.jose.errors import BadSignatureError
from authlib.oidc.core import UserInfo
from fastapi import (
HTTPException,
@ -139,6 +140,41 @@ auth_manager_config.OAUTH_UPDATE_EMAIL_ON_LOGIN = OAUTH_UPDATE_EMAIL_ON_LOGIN
auth_manager_config.OAUTH_AUDIENCE = OAUTH_AUDIENCE
# Conservative default when the provider omits both expires_in and expires_at.
# Matches the value recommended by Authlib's compliance_fix documentation.
DEFAULT_TOKEN_EXPIRY_SECONDS = 3600
def _normalize_token_expiry(token: dict) -> dict:
"""Ensure a token dict always has a numeric ``expires_at``.
Resolution order:
1. If *expires_at* is already present and non-None, trust it.
2. Else if *expires_in* is present and non-None, compute *expires_at*.
3. Otherwise fall back to ``DEFAULT_TOKEN_EXPIRY_SECONDS`` and log a
warning so operators can identify providers that omit expiration.
Also stamps *issued_at* for auditing.
"""
token['issued_at'] = datetime.now().timestamp()
if token.get('expires_at') is not None:
token['expires_at'] = int(token['expires_at'])
return token
if token.get('expires_in') is not None:
token['expires_at'] = int(datetime.now().timestamp() + token['expires_in'])
return token
# Neither field present — conservative fallback
log.warning(
"OAuth token response missing both 'expires_in' and 'expires_at'; "
f"defaulting to {DEFAULT_TOKEN_EXPIRY_SECONDS}s from now"
)
token['expires_at'] = int(datetime.now().timestamp() + DEFAULT_TOKEN_EXPIRY_SECONDS)
return token
FERNET = None
if len(OAUTH_CLIENT_INFO_ENCRYPTION_KEY) != 44:
@ -512,6 +548,25 @@ async def get_oauth_client_info_with_static_credentials(
raise e
def resolve_oauth_client_info(connection: dict) -> dict:
"""
Decrypt OAuth client info from a tool server connection config.
For oauth_2.1_static, overlays admin-provided credentials from
info.oauth_client_id and info.oauth_client_secret onto the blob.
"""
info = connection.get('info', {})
data = decrypt_data(info.get('oauth_client_info', ''))
if connection.get('auth_type') == 'oauth_2.1_static':
if info.get('oauth_client_id') and info.get('oauth_client_secret'):
data['client_id'] = info['oauth_client_id']
data['client_secret'] = info['oauth_client_secret']
return data
class OAuthClientManager:
def __init__(self, app):
self.oauth = OAuth()
@ -524,6 +579,7 @@ class OAuthClientManager:
'client_id': oauth_client_info.client_id,
'client_secret': oauth_client_info.client_secret,
'client_kwargs': {
'follow_redirects': True,
**({'scope': oauth_client_info.scope} if oauth_client_info.scope else {}),
**(
{'token_endpoint_auth_method': oauth_client_info.token_endpoint_auth_method}
@ -534,15 +590,20 @@ class OAuthClientManager:
'server_metadata_url': (oauth_client_info.issuer if oauth_client_info.issuer else None),
}
if oauth_client_info.server_metadata and oauth_client_info.server_metadata.code_challenge_methods_supported:
if (
isinstance(
oauth_client_info.server_metadata.code_challenge_methods_supported,
list,
)
and 'S256' in oauth_client_info.server_metadata.code_challenge_methods_supported
):
kwargs['code_challenge_method'] = 'S256'
# Default to S256 for OAuth 2.1 (PKCE is mandatory per RFC 9700)
kwargs['code_challenge_method'] = 'S256'
# Only remove PKCE if metadata explicitly excludes S256
if (
oauth_client_info.server_metadata
and oauth_client_info.server_metadata.code_challenge_methods_supported
and isinstance(
oauth_client_info.server_metadata.code_challenge_methods_supported,
list,
)
and 'S256' not in oauth_client_info.server_metadata.code_challenge_methods_supported
):
del kwargs['code_challenge_method']
self.clients[client_id] = {
'client': self.oauth.register(**kwargs),
@ -582,7 +643,7 @@ class OAuthClientManager:
continue
try:
oauth_client_info = decrypt_data(oauth_client_info)
oauth_client_info = resolve_oauth_client_info(connection)
return self.add_client(expected_client_id, OAuthClientInformationFull(**oauth_client_info))['client']
except Exception as e:
log.error(f'Failed to lazily add OAuth client {expected_client_id} from config: {e}')
@ -705,7 +766,7 @@ class OAuthClientManager:
log.warning(f'No OAuth session found for user {user_id}, client_id {client_id}')
return None
if force_refresh or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at):
if force_refresh or session.expires_at is None or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at):
log.debug(f'Token refresh needed for user {user_id}, client_id {session.provider}')
refreshed_token = await self._refresh_token(session)
if refreshed_token:
@ -816,14 +877,7 @@ class OAuthClientManager:
if 'refresh_token' not in new_token_data:
new_token_data['refresh_token'] = token_data['refresh_token']
# Add timestamp for tracking
new_token_data['issued_at'] = datetime.now().timestamp()
# Calculate expires_at if we have expires_in
if 'expires_in' in new_token_data and 'expires_at' not in new_token_data:
new_token_data['expires_at'] = int(
datetime.now().timestamp() + new_token_data['expires_in']
)
_normalize_token_expiry(new_token_data)
log.debug(f'Token refresh successful for client_id {client_id}')
return new_token_data
@ -876,12 +930,7 @@ class OAuthClientManager:
if token:
try:
# Add timestamp for tracking
token['issued_at'] = datetime.now().timestamp()
# Calculate expires_at if we have expires_in
if 'expires_in' in token and 'expires_at' not in token:
token['expires_at'] = datetime.now().timestamp() + token['expires_in']
_normalize_token_expiry(token)
# Clean up any existing sessions for this user/client_id first
sessions = await OAuthSessions.get_sessions_by_user_id(user_id)
@ -968,7 +1017,7 @@ class OAuthManager:
log.warning(f'No OAuth session found for user {user_id}, session {session_id}')
return None
if force_refresh or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at):
if force_refresh or session.expires_at is None or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at):
log.debug(f'Token refresh needed for user {user_id}, provider {session.provider}')
refreshed_token = await self._refresh_token(session)
if refreshed_token:
@ -1082,14 +1131,7 @@ class OAuthManager:
if 'refresh_token' not in new_token_data:
new_token_data['refresh_token'] = token_data['refresh_token']
# Add timestamp for tracking
new_token_data['issued_at'] = datetime.now().timestamp()
# Calculate expires_at if we have expires_in
if 'expires_in' in new_token_data and 'expires_at' not in new_token_data:
new_token_data['expires_at'] = int(
datetime.now().timestamp() + new_token_data['expires_in']
)
_normalize_token_expiry(new_token_data)
log.debug(f'Token refresh successful for provider {provider}')
return new_token_data
@ -1392,6 +1434,27 @@ class OAuthManager:
try:
token = await client.authorize_access_token(request, **auth_params)
except BadSignatureError:
# The IdP likely rotated its signing keys and the cached JWKS
# is stale. Evict the cached key set so the next attempt
# fetches fresh keys from the jwks_uri.
log.warning(
'OIDC bad_signature for provider %s — evicting cached JWKS and retrying',
provider,
)
if hasattr(client, 'server_metadata') and isinstance(client.server_metadata, dict):
client.server_metadata.pop('jwks', None)
try:
token = await client.authorize_access_token(request, **auth_params)
except Exception as retry_exc:
detailed_error = _build_oauth_callback_error_message(retry_exc)
log.warning(
'OAuth callback error during authorize_access_token retry for provider %s: %s',
provider,
detailed_error,
exc_info=True,
)
raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
except Exception as e:
detailed_error = _build_oauth_callback_error_message(e)
log.warning(
@ -1666,12 +1729,7 @@ class OAuthManager:
)
try:
# Add timestamp for tracking
token['issued_at'] = datetime.now().timestamp()
# Calculate expires_at if we have expires_in
if 'expires_in' in token and 'expires_at' not in token:
token['expires_at'] = datetime.now().timestamp() + token['expires_in']
_normalize_token_expiry(token)
# Enforce max concurrent sessions per user/provider to prevent
# unbounded growth while allowing multi-device usage
@ -1785,7 +1843,10 @@ class OAuthManager:
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'},
content={
'error': 'invalid_request',
'error_description': 'No configured provider matches token issuer',
},
)
# 4. Validate the logout_token signature and claims
@ -1886,5 +1947,7 @@ class OAuthManager:
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')
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={})

View file

@ -194,20 +194,12 @@ def get_redis_connection(
connection = None
connect_timeout_kwargs = (
{'socket_connect_timeout': REDIS_SOCKET_CONNECT_TIMEOUT}
if REDIS_SOCKET_CONNECT_TIMEOUT is not None
else {}
{'socket_connect_timeout': REDIS_SOCKET_CONNECT_TIMEOUT} if REDIS_SOCKET_CONNECT_TIMEOUT is not None else {}
)
keepalive_kwargs = (
{'socket_keepalive': True} if REDIS_SOCKET_KEEPALIVE else {}
)
keepalive_kwargs = {'socket_keepalive': True} if REDIS_SOCKET_KEEPALIVE else {}
health_check_kwargs = (
{'health_check_interval': REDIS_HEALTH_CHECK_INTERVAL}
if REDIS_HEALTH_CHECK_INTERVAL
else {}
)
health_check_kwargs = {'health_check_interval': REDIS_HEALTH_CHECK_INTERVAL} if REDIS_HEALTH_CHECK_INTERVAL else {}
if async_mode:
import redis.asyncio as redis

View file

@ -0,0 +1,119 @@
"""Shared aiohttp ClientSession pool.
Instead of creating a new ClientSession (and TCPConnector) per request,
callers acquire a long-lived session from this module. The pool manages
a single TCPConnector with configurable limits, enabling TCP/SSL connection
reuse, shared DNS cache, and bounded concurrency.
All pool parameters are configurable via environment variables:
- AIOHTTP_POOL_CONNECTIONS (default 100) max total connections
- AIOHTTP_POOL_CONNECTIONS_PER_HOST (default 30) per-host limit
- AIOHTTP_POOL_DNS_TTL (default 300) DNS cache TTL in seconds
Usage:
from open_webui.utils.session_pool import get_session, cleanup_response
session = await get_session()
r = await session.request(...)
# When done with the *response* (not the session):
await cleanup_response(r)
IMPORTANT: Callers must NOT close the shared session. Only the response
needs cleanup. The session is closed once during application shutdown
via ``close_session()``.
"""
import logging
from typing import Optional
import aiohttp
from open_webui.env import (
AIOHTTP_CLIENT_TIMEOUT,
AIOHTTP_POOL_CONNECTIONS,
AIOHTTP_POOL_CONNECTIONS_PER_HOST,
AIOHTTP_POOL_DNS_TTL,
)
log = logging.getLogger(__name__)
_session: Optional[aiohttp.ClientSession] = None
async def get_session() -> aiohttp.ClientSession:
"""Return the shared aiohttp ClientSession, creating it lazily."""
global _session
if _session is None or _session.closed:
connector_kwargs = {
'ttl_dns_cache': AIOHTTP_POOL_DNS_TTL,
'enable_cleanup_closed': True,
}
if AIOHTTP_POOL_CONNECTIONS is not None:
connector_kwargs['limit'] = AIOHTTP_POOL_CONNECTIONS
else:
connector_kwargs['limit'] = 0 # aiohttp: 0 = unlimited
if AIOHTTP_POOL_CONNECTIONS_PER_HOST is not None:
connector_kwargs['limit_per_host'] = AIOHTTP_POOL_CONNECTIONS_PER_HOST
else:
connector_kwargs['limit_per_host'] = 0 # aiohttp: 0 = unlimited
connector = aiohttp.TCPConnector(**connector_kwargs)
timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT)
_session = aiohttp.ClientSession(
connector=connector,
timeout=timeout,
trust_env=True,
)
log.info(
'Created shared aiohttp session pool (limit=%s, per_host=%s, dns_ttl=%d)',
AIOHTTP_POOL_CONNECTIONS or 'unlimited',
AIOHTTP_POOL_CONNECTIONS_PER_HOST or 'unlimited',
AIOHTTP_POOL_DNS_TTL,
)
return _session
async def close_session():
"""Close the shared session. Called during application shutdown."""
global _session
if _session and not _session.closed:
await _session.close()
log.info('Closed shared aiohttp session pool')
_session = None
async def cleanup_response(
response: Optional[aiohttp.ClientResponse],
session: Optional[aiohttp.ClientSession] = None,
):
"""Release and close an aiohttp response, optionally closing the session.
When using the shared pool, ``session`` should be ``None`` (the pool
session is never closed per-request). When a caller creates its own
one-off session, pass it here to close it after the response.
"""
if response:
if not response.closed:
# aiohttp 3.9+ made ClientResponse.close() synchronous (returns None).
# Older versions returned a coroutine. Handle both gracefully.
result = response.close()
if result is not None:
await result
if session:
if not session.closed:
result = session.close()
if result is not None:
await result
async def stream_wrapper(response, session=None, content_handler=None):
"""Wrap a stream to ensure cleanup happens even if streaming is interrupted.
This is more reliable than BackgroundTask which may not run if the client
disconnects. When using the shared pool, ``session`` should be ``None``.
"""
try:
stream = content_handler(response.content) if content_handler else response.content
async for chunk in stream:
yield chunk
finally:
await cleanup_response(response, session)

View file

@ -51,6 +51,7 @@ from open_webui.env import (
ENABLE_FORWARD_USER_INFO_HEADERS,
FORWARD_SESSION_INFO_HEADER_CHAT_ID,
FORWARD_SESSION_INFO_HEADER_MESSAGE_ID,
REDIS_KEY_PREFIX,
)
from open_webui.utils.headers import include_user_info_headers
from open_webui.tools.builtin import (
@ -85,7 +86,8 @@ from open_webui.tools.builtin import (
view_file,
view_knowledge_file,
view_skill,
tasks,
create_tasks,
update_task,
create_automation,
update_automation,
list_automations,
@ -101,7 +103,9 @@ log = logging.getLogger(__name__)
# Let no function be called without need, and let what
# it yields justify the cost of running it.
async def get_async_tool_function_and_apply_extra_params(function: Callable, extra_params: dict) -> Callable[..., Awaitable]:
async def get_async_tool_function_and_apply_extra_params(
function: Callable, extra_params: dict
) -> Callable[..., Awaitable]:
sig = inspect.signature(function)
extra_params = {k: v for k, v in extra_params.items() if k in sig.parameters}
partial_func = partial(function, **extra_params)
@ -540,11 +544,13 @@ async def get_builtin_tools(
# Task management - break down complex work into trackable steps
if is_builtin_tool_enabled('tasks'):
builtin_functions.append(tasks)
builtin_functions.extend([create_tasks, update_task])
# Automation tools - create and manage scheduled automations from chat
if is_builtin_tool_enabled('automations') and await has_user_permission('automations'):
builtin_functions.extend([create_automation, update_automation, list_automations, toggle_automation, delete_automation])
builtin_functions.extend(
[create_automation, update_automation, list_automations, toggle_automation, delete_automation]
)
for func in builtin_functions:
callable = await get_async_tool_function_and_apply_extra_params(
@ -806,7 +812,7 @@ def convert_openapi_to_tool_payload(openapi_spec):
if not description:
description = param.get('description') or ''
if param_schema.get('enum') and isinstance(param_schema.get('enum'), list):
description += f'. Possible values: {", ".join(param_schema.get("enum"))}'
description += f'. Possible values: {", ".join(str(v) for v in param_schema.get("enum"))}'
param_property = {
'type': param_schema.get('type') or 'string',
'description': description,
@ -849,7 +855,9 @@ async def set_tool_servers(request: Request):
request.app.state.TOOL_SERVERS = await get_tool_servers_data(request.app.state.config.TOOL_SERVER_CONNECTIONS)
if request.app.state.redis is not None:
await request.app.state.redis.set('tool_servers', json.dumps(request.app.state.TOOL_SERVERS))
await request.app.state.redis.set(
f'{REDIS_KEY_PREFIX}:tool_servers', json.dumps(request.app.state.TOOL_SERVERS)
)
return request.app.state.TOOL_SERVERS
@ -858,7 +866,7 @@ async def get_tool_servers(request: Request):
tool_servers = []
if request.app.state.redis is not None:
try:
tool_servers = json.loads(await request.app.state.redis.get('tool_servers'))
tool_servers = json.loads(await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:tool_servers'))
request.app.state.TOOL_SERVERS = tool_servers
except Exception as e:
log.error(f'Error fetching tool_servers from Redis: {e}')
@ -984,7 +992,9 @@ async def set_terminal_servers(request: Request):
)
if request.app.state.redis is not None:
await request.app.state.redis.set('terminal_servers', json.dumps(request.app.state.TERMINAL_SERVERS))
await request.app.state.redis.set(
f'{REDIS_KEY_PREFIX}:terminal_servers', json.dumps(request.app.state.TERMINAL_SERVERS)
)
return request.app.state.TERMINAL_SERVERS
@ -994,7 +1004,7 @@ async def get_terminal_servers(request: Request):
terminal_servers = []
if request.app.state.redis is not None:
try:
terminal_servers = json.loads(await request.app.state.redis.get('terminal_servers'))
terminal_servers = json.loads(await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:terminal_servers'))
request.app.state.TERMINAL_SERVERS = terminal_servers
except Exception as e:
log.error(f'Error fetching terminal_servers from Redis: {e}')

View file

@ -1,13 +1,31 @@
"""Validation utilities for user-supplied input."""
# Known static asset paths used as default profile images
_ALLOWED_STATIC_PATHS = (
'/user.png',
'/static/favicon.png',
)
import re
from urllib.parse import urlparse
# External URL prefixes that are explicitly trusted for profile images
_ALLOWED_URL_PREFIXES = ('https://www.gravatar.com/avatar/',)
# Matches the OWUI-generated profile image route. ``[^/?#]+`` accepts
# any user-ID without allowing path-traversal or query/fragment injection,
# and the ``$`` anchor rejects trailing path components.
_USER_PROFILE_IMAGE_RE = re.compile(r'^/api/v1/users/[^/?#]+/profile/image$')
# Validates MIME type and structure of base64 data URIs. Only the prefix
# is checked — validating the full base64 payload would mean running a
# regex across megabytes of data on every Pydantic instantiation for zero
# security benefit (corrupt base64 simply renders a broken image, same as
# a 404 URL). SVG is intentionally excluded: it can carry embedded scripts.
_SAFE_DATA_URI_RE = re.compile(r'^data:image/(png|jpeg|gif|webp);base64,', re.IGNORECASE)
# Exact relative paths accepted as profile images. These are the only
# static-asset paths OWUI itself assigns; no prefix/wildcard matching is
# used so that arbitrary relative paths cannot trigger authenticated GETs
# against internal endpoints when rendered as ``<img>`` sources.
_SAFE_STATIC_PATHS = frozenset(
{
'/user.png',
'/favicon.png',
'/static/favicon.png',
}
)
def validate_profile_image_url(url: str) -> str:
@ -16,28 +34,49 @@ def validate_profile_image_url(url: str) -> str:
Allowed formats:
- Empty string (falls back to default avatar)
- data:image/* URIs (base64-encoded uploads from the frontend)
- Known static asset paths (/user.png, /static/favicon.png)
- Trusted external URLs (e.g. Gravatar)
- Known static-asset paths assigned by OWUI (exact match)
- The OWUI profile-image API route ``/api/v1/users/{id}/profile/image``
- ``http://`` and ``https://`` URLs with a valid hostname
- ``data:image/{png,jpeg,gif,webp};base64,...`` URIs
Returns the url unchanged if valid, raises ValueError otherwise.
Everything else is rejected, including:
- Dangerous schemes (javascript:, file:, ftp:, )
- SVG data URIs (can contain embedded scripts)
- Arbitrary relative paths (prevents authenticated GET triggers)
- Scheme-relative URLs (``//host/path``)
"""
if not url:
return url
_ALLOWED_DATA_PREFIXES = (
'data:image/png',
'data:image/jpeg',
'data:image/gif',
'data:image/webp',
# --- Relative paths (exact match + anchored regex only) -----------
if url in _SAFE_STATIC_PATHS:
return url
if _USER_PROFILE_IMAGE_RE.match(url):
return url
# --- Absolute URLs -------------------------------------------------
# urlparse normalises the scheme to lowercase, giving us
# case-insensitive scheme matching for free.
parsed = urlparse(url)
# External images served over HTTP(S), e.g. OAuth provider avatars.
# Require a non-empty hostname (not just netloc, which can be ":80"
# for a URL like http://:80/path with no actual host).
if parsed.scheme in ('http', 'https'):
if not parsed.hostname:
raise ValueError('Invalid profile image URL: HTTP(S) URLs must include a host.')
return url
# Base64-encoded raster images uploaded via the frontend.
# The regex enforces the ;base64, boundary and is case-insensitive
# per the data-URI / MIME-type specs.
if _SAFE_DATA_URI_RE.match(url):
return url
raise ValueError(
'Invalid profile image URL: must be a known internal path, '
'an HTTP(S) URL with a host, or a data:image URI (png/jpeg/gif/webp).'
)
if any(url.startswith(prefix) for prefix in _ALLOWED_DATA_PREFIXES):
return url
if url in _ALLOWED_STATIC_PATHS:
return url
if any(url.startswith(prefix) for prefix in _ALLOWED_URL_PREFIXES):
return url
raise ValueError('Invalid profile image URL: only data URIs and default avatars are allowed.')

View file

@ -13,15 +13,15 @@ cryptography
bcrypt==5.0.0
argon2-cffi==25.1.0
PyJWT[crypto]==2.11.0
authlib==1.6.9
authlib==1.6.10
requests==2.32.5
aiohttp==3.13.2 # do not update to 3.13.3 - broken
requests==2.33.1
aiohttp==3.13.5 # do not update to 3.13.3 - broken
async-timeout
aiocache
aiofiles
starlette-compress==1.7.0
Brotli==1.1.0
Brotli==1.2.0
httpx[socks,http2,zstd,cli,brotli]==0.28.1
starsessions[redis]==2.2.1
@ -52,7 +52,7 @@ langchain-text-splitters==1.1.1
fake-useragent==2.2.0
chromadb==1.5.2
black==26.1.0
black==26.3.1
pydub
chardet==5.2.0
beautifulsoup4

View file

@ -10,15 +10,15 @@ cryptography==46.0.5
bcrypt==5.0.0
argon2-cffi==25.1.0
PyJWT[crypto]==2.11.0
authlib==1.6.9
authlib==1.6.10
requests==2.32.5
aiohttp==3.13.2 # do not update to 3.13.3 - broken
requests==2.33.1
aiohttp==3.13.5 # do not update to 3.13.3 - broken
async-timeout==5.0.1
aiocache==0.12.3
aiofiles==25.1.0
starlette-compress==1.7.0
Brotli==1.1.0
Brotli==1.2.0
httpx[socks,http2,zstd,cli,brotli]==0.28.1
starsessions[redis]==2.2.1
python-mimeparse==2.0.0
@ -58,8 +58,8 @@ chromadb==1.5.2
weaviate-client==4.20.3
opensearch-py==3.1.0
transformers==5.3.0
sentence-transformers==5.2.3
transformers==5.5.4
sentence-transformers==5.4.0
accelerate==1.13.0
pyarrow==20.0.0 # fix: pin pyarrow version to 20 for rpi compatibility #15897
einops==0.8.2
@ -95,7 +95,7 @@ rank-bm25==0.2.2
onnxruntime==1.24.3
faster-whisper==1.2.1
black==26.1.0
black==26.3.1
youtube-transcript-api==1.2.4
pytube==15.0.0

View file

@ -18,10 +18,10 @@ dependencies = [
"bcrypt==5.0.0",
"argon2-cffi==25.1.0",
"PyJWT[crypto]==2.11.0",
"authlib==1.6.9",
"authlib==1.6.10",
"requests==2.32.5",
"aiohttp==3.13.2", # do not update to 3.13.3 - broken
"requests==2.33.1",
"aiohttp==3.13.5", # do not update to 3.13.3 - broken
"async-timeout==5.0.1",
"aiocache==0.12.3",
"aiofiles==25.1.0",
@ -64,8 +64,8 @@ dependencies = [
"PyMySQL==1.1.2",
"boto3==1.42.62",
"transformers==5.3.0",
"sentence-transformers==5.2.3",
"transformers==5.5.4",
"sentence-transformers==5.4.0",
"accelerate==1.13.0",
"pyarrow==20.0.0", # fix: pin pyarrow version to 20 for rpi compatibility #15897
"einops==0.8.2",
@ -100,7 +100,7 @@ dependencies = [
"onnxruntime==1.24.3",
"faster-whisper==1.2.1",
"black==26.1.0",
"black==26.3.1",
"youtube-transcript-api==1.2.4",
"pytube==15.0.0",

View file

@ -273,10 +273,42 @@ export const stopTask = async (token: string, id: string) => {
return res;
};
export const stopTasksByChatId = async (token: string, chat_id: string) => {
let error = null;
const res = await fetch(`${WEBUI_BASE_URL}/api/tasks/chat/${encodeURIComponent(chat_id)}/stop`, {
method: 'POST',
headers: {
Accept: 'application/json',
'Content-Type': 'application/json',
...(token && { authorization: `Bearer ${token}` })
}
})
.then(async (res) => {
if (!res.ok) throw await res.json();
return res.json();
})
.catch((err) => {
console.error(err);
if ('detail' in err) {
error = err.detail;
} else {
error = err;
}
return null;
});
if (error) {
throw error;
}
return res;
};
export const getTaskIdsByChatId = async (token: string, chat_id: string) => {
let error = null;
const res = await fetch(`${WEBUI_BASE_URL}/api/tasks/chat/${chat_id}`, {
const res = await fetch(`${WEBUI_BASE_URL}/api/tasks/chat/${encodeURIComponent(chat_id)}`, {
method: 'GET',
headers: {
Accept: 'application/json',

View file

@ -78,21 +78,20 @@
return;
}
if (auth_type === 'oauth_2.1_static' && (!oauthClientId || !oauthClientSecret)) {
toast.error($i18n.t('Please enter Client ID and Client Secret'));
return;
}
// client_id is the tool server ID (used as the internal lookup key for both flows).
// For static, client_secret signals the backend to use the static credential path.
// The actual OAuth client_id/secret come from the connection info at save time.
const formData: { url: string; client_id: string; client_secret?: string } = {
url: url,
client_id: id
client_id: id,
...(auth_type === 'oauth_2.1_static' ? { client_secret: oauthClientSecret } : {})
};
// For static OAuth, include client credentials
if (auth_type === 'oauth_2.1_static') {
if (!oauthClientId || !oauthClientSecret) {
toast.error($i18n.t('Please enter Client ID and Client Secret'));
return;
}
formData.client_id = id;
formData.client_secret = oauthClientSecret;
}
const res = await registerOAuthClient(localStorage.token, formData, 'mcp').catch((err) => {
toast.error($i18n.t('Registration failed'));
return null;

View file

@ -8,7 +8,6 @@
import ScheduleDropdown from '$lib/components/automations/ScheduleDropdown.svelte';
import ModelDropdown from '$lib/components/automations/ModelDropdown.svelte';
import TerminalDropdown from '$lib/components/automations/TerminalDropdown.svelte';
import {
createAutomation,
@ -16,7 +15,6 @@
type AutomationForm,
type AutomationResponse
} from '$lib/apis/automations';
import { getTerminalServers, type TerminalServer } from '$lib/apis/terminal/index';
const i18n = getContext('i18n');
const dispatch = createEventDispatcher();
@ -31,11 +29,6 @@
let loading = false;
// Terminal state
let terminalServers: TerminalServer[] = [];
let terminalServerId = '';
let terminalCwd = '';
// Schedule dropdown ref
let scheduleDropdown: ScheduleDropdown;
@ -58,15 +51,7 @@
data: {
prompt: prompt.trim(),
model_id: model_id.trim(),
rrule: scheduleDropdown.buildRrule(),
...(terminalServerId
? {
terminal: {
server_id: terminalServerId,
...(terminalCwd.trim() ? { cwd: terminalCwd.trim() } : {})
}
}
: {})
rrule: scheduleDropdown.buildRrule()
},
is_active
};
@ -90,20 +75,11 @@
};
const init = async () => {
// Load terminal servers
try {
terminalServers = await getTerminalServers(localStorage.token);
} catch {
terminalServers = [];
}
if (automation) {
name = automation.name;
prompt = automation.data.prompt;
model_id = automation.data.model_id;
is_active = automation.is_active;
terminalServerId = automation.data.terminal?.server_id || '';
terminalCwd = automation.data.terminal?.cwd || '';
if (scheduleDropdown) {
scheduleDropdown.parseRrule(automation.data.rrule);
}
@ -112,8 +88,6 @@
prompt = '';
model_id = '';
is_active = true;
terminalServerId = '';
terminalCwd = '';
}
};
@ -158,14 +132,6 @@
<ScheduleDropdown bind:this={scheduleDropdown} side="top" align="start" />
<ModelDropdown bind:model_id side="top" align="start" />
<TerminalDropdown
{terminalServers}
bind:terminalServerId
bind:terminalCwd
side="top"
align="start"
/>
</div>
<div class="flex items-center gap-2 shrink-0">

View file

@ -19,7 +19,6 @@
type AutomationResponse,
type AutomationRunModel
} from '$lib/apis/automations';
import { getTerminalServers, type TerminalServer } from '$lib/apis/terminal/index';
import Spinner from '$lib/components/common/Spinner.svelte';
import Tooltip from '$lib/components/common/Tooltip.svelte';
@ -29,7 +28,6 @@
import ScheduleDropdown from '$lib/components/automations/ScheduleDropdown.svelte';
import ModelDropdown from '$lib/components/automations/ModelDropdown.svelte';
import TerminalDropdown from '$lib/components/automations/TerminalDropdown.svelte';
dayjs.extend(relativeTime);
dayjs.extend(localizedFormat);
@ -43,10 +41,6 @@
let model_id = '';
let is_active = true;
let terminalServers: TerminalServer[] = [];
let terminalServerId = '';
let terminalCwd = '';
let loading = false;
let saving = false;
let showDeleteConfirm = false;
@ -97,15 +91,7 @@
data: {
prompt: prompt.trim(),
model_id: model_id.trim(),
rrule: scheduleDropdown.buildRrule(),
...(terminalServerId
? {
terminal: {
server_id: terminalServerId,
...(terminalCwd.trim() ? { cwd: terminalCwd.trim() } : {})
}
}
: {})
rrule: scheduleDropdown.buildRrule()
},
is_active
};
@ -204,18 +190,11 @@
prompt = automation.data.prompt;
model_id = automation.data.model_id;
is_active = automation.is_active;
terminalServerId = automation.data.terminal?.server_id || '';
terminalCwd = automation.data.terminal?.cwd || '';
if (scheduleDropdown) {
scheduleDropdown.parseRrule(automation.data.rrule);
}
try {
terminalServers = await getTerminalServers(localStorage.token);
} catch {
terminalServers = [];
}
await loadRuns();
});
</script>
@ -354,21 +333,6 @@
<span class="text-gray-600 dark:text-gray-400">{$i18n.t('Model')}</span>
<ModelDropdown bind:model_id side="bottom" align="end" onChange={markDirty} />
</div>
<!-- Terminal -->
{#if terminalServers.length > 0}
<div class="flex items-center justify-between text-xs">
<span class="text-gray-600 dark:text-gray-400">{$i18n.t('Terminal')}</span>
<TerminalDropdown
{terminalServers}
bind:terminalServerId
bind:terminalCwd
side="bottom"
align="end"
onChange={markDirty}
/>
</div>
{/if}
</div>
</div>

View file

@ -51,6 +51,7 @@
type="button"
on:click={() => {
uploadFilesHandler();
show = false;
}}
>
<Clip />
@ -62,6 +63,7 @@
type="button"
on:click={() => {
screenCaptureHandler();
show = false;
}}
>
<Camera />

View file

@ -81,11 +81,11 @@
import { processWeb, processWebSearch, processYoutubeVideo } from '$lib/apis/retrieval';
import { getAndUpdateUserLocation, getUserSettings } from '$lib/apis/users';
import {
chatCompleted,
generateQueries,
chatAction,
generateMoACompletion,
stopTask,
stopTasksByChatId,
getTaskIdsByChatId
} from '$lib/apis';
import { getTools } from '$lib/apis/tools';
@ -180,6 +180,12 @@
}
const navigateHandler = async () => {
// Mark the outgoing chat as read before loading the new one.
// $chatId still holds the previous chat here — loadChat() updates it.
if ($chatId && $chatId !== chatIdProp && !$temporaryChatEnabled) {
updateLastReadAt($chatId);
}
loading = true;
prompt = '';
@ -370,6 +376,11 @@
codeInterpreterEnabled = model.info.meta.defaultFeatureIds.includes('code_interpreter');
}
}
// Set Default Terminal
if (model?.info?.meta?.terminalId) {
selectedTerminalId.set(model.info.meta.terminalId);
}
}
};
@ -485,6 +496,22 @@
if (autoScroll) {
scrollToBottom('smooth');
}
} else if (type === 'chat:outlet') {
// Outlet filter ran on backend — sync in-memory state
const outletMessages = data.messages ?? [];
for (const msg of outletMessages) {
if (msg?.id && history.messages[msg.id]) {
const existing = history.messages[msg.id];
if (existing.content !== msg.content) {
history.messages[msg.id] = {
...existing,
originalContent: existing.content,
...msg
};
}
}
}
history = history;
} else if (type === 'chat:message:favorite') {
// Update message favorite status
message.favorite = data.favorite;
@ -1355,6 +1382,17 @@
taskIds = taskRes.task_ids;
}
// If no active tasks and current message is incomplete, generation was interrupted
const currentMessage = history.currentId ? history.messages[history.currentId] : null;
if (
currentMessage &&
currentMessage.role === 'assistant' &&
!currentMessage.done &&
(!taskIds || taskIds.length === 0)
) {
currentMessage.done = true;
}
await tick();
return true;
@ -1385,87 +1423,37 @@
}
};
let processingQueueChats = new Set<string>();
const processNextInQueue = async (targetChatId: string) => {
if (processingQueueChats.has(targetChatId)) return;
const queue = $chatRequestQueues[targetChatId];
if (!queue || queue.length === 0) return;
const combinedPrompt = queue.map((m) => m.prompt).join('\n\n');
const combinedFiles = queue.flatMap((m) => m.files);
processingQueueChats.add(targetChatId);
try {
const combinedPrompt = queue.map((m) => m.prompt).join('\n\n');
const combinedFiles = queue.flatMap((m) => m.files);
chatRequestQueues.update((q) => {
const { [targetChatId]: _, ...rest } = q;
return rest;
});
chatRequestQueues.update((q) => {
const { [targetChatId]: _, ...rest } = q;
return rest;
});
await submitPrompt(combinedPrompt, combinedFiles);
await submitPrompt(combinedPrompt, combinedFiles);
} finally {
processingQueueChats.delete(targetChatId);
}
};
const chatCompletedHandler = async (_chatId, modelId, responseMessageId, messages) => {
if (!responseMessageId) {
console.error('chatCompleted: missing message id', {
chatId: _chatId,
modelId,
messageCount: messages?.length ?? 0
});
return;
// Backend handles outlet filters and persistence inline.
// Just refresh the sidebar chat list.
if ($chatId == _chatId && !$temporaryChatEnabled) {
currentChatPage.set(1);
await chats.set(await getChatList(localStorage.token, $currentChatPage));
}
const res = await chatCompleted(localStorage.token, {
model: modelId,
messages: messages.map((m) => ({
id: m.id,
role: m.role,
content: m.content,
info: m.info ? m.info : undefined,
timestamp: m.timestamp,
...(m.usage ? { usage: m.usage } : {}),
...(m.sources ? { sources: m.sources } : {})
})),
filter_ids: selectedFilterIds.length > 0 ? selectedFilterIds : undefined,
model_item: $models.find((m) => m.id === modelId),
chat_id: _chatId,
session_id: $socket?.id,
id: responseMessageId
}).catch((error) => {
toast.error(`${error}`);
messages.at(-1).error = { content: error };
return null;
});
if (res !== null && res.messages) {
// Update chat history with the new messages
for (const message of res.messages) {
if (message?.id) {
// Add null check for message and message.id
history.messages[message.id] = {
...history.messages[message.id],
...(history.messages[message.id].content !== message.content
? { originalContent: history.messages[message.id].content }
: {}),
...message
};
}
}
}
await tick();
if ($chatId == _chatId) {
if (!$temporaryChatEnabled) {
chat = await updateChatById(localStorage.token, _chatId, {
models: selectedModels,
messages: messages,
history: history,
params: params,
files: chatFiles
});
currentChatPage.set(1);
await chats.set(await getChatList(localStorage.token, $currentChatPage));
}
}
taskIds = null;
};
@ -1873,13 +1861,15 @@
history.currentId = userMessageId;
// focus on chat input
const chatInput = document.getElementById('chat-input');
chatInput?.focus();
// focus on chat input (skip during voice call to avoid triggering mobile keyboard)
if (!$showCallOverlay) {
const chatInput = document.getElementById('chat-input');
chatInput?.focus();
}
saveSessionSelectedModels();
await sendMessage(history, userMessageId, { newChat: true });
await sendMessage(history, userMessageId);
};
const submitHandler = async (userPrompt, { _raw = false } = {}) => {
@ -1979,13 +1969,11 @@
{
messages = null,
modelId = null,
modelIdx = null,
newChat = false
modelIdx = null
}: {
messages?: any[] | null;
modelId?: string | null;
modelIdx?: number | null;
newChat?: boolean;
} = {}
) => {
if (autoScroll) {
@ -2004,6 +1992,8 @@
: selectedModels;
// Create response messages for each selected model
// Build message_ids map: {model_id: assistant_message_id}
const messageIdsMap: Record<string, string> = {};
for (const [_modelIdx, modelId] of selectedModelIds.entries()) {
const model = $models.filter((m) => m.id === modelId).at(0);
@ -2015,6 +2005,7 @@
childrenIds: [],
role: 'assistant',
content: '',
done: false,
model: model.id,
modelName: model.name ?? model.id,
modelIdx: modelIdx ? modelIdx : _modelIdx,
@ -2027,7 +2018,6 @@
// Append messageId to childrenIds of parent message
if (parentId !== null && history.messages[parentId]) {
// Add null check before accessing childrenIds
history.messages[parentId].childrenIds = [
...history.messages[parentId].childrenIds,
responseMessageId
@ -2035,68 +2025,71 @@
}
responseMessageIds[`${modelId}-${modelIdx ? modelIdx : _modelIdx}`] = responseMessageId;
messageIdsMap[modelId] = responseMessageId;
}
}
history = history;
// Create new chat if newChat is true and first user message
if (newChat && _history.messages[_history.currentId].parentId === null) {
_chatId = await initChatHandler(_history);
// New chat — backend generates the chat_id on first request
if (!_chatId) {
if ($temporaryChatEnabled) {
_chatId = `local:${$socket?.id}`;
await chatId.set(_chatId);
}
await tick();
}
await tick();
// Re-clone history so sendMessageSocket gets the response messages we just added
_history = structuredClone(history);
// Save chat after all messages have been created
await saveChatHandler(_chatId, _history);
await Promise.all(
selectedModelIds.map(async (modelId, _modelIdx) => {
console.log('modelId', modelId);
const model = $models.filter((m) => m.id === modelId).at(0);
// Vision capability check
for (const mid of selectedModelIds) {
const model = $models.filter((m) => m.id === mid).at(0);
if (model) {
const hasImages = createMessagesList(_history, parentId).some((message) =>
message.files?.some(
(file) => file.type === 'image' || (file?.content_type ?? '').startsWith('image/')
)
);
if (model) {
// If there are image files, check if model is vision capable
// Skip this check if image generation is enabled, as images may be for editing or are generated outputs in the history
const hasImages = createMessagesList(_history, parentId).some((message) =>
message.files?.some(
(file) => file.type === 'image' || (file?.content_type ?? '').startsWith('image/')
)
if (
hasImages &&
!(model.info?.meta?.capabilities?.vision ?? true) &&
!imageGenerationEnabled
) {
toast.error(
$i18n.t('Model {{modelName}} is not vision capable', {
modelName: model.name ?? model.id
})
);
if (
hasImages &&
!(model.info?.meta?.capabilities?.vision ?? true) &&
!imageGenerationEnabled
) {
toast.error(
$i18n.t('Model {{modelName}} is not vision capable', {
modelName: model.name ?? model.id
})
);
}
let responseMessageId =
responseMessageIds[`${modelId}-${modelIdx ? modelIdx : _modelIdx}`];
const chatEventEmitter = await getChatEventEmitter(model.id, _chatId);
scrollToBottom();
await sendMessageSocket(
model,
messages && messages.length > 0
? messages
: createMessagesList(_history, responseMessageId),
_history,
responseMessageId,
_chatId
);
if (chatEventEmitter) clearInterval(chatEventEmitter);
} else {
toast.error($i18n.t(`Model {{modelId}} not found`, { modelId }));
}
})
);
}
}
// Single request — backend fans out to all models
const primaryModelId = selectedModelIds[0];
const primaryModel = $models.filter((m) => m.id === primaryModelId).at(0);
const primaryResponseMessageId = messageIdsMap[primaryModelId];
if (primaryModel && primaryResponseMessageId) {
const chatEventEmitter = await getChatEventEmitter(primaryModel.id, _chatId);
scrollToBottom();
await sendMessageSocket(
primaryModel,
messages && messages.length > 0
? messages
: createMessagesList(_history, primaryResponseMessageId),
_history,
primaryResponseMessageId,
_chatId,
selectedModelIds.length > 1 ? messageIdsMap : undefined
);
if (chatEventEmitter) clearInterval(chatEventEmitter);
}
};
const getFeatures = () => {
@ -2151,7 +2144,7 @@
.map((token) => decodeURIComponent(JSON.parse(`"${token.replace(/"/g, '\\"')}"`)));
};
const sendMessageSocket = async (model, _messages, _history, responseMessageId, _chatId) => {
const sendMessageSocket = async (model, _messages, _history, responseMessageId, _chatId, messageIdsMap?: Record<string, string>) => {
const responseMessage = _history.messages[responseMessageId];
const userMessage = _history.messages[responseMessage.parentId];
@ -2202,53 +2195,54 @@
$settings?.params?.stream_response ??
params?.stream_response ??
true;
// Always include system prompt — backend extracts it and prepends to DB messages.
// Only temp chats need conversation messages (persisted chats load from DB).
let messages = [
params?.system || $settings.system
? {
role: 'system',
content: `${params?.system ?? $settings?.system ?? ''}`
}
: undefined,
..._messages.map((message) => ({
...message,
content: processDetails(message.content),
// Include output for temp chats (backend will use it and strip before LLM)
...(message.output ? { output: message.output } : {})
}))
].filter((message) => message);
? { role: 'system', content: `${params?.system ?? $settings?.system ?? ''}` }
: undefined
].filter(Boolean);
if ($temporaryChatEnabled) {
messages = [
...messages,
..._messages.map((message) => ({
...message,
content: processDetails(message.content),
...(message.output ? { output: message.output } : {})
}))
].filter((message) => message);
messages = messages
.map((message, idx, arr) => {
const imageFiles = (message?.files ?? []).filter(
(file) => file.type === 'image' || (file?.content_type ?? '').startsWith('image/')
);
messages = messages
.map((message, idx, arr) => {
const imageFiles = (message?.files ?? []).filter(
(file) => file.type === 'image' || (file?.content_type ?? '').startsWith('image/')
);
return {
role: message.role,
// Preserve output items so backend can reconstruct tool_calls/tool-role messages (temp chats)
...(message.output ? { output: message.output } : {}),
...(message.role === 'user' && imageFiles.length > 0
? {
content: [
{
type: 'text',
text: message?.merged?.content ?? message.content
},
...imageFiles.map((file) => ({
type: 'image_url',
image_url: {
url: file.url
}
}))
]
}
: {
content: message?.merged?.content ?? message.content
})
};
})
.filter((message) => message?.role === 'user' || message?.content?.trim());
return {
role: message.role,
...(message.output ? { output: message.output } : {}),
...(message.role === 'user' && imageFiles.length > 0
? {
content: [
{
type: 'text',
text: message?.merged?.content ?? message.content
},
...imageFiles.map((file) => ({
type: 'image_url',
image_url: {
url: file.url
}
}))
]
}
: {
content: message?.merged?.content ?? message.content
})
};
})
.filter((message) => message?.role === 'user' || message?.content?.trim());
}
const toolIds = [];
const toolServerIds = [];
@ -2310,7 +2304,7 @@
{
stream: stream,
model: model.id,
messages: messages,
...(messages.length > 0 ? { messages } : {}),
params: {
...$settings?.params,
...params,
@ -2341,12 +2335,13 @@
model_item: $models.find((m) => m.id === model.id),
session_id: $socket?.id,
chat_id: $chatId,
chat_id: _chatId || undefined,
folder_id: $selectedFolder?.id ?? undefined,
id: responseMessageId,
parent_id: userMessage?.id ?? null,
parent_message: userMessage,
...(messageIdsMap ? { message_ids: messageIdsMap } : {}),
parent_id: userMessage?.parentId ?? null,
user_message: userMessage,
background_tasks: {
...(!$temporaryChatEnabled &&
@ -2403,10 +2398,22 @@
if (res.error) {
await handleOpenAIError(res.error, responseMessage);
} else {
// Backend returns task_ids (multi-model) or task_id (single model)
const newTaskIds = res.task_ids ?? (res.task_id ? [res.task_id] : []);
if (taskIds) {
taskIds.push(res.task_id);
taskIds.push(...newTaskIds);
} else {
taskIds = [res.task_id];
taskIds = newTaskIds;
}
// Backend returns chat_id for new chats — set store + URL
if (res.chat_id && $chatId !== res.chat_id) {
await chatId.set(res.chat_id);
if (!$temporaryChatEnabled) {
window.history.replaceState(history.state, '', `/c/${res.chat_id}`);
currentChatPage.set(1);
await chats.set(await getChatList(localStorage.token, $currentChatPage));
}
}
}
}
@ -2459,11 +2466,18 @@
const stopResponse = async (processQueue = true) => {
if (taskIds) {
for (const taskId of taskIds) {
const res = await stopTask(localStorage.token, taskId).catch((error) => {
if ($chatId) {
await stopTasksByChatId(localStorage.token, $chatId).catch((error) => {
toast.error(`${error}`);
return null;
});
} else {
for (const taskId of taskIds) {
const res = await stopTask(localStorage.token, taskId).catch((error) => {
toast.error(`${error}`);
return null;
});
}
}
taskIds = null;

View file

@ -46,11 +46,13 @@
}}
>
<Tooltip
content={DOMPurify.sanitize(marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description ?? ''
).replaceAll('\n', '<br>')
))}
content={DOMPurify.sanitize(
marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description ?? ''
).replaceAll('\n', '<br>')
)
)}
placement="right"
>
<img
@ -97,11 +99,13 @@
<div
class="mt-0.5 text-base font-normal text-gray-500 dark:text-gray-400 line-clamp-3 markdown"
>
{@html DOMPurify.sanitize(marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description
).replaceAll('\n', '<br>')
))}
{@html DOMPurify.sanitize(
marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description
).replaceAll('\n', '<br>')
)
)}
</div>
{#if models[selectedModelIdx]?.info?.meta?.user}
<div class="mt-0.5 text-sm font-normal text-gray-400 dark:text-gray-500">

View file

@ -399,7 +399,8 @@
{/if}
<iframe
src={serveUrl}
sandbox="allow-scripts allow-same-origin allow-downloads{($settings?.iframeSandboxAllowForms ?? false)
sandbox="allow-scripts allow-same-origin allow-downloads{($settings?.iframeSandboxAllowForms ??
false)
? ' allow-forms'
: ''}"
class="w-full h-full border-none bg-white"

View file

@ -343,7 +343,9 @@
}
chatInputElement?.setText(text);
chatInputElement?.focus();
if (!$showCallOverlay) {
chatInputElement?.focus();
}
if (text !== '') {
text = await inputVariableHandler(text);

View file

@ -102,9 +102,7 @@
<Cloud className="size-3.5" strokeWidth="2" />
{#if $selectedTerminalId && selectedLabel}
<span class="truncate text-[13px] max-w-[100px] sm:max-w-[150px]"
>{selectedLabel}</span
>
<span class="truncate text-[13px] max-w-[100px] sm:max-w-[150px]">{selectedLabel}</span>
{/if}
</button>
</Tooltip>

View file

@ -57,18 +57,119 @@
export let onSelect = (e) => {};
export let messagesCount: number | null = 20;
export let messagesCount: number | null = 8;
let messagesLoading = false;
// Off-screen message unloading. Heights are measured on scroll so spacers
// always match real sizes — no scroll jumps, no feedback loops needed.
const OVERSCAN = 3;
const DEFAULT_HEIGHT = 150;
let visibleStart = 0;
let visibleEnd = 0;
let messageHeights = new Map();
let topSpacerHeight = 0;
let bottomSpacerHeight = 0;
let pendingCull = null;
// Helper: get height for a message (cached or default)
const heightOf = (id) => messageHeights.get(id) ?? DEFAULT_HEIGHT;
/** Measure all currently rendered message elements and cache their heights */
const measureMessageHeights = () => {
const elements = document.getElementById('messages-container')?.querySelectorAll('[role="listitem"]');
if (!elements) return;
messageHeights = new Map([
...messageHeights,
...Array.from(elements)
.map((el, i) => [messages[visibleStart + i]?.id, el.getBoundingClientRect().height])
.filter(([id]) => id != null)
]);
};
/** Compute visible range from current scroll position and apply */
const updateVisibleRange = () => {
const container = document.getElementById('messages-container');
if (!container || messages.length === 0) return;
const st = container.scrollTop;
const ch = container.clientHeight;
// Build prefix sums from measured heights
const prefixSums = messages.reduce(
(acc, m) => [...acc, acc[acc.length - 1] + heightOf(m.id)],
[0]
);
const firstVisible = Math.max(0, prefixSums.findIndex((h) => h > st) - 1);
const lastVisible = prefixSums.findIndex((h) => h > st + ch);
// Only cull messages that have been measured (so spacer height is accurate)
// findIndex returns -1 when all are measured → no limit on culling
const firstUnmeasured = messages.findIndex((m) => !messageHeights.has(m.id));
const cullLimit = firstUnmeasured === -1 ? messages.length : firstUnmeasured;
visibleStart = Math.max(0, Math.min(firstVisible - OVERSCAN, cullLimit));
visibleEnd = Math.min(messages.length, (lastVisible === -1 ? messages.length : lastVisible) + OVERSCAN);
topSpacerHeight = prefixSums[visibleStart] ?? 0;
bottomSpacerHeight = (prefixSums[messages.length] ?? 0) - (prefixSums[visibleEnd] ?? 0);
};
/** Scroll handler: measure every frame, cull via rAF (same throttle as pendingRebuild) */
const handleContainerScroll = () => {
measureMessageHeights();
// Don't cull during progressive loading
if (messagesLoading) return;
if (!pendingCull) {
pendingCull = requestAnimationFrame(() => {
pendingCull = null;
updateVisibleRange();
});
}
};
let scrollListenerAttached = false;
const attachScrollListener = () => {
if (scrollListenerAttached) return;
const container = document.getElementById('messages-container');
if (!container) return;
container.addEventListener('scroll', handleContainerScroll, { passive: true });
scrollListenerAttached = true;
};
onMount(() => {
attachScrollListener();
});
onDestroy(() => {
const container = document.getElementById('messages-container');
if (container && scrollListenerAttached) {
container.removeEventListener('scroll', handleContainerScroll);
}
cancelAnimationFrame(pendingCull);
cancelAnimationFrame(pendingRebuild);
});
const loadMoreMessages = async () => {
// scroll slightly down to disable continuous loading
const element = document.getElementById('messages-container');
element.scrollTop = element.scrollTop + 100;
messagesLoading = true;
messagesCount += 20;
messagesCount += 8;
buildMessages();
// Show all messages during progressive loading (no culling)
visibleStart = 0;
visibleEnd = messages.length;
topSpacerHeight = 0;
bottomSpacerHeight = 0;
await tick();
messagesLoading = false;
@ -95,6 +196,7 @@
}
messages = _messages.reverse();
visibleEnd = messages.length;
};
// Throttle message list rebuilds to once per animation frame during streaming.
@ -113,6 +215,8 @@
cancelAnimationFrame(pendingRebuild);
pendingRebuild = null;
buildMessages();
// No explicit culling needed — scrollToBottom will fire a scroll event,
// which triggers handleContainerScroll → rAF → updateVisibleRange
} else if (_messages) {
// Content update (streaming) — throttle to once per frame
if (!pendingRebuild) {
@ -426,9 +530,7 @@
showMessage({ id: parentMessageId }, false);
};
onDestroy(() => {
cancelAnimationFrame(pendingRebuild);
});
const triggerScroll = () => {
if (autoScroll) {
@ -465,7 +567,13 @@
</Loader>
{/if}
<ul role="log" aria-live="polite" aria-relevant="additions" aria-atomic="false">
{#each messages as message, messageIdx (message.id)}
<!-- Top spacer: sum of cached heights for messages above visible range -->
{#if topSpacerHeight > 0}
<div style="height: {topSpacerHeight}px" aria-hidden="true" />
{/if}
{#each messages.slice(visibleStart, visibleEnd) as message, i (message.id)}
{@const messageIdx = visibleStart + i}
<Message
{chatId}
bind:history
@ -494,6 +602,11 @@
{topPadding}
/>
{/each}
<!-- Bottom spacer: sum of cached heights for messages below visible range -->
{#if bottomSpacerHeight > 0}
<div style="height: {bottomSpacerHeight}px" aria-hidden="true" />
{/if}
</ul>
</section>
<div class="pb-18" />

View file

@ -109,7 +109,7 @@
<del><svelte:self id={`${id}-del`} tokens={token.tokens} {onSourceClick} /></del>
{:else if token.type === 'inlineKatex'}
{#if token.text}
<KatexRenderer content={token.text} displayMode={false} />
<KatexRenderer content={token.text} displayMode={token?.displayMode ?? false} />
{/if}
{:else if token.type === 'iframe'}
<iframe

View file

@ -380,6 +380,7 @@
<ToolCallDisplay
id={`${id}-${tokenIdx}-${detailIdx}-tc`}
attributes={detailToken.attributes}
resultContent={getDetailTextContent(detailToken)}
grouped={true}
open={$settings?.expandDetails ?? false}
className="w-full space-y-1"
@ -428,6 +429,7 @@
<ToolCallDisplay
id={`${id}-${tokenIdx}-tc`}
attributes={token.attributes}
resultContent={getDetailTextContent(token)}
open={$settings?.expandDetails ?? false}
className="w-full space-y-1"
/>

View file

@ -1430,6 +1430,7 @@
: ''}"
style="fill: currentColor;"
alt={action.name}
draggable="false"
/>
</div>
{:else}

View file

@ -165,21 +165,25 @@
{#if models[selectedModelIdx]?.info?.meta?.description ?? null}
<Tooltip
className=" w-fit"
content={DOMPurify.sanitize(marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description ?? ''
).replaceAll('\n', '<br>')
))}
content={DOMPurify.sanitize(
marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description ?? ''
).replaceAll('\n', '<br>')
)
)}
placement="top"
>
<div
class="mt-0.5 px-2 text-sm font-normal text-gray-500 dark:text-gray-400 line-clamp-2 max-w-xl markdown"
>
{@html DOMPurify.sanitize(marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description ?? ''
).replaceAll('\n', '<br>')
))}
{@html DOMPurify.sanitize(
marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description ?? ''
).replaceAll('\n', '<br>')
)
)}
</div>
</Tooltip>

View file

@ -248,7 +248,7 @@
<hr class="border-gray-50 dark:border-gray-850/30 my-4" />
{#if $config?.features.enable_login_form}
{#if $config?.features.enable_login_form && $config?.features.enable_password_change_form}
<div class="mt-2">
<UpdatePassword />
</div>

View file

@ -77,7 +77,9 @@
}
$: args = decode(attributes?.arguments ?? '');
$: result = decode(attributes?.result ?? '');
export let resultContent: string = '';
$: result = resultContent || decode(attributes?.result ?? '');
$: files = parseJSONString(decode(attributes?.files ?? ''));
$: embeds = parseJSONString(decode(attributes?.embeds ?? ''));
$: isDone = attributes?.done === 'true';
@ -198,10 +200,10 @@
</div>
{:else}
<div class="tool-call-body w-full max-w-none!">
<Markdown
id={`${componentId}-tool-call-args`}
content={`\`\`\`json\n${formatJSONString(args)}\n\`\`\``}
/>
<pre
class="text-xs text-gray-600 dark:text-gray-300 whitespace-pre font-mono bg-gray-50 dark:bg-gray-900 rounded-lg p-2.5 overflow-x-auto">{formatJSONString(
args
)}</pre>
</div>
{/if}
</div>
@ -217,10 +219,12 @@
</div>
<div class="w-full max-w-none!">
{#if typeof parsedResult === 'object' && parsedResult !== null}
<Markdown
id={`${componentId}-tool-call-result`}
content={`\`\`\`json\n${JSON.stringify(parsedResult, null, 2)}\n\`\`\``}
/>
<pre
class="text-xs text-gray-600 dark:text-gray-300 whitespace-pre font-mono bg-gray-50 dark:bg-gray-900 rounded-lg p-2.5 overflow-x-auto">{JSON.stringify(
parsedResult,
null,
2
)}</pre>
{:else}
{@const resultStr = String(parsedResult)}
{@const isTruncated = resultStr.length > RESULT_PREVIEW_LIMIT && !expandedResult}

View file

@ -88,10 +88,21 @@
let mouseOver = false;
// Local state: tracks the last updatedAt seen while the user was viewing
// this chat. Survives prop refreshes from sidebar data re-fetches that
// would overwrite the `lastReadAt` prop with a stale server value.
let viewedAt: number | null = null;
$: if (id === $chatId) {
viewedAt = updatedAt ?? Date.now() / 1000;
}
$: effectiveReadAt = Math.max(lastReadAt ?? 0, viewedAt ?? 0) || null;
$: unread =
id !== $chatId &&
!$activeChatIds.has(id) &&
(lastReadAt === null || (updatedAt !== null && updatedAt > lastReadAt));
(effectiveReadAt === null || (updatedAt !== null && updatedAt > effectiveReadAt));
const loadChat = async () => {
if (!chat) {

View file

@ -10,7 +10,6 @@
import { splitStream } from '$lib/utils';
import Spinner from '$lib/components/common/Spinner.svelte';
const i18n = getContext('i18n');

View file

@ -7,7 +7,6 @@
const dispatch = createEventDispatcher();
import Modal from '$lib/components/common/Modal.svelte';
import RichTextInput from '$lib/components/common/RichTextInput.svelte';
import XMark from '$lib/components/icons/XMark.svelte';
import MicSolid from '$lib/components/icons/MicSolid.svelte';
import Tooltip from '$lib/components/common/Tooltip.svelte';
@ -57,7 +56,7 @@
<div class="shrink-0 w-full flex justify-between items-center">
<div class="w-full">
<input
class="w-full text-3xl font-medium bg-transparent outline-hidden"
class="w-full text-3xl bg-transparent outline-hidden"
type="text"
bind:value={name}
placeholder={$i18n.t('Title')}
@ -67,17 +66,17 @@
</div>
<div class=" flex-1 w-full h-full">
<RichTextInput
<textarea
class="w-full h-full min-h-[200px] bg-transparent outline-none resize-none text-base leading-relaxed placeholder:text-gray-300 dark:placeholder:text-gray-600"
bind:value={content}
placeholder={$i18n.t('Write something...')}
preserveBreaks={true}
/>
</div>
</div>
</div>
<div
class="flex flex-row items-center justify-end text-sm font-medium shrink-0 mt-1 p-4 gap-1.5"
class="flex flex-row items-center justify-end text-sm shrink-0 mt-1 p-4 gap-1.5"
>
<div class="">
{#if voiceInput}

View file

@ -25,6 +25,7 @@
import DefaultFeatures from './DefaultFeatures.svelte';
import BuiltinTools from './BuiltinTools.svelte';
import PromptSuggestions from './PromptSuggestions.svelte';
import TerminalSelector from './TerminalSelector.svelte';
import AccessControlModal from '../common/AccessControlModal.svelte';
import LockClosed from '$lib/components/icons/LockClosed.svelte';
import { updateModelAccessGrants } from '$lib/apis/models';
@ -102,6 +103,7 @@
let actionIds = [];
let accessGrants = [];
let terminalId = '';
let tts = { voice: '' };
const submitHandler = async () => {
@ -206,6 +208,14 @@
}
}
if (terminalId) {
info.meta.terminalId = terminalId;
} else {
if (info.meta.terminalId) {
delete info.meta.terminalId;
}
}
if (tts.voice !== '') {
if (!info.meta.tts) info.meta.tts = {};
info.meta.tts.voice = tts.voice;
@ -316,6 +326,7 @@
capabilities = { ...capabilities, ...(model?.meta?.capabilities ?? {}) };
defaultFeatureIds = model?.meta?.defaultFeatureIds ?? defaultFeatureIds;
builtinTools = model?.meta?.builtinTools ?? builtinTools;
terminalId = model?.meta?.terminalId ?? '';
tts = { voice: model?.meta?.tts?.voice ?? '' };
accessGrants = model?.access_grants ?? [];
@ -828,6 +839,10 @@
</div>
{/if}
<div class="my-4">
<TerminalSelector bind:terminalId />
</div>
<div class="my-4">
<div class="flex w-full justify-between mb-1">
<div class="self-center text-xs font-medium text-gray-500">

View file

@ -0,0 +1,30 @@
<script lang="ts">
import { getContext, onMount } from 'svelte';
import { getTerminalServers, type TerminalServer } from '$lib/apis/terminal';
const i18n = getContext('i18n');
export let terminalId: string = '';
let terminals: TerminalServer[] = [];
onMount(async () => {
terminals = await getTerminalServers(localStorage.token);
});
</script>
{#if terminals.length > 0}
<div class="flex w-full justify-between mb-1">
<div class="self-center text-xs font-medium text-gray-500">{$i18n.t('Terminal')}</div>
</div>
<select
class="w-full text-sm bg-transparent outline-hidden cursor-pointer"
bind:value={terminalId}
>
<option value="">{$i18n.t('None')}</option>
{#each terminals as terminal (terminal.id)}
<option value={terminal.id}>{terminal.name || terminal.id}</option>
{/each}
</select>
{/if}

View file

@ -183,7 +183,7 @@
"Are you sure you want to archive all chats? This action cannot be undone.": "Tem certeza de que deseja arquivar todos os chats? Esta ação não pode ser desfeita.",
"Are you sure you want to clear all memories? This action cannot be undone.": "Tem certeza de que deseja apagar todas as memórias? Esta ação não pode ser desfeita.",
"Are you sure you want to delete \"{{NAME}}\"?": "Tem certeza de que deseja excluir \"{{NAME}}\"?",
"Are you sure you want to delete **{{modelName}}**?": "",
"Are you sure you want to delete **{{modelName}}**?": "Tem certeza de que deseja excluir **{{modelName}}**?",
"Are you sure you want to delete all chats? This action cannot be undone.": "Tem certeza de que deseja excluir todas as conversas? Esta ação não pode ser desfeita.",
"Are you sure you want to delete this channel?": "Tem certeza de que deseja excluir este canal?",
"Are you sure you want to delete this connection? This action cannot be undone.": "Tem certeza de que deseja excluir esta conexão? Esta ação não pode ser desfeita.",
@ -200,7 +200,7 @@
"Assistant": "Assistente",
"Async Embedding Processing": "Processamento de Embedding assíncrono",
"Attach File From Knowledge": "Anexar arquivo da base de conhecimento",
"Attach Files": "",
"Attach Files": "Anexar arquivos",
"Attach Knowledge": "Anexar Base de Conhecimento",
"Attach Notes": "Anexar Notas",
"Attach Webpage": "Anexar Página Web",
@ -223,13 +223,13 @@
"AUTOMATIC1111 Base URL": "URL Base AUTOMATIC1111",
"AUTOMATIC1111 Base URL is required.": "URL Base AUTOMATIC1111 é necessária.",
"Automatically inject system tools in native function calling mode (e.g., timestamps, memory, chat history, notes, etc.)": "Injetar automaticamente ferramentas do sistema no modo de chamada de função nativa (por exemplo, carimbos de data/hora, memória, histórico de chat, notas, etc.)",
"Automation": "",
"Automation created": "",
"Automation Name": "",
"Automation title": "",
"Automation triggered": "",
"Automation updated": "",
"Automations": "",
"Automation": "Automação",
"Automation created": "Automação criada",
"Automation Name": "Nome da automação",
"Automation title": "Título de automação",
"Automation triggered": "Automação acionada",
"Automation updated": "Automação atualizada",
"Automations": "Automações",
"Available list": "Lista disponível",
"Available models": "Modelos disponíveis",
"Available Tools": "Ferramentas disponíveis",
@ -260,7 +260,7 @@
"Boosting or penalizing specific tokens for constrained responses. Bias values will be clamped between -100 and 100 (inclusive). (Default: none)": "Aumentar ou penalizar tokens específicos para respostas restritas. Os valores de viés serão fixados entre -100 e 100 (inclusive). (Padrão: nenhum)",
"Brave": "Brave",
"Brave Search API Key": "Chave API do Brave Search",
"Break down complex requests into trackable steps": "",
"Break down complex requests into trackable steps": "Divida solicitações complexas em etapas rastreáveis.",
"Browse and query knowledge bases": "Navegue e consulte bases de conhecimento.",
"Builtin Tools": "Ferramentas integradas",
"Bullet List": "Lista com marcadores",
@ -394,7 +394,7 @@
"Concurrent Requests": "Solicitações simultâneas",
"Config": "Configuração",
"Config imported successfully": "Configuração importada com sucesso",
"Configuration": "",
"Configuration": "Configuração",
"Configure": "Configurar",
"Confirm": "Confirmar",
"Confirm Password": "Confirmar Senha",
@ -464,7 +464,7 @@
"Create new secret key": "Criar nova chave secreta",
"Create note": "Criar nota",
"Create Note": "Criar Nota",
"Create scheduled prompts that run automatically on a recurring basis.": "",
"Create scheduled prompts that run automatically on a recurring basis.": "Crie prompts agendados que sejam executados automaticamente de forma recorrente.",
"Create your first note by clicking on the plus button below.": "Crie sua primeira nota clicando no botão de adição abaixo.",
"Created at": "Criado em",
"Created At": "Criado Em",
@ -486,7 +486,7 @@
"Data Controls": "Controle de Dados",
"Database": "Banco de Dados",
"Datalab Marker API": "API do Marcador do Datalab",
"Day": "",
"Day": "Dia",
"DD/MM/YYYY": "DD/MM/AAAA",
"DDGS Backend": "Backend DDGS",
"December": "Dezembro",
@ -517,7 +517,7 @@
"Delete All": "Excluir tudo",
"Delete All Chats": "Excluir Todos os Chats",
"Delete all contents inside this folder": "Apague todo o conteúdo desta pasta.",
"Delete automation?": "",
"Delete automation?": "Excluir automação?",
"Delete Chat": "Excluir Chat",
"Delete chat?": "Excluir chat?",
"Delete File": "Excluir arquivo",
@ -662,7 +662,7 @@
"Embedding Concurrent Requests": "Solicitações Simultâneas de Embedding",
"Embedding Model": "Modelo de Embedding",
"Embedding Model Engine": "Motor do Modelo de Embedding",
"Emojis": "",
"Emojis": "Emojis",
"Empty message": "Mensagem vazia",
"Enable All": "Ativar tudo",
"Enable API Keys": "Habilitar Chaves de API",
@ -754,7 +754,7 @@
"Enter Perplexity Search API URL": "Insira a URL da API de pesquisa Perplexity",
"Enter Playwright Timeout": "Insira o tempo limite do Playwright",
"Enter Playwright WebSocket URL": "Insira a URL do WebSocket do Playwright",
"Enter prompt here.": "",
"Enter prompt here.": "Insira o prompt aqui.",
"Enter proxy URL (e.g. https://user:password@host:port)": "Insira a URL do proxy (por exemplo, https://usuário:senha@host:porta)",
"Enter reasoning effort": "Insira o esforço de raciocínio",
"Enter Score": "Digite a Pontuação",
@ -779,7 +779,7 @@
"Enter system prompt here": "Insira o prompt do sistema aqui",
"Enter Tavily API Key": "Digite a Chave API do Tavily",
"Enter Tavily Extract Depth": "Insira a profundidade de extração do Tavily",
"Enter the prompt instructions for this automation...": "",
"Enter the prompt instructions for this automation...": "Insira as instruções para esta automação...",
"Enter the public URL of your WebUI. This URL will be used to generate links in the notifications.": "Insira a URL pública da sua WebUI. Esta URL será usada para gerar links nas notificações.",
"Enter the URL of the function to import": "Digite a URL da função a ser importada",
"Enter the URL to import": "Digite a URL para importar",
@ -819,7 +819,7 @@
"Error accessing directory": "Erro ao acessar o diretório",
"Error accessing Google Drive: {{error}}": "Erro ao acessar o Google Drive: {{error}}",
"Error accessing media devices.": "Erro ao acessar dispositivos de mídia.",
"Error deleting model: {{error}}": "",
"Error deleting model: {{error}}": "Erro ao excluir o modelo: {{error}}",
"Error starting recording.": "Erro ao iniciar a gravação.",
"Error unloading model: {{error}}": "Erro ao descarregar modelo: {{error}}",
"Error uploading file: {{error}}": "Erro ao carregar o arquivo: {{error}}",
@ -837,7 +837,7 @@
"Execute code": "Executar código",
"Execute code for analysis": "Executar código para análise",
"Executing **{{NAME}}**...": "Executando **{{NAME}}**...",
"Execution Logs": "",
"Execution Logs": "Registros de execução",
"Expand": "Expandir",
"Experimental": "Experimental",
"Explain": "Explicar",
@ -847,8 +847,8 @@
"Export": "Exportar",
"Export All Archived Chats": "Exportar todos os chats arquivados",
"Export All Chats (All Users)": "Exportar Todos os Chats (Todos os Usuários)",
"Export as CSV": "",
"Export as JSON": "",
"Export as CSV": "Exportar como CSV",
"Export as JSON": "Exportar como JSON",
"Export chat (.json)": "Exportar chat (.json)",
"Export Chats": "Exportar Chats",
"Export Config": "Exportar Configuração",
@ -1103,7 +1103,7 @@
"Insert Suggestion Prompt to Input": "Inserir prompt de sugestão para entrada",
"Install from Github URL": "Instalar da URL do Github",
"Instant Auto-Send After Voice Transcription": "Envio Automático Instantâneo Após Transcrição de Voz",
"Instructions": "",
"Instructions": "Instruções",
"Integration": "Integração",
"Integrations": "Integrações",
"Interface": "Interface",
@ -1163,7 +1163,7 @@
"Last 90 days": "Últimos 90 dias",
"Last Active": "Última Atividade",
"Last Modified": "Última Modificação",
"Last ran": "",
"Last ran": "Última execução",
"Last reply": "Última resposta",
"LDAP": "LDAP",
"LDAP server updated": "Servidor LDAP atualizado",
@ -1270,7 +1270,7 @@
"Model '{{modelName}}' has been successfully downloaded.": "Modelo '{{modelName}}' foi baixado com sucesso.",
"Model '{{modelTag}}' is already in queue for downloading.": "Modelo '{{modelTag}}' já está na fila para download.",
"Model {{modelId}} not found": "Modelo {{modelId}} não encontrado",
"Model {{modelName}} deleted successfully": "",
"Model {{modelName}} deleted successfully": "Modelo {{modelName}} excluído com sucesso",
"Model {{modelName}} is not vision capable": "Modelo {{modelName}} não é capaz de visão",
"Model {{name}} is now {{status}}": "Modelo {{name}} está agora {{status}}",
"Model {{name}} is now hidden": "O modelo {{name}} agora está oculto",
@ -1320,11 +1320,11 @@
"Name": "Nome",
"Name and ID are required, please fill them out": "Nome e ID são obrigatórios, por favor preencha-os",
"Name your knowledge base": "Nome da sua base de conhecimento",
"Name, prompt, and model are required": "",
"Name, prompt, and model are required": "Nome, prompt e modelo são obrigatórios.",
"Native": "Nativo",
"Never": "",
"Never": "Nunca",
"New": "Novo",
"New Automation": "",
"New Automation": "Nova Automação",
"New Button": "Novo Botão",
"New Chat": "Novo Chat",
"New File": "Novo Arquivo",
@ -1343,11 +1343,11 @@
"New Webhook": "Novo Webhook",
"new-channel": "novo-canal",
"Next message": "Próxima mensagem",
"Next run": "",
"Next run": "Próxima execução",
"No access grants. Private to you.": "Sem permissões de acesso. Privacidade exclusiva para você.",
"No activity data": "Sem dados de atividade",
"No authentication": "Sem autenticação",
"No automations found": "",
"No automations found": "Nenhuma automação encontrada",
"No chats found": "Nenhum chat encontrado",
"No chats found for this user.": "Nenhum chat encontrado para este usuário.",
"No chats found.": "Nenhum chat encontrado.",
@ -1358,7 +1358,7 @@
"No data": "Sem dados",
"No data found": "Nenhum dado encontrado",
"No distance available": "Sem distância disponível",
"No execution logs available yet": "",
"No execution logs available yet": "Ainda não há registros de execução disponíveis.",
"No expiration can pose security risks.": "A ausência de expiração pode representar riscos de segurança.",
"No feedback found": "Nenhum feedback encontrado",
"No file selected": "Nenhum arquivo selecionado",
@ -1405,7 +1405,7 @@
"Not factually correct": "Não está factualmente correto",
"Not helpful": "Não é útil",
"Not Registered": "Não registrado",
"Not scheduled": "",
"Not scheduled": "Não agendado",
"Note": "Nota",
"Note deleted successfully": "Nota excluída com sucesso",
"Note: If you set a minimum score, the search will only return documents with a score greater than or equal to the minimum score.": "Nota: Se você definir uma pontuação mínima, a pesquisa retornará apenas documentos com pontuação igual ou superior à pontuação mínima.",
@ -1495,7 +1495,7 @@
"Password": "Senha",
"Passwords do not match.": "As senhas não coincidem.",
"Paste Large Text as File": "Cole Textos Longos como Arquivo",
"Paused": "",
"Paused": "Em pausa",
"PDF document (.pdf)": "Documento PDF (.pdf)",
"PDF Extract Images (OCR)": "Extrair Imagens do PDF (OCR)",
"PDF Loader Mode": "Modo de carregamento de PDF",
@ -1600,7 +1600,7 @@
"Reason": "Razão",
"Reasoning Effort": "Esforço de raciocínio",
"Reasoning Tags": "Tags de raciocínio",
"Recently Used": "",
"Recently Used": "Usado recentemente",
"Record": "Gravar",
"Record voice": "Gravar voz",
"Redirecting you to Open WebUI Community": "Redirecionando você para a Comunidade OpenWebUI",
@ -1636,7 +1636,7 @@
"Renamed to {{name}}": "Renomeado para {{name}}",
"Render Markdown in Previews": "Renderizar Markdown nas Pré-visualizações",
"Reorder Models": "Reordenar modelos",
"Repeats": "",
"Repeats": "Repetições",
"Reply": "Responder",
"Reply in Thread": "Responder no tópico",
"Reply to thread...": "Responder ao tópico...",
@ -1670,8 +1670,8 @@
"RTL": "Direita para Esquerda",
"Run": "Executar",
"Run All": "Executar Tudo",
"Run now": "",
"Run Now": "",
"Run now": "Executar agora",
"Run Now": "Executar Agora",
"Running": "Executando",
"Running...": "Executando...",
"Runs embedding tasks concurrently to speed up processing. Turn off if rate limits become an issue.": "Executa tarefas de incorporação simultaneamente para acelerar o processamento. Desative se os limites de taxa se tornarem um problema.",
@ -1682,15 +1682,15 @@
"Save Chat": "Salvar Chat",
"Saved": "Armazenado",
"Saving chat logs directly to your browser's storage is no longer supported. Please take a moment to download and delete your chat logs by clicking the button below. Don't worry, you can easily re-import your chat logs to the backend through": "Salvar registros de chat diretamente no armazenamento do seu navegador não é mais suportado. Por favor, reserve um momento para baixar e excluir seus registros de chat clicando no botão abaixo. Não se preocupe, você pode facilmente reimportar seus registros de chat para o backend através de",
"Schedule": "",
"Scheduled time must be in the future": "",
"Schedule": "Agendar",
"Scheduled time must be in the future": "O horário agendado deve ser no futuro.",
"Scroll On Branch Change": "Rolar na mudança de ramo",
"Search": "Pesquisar",
"Search a model": "Pesquisar um modelo",
"Search all emojis": "Pesquisar todos os emojis",
"Search and manage user memories": "Pesquisar e gerenciar memórias de usuários",
"Search and view user chat history": "Pesquise e visualize o histórico de chat do usuário",
"Search Automations": "",
"Search Automations": "Pesquisar Automações",
"Search Base": "Pesquisar Base",
"Search channels and channel messages": "Pesquisar canais e mensagens de canais",
"Search Chats": "Pesquisar Chats",
@ -1760,7 +1760,7 @@
"Select how to split message text for TTS requests": "Selecione como dividir o texto da mensagem para solicitações TTS",
"Select Knowledge": "Selecionar Conhecimento",
"Select Method": "Selecione o método",
"Select model": "",
"Select model": "Selecione o modelo",
"Select only one model to call": "Selecione apenas um modelo para chamar",
"Select view": "Selecionar visualização",
"Selected model: {{modelName}}": "Modelo selecionado: {{modelName}}",
@ -1870,7 +1870,7 @@
"Start of the channel": "Início do canal",
"Start Tag": "Tag inicial",
"Starting kernel...": "Iniciando kernel...",
"State": "",
"State": "Estado",
"Status": "Status",
"Status cleared successfully": "Status liberado com sucesso",
"Status updated successfully": "Status atualizado com sucesso",
@ -1921,10 +1921,10 @@
"Talk to Model": "Fale com o modelo",
"Tap to interrupt": "Toque para interromper",
"Task List": "Lista de tarefas",
"Task Management": "",
"Task Management": "Gerenciamento de Tarefas",
"Task Model": "Modelo de Tarefa",
"Tasks": "Tarefas",
"tasks completed": "",
"tasks completed": "tarefas concluídas",
"Tavily API Key": "Chave da API Tavily",
"Tavily Extract Depth": "Profundidade de extração do Tavily",
"Tell us more:": "Conte-nos mais:",
@ -1991,7 +1991,7 @@
"Tika": "Tika",
"Tika Server URL required.": "URL do servidor Tika necessária.",
"Tiktoken": "Tiktoken",
"Time": "",
"Time": "Tempo",
"Time & Calculation": "Tempo e Cálculo",
"Timeout": "Tempo limite",
"Title": "Título",
@ -2009,7 +2009,7 @@
"To select toolkits here, add them to the \"Tools\" workspace first.": "Para selecionar kits de ferramentas aqui, adicione-os ao espaço de trabalho \"Ferramentas\" primeiro.",
"Toast notifications for new updates": "Notificações de alerta para novas atualizações",
"Today": "Hoje",
"Today at": "",
"Today at": "Hoje em",
"Today at {{LOCALIZED_TIME}}": "Hoje às {{LOCALIZED_TIME}}",
"Toggle {{COUNT}} sources": "Alternar {{COUNT}} origens",
"Toggle 1 source": "Alternar 1 origem",
@ -2148,7 +2148,7 @@
"Waiting for upload...": "Aguardando upload...",
"Warning": "Aviso",
"Warning:": "Aviso:",
"Warning: Enabling this will allow users to run scheduled prompts automatically.": "",
"Warning: Enabling this will allow users to run scheduled prompts automatically.": "Aviso: Habilitar esta opção permitirá que os usuários executem solicitações agendadas automaticamente.",
"Warning: Enabling this will allow users to upload arbitrary code on the server.": "Aviso: Habilitar isso permitirá que os usuários façam upload de código arbitrário no servidor.",
"Warning: Jupyter execution enables arbitrary code execution, posing severe security risks—proceed with extreme caution.": "Aviso: a execução do Jupyter permite a execução de código arbitrário, o que representa sérios riscos de segurança. Prossiga com extremo cuidado.",
"Web": "Web",
@ -2183,7 +2183,7 @@
"Width": "Largura",
"Wikipedia": "Wikipédia",
"Won": "Ganhou",
"Working Directory": "",
"Working Directory": "Diretório de Trabalho",
"Works together with top-k. A higher value (e.g., 0.95) will lead to more diverse text, while a lower value (e.g., 0.5) will generate more focused and conservative text.": "Funciona em conjunto com o top-k. Um valor mais alto (por exemplo, 0,95) resultará em um texto mais diverso, enquanto um valor mais baixo (por exemplo, 0,5) gerará um texto mais focado e conservador.",
"Workspace": "Espaço de Trabalho",
"Workspace Permissions": "Permissões do espaço de trabalho",

View file

@ -520,7 +520,7 @@ export const copyToClipboard = async (text, html = null, formatted = false) => {
textArea.style.position = 'fixed';
document.body.appendChild(textArea);
textArea.focus();
textArea.focus({ preventScroll: true });
textArea.select();
try {
@ -923,8 +923,19 @@ export const processDetails = (content) => {
attributes[attributeMatch[1]] = attributeMatch[2];
}
// New format: result in body content; Old format: result in attribute
let resultText = '';
if (attributes.result) {
content = content.replace(match, unescapeHtml(attributes.result));
resultText = unescapeHtml(attributes.result);
} else {
// Extract body content (strip <summary>...</summary>)
const bodyMatch = match.match(/<summary>[\s\S]*?<\/summary>\s*([\s\S]*?)\s*<\/details>/i);
if (bodyMatch && bodyMatch[1].trim()) {
resultText = unescapeHtml(bodyMatch[1].trim());
}
}
if (resultText) {
content = content.replace(match, resultText);
}
}
}

View file

@ -60,7 +60,7 @@ function detailsTokenizer(src: string) {
}
function detailsStart(src: string) {
return src.match(/^<details>/) ? 0 : -1;
return src.match(/^<details[\s>]/) ? 0 : -1;
}
function detailsRenderer(token: any) {

View file

@ -66,6 +66,44 @@ function generateRegexRules(delimiters) {
const { inlineRule, blockRule } = generateRegexRules(DELIMITER_LIST);
const isAllowedTrailing = (src: string, i: number): boolean =>
i >= src.length || ALLOWED_SURROUNDING_CHARS_REGEX.test(src.charAt(i));
const isBlockBoundary = (src: string, i: number): boolean =>
/^(?:[ \t]*\r?\n|$)/.test(src.slice(i));
const findClosingDelimiter = (src: string, i: number): number =>
i >= src.length - 1
? -1
: src[i] === '\\'
? findClosingDelimiter(src, i + 2)
: src[i] === '$' && src[i + 1] === '$'
? i
: findClosingDelimiter(src, i + 1);
export const tokenizeDisplayMath = (
src: string,
type: 'inlineKatex' | 'blockKatex',
requireBlockBoundary = false
) => {
if (!src.startsWith('$$')) return;
const endIndex = findClosingDelimiter(src, 2);
if (endIndex === -1) return;
const raw = src.slice(0, endIndex + 2);
const text = raw.slice(2, -2);
const afterClose = endIndex + 2;
const validators: Array<() => boolean> = [
() => text.trim().length > 0,
() => isAllowedTrailing(src, afterClose),
() => !requireBlockBoundary || isBlockBoundary(src, afterClose)
];
return validators.every((v) => v()) ? { type, raw, text, displayMode: true } : undefined;
};
export default function (options = {}) {
return {
extensions: [inlineKatex(options), blockKatex(options)]
@ -102,6 +140,17 @@ function katexStart(src, displayMode: boolean) {
}
function katexTokenizer(src, tokens, displayMode: boolean) {
if (src.startsWith('$$')) {
const displayToken = tokenizeDisplayMath(
src,
displayMode ? 'blockKatex' : 'inlineKatex',
displayMode
);
if (displayToken) {
return displayToken;
}
}
const ruleReg = displayMode ? blockRule : inlineRule;
const type = displayMode ? 'blockKatex' : 'inlineKatex';

View file

@ -41,7 +41,11 @@ export function mentionExtension(opts: MentionOptions = {}) {
// mentionStart fires on every '<' in the document, making the tokenizer a hot path.
const trigger = opts.triggerChar ?? '@';
const re = new RegExp(`^<\\${trigger}([\\w.\\-:/]+)(?:\\|([^>]*))?>`);
const snapshot: MentionOptions = { triggerChar: trigger, className: opts.className, extraAttrs: opts.extraAttrs };
const snapshot: MentionOptions = {
triggerChar: trigger,
className: opts.className,
extraAttrs: opts.extraAttrs
};
return {
name: 'mention',

View file

@ -78,7 +78,8 @@ class OneDriveConfig {
const msalParams = {
auth: {
authority: `https://login.microsoftonline.com/${authorityEndpoint}`,
clientId: clientId
clientId: clientId,
redirectUri: window.location.origin
}
};

View file

@ -728,7 +728,8 @@
// Apply theme classes (mirrors logic from chat/Settings/General.svelte)
const themes = ['dark', 'light', 'oled-dark'];
let themeToApply = newTheme === 'oled-dark' ? 'dark' : newTheme === 'her' ? 'light' : newTheme;
let themeToApply =
newTheme === 'oled-dark' ? 'dark' : newTheme === 'her' ? 'light' : newTheme;
if (newTheme === 'system') {
themeToApply = window.matchMedia('(prefers-color-scheme: dark)').matches ? 'dark' : 'light';
}
@ -988,13 +989,15 @@
console.error('Error refreshing backend config:', error);
}
// Relay auth token to desktop app for API access
if (window.electronAPI?.send) {
window.electronAPI.send({
type: 'token:update',
token: localStorage.token
}).catch(() => {});
}
// Relay auth token to desktop app for API access
if (window.electronAPI?.send) {
window.electronAPI
.send({
type: 'token:update',
token: localStorage.token
})
.catch(() => {});
}
} else {
// Redirect Invalid Session User to /auth Page
localStorage.removeItem('token');