mirror of
https://github.com/open-webui/open-webui.git
synced 2026-09-12 23:02:35 +00:00
Merge branch 'dev' into refactor/serialize-output-on
This commit is contained in:
commit
8893857461
178 changed files with 5005 additions and 2158 deletions
24
LICENSE
24
LICENSE
|
|
@ -1,3 +1,5 @@
|
|||
Open WebUI License
|
||||
|
||||
Copyright (c) 2023- Open WebUI Inc. [Created by Timothy Jaeryang Baek]
|
||||
All rights reserved.
|
||||
|
||||
|
|
@ -15,11 +17,27 @@ modification, are permitted provided that the following conditions are met:
|
|||
contributors may be used to endorse or promote products derived from
|
||||
this software without specific prior written permission.
|
||||
|
||||
4. Notwithstanding any other provision of this License, and as a material condition of the rights granted herein, licensees are strictly prohibited from altering, removing, obscuring, or replacing any "Open WebUI" branding, including but not limited to the name, logo, or any visual, textual, or symbolic identifiers that distinguish the software and its interfaces, in any deployment or distribution, regardless of the number of users, except as explicitly set forth in Clauses 5 and 6 below.
|
||||
4. Notwithstanding any other provision of this License, and as a material
|
||||
condition of the rights granted herein, licensees are strictly prohibited
|
||||
from altering, removing, obscuring, or replacing any "Open WebUI"
|
||||
branding, including but not limited to the name, logo, or any visual,
|
||||
textual, or symbolic identifiers that distinguish the software and its
|
||||
interfaces, in any deployment or distribution, except in the following
|
||||
circumstances: (i) deployments or distributions where the total number
|
||||
of end users (defined as individual natural persons with direct access
|
||||
to the application) does not exceed fifty (50) within any rolling
|
||||
thirty (30) day period; (ii) the licensee has obtained specific prior
|
||||
written permission from the copyright holder; or (iii) where the
|
||||
licensee has obtained a duly executed enterprise license expressly
|
||||
permitting such modification. For all other cases, any removal or
|
||||
alteration of the "Open WebUI" branding shall constitute a material
|
||||
breach of license.
|
||||
|
||||
5. The branding restriction enumerated in Clause 4 shall not apply in the following limited circumstances: (i) deployments or distributions where the total number of end users (defined as individual natural persons with direct access to the application) does not exceed fifty (50) within any rolling thirty (30) day period; (ii) cases in which the licensee is an official contributor to the codebase—with a substantive code change successfully merged into the main branch of the official codebase maintained by the copyright holder—who has obtained specific prior written permission for branding adjustment from the copyright holder; or (iii) where the licensee has obtained a duly executed enterprise license expressly permitting such modification. For all other cases, any removal or alteration of the "Open WebUI" branding shall constitute a material breach of license.
|
||||
Materials governed by prior licenses retain those original license
|
||||
terms, as specified in LICENSE_HISTORY.
|
||||
|
||||
6. All code, modifications, or derivative works incorporated into this project prior to the incorporation of this branding clause remain licensed under the BSD 3-Clause License, and prior contributors retain all BSD-3 rights therein; if any such contributor requests the removal of their BSD-3-licensed code, the copyright holder will do so, and any replacement code will be licensed under the project's primary license then in effect. By contributing after this clause's adoption, you agree to the project's Contributor License Agreement (CLA) and to these updated terms for all new contributions.
|
||||
By contributing to this project, you agree to the project's Contributor
|
||||
License Agreement (CONTRIBUTOR_LICENSE_AGREEMENT).
|
||||
|
||||
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -91,6 +91,17 @@ 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:
|
||||
|
|
|
|||
|
|
@ -375,7 +375,7 @@ else:
|
|||
except Exception:
|
||||
DATABASE_POOL_RECYCLE = 3600
|
||||
|
||||
DATABASE_ENABLE_SQLITE_WAL = os.environ.get('DATABASE_ENABLE_SQLITE_WAL', 'False').lower() == 'true'
|
||||
DATABASE_ENABLE_SQLITE_WAL = os.environ.get('DATABASE_ENABLE_SQLITE_WAL', 'True').lower() == 'true'
|
||||
|
||||
DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL = os.environ.get('DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL', None)
|
||||
if DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL is not None:
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ from typing import Optional
|
|||
from aiocache import cached
|
||||
import aiohttp
|
||||
import anyio.to_thread
|
||||
import requests
|
||||
|
||||
from redis import Redis
|
||||
|
||||
|
||||
|
|
@ -46,7 +46,6 @@ from fastapi.staticfiles import StaticFiles
|
|||
from starlette_compress import CompressMiddleware
|
||||
|
||||
from starlette.exceptions import HTTPException as StarletteHTTPException
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.middleware.sessions import SessionMiddleware
|
||||
from starlette.responses import Response, StreamingResponse
|
||||
from starlette.datastructures import Headers
|
||||
|
|
@ -58,8 +57,15 @@ from starsessions import (
|
|||
from starsessions.stores.redis import RedisStore
|
||||
|
||||
from open_webui.utils import logger
|
||||
from open_webui.utils.asgi_middleware import (
|
||||
AuthTokenMiddleware,
|
||||
CommitSessionMiddleware,
|
||||
RedirectMiddleware,
|
||||
WebsocketUpgradeGuardMiddleware,
|
||||
)
|
||||
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 +73,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 +121,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 +385,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 +477,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 +563,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,13 +576,14 @@ from open_webui.tasks import (
|
|||
list_task_ids_by_item_id,
|
||||
create_task,
|
||||
stop_task,
|
||||
stop_item_tasks,
|
||||
list_tasks,
|
||||
) # Import from tasks.py
|
||||
|
||||
from open_webui.utils.redis import get_sentinels_from_env
|
||||
|
||||
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.constants import ERROR_MESSAGES, TASKS
|
||||
|
||||
if SAFE_MODE:
|
||||
print('SAFE MODE ENABLED')
|
||||
|
|
@ -622,7 +634,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 +730,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 +862,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
|
||||
|
|
@ -1346,149 +1364,19 @@ if ENABLE_COMPRESSION_MIDDLEWARE:
|
|||
app.add_middleware(CompressMiddleware)
|
||||
|
||||
|
||||
class RedirectMiddleware(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next):
|
||||
# Check if the request is a GET request
|
||||
if request.method == 'GET':
|
||||
path = request.url.path
|
||||
query_params = dict(parse_qs(urlparse(str(request.url)).query))
|
||||
|
||||
redirect_params = {}
|
||||
|
||||
# Check for the specific watch path and the presence of 'v' parameter
|
||||
if path.endswith('/watch') and 'v' in query_params:
|
||||
# Extract the first 'v' parameter
|
||||
youtube_video_id = query_params['v'][0]
|
||||
redirect_params['youtube'] = youtube_video_id
|
||||
|
||||
if 'shared' in query_params and len(query_params['shared']) > 0:
|
||||
# PWA share_target support
|
||||
|
||||
text = query_params['shared'][0]
|
||||
if text:
|
||||
urls = re.match(r'https://\S+', text)
|
||||
if urls:
|
||||
from open_webui.retrieval.loaders.youtube import _parse_video_id
|
||||
|
||||
if youtube_video_id := _parse_video_id(urls[0]):
|
||||
redirect_params['youtube'] = youtube_video_id
|
||||
else:
|
||||
redirect_params['load-url'] = urls[0]
|
||||
else:
|
||||
redirect_params['q'] = text
|
||||
|
||||
if redirect_params:
|
||||
redirect_url = f'/?{urlencode(redirect_params)}'
|
||||
return RedirectResponse(url=redirect_url)
|
||||
|
||||
# Proceed with the normal flow of other requests
|
||||
response = await call_next(request)
|
||||
return response
|
||||
|
||||
|
||||
# All HTTP middlewares below are pure-ASGI implementations. The previous
|
||||
# `BaseHTTPMiddleware` / `@app.middleware('http')` versions wrapped the
|
||||
# downstream app in an anyio task group whose cancel scope cancelled
|
||||
# in-flight DB calls (and any other awaits) on client disconnect /
|
||||
# response completion — which surfaced as noisy SQLAlchemy
|
||||
# `terminate_force_close` tracebacks under aiosqlite and as random
|
||||
# CancelledError storms across the request path. See
|
||||
# `open_webui.utils.asgi_middleware` for the rationale.
|
||||
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)
|
||||
# log.debug("Commit session after request")
|
||||
try:
|
||||
ScopedSession.commit()
|
||||
finally:
|
||||
# CRITICAL: remove() returns the connection to the pool.
|
||||
# Without this, connections remain "checked out" and accumulate
|
||||
# as "idle in transaction" in PostgreSQL.
|
||||
ScopedSession.remove()
|
||||
return response
|
||||
|
||||
|
||||
@app.middleware('http')
|
||||
async def check_url(request: Request, call_next):
|
||||
start_time = int(time.time())
|
||||
request.state.token = get_http_authorization_cred(request.headers.get('Authorization'))
|
||||
# Fallback to cookie token for browser sessions
|
||||
if request.state.token is None and request.cookies.get('token'):
|
||||
from fastapi.security import HTTPAuthorizationCredentials
|
||||
|
||||
request.state.token = HTTPAuthorizationCredentials(scheme='Bearer', credentials=request.cookies.get('token'))
|
||||
|
||||
# Fallback to x-api-key header for Anthropic Messages API routes
|
||||
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
|
||||
|
||||
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)
|
||||
process_time = int(time.time()) - start_time
|
||||
response.headers['X-Process-Time'] = str(process_time)
|
||||
return response
|
||||
|
||||
|
||||
@app.middleware('http')
|
||||
async def inspect_websocket(request: Request, call_next):
|
||||
if '/ws/socket.io' in request.url.path and request.query_params.get('transport') == 'websocket':
|
||||
upgrade = (request.headers.get('Upgrade') or '').lower()
|
||||
connection = (request.headers.get('Connection') or '').lower().split(',')
|
||||
# Check that there's the correct headers for an upgrade, else reject the connection
|
||||
# This is to work around this upstream issue: https://github.com/miguelgrinberg/python-engineio/issues/367
|
||||
if upgrade != 'websocket' or 'upgrade' not in connection:
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
content={'detail': 'Invalid WebSocket upgrade request'},
|
||||
)
|
||||
return await call_next(request)
|
||||
app.add_middleware(CommitSessionMiddleware)
|
||||
app.add_middleware(AuthTokenMiddleware, fastapi_app=app)
|
||||
app.add_middleware(WebsocketUpgradeGuardMiddleware)
|
||||
|
||||
|
||||
app.add_middleware(
|
||||
|
|
@ -1560,6 +1448,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 +1616,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 +1662,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
|
||||
|
|
@ -1797,24 +1827,26 @@ async def chat_completion(
|
|||
detail=str(e),
|
||||
)
|
||||
|
||||
async def process_chat(request, form_data, user, metadata, model):
|
||||
async def process_chat(request, form_data, user, metadata, model, tasks=None):
|
||||
try:
|
||||
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 +1855,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 +1875,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,20 +1926,72 @@ 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 idx, (target_model_id, assistant_message_id) in enumerate(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)
|
||||
|
||||
# Only the first model runs title/tags generation;
|
||||
# subsequent models only run follow-ups.
|
||||
task_id, _ = await create_task(
|
||||
request.app.state.redis,
|
||||
process_chat(
|
||||
request,
|
||||
model_form_data,
|
||||
user,
|
||||
per_model_metadata,
|
||||
resolved_model,
|
||||
tasks
|
||||
if idx == 0
|
||||
else {
|
||||
k: v
|
||||
for k, v in (tasks or {}).items()
|
||||
if k not in (TASKS.TITLE_GENERATION, TASKS.TAGS_GENERATION)
|
||||
}
|
||||
or None,
|
||||
),
|
||||
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:
|
||||
return await process_chat(request, form_data, user, metadata, model)
|
||||
# Legacy/direct: single model, synchronous
|
||||
metadata['message_id'] = list(message_ids.values())[0]
|
||||
return await process_chat(request, form_data, user, metadata, model, tasks)
|
||||
|
||||
|
||||
# Alias for chat_completion (Legacy)
|
||||
|
|
@ -1979,6 +2065,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 +2100,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 +2109,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 +2131,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 +2200,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 +2411,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 +2473,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 +2624,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,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,23 @@
|
|||
"""Add is_pinned to note table
|
||||
|
||||
Revision ID: e1f2a3b4c5d6
|
||||
Revises: b7c8d9e0f1a2
|
||||
Create Date: 2026-04-14 22:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = 'e1f2a3b4c5d6'
|
||||
down_revision = 'b7c8d9e0f1a2'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.add_column('note', sa.Column('is_pinned', sa.Boolean(), nullable=True))
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_column('note', 'is_pinned')
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
@ -476,9 +464,8 @@ class ChatMessageTable:
|
|||
async with get_async_db_context(db) as db:
|
||||
from open_webui.models.groups import GroupMember
|
||||
|
||||
stmt = (
|
||||
select(ChatMessage.user_id, func.count(ChatMessage.id).label('count'))
|
||||
.filter(~ChatMessage.user_id.like('shared-%'))
|
||||
stmt = select(ChatMessage.user_id, func.count(ChatMessage.id).label('count')).filter(
|
||||
~ChatMessage.user_id.like('shared-%')
|
||||
)
|
||||
|
||||
if start_date:
|
||||
|
|
@ -503,9 +490,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:
|
||||
|
|
@ -532,13 +518,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:
|
||||
|
|
@ -582,13 +565,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:
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()),
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import uuid
|
|||
from typing import Optional
|
||||
from functools import lru_cache
|
||||
|
||||
from sqlalchemy import select, delete, update, or_, func, cast
|
||||
from sqlalchemy import Boolean, select, delete, update, or_, func, cast
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, get_async_db_context
|
||||
from open_webui.models.groups import Groups
|
||||
|
|
@ -29,6 +29,7 @@ class Note(Base):
|
|||
title = Column(Text)
|
||||
data = Column(JSON, nullable=True)
|
||||
meta = Column(JSON, nullable=True)
|
||||
is_pinned = Column(Boolean, default=False, nullable=True)
|
||||
|
||||
created_at = Column(BigInteger)
|
||||
updated_at = Column(BigInteger)
|
||||
|
|
@ -43,6 +44,7 @@ class NoteModel(BaseModel):
|
|||
title: str
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
is_pinned: Optional[bool] = False
|
||||
|
||||
access_grants: list[AccessGrantModel] = Field(default_factory=list)
|
||||
|
||||
|
|
@ -77,6 +79,7 @@ class NoteItemResponse(BaseModel):
|
|||
id: str
|
||||
title: str
|
||||
data: Optional[dict]
|
||||
is_pinned: Optional[bool] = False
|
||||
updated_at: int
|
||||
created_at: int
|
||||
user: Optional[UserResponse] = None
|
||||
|
|
@ -113,7 +116,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 +164,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 +225,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 +243,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),
|
||||
)
|
||||
)
|
||||
|
|
@ -305,6 +314,39 @@ class NoteTable:
|
|||
await db.commit()
|
||||
return await self._to_note_model(note, db=db) if note else None
|
||||
|
||||
async def toggle_note_pinned_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[NoteModel]:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Note).filter(Note.id == id))
|
||||
note = result.scalars().first()
|
||||
if not note:
|
||||
return None
|
||||
note.is_pinned = not note.is_pinned
|
||||
note.updated_at = int(time.time_ns())
|
||||
await db.commit()
|
||||
return await self._to_note_model(note, db=db)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def get_pinned_notes_by_user_id(
|
||||
self,
|
||||
user_id: str,
|
||||
permission: str = 'read',
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[NoteModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = [group.id for group in user_groups]
|
||||
|
||||
stmt = select(Note).filter(Note.is_pinned == True).order_by(Note.updated_at.desc())
|
||||
stmt = self._has_permission(db, stmt, {'user_id': user_id, 'group_ids': user_group_ids}, permission)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
notes = result.scalars().all()
|
||||
note_ids = [note.id for note in notes]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db)
|
||||
return [await self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
|
||||
|
||||
async def delete_note_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import asyncio
|
||||
import requests
|
||||
import logging
|
||||
import ftfy
|
||||
|
|
@ -238,6 +239,18 @@ class Loader:
|
|||
|
||||
return [Document(page_content=ftfy.fix_text(doc.page_content), metadata=doc.metadata) for doc in docs]
|
||||
|
||||
async def aload(self, filename: str, file_content_type: str, file_path: str) -> list[Document]:
|
||||
"""
|
||||
Async wrapper around `load`.
|
||||
|
||||
Document loaders dispatched by `_get_loader` (PyMuPDF, Unstructured,
|
||||
python-docx, Tika, etc.) are uniformly synchronous and CPU/IO-bound.
|
||||
Calling `load` directly from an async handler would block the event
|
||||
loop for the entire parse — minutes for large PDFs. This offloads
|
||||
the work to a worker thread so the loop stays responsive.
|
||||
"""
|
||||
return await asyncio.to_thread(self.load, filename, file_content_type, file_path)
|
||||
|
||||
def _is_text_file(self, file_ext: str, file_content_type: str) -> bool:
|
||||
return file_ext in known_source_ext or (
|
||||
file_content_type
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from langchain_community.retrievers import BM25Retriever
|
|||
from langchain_core.documents import Document
|
||||
|
||||
from open_webui.config import VECTOR_DB
|
||||
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
||||
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
|
||||
|
||||
|
||||
|
|
@ -121,7 +122,7 @@ class VectorSearchRetriever(BaseRetriever):
|
|||
run_manager: CallbackManagerForRetrieverRun,
|
||||
) -> list[Document]:
|
||||
embedding = await self.embedding_function(query, RAG_EMBEDDING_QUERY_PREFIX)
|
||||
result = VECTOR_DB_CLIENT.search(
|
||||
result = await ASYNC_VECTOR_DB_CLIENT.search(
|
||||
collection_name=self.collection_name,
|
||||
vectors=[embedding],
|
||||
limit=self.top_k,
|
||||
|
|
@ -488,16 +489,24 @@ async def query_collection_with_hybrid_search(
|
|||
) -> dict:
|
||||
results = []
|
||||
error = False
|
||||
# Fetch collection data once per collection sequentially
|
||||
# Avoid fetching the same data multiple times later
|
||||
collection_results = {}
|
||||
for collection_name in collection_names:
|
||||
# Fetch every collection's contents once up front so the
|
||||
# per-query/per-document loop below can reuse them. Each fetch
|
||||
# offloads to a worker thread, so run them concurrently with
|
||||
# `asyncio.gather` instead of awaiting them serially — otherwise
|
||||
# latency scales linearly with `len(collection_names)`.
|
||||
log.debug(
|
||||
'query_collection_with_hybrid_search: prefetching %d collections',
|
||||
len(collection_names),
|
||||
)
|
||||
|
||||
async def _fetch_collection(name: str):
|
||||
try:
|
||||
log.debug(f'query_collection_with_hybrid_search:VECTOR_DB_CLIENT.get:collection {collection_name}')
|
||||
collection_results[collection_name] = VECTOR_DB_CLIENT.get(collection_name=collection_name)
|
||||
return name, await ASYNC_VECTOR_DB_CLIENT.get(collection_name=name)
|
||||
except Exception as e:
|
||||
log.exception(f'Failed to fetch collection {collection_name}: {e}')
|
||||
collection_results[collection_name] = None
|
||||
log.exception(f'Failed to fetch collection {name}: {e}')
|
||||
return name, None
|
||||
|
||||
collection_results = dict(await asyncio.gather(*(_fetch_collection(name) for name in collection_names)))
|
||||
|
||||
log.info(f'Starting hybrid search for {len(queries)} queries in {len(collection_names)} collections...')
|
||||
|
||||
|
|
@ -1140,7 +1149,9 @@ async def get_sources_from_items(
|
|||
|
||||
try:
|
||||
if full_context:
|
||||
query_result = get_all_items_from_collections(collection_names)
|
||||
# Sync helper makes blocking VECTOR_DB_CLIENT calls;
|
||||
# offload so the async caller's event loop stays free.
|
||||
query_result = await asyncio.to_thread(get_all_items_from_collections, collection_names)
|
||||
else:
|
||||
query_result = await query_collection(
|
||||
request,
|
||||
|
|
|
|||
129
backend/open_webui/retrieval/vector/async_client.py
Normal file
129
backend/open_webui/retrieval/vector/async_client.py
Normal file
|
|
@ -0,0 +1,129 @@
|
|||
"""
|
||||
Async facade over the synchronous VECTOR_DB_CLIENT.
|
||||
|
||||
The vector DB backends bundled with Open WebUI (Chroma, pgvector, Qdrant,
|
||||
Milvus, OpenSearch, Pinecone, Weaviate, …) all expose a uniformly
|
||||
synchronous API. Each method performs blocking network or disk I/O — and
|
||||
some, like `insert`/`upsert`, can run for several seconds.
|
||||
|
||||
When such a sync method is awaited from an async route handler, it blocks
|
||||
the event loop for its entire duration, freezing every other in-flight
|
||||
HTTP request, websocket message and background task.
|
||||
|
||||
This module wraps the sync client in an `AsyncVectorDBClient` that
|
||||
transparently dispatches each call to a worker thread via
|
||||
`asyncio.to_thread`. Async callers can `await ASYNC_VECTOR_DB_CLIENT.x(...)`
|
||||
in place of `VECTOR_DB_CLIENT.x(...)` and the loop stays responsive.
|
||||
|
||||
The original `VECTOR_DB_CLIENT` is unchanged, so callers already running
|
||||
inside `run_in_threadpool` (e.g. `save_docs_to_vector_db`) are not
|
||||
affected.
|
||||
|
||||
Thread-safety expectations
|
||||
--------------------------
|
||||
Every async caller now invokes `VECTOR_DB_CLIENT` from a worker thread
|
||||
rather than the event-loop thread, and many can run concurrently. The
|
||||
sync client (and its underlying backend driver) is therefore expected
|
||||
to be safe for concurrent use across threads, which is the standard
|
||||
contract for the bundled drivers (chroma, pgvector via SQLAlchemy
|
||||
pool, qdrant-client, opensearch-py, …). This is *not* a new exposure
|
||||
introduced by this facade — `save_docs_to_vector_db` already called
|
||||
the sync client from `run_in_threadpool`, so concurrent threaded
|
||||
access has always been a requirement of the codebase. Adding a global
|
||||
serialization lock here would defeat the responsiveness this facade
|
||||
exists to provide; any backend that genuinely cannot tolerate
|
||||
concurrent access should grow its own internal serialization.
|
||||
|
||||
API surface
|
||||
-----------
|
||||
Method signatures mirror `VectorDBBase` exactly. This is deliberate:
|
||||
permissive `*args/**kwargs` forwarding hides typos at the call site
|
||||
(an earlier revision of this file shipped that, and a `metadata=`
|
||||
typo silently broke an entire endpoint until explicit signatures
|
||||
surfaced it). Callers that need a backend-specific parameter not on
|
||||
`VectorDBBase` should reach for the `.sync` escape hatch and wrap
|
||||
their own `asyncio.to_thread`, e.g. ::
|
||||
|
||||
await asyncio.to_thread(
|
||||
ASYNC_VECTOR_DB_CLIENT.sync.some_backend_specific_op,
|
||||
collection_name, special_kwarg=value,
|
||||
)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
|
||||
from open_webui.retrieval.vector.main import (
|
||||
GetResult,
|
||||
SearchResult,
|
||||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
|
||||
|
||||
class AsyncVectorDBClient:
|
||||
"""Awaitable mirror of `VectorDBBase` that off-loads each call to a thread.
|
||||
|
||||
Method signatures mirror `VectorDBBase` exactly so static analysis
|
||||
catches bad kwargs at the call site instead of letting them surface
|
||||
deep inside the worker thread (where the resulting ``TypeError`` is
|
||||
typically swallowed by surrounding ``try/except``).
|
||||
"""
|
||||
|
||||
def __init__(self, sync_client: VectorDBBase) -> None:
|
||||
self._sync = sync_client
|
||||
|
||||
@property
|
||||
def sync(self) -> VectorDBBase:
|
||||
"""Escape hatch for code that must call the sync client directly
|
||||
(e.g. already inside a worker thread)."""
|
||||
return self._sync
|
||||
|
||||
async def has_collection(self, collection_name: str) -> bool:
|
||||
return await asyncio.to_thread(self._sync.has_collection, collection_name)
|
||||
|
||||
async def delete_collection(self, collection_name: str) -> None:
|
||||
return await asyncio.to_thread(self._sync.delete_collection, collection_name)
|
||||
|
||||
async def insert(self, collection_name: str, items: List[VectorItem]) -> None:
|
||||
return await asyncio.to_thread(self._sync.insert, collection_name, items)
|
||||
|
||||
async def upsert(self, collection_name: str, items: List[VectorItem]) -> None:
|
||||
return await asyncio.to_thread(self._sync.upsert, collection_name, items)
|
||||
|
||||
async def search(
|
||||
self,
|
||||
collection_name: str,
|
||||
vectors: List[List[Union[float, int]]],
|
||||
filter: Optional[Dict] = None,
|
||||
limit: int = 10,
|
||||
) -> Optional[SearchResult]:
|
||||
return await asyncio.to_thread(self._sync.search, collection_name, vectors, filter, limit)
|
||||
|
||||
async def query(
|
||||
self,
|
||||
collection_name: str,
|
||||
filter: Dict,
|
||||
limit: Optional[int] = None,
|
||||
) -> Optional[GetResult]:
|
||||
return await asyncio.to_thread(self._sync.query, collection_name, filter, limit)
|
||||
|
||||
async def get(self, collection_name: str) -> Optional[GetResult]:
|
||||
return await asyncio.to_thread(self._sync.get, collection_name)
|
||||
|
||||
async def delete(
|
||||
self,
|
||||
collection_name: str,
|
||||
ids: Optional[List[str]] = None,
|
||||
filter: Optional[Dict] = None,
|
||||
) -> None:
|
||||
return await asyncio.to_thread(self._sync.delete, collection_name, ids, filter)
|
||||
|
||||
async def reset(self) -> None:
|
||||
return await asyncio.to_thread(self._sync.reset)
|
||||
|
||||
|
||||
ASYNC_VECTOR_DB_CLIENT = AsyncVectorDBClient(VECTOR_DB_CLIENT)
|
||||
|
|
@ -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,
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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,7 @@ 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
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
|
|
|||
|
|
@ -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}')
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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', [])
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -22,10 +22,10 @@ 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
|
||||
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
||||
|
||||
from open_webui.models.channels import Channels
|
||||
from open_webui.models.users import Users
|
||||
|
|
@ -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,
|
||||
|
|
@ -110,10 +124,10 @@ async def process_uploaded_file(
|
|||
stt_supported_content_types = getattr(request.app.state.config, 'STT_SUPPORTED_CONTENT_TYPES', [])
|
||||
|
||||
if strict_match_mime_type(stt_supported_content_types, content_type):
|
||||
file_path_processed = Storage.get_file(file_path)
|
||||
file_path_processed = await asyncio.to_thread(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)
|
||||
|
|
@ -225,7 +242,8 @@ async def upload_file_handler(
|
|||
id = str(uuid.uuid4())
|
||||
name = filename
|
||||
filename = f'{id}_{filename}'
|
||||
contents, file_path = Storage.upload_file(
|
||||
contents, file_path = await asyncio.to_thread(
|
||||
Storage.upload_file,
|
||||
file.file,
|
||||
filename,
|
||||
{
|
||||
|
|
@ -389,8 +407,8 @@ async def delete_all_files(user=Depends(get_admin_user), db: AsyncSession = Depe
|
|||
result = await Files.delete_all_files(db=db)
|
||||
if result:
|
||||
try:
|
||||
Storage.delete_all_files()
|
||||
VECTOR_DB_CLIENT.reset()
|
||||
await asyncio.to_thread(Storage.delete_all_files)
|
||||
await ASYNC_VECTOR_DB_CLIENT.reset()
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
log.error('Error deleting files')
|
||||
|
|
@ -495,7 +513,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 +560,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,
|
||||
|
|
@ -558,9 +578,9 @@ async def update_file_data_content_by_id(
|
|||
for knowledge in knowledges:
|
||||
try:
|
||||
# Remove old embeddings for this file from the KB collection
|
||||
VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'file_id': id})
|
||||
await ASYNC_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,
|
||||
|
|
@ -599,7 +619,7 @@ async def get_file_content_by_id(
|
|||
|
||||
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db):
|
||||
try:
|
||||
file_path = Storage.get_file(file.path)
|
||||
file_path = await asyncio.to_thread(Storage.get_file, file.path)
|
||||
file_path = Path(file_path)
|
||||
|
||||
# Check if the file already exists in the cache
|
||||
|
|
@ -646,7 +666,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:
|
||||
|
|
@ -664,7 +686,7 @@ async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user),
|
|||
|
||||
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db):
|
||||
try:
|
||||
file_path = Storage.get_file(file.path)
|
||||
file_path = await asyncio.to_thread(Storage.get_file, file.path)
|
||||
file_path = Path(file_path)
|
||||
|
||||
# Check if the file already exists in the cache
|
||||
|
|
@ -693,7 +715,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:
|
||||
|
|
@ -711,7 +735,7 @@ async def get_file_content_by_id(id: str, user=Depends(get_verified_user), db: A
|
|||
headers = {'Content-Disposition': f"attachment; filename*=UTF-8''{encoded_filename}"}
|
||||
|
||||
if file_path:
|
||||
file_path = Storage.get_file(file_path)
|
||||
file_path = await asyncio.to_thread(Storage.get_file, file_path)
|
||||
file_path = Path(file_path)
|
||||
|
||||
# Check if the file already exists in the cache
|
||||
|
|
@ -766,17 +790,17 @@ async def delete_file_by_id(id: str, user=Depends(get_verified_user), db: AsyncS
|
|||
await Knowledges.remove_file_from_knowledge_by_id(knowledge.id, id, db=db)
|
||||
# Clean KB embeddings (same logic as /knowledge/{id}/file/remove)
|
||||
try:
|
||||
VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'file_id': id})
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'file_id': id})
|
||||
if file.hash:
|
||||
VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'hash': file.hash})
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'hash': file.hash})
|
||||
except Exception as e:
|
||||
log.debug(f'KB embedding cleanup for {knowledge.id}: {e}')
|
||||
|
||||
result = await Files.delete_file_by_id(id, db=db)
|
||||
if result:
|
||||
try:
|
||||
Storage.delete_file(file.path)
|
||||
VECTOR_DB_CLIENT.delete(collection_name=f'file-{id}')
|
||||
await asyncio.to_thread(Storage.delete_file, file.path)
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=f'file-{id}')
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
log.error('Error deleting files')
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,6 @@ async def image_generations(
|
|||
|
||||
model = get_image_model(request)
|
||||
|
||||
r = None
|
||||
try:
|
||||
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
headers = {
|
||||
|
|
@ -552,7 +561,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 +581,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 +631,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 +738,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 +763,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 +806,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 +841,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 +855,6 @@ async def image_edits(
|
|||
),
|
||||
)
|
||||
|
||||
r = None
|
||||
try:
|
||||
if request.app.state.config.IMAGE_EDIT_ENGINE == 'openai':
|
||||
headers = {
|
||||
|
|
@ -883,17 +891,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 +961,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 +1068,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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -19,7 +19,7 @@ from open_webui.models.knowledge import (
|
|||
KnowledgeUserResponse,
|
||||
)
|
||||
from open_webui.models.files import Files, FileModel, FileMetadataResponse
|
||||
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
|
||||
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
||||
from open_webui.routers.retrieval import (
|
||||
process_file,
|
||||
ProcessFileForm,
|
||||
|
|
@ -66,7 +66,7 @@ async def embed_knowledge_base_metadata(
|
|||
try:
|
||||
content = f'{name}\n\n{description}' if description else name
|
||||
embedding = await request.app.state.EMBEDDING_FUNCTION(content)
|
||||
VECTOR_DB_CLIENT.upsert(
|
||||
await ASYNC_VECTOR_DB_CLIENT.upsert(
|
||||
collection_name=KNOWLEDGE_BASES_COLLECTION,
|
||||
items=[
|
||||
{
|
||||
|
|
@ -85,10 +85,10 @@ async def embed_knowledge_base_metadata(
|
|||
return False
|
||||
|
||||
|
||||
def remove_knowledge_base_metadata_embedding(knowledge_base_id: str) -> bool:
|
||||
async def remove_knowledge_base_metadata_embedding(knowledge_base_id: str) -> bool:
|
||||
"""Remove knowledge base embedding."""
|
||||
try:
|
||||
VECTOR_DB_CLIENT.delete(
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(
|
||||
collection_name=KNOWLEDGE_BASES_COLLECTION,
|
||||
ids=[knowledge_base_id],
|
||||
)
|
||||
|
|
@ -310,8 +310,8 @@ async def reindex_knowledge_files(
|
|||
try:
|
||||
files = await Knowledges.get_files_by_id(knowledge_base.id, db=db)
|
||||
try:
|
||||
if VECTOR_DB_CLIENT.has_collection(collection_name=knowledge_base.id):
|
||||
VECTOR_DB_CLIENT.delete_collection(collection_name=knowledge_base.id)
|
||||
if await ASYNC_VECTOR_DB_CLIENT.has_collection(collection_name=knowledge_base.id):
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete_collection(collection_name=knowledge_base.id)
|
||||
except Exception as e:
|
||||
log.error(f'Error deleting collection {knowledge_base.id}: {str(e)}')
|
||||
continue # Skip, don't raise
|
||||
|
|
@ -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,
|
||||
|
|
@ -733,11 +732,11 @@ async def update_file_from_knowledge_by_id(
|
|||
)
|
||||
|
||||
# Remove content from the vector database
|
||||
VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'file_id': form_data.file_id})
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'file_id': form_data.file_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,
|
||||
|
|
@ -815,11 +814,11 @@ async def remove_file_from_knowledge_by_id(
|
|||
|
||||
# Remove content from the vector database
|
||||
try:
|
||||
VECTOR_DB_CLIENT.delete(
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(
|
||||
collection_name=knowledge.id, filter={'file_id': form_data.file_id}
|
||||
) # Remove by file_id first
|
||||
|
||||
VECTOR_DB_CLIENT.delete(
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(
|
||||
collection_name=knowledge.id, filter={'hash': file.hash}
|
||||
) # Remove by hash as well in case of duplicates
|
||||
except Exception as e:
|
||||
|
|
@ -831,8 +830,8 @@ async def remove_file_from_knowledge_by_id(
|
|||
try:
|
||||
# Remove the file's collection from vector database
|
||||
file_collection = f'file-{form_data.file_id}'
|
||||
if VECTOR_DB_CLIENT.has_collection(collection_name=file_collection):
|
||||
VECTOR_DB_CLIENT.delete_collection(collection_name=file_collection)
|
||||
if await ASYNC_VECTOR_DB_CLIENT.has_collection(collection_name=file_collection):
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete_collection(collection_name=file_collection)
|
||||
except Exception as e:
|
||||
log.debug('This was most likely caused by bypassing embedding processing')
|
||||
log.debug(e)
|
||||
|
|
@ -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(
|
||||
|
|
@ -914,13 +915,13 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: A
|
|||
|
||||
# Clean up vector DB
|
||||
try:
|
||||
VECTOR_DB_CLIENT.delete_collection(collection_name=id)
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete_collection(collection_name=id)
|
||||
except Exception as e:
|
||||
log.debug(e)
|
||||
pass
|
||||
|
||||
# Remove knowledge base embedding
|
||||
remove_knowledge_base_metadata_embedding(id)
|
||||
await remove_knowledge_base_metadata_embedding(id)
|
||||
|
||||
result = await Knowledges.delete_knowledge_by_id(id=id, db=db)
|
||||
return result
|
||||
|
|
@ -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(
|
||||
|
|
@ -957,12 +960,12 @@ async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: As
|
|||
)
|
||||
|
||||
try:
|
||||
VECTOR_DB_CLIENT.delete_collection(collection_name=id)
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete_collection(collection_name=id)
|
||||
except Exception as e:
|
||||
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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ import asyncio
|
|||
from typing import Optional
|
||||
|
||||
from open_webui.models.memories import Memories, MemoryModel
|
||||
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
|
||||
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
||||
from open_webui.utils.auth import get_verified_user
|
||||
from open_webui.internal.db import get_async_session
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
|
@ -85,7 +85,7 @@ async def add_memory(
|
|||
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)
|
||||
|
||||
VECTOR_DB_CLIENT.upsert(
|
||||
await ASYNC_VECTOR_DB_CLIENT.upsert(
|
||||
collection_name=f'user-memory-{user.id}',
|
||||
items=[
|
||||
{
|
||||
|
|
@ -138,7 +138,7 @@ async def query_memory(
|
|||
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(form_data.content, user=user)
|
||||
|
||||
results = VECTOR_DB_CLIENT.search(
|
||||
results = await ASYNC_VECTOR_DB_CLIENT.search(
|
||||
collection_name=f'user-memory-{user.id}',
|
||||
vectors=[vector],
|
||||
limit=form_data.k,
|
||||
|
|
@ -175,7 +175,7 @@ async def reset_memory_from_vector_db(
|
|||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
VECTOR_DB_CLIENT.delete_collection(f'user-memory-{user.id}')
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete_collection(f'user-memory-{user.id}')
|
||||
|
||||
memories = await Memories.get_memories_by_user_id(user.id)
|
||||
|
||||
|
|
@ -184,7 +184,7 @@ async def reset_memory_from_vector_db(
|
|||
*[request.app.state.EMBEDDING_FUNCTION(memory.content, user=user) for memory in memories]
|
||||
)
|
||||
|
||||
VECTOR_DB_CLIENT.upsert(
|
||||
await ASYNC_VECTOR_DB_CLIENT.upsert(
|
||||
collection_name=f'user-memory-{user.id}',
|
||||
items=[
|
||||
{
|
||||
|
|
@ -230,7 +230,7 @@ async def delete_memory_by_user_id(
|
|||
|
||||
if result:
|
||||
try:
|
||||
VECTOR_DB_CLIENT.delete_collection(f'user-memory-{user.id}')
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete_collection(f'user-memory-{user.id}')
|
||||
except Exception as e:
|
||||
log.error(e)
|
||||
return True
|
||||
|
|
@ -268,12 +268,12 @@ 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)
|
||||
|
||||
VECTOR_DB_CLIENT.upsert(
|
||||
await ASYNC_VECTOR_DB_CLIENT.upsert(
|
||||
collection_name=f'user-memory-{user.id}',
|
||||
items=[
|
||||
{
|
||||
|
|
@ -318,7 +318,7 @@ async def delete_memory_by_id(
|
|||
result = await Memories.delete_memory_by_id_and_user_id(memory_id, user.id, db=db)
|
||||
|
||||
if result:
|
||||
VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=[memory_id])
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=[memory_id])
|
||||
return True
|
||||
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -58,6 +58,7 @@ class NoteItemResponse(BaseModel):
|
|||
id: str
|
||||
title: str
|
||||
data: Optional[dict]
|
||||
is_pinned: Optional[bool] = False
|
||||
updated_at: int
|
||||
created_at: int
|
||||
user: Optional[UserResponse] = None
|
||||
|
|
@ -104,6 +105,45 @@ async def get_notes(
|
|||
]
|
||||
|
||||
|
||||
############################
|
||||
# GetPinnedNotes
|
||||
############################
|
||||
|
||||
|
||||
@router.get('/pinned', response_model=list[NoteItemResponse])
|
||||
async def get_pinned_notes(
|
||||
request: Request,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
)
|
||||
|
||||
notes = await Notes.get_pinned_notes_by_user_id(user.id, 'read', db=db)
|
||||
if not notes:
|
||||
return []
|
||||
|
||||
user_ids = list(set(note.user_id for note in notes))
|
||||
users = {user.id: user for user in await Users.get_users_by_user_ids(user_ids, db=db)}
|
||||
|
||||
return [
|
||||
NoteUserResponse(
|
||||
**{
|
||||
**note.model_dump(),
|
||||
'data': _truncate_note_data(note.data),
|
||||
'user': UserResponse(**users[note.user_id].model_dump()),
|
||||
}
|
||||
)
|
||||
for note in notes
|
||||
if note.user_id in users
|
||||
]
|
||||
|
||||
|
||||
@router.get('/search', response_model=NoteListResponse)
|
||||
async def search_notes(
|
||||
request: Request,
|
||||
|
|
@ -364,6 +404,46 @@ async def update_note_access_by_id(
|
|||
return await Notes.get_note_by_id(id, db=db)
|
||||
|
||||
|
||||
############################
|
||||
# PinNoteById
|
||||
############################
|
||||
|
||||
|
||||
@router.post('/{id}/pin', response_model=Optional[NoteModel])
|
||||
async def pin_note_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
)
|
||||
|
||||
note = await Notes.get_note_by_id(id, db=db)
|
||||
if not note:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
|
||||
if user.role != 'admin' and (
|
||||
user.id != note.user_id
|
||||
and not await AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type='note',
|
||||
resource_id=note.id,
|
||||
permission='read',
|
||||
db=db,
|
||||
)
|
||||
):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
note = await Notes.toggle_note_pinned_by_id(id, db=db)
|
||||
return note
|
||||
|
||||
|
||||
############################
|
||||
# DeleteNoteById
|
||||
############################
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -1015,6 +1020,8 @@ class ChatMessage(BaseModel):
|
|||
tool_calls: Optional[list[dict]] = None
|
||||
images: Optional[list[str]] = None
|
||||
|
||||
model_config = ConfigDict(extra='allow')
|
||||
|
||||
@validator('content', pre=True)
|
||||
@classmethod
|
||||
def check_at_least_one_field(cls, field_value, values, **kwargs):
|
||||
|
|
@ -1062,7 +1069,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 +1315,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 +1373,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 +1398,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 +1656,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.')
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -40,11 +40,12 @@ 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
|
||||
|
||||
|
||||
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
|
||||
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
||||
|
||||
# Document loaders
|
||||
from open_webui.retrieval.loaders.main import Loader
|
||||
|
|
@ -151,7 +152,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
|
||||
|
||||
|
|
@ -1556,7 +1557,7 @@ async def process_file(
|
|||
|
||||
try:
|
||||
# /files/{file_id}/data/content/update
|
||||
VECTOR_DB_CLIENT.delete_collection(collection_name=f'file-{file.id}')
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete_collection(collection_name=f'file-{file.id}')
|
||||
except Exception:
|
||||
# Audio file upload pipeline
|
||||
pass
|
||||
|
|
@ -1579,7 +1580,9 @@ async def process_file(
|
|||
# Check if the file has already been processed and save the content
|
||||
# Usage: /knowledge/{id}/file/add, /knowledge/{id}/file/update
|
||||
|
||||
result = VECTOR_DB_CLIENT.query(collection_name=f'file-{file.id}', filter={'file_id': file.id})
|
||||
result = await ASYNC_VECTOR_DB_CLIENT.query(
|
||||
collection_name=f'file-{file.id}', filter={'file_id': file.id}
|
||||
)
|
||||
|
||||
if result is not None and len(result.ids[0]) > 0:
|
||||
docs = [
|
||||
|
|
@ -1609,7 +1612,7 @@ async def process_file(
|
|||
# Usage: /files/
|
||||
file_path = file.path
|
||||
if file_path:
|
||||
file_path = Storage.get_file(file_path)
|
||||
file_path = await asyncio.to_thread(Storage.get_file, file_path)
|
||||
loader = Loader(
|
||||
engine=request.app.state.config.CONTENT_EXTRACTION_ENGINE,
|
||||
user=user,
|
||||
|
|
@ -1643,7 +1646,7 @@ async def process_file(
|
|||
MINERU_API_TIMEOUT=request.app.state.config.MINERU_API_TIMEOUT,
|
||||
MINERU_PARAMS=request.app.state.config.MINERU_PARAMS,
|
||||
)
|
||||
docs = loader.load(file.filename, file.meta.get('content_type'), file_path)
|
||||
docs = await loader.aload(file.filename, file.meta.get('content_type'), file_path)
|
||||
|
||||
docs = [
|
||||
Document(
|
||||
|
|
@ -1694,11 +1697,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 +1722,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 +1752,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'},
|
||||
|
|
@ -2375,7 +2383,7 @@ async def query_doc_handler(
|
|||
try:
|
||||
if request.app.state.config.ENABLE_RAG_HYBRID_SEARCH and (form_data.hybrid is None or form_data.hybrid):
|
||||
collection_results = {}
|
||||
collection_results[form_data.collection_name] = VECTOR_DB_CLIENT.get(
|
||||
collection_results[form_data.collection_name] = await ASYNC_VECTOR_DB_CLIENT.get(
|
||||
collection_name=form_data.collection_name
|
||||
)
|
||||
return await query_doc_with_hybrid_search(
|
||||
|
|
@ -2404,7 +2412,10 @@ async def query_doc_handler(
|
|||
query_embedding = await request.app.state.EMBEDDING_FUNCTION(
|
||||
form_data.query, prefix=RAG_EMBEDDING_QUERY_PREFIX, user=user
|
||||
)
|
||||
return query_doc(
|
||||
# query_doc wraps a blocking VECTOR_DB_CLIENT.search call;
|
||||
# offload so the request's event loop stays responsive.
|
||||
return await asyncio.to_thread(
|
||||
query_doc,
|
||||
collection_name=form_data.collection_name,
|
||||
query_embedding=query_embedding,
|
||||
k=form_data.k if form_data.k else request.app.state.config.TOP_K,
|
||||
|
|
@ -2502,7 +2513,7 @@ async def delete_entries_from_collection(
|
|||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
try:
|
||||
if VECTOR_DB_CLIENT.has_collection(collection_name=form_data.collection_name):
|
||||
if await ASYNC_VECTOR_DB_CLIENT.has_collection(collection_name=form_data.collection_name):
|
||||
file = await Files.get_file_by_id(form_data.file_id, db=db)
|
||||
if not file:
|
||||
raise HTTPException(
|
||||
|
|
@ -2511,13 +2522,37 @@ async def delete_entries_from_collection(
|
|||
)
|
||||
hash = file.hash
|
||||
|
||||
VECTOR_DB_CLIENT.delete(
|
||||
# Refuse to issue a `filter={'hash': None}` query — the
|
||||
# match semantics of a null filter value are
|
||||
# backend-dependent (some backends ignore the key, some
|
||||
# match every row whose metadata lacks `hash`) and risk
|
||||
# deleting unrelated entries. Files without a hash are
|
||||
# typically unprocessed / failed / legacy records that
|
||||
# can't be targeted by hash anyway.
|
||||
if hash is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('File has no hash; cannot delete vector entries by hash.'),
|
||||
)
|
||||
|
||||
# Pre-existing bug: this used `metadata=` which is not a
|
||||
# parameter on `VectorDBBase.delete` nor on any backend
|
||||
# implementation, so the call always raised TypeError that
|
||||
# was silently swallowed by the surrounding `except
|
||||
# Exception` and the endpoint reported `{'status': False}`
|
||||
# for every request. Use `filter` to actually do what the
|
||||
# endpoint name promises.
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(
|
||||
collection_name=form_data.collection_name,
|
||||
metadata={'hash': hash},
|
||||
filter={'hash': hash},
|
||||
)
|
||||
return {'status': True}
|
||||
else:
|
||||
return {'status': False}
|
||||
except HTTPException:
|
||||
# Caller-meaningful errors (404/400 above) must not be
|
||||
# swallowed and re-shaped as `{'status': False}`.
|
||||
raise
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
return {'status': False}
|
||||
|
|
@ -2525,7 +2560,7 @@ async def delete_entries_from_collection(
|
|||
|
||||
@router.post('/reset/db')
|
||||
async def reset_vector_db(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
|
||||
VECTOR_DB_CLIENT.reset()
|
||||
await ASYNC_VECTOR_DB_CLIENT.reset()
|
||||
await Knowledges.delete_all_knowledge(db=db)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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'),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -37,7 +37,7 @@ from open_webui.models.channels import Channels, ChannelMember, Channel
|
|||
from open_webui.models.messages import Messages, Message
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.memories import Memories
|
||||
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
|
||||
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
||||
from open_webui.utils.sanitize import sanitize_code
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
|
@ -653,7 +653,7 @@ async def delete_memory(
|
|||
result = await Memories.delete_memory_by_id_and_user_id(memory_id, user.id)
|
||||
|
||||
if result:
|
||||
VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=[memory_id])
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=[memory_id])
|
||||
return json.dumps(
|
||||
{'status': 'success', 'message': f'Memory {memory_id} deleted'},
|
||||
ensure_ascii=False,
|
||||
|
|
@ -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 = []
|
||||
|
|
@ -2190,7 +2202,7 @@ async def query_knowledge_bases(
|
|||
import heapq
|
||||
from open_webui.models.knowledge import Knowledges
|
||||
from open_webui.routers.knowledge import KNOWLEDGE_BASES_COLLECTION
|
||||
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
|
||||
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
||||
|
||||
user_id = __user__.get('id')
|
||||
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
|
||||
|
|
@ -2215,7 +2227,7 @@ async def query_knowledge_bases(
|
|||
|
||||
accessible_ids = [kb.id for kb in accessible_knowledge_bases.items]
|
||||
|
||||
search_results = VECTOR_DB_CLIENT.search(
|
||||
search_results = await ASYNC_VECTOR_DB_CLIENT.search(
|
||||
collection_name=KNOWLEDGE_BASES_COLLECTION,
|
||||
vectors=[query_embedding],
|
||||
filter={'knowledge_base_id': {'$in': accessible_ids}},
|
||||
|
|
@ -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
|
||||
# =============================================================================
|
||||
|
|
|
|||
|
|
@ -181,17 +181,99 @@ 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(
|
||||
{
|
||||
|
|
|
|||
268
backend/open_webui/utils/asgi_middleware.py
Normal file
268
backend/open_webui/utils/asgi_middleware.py
Normal file
|
|
@ -0,0 +1,268 @@
|
|||
"""
|
||||
Pure-ASGI replacements for the project's previous
|
||||
`@app.middleware('http')` / `BaseHTTPMiddleware` middlewares.
|
||||
|
||||
Why this matters
|
||||
----------------
|
||||
Starlette's `BaseHTTPMiddleware` (which `@app.middleware('http')` is
|
||||
sugar for) runs the downstream app inside an `anyio` task group. When
|
||||
the wrapper exits — for any reason: response complete, client
|
||||
disconnect, an outer middleware bailing out — the task group cancels
|
||||
the inner task. That `CancelledError` then propagates into whatever
|
||||
the inner task was doing, including in-flight DB queries, embedding
|
||||
calls and disk I/O.
|
||||
|
||||
In Open WebUI this surfaces as:
|
||||
|
||||
* SQLAlchemy logging multi-page `NotImplementedError:
|
||||
terminate_force_close()` tracebacks at ERROR every time a request is
|
||||
cancelled mid-DB-call (the aiosqlite connector cleanup path).
|
||||
* Spurious cancellations cascading through the four stacked
|
||||
`@app.middleware('http')` wrappers.
|
||||
|
||||
Pure ASGI middleware does not introduce a cancel scope around the
|
||||
downstream app, so client disconnects propagate the way ASGI was
|
||||
designed to (via `receive()` returning `http.disconnect`) instead of
|
||||
being injected as `CancelledError` into arbitrary `await` points.
|
||||
|
||||
Reference: https://www.starlette.io/middleware/#limitations
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
from urllib.parse import parse_qs, urlencode
|
||||
|
||||
from fastapi.responses import JSONResponse, RedirectResponse
|
||||
from fastapi.security import HTTPAuthorizationCredentials
|
||||
from starlette.datastructures import MutableHeaders
|
||||
from starlette.requests import Request
|
||||
from starlette.types import ASGIApp, Message, Receive, Scope, Send
|
||||
|
||||
from open_webui.internal.db import ScopedSession
|
||||
from open_webui.utils.auth import get_http_authorization_cred
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CommitSessionMiddleware:
|
||||
"""Commit and release the thread-local sync `ScopedSession` after each
|
||||
HTTP request.
|
||||
|
||||
Most requests now use the async session; the sync ScopedSession is
|
||||
only touched by startup, healthchecks, and a handful of legacy
|
||||
helpers (notably the pgvector / opengauss vector-DB clients). The
|
||||
middleware exists so that PostgreSQL connections do not accumulate
|
||||
as "idle in transaction" and so that any pending sync work made
|
||||
inside the request is durably persisted.
|
||||
|
||||
Failure semantics
|
||||
-----------------
|
||||
* Downstream raised → roll back any pending sync work, release the
|
||||
connection, and re-raise so the outer exception middleware can
|
||||
turn it into an error response. We never commit work on a
|
||||
request that did not complete successfully.
|
||||
* Downstream returned → commit pending sync work; on commit
|
||||
failure, log loudly, roll back, and re-raise. Note that in pure
|
||||
ASGI the response messages have already been emitted by the
|
||||
time `await self.app(...)` returns, so a commit failure cannot
|
||||
retroactively change what the client sees on the wire — but
|
||||
re-raising still surfaces the error in logs and to ASGI servers
|
||||
that expose it. We deliberately do not buffer the response to
|
||||
gate it on commit success, because that would defeat streaming
|
||||
responses (chat completions, SSE) which are core to the app.
|
||||
|
||||
For request paths where commit-before-send is required, manage the
|
||||
sync session explicitly inside the handler instead of relying on
|
||||
this middleware.
|
||||
"""
|
||||
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope['type'] != 'http':
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
try:
|
||||
await self.app(scope, receive, send)
|
||||
except BaseException:
|
||||
# Downstream did not complete successfully. Roll back any
|
||||
# pending sync writes, release the connection, and let the
|
||||
# exception propagate.
|
||||
try:
|
||||
ScopedSession.rollback()
|
||||
except Exception:
|
||||
log.exception('CommitSessionMiddleware: rollback failed after downstream error')
|
||||
finally:
|
||||
ScopedSession.remove()
|
||||
raise
|
||||
|
||||
# Downstream completed. Commit pending sync work.
|
||||
try:
|
||||
ScopedSession.commit()
|
||||
except Exception:
|
||||
log.exception('CommitSessionMiddleware: post-request commit failed; response was already sent to client')
|
||||
try:
|
||||
ScopedSession.rollback()
|
||||
except Exception:
|
||||
log.exception('CommitSessionMiddleware: rollback failed after commit failure')
|
||||
raise
|
||||
finally:
|
||||
# CRITICAL: remove() returns the connection to the pool.
|
||||
# Without this, connections remain "checked out" and
|
||||
# accumulate as "idle in transaction" in PostgreSQL.
|
||||
ScopedSession.remove()
|
||||
|
||||
|
||||
class AuthTokenMiddleware:
|
||||
"""Extract the bearer/cookie/x-api-key credential and stash it on
|
||||
`request.state.token`.
|
||||
|
||||
Routes that depend on `get_verified_user` etc. read this state.
|
||||
Also exposes `request.state.enable_api_keys` (snapshotted at request
|
||||
entry from runtime config) and stamps an `X-Process-Time` response
|
||||
header.
|
||||
"""
|
||||
|
||||
def __init__(self, app: ASGIApp, *, fastapi_app) -> None:
|
||||
self.app = app
|
||||
self._fastapi_app = fastapi_app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope['type'] != 'http':
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
start_time = time.monotonic()
|
||||
request = Request(scope)
|
||||
|
||||
token = get_http_authorization_cred(request.headers.get('Authorization'))
|
||||
if token is None:
|
||||
cookie_token = request.cookies.get('token')
|
||||
if cookie_token:
|
||||
token = HTTPAuthorizationCredentials(scheme='Bearer', credentials=cookie_token)
|
||||
if token is None:
|
||||
api_key = request.headers.get('x-api-key')
|
||||
if api_key:
|
||||
token = HTTPAuthorizationCredentials(scheme='Bearer', credentials=api_key)
|
||||
|
||||
request.state.token = token
|
||||
request.state.enable_api_keys = self._fastapi_app.state.config.ENABLE_API_KEYS
|
||||
|
||||
async def send_with_timing(message: Message) -> None:
|
||||
if message['type'] == 'http.response.start':
|
||||
process_time = int(time.monotonic() - start_time)
|
||||
headers = MutableHeaders(scope=message)
|
||||
headers['X-Process-Time'] = str(process_time)
|
||||
await send(message)
|
||||
|
||||
await self.app(scope, receive, send_with_timing)
|
||||
|
||||
|
||||
class WebsocketUpgradeGuardMiddleware:
|
||||
"""Reject HTTP requests to `/ws/socket.io` that claim
|
||||
`transport=websocket` but lack the proper `Upgrade`/`Connection`
|
||||
headers.
|
||||
|
||||
Works around https://github.com/miguelgrinberg/python-engineio/issues/367
|
||||
where engineio mishandles such requests.
|
||||
"""
|
||||
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope['type'] != 'http':
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
path = scope.get('path', '')
|
||||
if '/ws/socket.io' in path:
|
||||
query_string = scope.get('query_string', b'').decode('latin-1', errors='replace')
|
||||
query_params = parse_qs(query_string)
|
||||
if query_params.get('transport', [''])[0] == 'websocket':
|
||||
headers = _scope_headers(scope)
|
||||
upgrade = headers.get('upgrade', '').lower()
|
||||
connection_tokens = [token.strip() for token in headers.get('connection', '').lower().split(',')]
|
||||
if upgrade != 'websocket' or 'upgrade' not in connection_tokens:
|
||||
response = JSONResponse(
|
||||
status_code=400,
|
||||
content={'detail': 'Invalid WebSocket upgrade request'},
|
||||
)
|
||||
await response(scope, receive, send)
|
||||
return
|
||||
|
||||
await self.app(scope, receive, send)
|
||||
|
||||
|
||||
class RedirectMiddleware:
|
||||
"""Rewrites a couple of legacy entry-points to the SPA's own routes:
|
||||
|
||||
* ``GET /watch?v=ID`` (YouTube) → ``/?youtube=ID``
|
||||
* ``GET /?shared=…`` (PWA share-target) → ``/?youtube=…`` /
|
||||
``/?load-url=…`` / ``/?q=…``
|
||||
"""
|
||||
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope['type'] != 'http' or scope.get('method', '').upper() != 'GET':
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
path = scope.get('path', '')
|
||||
query_string = scope.get('query_string', b'').decode('latin-1', errors='replace')
|
||||
query_params = parse_qs(query_string)
|
||||
|
||||
redirect_params: dict[str, str] = {}
|
||||
if path.endswith('/watch') and 'v' in query_params and query_params['v']:
|
||||
redirect_params['youtube'] = query_params['v'][0]
|
||||
|
||||
if 'shared' in query_params and query_params['shared']:
|
||||
text = query_params['shared'][0]
|
||||
if text:
|
||||
url_match = re.match(r'https://\S+', text)
|
||||
if url_match:
|
||||
# Local import: youtube loader pulls heavy deps and is
|
||||
# only needed when a share-target actually contains a
|
||||
# YouTube URL.
|
||||
from open_webui.retrieval.loaders.youtube import _parse_video_id
|
||||
|
||||
youtube_video_id = _parse_video_id(url_match[0])
|
||||
if youtube_video_id:
|
||||
redirect_params['youtube'] = youtube_video_id
|
||||
else:
|
||||
redirect_params['load-url'] = url_match[0]
|
||||
else:
|
||||
redirect_params['q'] = text
|
||||
|
||||
if redirect_params:
|
||||
redirect_url = f'/?{urlencode(redirect_params)}'
|
||||
response = RedirectResponse(url=redirect_url)
|
||||
await response(scope, receive, send)
|
||||
return
|
||||
|
||||
await self.app(scope, receive, send)
|
||||
|
||||
|
||||
def _scope_headers(scope: Scope) -> dict[str, str]:
|
||||
"""Return ASGI scope headers as a lower-cased str→str dict.
|
||||
|
||||
ASGI delivers headers as a list of (bytes, bytes) pairs. For
|
||||
convenience, fold duplicate keys with comma-joining (matching
|
||||
HTTP/1.1 semantics).
|
||||
"""
|
||||
decoded: dict[str, str] = {}
|
||||
for raw_key, raw_value in scope.get('headers', []):
|
||||
key = raw_key.decode('latin-1').lower()
|
||||
value = raw_value.decode('latin-1')
|
||||
if key in decoded:
|
||||
decoded[key] = f'{decoded[key]}, {value}'
|
||||
else:
|
||||
decoded[key] = value
|
||||
return decoded
|
||||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -20,12 +20,14 @@ from open_webui.models.files import Files
|
|||
from open_webui.routers.files import upload_file_handler
|
||||
from open_webui.retrieval.web.utils import validate_url
|
||||
|
||||
import asyncio
|
||||
import mimetypes
|
||||
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,19 +39,20 @@ 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)
|
||||
|
||||
if not file:
|
||||
return None
|
||||
|
||||
file_path = Storage.get_file(file.path)
|
||||
file_path = await asyncio.to_thread(Storage.get_file, file.path)
|
||||
file_path = Path(file_path)
|
||||
|
||||
if file_path.is_file():
|
||||
|
|
@ -88,7 +91,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)
|
||||
|
|
@ -170,7 +173,7 @@ async def get_image_base64_from_file_id(id: str) -> Optional[str]:
|
|||
return None
|
||||
|
||||
try:
|
||||
file_path = Storage.get_file(file.path)
|
||||
file_path = await asyncio.to_thread(Storage.get_file, file.path)
|
||||
file_path = Path(file_path)
|
||||
|
||||
# Check if the file already exists in the cache
|
||||
|
|
|
|||
|
|
@ -1,49 +1,51 @@
|
|||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import random
|
||||
import requests
|
||||
import aiohttp
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
from typing import Optional
|
||||
|
||||
import websocket # NOTE: websocket-client (https://github.com/websocket-client/websocket-client)
|
||||
import aiohttp
|
||||
from pydantic import BaseModel
|
||||
|
||||
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL
|
||||
from open_webui.utils.session_pool import get_session
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
default_headers = {'User-Agent': 'Mozilla/5.0'}
|
||||
|
||||
|
||||
def queue_prompt(prompt, client_id, base_url, api_key):
|
||||
async def queue_prompt(prompt, client_id, base_url, api_key):
|
||||
log.info('queue_prompt')
|
||||
p = {'prompt': prompt, 'client_id': client_id}
|
||||
data = json.dumps(p).encode('utf-8')
|
||||
log.debug(f'queue_prompt data: {data}')
|
||||
log.debug(f'queue_prompt data: {p}')
|
||||
try:
|
||||
req = urllib.request.Request(
|
||||
session = await get_session()
|
||||
async with session.post(
|
||||
f'{base_url}/prompt',
|
||||
data=data,
|
||||
json=p,
|
||||
headers={**default_headers, 'Authorization': f'Bearer {api_key}'},
|
||||
)
|
||||
response = urllib.request.urlopen(req).read()
|
||||
return json.loads(response)
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
r.raise_for_status()
|
||||
return await r.json()
|
||||
except Exception as e:
|
||||
log.exception(f'Error while queuing prompt: {e}')
|
||||
raise e
|
||||
raise
|
||||
|
||||
|
||||
def get_image(filename, subfolder, folder_type, base_url, api_key):
|
||||
async def get_image(filename, subfolder, folder_type, base_url, api_key):
|
||||
log.info('get_image')
|
||||
data = {'filename': filename, 'subfolder': subfolder, 'type': folder_type}
|
||||
url_values = urllib.parse.urlencode(data)
|
||||
req = urllib.request.Request(
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
f'{base_url}/view?{url_values}',
|
||||
headers={**default_headers, 'Authorization': f'Bearer {api_key}'},
|
||||
)
|
||||
with urllib.request.urlopen(req) as response:
|
||||
return response.read()
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
r.raise_for_status()
|
||||
return await r.read()
|
||||
|
||||
|
||||
def get_image_url(filename, subfolder, folder_type, base_url):
|
||||
|
|
@ -53,32 +55,39 @@ def get_image_url(filename, subfolder, folder_type, base_url):
|
|||
return f'{base_url}/view?{url_values}'
|
||||
|
||||
|
||||
def get_history(prompt_id, base_url, api_key):
|
||||
async def get_history(prompt_id, base_url, api_key):
|
||||
log.info('get_history')
|
||||
|
||||
req = urllib.request.Request(
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
f'{base_url}/history/{prompt_id}',
|
||||
headers={**default_headers, 'Authorization': f'Bearer {api_key}'},
|
||||
)
|
||||
with urllib.request.urlopen(req) as response:
|
||||
return json.loads(response.read())
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
r.raise_for_status()
|
||||
return await r.json()
|
||||
|
||||
|
||||
def get_images(ws, workflow, client_id, base_url, api_key):
|
||||
prompt_id = queue_prompt(workflow, client_id, base_url, api_key)['prompt_id']
|
||||
async def _ws_get_images(ws, workflow, client_id, base_url, api_key):
|
||||
"""Queue a prompt and wait on *ws* for ComfyUI to finish executing it.
|
||||
|
||||
Returns a dict of ``{'data': [{'url': ...}, ...]}``.
|
||||
"""
|
||||
prompt_id = (await queue_prompt(workflow, client_id, base_url, api_key))['prompt_id']
|
||||
output_images = []
|
||||
while True:
|
||||
out = ws.recv()
|
||||
if isinstance(out, str):
|
||||
message = json.loads(out)
|
||||
|
||||
async for msg in ws:
|
||||
if msg.type == aiohttp.WSMsgType.TEXT:
|
||||
message = json.loads(msg.data)
|
||||
if message['type'] == 'executing':
|
||||
data = message['data']
|
||||
if data['node'] is None and data['prompt_id'] == prompt_id:
|
||||
break # Execution is done
|
||||
else:
|
||||
continue # previews are binary data
|
||||
elif msg.type in (aiohttp.WSMsgType.CLOSED, aiohttp.WSMsgType.ERROR):
|
||||
log.error(f'WebSocket closed unexpectedly: {msg.type}')
|
||||
break
|
||||
# binary messages (previews) are silently skipped
|
||||
|
||||
history = get_history(prompt_id, base_url, api_key)[prompt_id]
|
||||
history = (await get_history(prompt_id, base_url, api_key))[prompt_id]
|
||||
for node_id in history['outputs']:
|
||||
node_output = history['outputs'][node_id]
|
||||
if node_id in workflow and workflow[node_id].get('class_type') in [
|
||||
|
|
@ -105,10 +114,10 @@ async def comfyui_upload_image(image_file_item, base_url, api_key):
|
|||
form.add_field('image', file_bytes, filename=filename, content_type=mime_type)
|
||||
form.add_field('type', 'input') # required by ComfyUI
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(url, data=form, headers=headers) as resp:
|
||||
resp.raise_for_status()
|
||||
return await resp.json()
|
||||
session = await get_session()
|
||||
async with session.post(url, data=form, headers=headers, ssl=AIOHTTP_CLIENT_SESSION_SSL) as resp:
|
||||
resp.raise_for_status()
|
||||
return await resp.json()
|
||||
|
||||
|
||||
class ComfyUINodeInput(BaseModel):
|
||||
|
|
@ -136,11 +145,9 @@ class ComfyUICreateImageForm(BaseModel):
|
|||
seed: Optional[int] = None
|
||||
|
||||
|
||||
async def comfyui_create_image(model: str, payload: ComfyUICreateImageForm, client_id, base_url, api_key):
|
||||
ws_url = base_url.replace('http://', 'ws://').replace('https://', 'wss://')
|
||||
workflow = json.loads(payload.workflow.workflow)
|
||||
|
||||
for node in payload.workflow.nodes:
|
||||
def _apply_workflow_nodes(workflow, nodes, model, payload):
|
||||
"""Mutate *workflow* dict in-place based on typed node definitions."""
|
||||
for node in nodes:
|
||||
if node.type:
|
||||
if node.type == 'model':
|
||||
for node_id in node.node_ids:
|
||||
|
|
@ -151,6 +158,14 @@ async def comfyui_create_image(model: str, payload: ComfyUICreateImageForm, clie
|
|||
elif node.type == 'negative_prompt':
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key if node.key else 'text'] = payload.negative_prompt
|
||||
elif node.type == 'image':
|
||||
if isinstance(payload.image, list):
|
||||
for idx, node_id in enumerate(node.node_ids):
|
||||
if idx < len(payload.image):
|
||||
workflow[node_id]['inputs'][node.key] = payload.image[idx]
|
||||
else:
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key] = payload.image
|
||||
elif node.type == 'width':
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key if node.key else 'width'] = payload.width
|
||||
|
|
@ -171,24 +186,31 @@ async def comfyui_create_image(model: str, payload: ComfyUICreateImageForm, clie
|
|||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key] = node.value
|
||||
|
||||
|
||||
async def comfyui_create_image(model: str, payload: ComfyUICreateImageForm, client_id, base_url, api_key):
|
||||
ws_url = base_url.replace('http://', 'ws://').replace('https://', 'wss://')
|
||||
workflow = json.loads(payload.workflow.workflow)
|
||||
_apply_workflow_nodes(workflow, payload.workflow.nodes, model, payload)
|
||||
|
||||
headers = {'Authorization': f'Bearer {api_key}'}
|
||||
session = await get_session()
|
||||
|
||||
try:
|
||||
ws = websocket.WebSocket()
|
||||
headers = {'Authorization': f'Bearer {api_key}'}
|
||||
ws.connect(f'{ws_url}/ws?clientId={client_id}', header=headers)
|
||||
log.info('WebSocket connection established.')
|
||||
except Exception as e:
|
||||
async with session.ws_connect(
|
||||
f'{ws_url}/ws?clientId={client_id}',
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as ws:
|
||||
log.info('WebSocket connection established.')
|
||||
log.info('Sending workflow to WebSocket server.')
|
||||
log.info(f'Workflow: {workflow}')
|
||||
images = await _ws_get_images(ws, workflow, client_id, base_url, api_key)
|
||||
except aiohttp.WSServerHandshakeError as e:
|
||||
log.exception(f'Failed to connect to WebSocket server: {e}')
|
||||
return None
|
||||
|
||||
try:
|
||||
log.info('Sending workflow to WebSocket server.')
|
||||
log.info(f'Workflow: {workflow}')
|
||||
images = await asyncio.to_thread(get_images, ws, workflow, client_id, base_url, api_key)
|
||||
except Exception as e:
|
||||
log.exception(f'Error while receiving images: {e}')
|
||||
images = None
|
||||
|
||||
ws.close()
|
||||
log.exception(f'Error during image generation: {e}')
|
||||
return None
|
||||
|
||||
return images
|
||||
|
||||
|
|
@ -209,64 +231,26 @@ class ComfyUIEditImageForm(BaseModel):
|
|||
async def comfyui_edit_image(model: str, payload: ComfyUIEditImageForm, client_id, base_url, api_key):
|
||||
ws_url = base_url.replace('http://', 'ws://').replace('https://', 'wss://')
|
||||
workflow = json.loads(payload.workflow.workflow)
|
||||
_apply_workflow_nodes(workflow, payload.workflow.nodes, model, payload)
|
||||
|
||||
for node in payload.workflow.nodes:
|
||||
if node.type:
|
||||
if node.type == 'model':
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key] = model
|
||||
elif node.type == 'image':
|
||||
if isinstance(payload.image, list):
|
||||
# check if multiple images are provided
|
||||
for idx, node_id in enumerate(node.node_ids):
|
||||
if idx < len(payload.image):
|
||||
workflow[node_id]['inputs'][node.key] = payload.image[idx]
|
||||
else:
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key] = payload.image
|
||||
elif node.type == 'prompt':
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key if node.key else 'text'] = payload.prompt
|
||||
elif node.type == 'negative_prompt':
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key if node.key else 'text'] = payload.negative_prompt
|
||||
elif node.type == 'width':
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key if node.key else 'width'] = payload.width
|
||||
elif node.type == 'height':
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key if node.key else 'height'] = payload.height
|
||||
elif node.type == 'n':
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key if node.key else 'batch_size'] = payload.n
|
||||
elif node.type == 'steps':
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key if node.key else 'steps'] = payload.steps
|
||||
elif node.type == 'seed':
|
||||
seed = payload.seed if payload.seed else random.randint(0, 1125899906842624)
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key] = seed
|
||||
else:
|
||||
for node_id in node.node_ids:
|
||||
workflow[node_id]['inputs'][node.key] = node.value
|
||||
headers = {'Authorization': f'Bearer {api_key}'}
|
||||
session = await get_session()
|
||||
|
||||
try:
|
||||
ws = websocket.WebSocket()
|
||||
headers = {'Authorization': f'Bearer {api_key}'}
|
||||
ws.connect(f'{ws_url}/ws?clientId={client_id}', header=headers)
|
||||
log.info('WebSocket connection established.')
|
||||
except Exception as e:
|
||||
async with session.ws_connect(
|
||||
f'{ws_url}/ws?clientId={client_id}',
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as ws:
|
||||
log.info('WebSocket connection established.')
|
||||
log.info('Sending workflow to WebSocket server.')
|
||||
log.info(f'Workflow: {workflow}')
|
||||
images = await _ws_get_images(ws, workflow, client_id, base_url, api_key)
|
||||
except aiohttp.WSServerHandshakeError as e:
|
||||
log.exception(f'Failed to connect to WebSocket server: {e}')
|
||||
return None
|
||||
|
||||
try:
|
||||
log.info('Sending workflow to WebSocket server.')
|
||||
log.info(f'Workflow: {workflow}')
|
||||
images = await asyncio.to_thread(get_images, ws, workflow, client_id, base_url, api_key)
|
||||
except Exception as e:
|
||||
log.exception(f'Error while receiving images: {e}')
|
||||
images = None
|
||||
|
||||
ws.close()
|
||||
log.exception(f'Error during image editing: {e}')
|
||||
return None
|
||||
|
||||
return images
|
||||
|
|
|
|||
|
|
@ -450,7 +450,8 @@ def serialize_output(output: list) -> str:
|
|||
files = result_item.get('files')
|
||||
embeds = result_item.get('embeds', '')
|
||||
|
||||
append_part(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')
|
||||
append_part(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:
|
||||
append_part(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')
|
||||
|
||||
|
|
@ -907,10 +908,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
|
||||
|
|
@ -2068,13 +2073,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)
|
||||
|
|
@ -2164,10 +2172,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
|
||||
|
|
@ -3072,6 +3080,145 @@ 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.
|
||||
|
||||
For temp chats (local: prefix), messages are built from form_data
|
||||
plus the assistant response message stored in ctx['assistant_message'],
|
||||
since temp chats have no DB-persisted history.
|
||||
"""
|
||||
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 not message_id:
|
||||
return
|
||||
|
||||
is_temp_chat = chat_id.startswith('local:')
|
||||
|
||||
try:
|
||||
messages_map = None
|
||||
|
||||
if is_temp_chat:
|
||||
# Temp chats have no DB record — build message list from
|
||||
# the in-memory form_data plus the assistant response.
|
||||
form_messages = ctx.get('form_data', {}).get('messages', [])
|
||||
assistant_message = ctx.get('assistant_message', {})
|
||||
|
||||
message_list = [
|
||||
{
|
||||
'role': m.get('role'),
|
||||
'content': m.get('content', ''),
|
||||
}
|
||||
for m in form_messages
|
||||
]
|
||||
|
||||
# Append the full assistant message (content, output, usage, etc.)
|
||||
if assistant_message:
|
||||
message_list.append({
|
||||
'id': message_id,
|
||||
'role': 'assistant',
|
||||
**assistant_message,
|
||||
})
|
||||
else:
|
||||
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'),
|
||||
**({'output': m['output']} if m.get('output') else {}),
|
||||
**({'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
|
||||
# (skip DB persistence for temp chats — they have no DB record)
|
||||
if outlet_result and outlet_result.get('messages'):
|
||||
if not is_temp_chat and messages_map:
|
||||
for message in outlet_result['messages']:
|
||||
outlet_message_id = message.get('id')
|
||||
if outlet_message_id and outlet_message_id in messages_map:
|
||||
original_message = messages_map[outlet_message_id]
|
||||
if original_message.get('content') != message.get('content'):
|
||||
await Chats.upsert_message_to_chat_by_id_and_message_id(
|
||||
chat_id,
|
||||
outlet_message_id,
|
||||
{
|
||||
'content': message['content'],
|
||||
'originalContent': original_message.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']
|
||||
|
||||
|
|
@ -3095,6 +3242,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'],
|
||||
|
|
@ -3191,6 +3340,12 @@ async def non_streaming_chat_response_handler(response, ctx):
|
|||
)
|
||||
|
||||
await background_tasks_handler(ctx)
|
||||
ctx['assistant_message'] = {
|
||||
'content': content,
|
||||
'output': response_output,
|
||||
**({'usage': usage} if usage else {}),
|
||||
}
|
||||
await outlet_filter_handler(ctx)
|
||||
|
||||
response = build_response_object(response, merge_events_into_response(response_data, events))
|
||||
except Exception as e:
|
||||
|
|
@ -3663,6 +3818,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'],
|
||||
|
|
@ -4078,11 +4234,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
|
||||
|
|
@ -4700,27 +4857,41 @@ async def streaming_chat_response_handler(response, ctx):
|
|||
)
|
||||
|
||||
await background_tasks_handler(ctx)
|
||||
ctx['assistant_message'] = {
|
||||
'content': serialize_output(output),
|
||||
'output': output,
|
||||
**({'usage': usage} if usage else {}),
|
||||
}
|
||||
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()
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,24 @@ 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 +578,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 +589,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 +642,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 +765,11 @@ 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 +880,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 +933,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 +1020,11 @@ 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 +1138,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 +1441,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 +1736,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 +1850,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 +1954,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={})
|
||||
|
|
|
|||
|
|
@ -204,6 +204,11 @@ def convert_messages_openai_to_ollama(messages: list[dict]) -> list[dict]:
|
|||
# Initialize the new message structure with the role
|
||||
new_message = {'role': message['role']}
|
||||
|
||||
# Preserve Ollama-native 'thinking' field (used by reasoning models,
|
||||
# may be injected by filter inlet functions).
|
||||
if 'thinking' in message:
|
||||
new_message['thinking'] = message['thinking']
|
||||
|
||||
content = message.get('content', [])
|
||||
tool_calls = message.get('tool_calls', None)
|
||||
tool_call_id = message.get('tool_call_id', None)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
119
backend/open_webui/utils/session_pool.py
Normal file
119
backend/open_webui/utils/session_pool.py
Normal 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)
|
||||
|
|
@ -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}')
|
||||
|
|
|
|||
|
|
@ -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.')
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
||||
|
|
|
|||
|
|
@ -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',
|
||||
|
|
|
|||
|
|
@ -313,3 +313,65 @@ export const deleteNoteById = async (token: string, id: string) => {
|
|||
|
||||
return res;
|
||||
};
|
||||
|
||||
export const getPinnedNoteList = async (token: string = '') => {
|
||||
let error = null;
|
||||
|
||||
const res = await fetch(`${WEBUI_API_BASE_URL}/notes/pinned`, {
|
||||
method: 'GET',
|
||||
headers: {
|
||||
Accept: 'application/json',
|
||||
'Content-Type': 'application/json',
|
||||
authorization: `Bearer ${token}`
|
||||
}
|
||||
})
|
||||
.then(async (res) => {
|
||||
if (!res.ok) throw await res.json();
|
||||
return res.json();
|
||||
})
|
||||
.then((json) => {
|
||||
return json;
|
||||
})
|
||||
.catch((err) => {
|
||||
error = err.detail;
|
||||
console.error(err);
|
||||
return null;
|
||||
});
|
||||
|
||||
if (error) {
|
||||
throw error;
|
||||
}
|
||||
|
||||
return res ?? [];
|
||||
};
|
||||
|
||||
export const toggleNotePinnedStatusById = async (token: string, id: string) => {
|
||||
let error = null;
|
||||
|
||||
const res = await fetch(`${WEBUI_API_BASE_URL}/notes/${id}/pin`, {
|
||||
method: 'POST',
|
||||
headers: {
|
||||
Accept: 'application/json',
|
||||
'Content-Type': 'application/json',
|
||||
authorization: `Bearer ${token}`
|
||||
}
|
||||
})
|
||||
.then(async (res) => {
|
||||
if (!res.ok) throw await res.json();
|
||||
return res.json();
|
||||
})
|
||||
.then((json) => {
|
||||
return json;
|
||||
})
|
||||
.catch((err) => {
|
||||
error = err.detail;
|
||||
console.error(err);
|
||||
return null;
|
||||
});
|
||||
|
||||
if (error) {
|
||||
throw error;
|
||||
}
|
||||
|
||||
return res;
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
||||
|
|
|
|||
|
|
@ -1,9 +1,10 @@
|
|||
<script lang="ts">
|
||||
import type i18nType from '$lib/i18n';
|
||||
import { getContext } from 'svelte';
|
||||
|
||||
import Dropdown from '$lib/components/common/Dropdown.svelte';
|
||||
|
||||
const i18n = getContext('i18n');
|
||||
const i18n: typeof i18nType = getContext('i18n');
|
||||
|
||||
export let frequency = 'DAILY';
|
||||
export let interval = 1;
|
||||
|
|
@ -23,23 +24,23 @@
|
|||
|
||||
let showDropdown = false;
|
||||
|
||||
const FREQUENCIES = [
|
||||
{ key: 'ONCE', label: 'Once' },
|
||||
{ key: 'HOURLY', label: 'Hourly' },
|
||||
{ key: 'DAILY', label: 'Daily' },
|
||||
{ key: 'WEEKLY', label: 'Weekly' },
|
||||
{ key: 'MONTHLY', label: 'Monthly' },
|
||||
{ key: 'CUSTOM', label: 'Custom' }
|
||||
$: FREQUENCIES = [
|
||||
{ key: 'ONCE', label: $i18n.t('Once') },
|
||||
{ key: 'HOURLY', label: $i18n.t('Hourly') },
|
||||
{ key: 'DAILY', label: $i18n.t('Daily') },
|
||||
{ key: 'WEEKLY', label: $i18n.t('Weekly') },
|
||||
{ key: 'MONTHLY', label: $i18n.t('Monthly') },
|
||||
{ key: 'CUSTOM', label: $i18n.t('Custom') }
|
||||
];
|
||||
|
||||
const DAYS = [
|
||||
{ key: 'MO', label: 'Mo' },
|
||||
{ key: 'TU', label: 'Tu' },
|
||||
{ key: 'WE', label: 'We' },
|
||||
{ key: 'TH', label: 'Th' },
|
||||
{ key: 'FR', label: 'Fr' },
|
||||
{ key: 'SA', label: 'Sa' },
|
||||
{ key: 'SU', label: 'Su' }
|
||||
$: DAYS = [
|
||||
{ key: 'MO', label: $i18n.t('Mo', { context: 'day_of_week' }) },
|
||||
{ key: 'TU', label: $i18n.t('Tu', { context: 'day_of_week' }) },
|
||||
{ key: 'WE', label: $i18n.t('We', { context: 'day_of_week' }) },
|
||||
{ key: 'TH', label: $i18n.t('Th', { context: 'day_of_week' }) },
|
||||
{ key: 'FR', label: $i18n.t('Fr', { context: 'day_of_week' }) },
|
||||
{ key: 'SA', label: $i18n.t('Sa', { context: 'day_of_week' }) },
|
||||
{ key: 'SU', label: $i18n.t('Su', { context: 'day_of_week' }) }
|
||||
];
|
||||
|
||||
let lastVisualFrequency = 'DAILY';
|
||||
|
|
|
|||
|
|
@ -51,6 +51,7 @@
|
|||
type="button"
|
||||
on:click={() => {
|
||||
uploadFilesHandler();
|
||||
show = false;
|
||||
}}
|
||||
>
|
||||
<Clip />
|
||||
|
|
@ -62,6 +63,7 @@
|
|||
type="button"
|
||||
on:click={() => {
|
||||
screenCaptureHandler();
|
||||
show = false;
|
||||
}}
|
||||
>
|
||||
<Camera />
|
||||
|
|
|
|||
|
|
@ -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,14 @@
|
|||
.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 +2202,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 +2311,7 @@
|
|||
{
|
||||
stream: stream,
|
||||
model: model.id,
|
||||
messages: messages,
|
||||
...(messages.length > 0 ? { messages } : {}),
|
||||
params: {
|
||||
...$settings?.params,
|
||||
...params,
|
||||
|
|
@ -2341,20 +2342,16 @@
|
|||
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 &&
|
||||
(messages.length == 1 ||
|
||||
(messages.length == 2 &&
|
||||
messages.at(0)?.role === 'system' &&
|
||||
messages.at(1)?.role === 'user')) &&
|
||||
(selectedModels[0] === model.id || atSelectedModel !== undefined)
|
||||
...(!$temporaryChatEnabled && !_chatId && (userMessage?.parentId ?? null) === null
|
||||
? {
|
||||
title_generation: $settings?.title?.auto ?? true,
|
||||
tags_generation: $settings?.autoTags ?? true
|
||||
|
|
@ -2403,10 +2400,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 +2468,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;
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -343,7 +343,9 @@
|
|||
}
|
||||
|
||||
chatInputElement?.setText(text);
|
||||
chatInputElement?.focus();
|
||||
if (!$showCallOverlay) {
|
||||
chatInputElement?.focus();
|
||||
}
|
||||
|
||||
if (text !== '') {
|
||||
text = await inputVariableHandler(text);
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -57,18 +57,124 @@
|
|||
|
||||
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 +201,7 @@
|
|||
}
|
||||
|
||||
messages = _messages.reverse();
|
||||
visibleEnd = messages.length;
|
||||
};
|
||||
|
||||
// Throttle message list rebuilds to once per animation frame during streaming.
|
||||
|
|
@ -113,6 +220,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,10 +535,6 @@
|
|||
showMessage({ id: parentMessageId }, false);
|
||||
};
|
||||
|
||||
onDestroy(() => {
|
||||
cancelAnimationFrame(pendingRebuild);
|
||||
});
|
||||
|
||||
const triggerScroll = () => {
|
||||
if (autoScroll) {
|
||||
const element = document.getElementById('messages-container');
|
||||
|
|
@ -465,7 +570,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 +605,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" />
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -1430,6 +1430,7 @@
|
|||
: ''}"
|
||||
style="fill: currentColor;"
|
||||
alt={action.name}
|
||||
draggable="false"
|
||||
/>
|
||||
</div>
|
||||
{:else}
|
||||
|
|
|
|||
|
|
@ -450,6 +450,7 @@
|
|||
});
|
||||
|
||||
if (res) {
|
||||
// $i18n.t('Model {{modelId}} not found')
|
||||
toast.success(
|
||||
$i18n.t('Model {{modelName}} deleted successfully', { modelName: model.name ?? model.id })
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@
|
|||
}
|
||||
|
||||
onDestroy(() => {
|
||||
window.removeEventListener('keydown', handleKeyDown);
|
||||
show = false;
|
||||
if (modalElement) {
|
||||
if (document.body.contains(modalElement)) {
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue