From a59c967d7eede4882bcc013f541780967730b2dc Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Tue, 12 May 2026 06:30:38 +0900 Subject: [PATCH] refac: modernize imports, standardize type hints and docstrings --- backend/open_webui/constants.py | 6 +- backend/open_webui/migrations/env.py | 3 +- backend/open_webui/migrations/util.py | 9 +- backend/open_webui/retrieval/utils.py | 87 ++++++++----------- backend/open_webui/retrieval/web/brave.py | 6 +- .../open_webui/retrieval/web/duckduckgo.py | 12 +-- .../open_webui/retrieval/web/google_pse.py | 10 +-- .../open_webui/retrieval/web/jina_search.py | 7 +- backend/open_webui/retrieval/web/main.py | 13 ++- backend/open_webui/retrieval/web/mojeek.py | 6 +- backend/open_webui/retrieval/web/searxng.py | 10 +-- backend/open_webui/retrieval/web/serper.py | 6 +- backend/open_webui/retrieval/web/serply.py | 6 +- backend/open_webui/retrieval/web/serpstack.py | 6 +- backend/open_webui/retrieval/web/tavily.py | 8 +- backend/open_webui/socket/main.py | 65 +++++++------- 16 files changed, 118 insertions(+), 142 deletions(-) diff --git a/backend/open_webui/constants.py b/backend/open_webui/constants.py index ad1bdf4a20..c767e72520 100644 --- a/backend/open_webui/constants.py +++ b/backend/open_webui/constants.py @@ -72,11 +72,11 @@ class ERROR_MESSAGES(str, Enum): EMPTY_CONTENT = 'The content provided is empty. Please ensure that there is text or data present before proceeding.' - DB_NOT_SQLITE = 'This feature is only available when running with SQLite databases.' + DB_NOT_SQLITE = 'This feature is only available with SQLite databases.' - INVALID_URL = 'Oops! The URL you provided is invalid. Please double-check and try again.' + INVALID_URL = 'The URL you provided is invalid. Please double-check and try again.' - WEB_SEARCH_ERROR = lambda err='': f'{err if err else "Oops! Something went wrong while searching the web."}' + WEB_SEARCH_ERROR = lambda err='': err if err else 'Something went wrong while searching the web.' OLLAMA_API_DISABLED = 'The Ollama API is disabled. Please enable it to use this feature.' diff --git a/backend/open_webui/migrations/env.py b/backend/open_webui/migrations/env.py index dc581c6bb3..00c9e9569e 100644 --- a/backend/open_webui/migrations/env.py +++ b/backend/open_webui/migrations/env.py @@ -9,12 +9,11 @@ import logging from logging.config import fileConfig from alembic import context -from sqlalchemy import create_engine, engine_from_config, pool - from open_webui.env import DATABASE_PASSWORD, DATABASE_URL, LOG_FORMAT from open_webui.internal.db import extract_ssl_params_from_url, reattach_ssl_params_to_url from open_webui.models.auths import Auth from open_webui.models.calendar import Calendar, CalendarEvent, CalendarEventAttendee # noqa: F401 +from sqlalchemy import create_engine, engine_from_config, pool # ── Alembic config & logging ───────────────────────────────────────────────── diff --git a/backend/open_webui/migrations/util.py b/backend/open_webui/migrations/util.py index e9f3aa7dd4..807baad4ba 100644 --- a/backend/open_webui/migrations/util.py +++ b/backend/open_webui/migrations/util.py @@ -1,14 +1,13 @@ """Alembic migration utilities.""" from alembic import op -from sqlalchemy import inspect as sa_inspect +from sqlalchemy import inspect def get_existing_tables() -> set[str]: - """Return the set of table names already present in the database.""" - bind = op.get_bind() - inspector = sa_inspect(bind) - return set(inspector.get_table_names()) + """Return table names already present in the database.""" + conn = op.get_bind() + return set(inspect(conn).get_table_names()) def get_revision_id() -> str: diff --git a/backend/open_webui/retrieval/utils.py b/backend/open_webui/retrieval/utils.py index 8e672b7a8f..9635f5b2e1 100644 --- a/backend/open_webui/retrieval/utils.py +++ b/backend/open_webui/retrieval/utils.py @@ -1,16 +1,15 @@ -import logging -import os -from typing import Awaitable, Optional, Union - -import requests -import aiohttp import asyncio import hashlib -from concurrent.futures import ThreadPoolExecutor -import time +import logging +import os import re - +import time +from concurrent.futures import ThreadPoolExecutor +from typing import Awaitable, Optional, Union from urllib.parse import quote + +import aiohttp +import requests from huggingface_hub import snapshot_download from langchain_classic.retrievers import ( ContextualCompressionRetriever, @@ -18,41 +17,33 @@ from langchain_classic.retrievers import ( ) 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 - - -from open_webui.models.users import UserModel -from open_webui.models.files import Files -from open_webui.models.knowledge import Knowledges - -from open_webui.models.chats import Chats -from open_webui.models.notes import Notes -from open_webui.models.access_grants import AccessGrants -from open_webui.utils.access_control.files import has_access_to_file - -from open_webui.retrieval.vector.main import GetResult -from open_webui.utils.headers import include_user_info_headers -from open_webui.utils.misc import get_message_list - -from open_webui.retrieval.web.utils import get_web_loader -from open_webui.retrieval.loaders.youtube import YoutubeLoader - - -from open_webui.env import ( - AIOHTTP_CLIENT_TIMEOUT, - AIOHTTP_CLIENT_ALLOW_REDIRECTS, - OFFLINE_MODE, - ENABLE_FORWARD_USER_INFO_HEADERS, - AIOHTTP_CLIENT_SESSION_SSL, -) from open_webui.config import ( - RAG_EMBEDDING_QUERY_PREFIX, RAG_EMBEDDING_CONTENT_PREFIX, RAG_EMBEDDING_PREFIX_FIELD_NAME, + RAG_EMBEDDING_QUERY_PREFIX, + VECTOR_DB, ) +from open_webui.env import ( + AIOHTTP_CLIENT_ALLOW_REDIRECTS, + AIOHTTP_CLIENT_SESSION_SSL, + AIOHTTP_CLIENT_TIMEOUT, + ENABLE_FORWARD_USER_INFO_HEADERS, + OFFLINE_MODE, +) +from open_webui.models.access_grants import AccessGrants +from open_webui.models.chats import Chats +from open_webui.models.files import Files +from open_webui.models.knowledge import Knowledges +from open_webui.models.notes import Notes +from open_webui.models.users import UserModel +from open_webui.retrieval.loaders.youtube import YoutubeLoader +from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT +from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT +from open_webui.retrieval.vector.main import GetResult +from open_webui.retrieval.web.utils import get_web_loader +from open_webui.utils.access_control.files import has_access_to_file +from open_webui.utils.headers import include_user_info_headers +from open_webui.utils.misc import get_message_list log = logging.getLogger(__name__) @@ -525,11 +516,8 @@ def get_all_items_from_collections(collection_names: list[str]) -> dict: async def query_collection( - request, - collection_names: list[str], - queries: list[str], - embedding_function, - k: int, + request, collection_names: list[str], queries: list[str], + embedding_function, k: int, ) -> dict: # When request is provided, try hybrid search + reranking if enabled if request and request.app.state.config.ENABLE_RAG_HYBRID_SEARCH: @@ -1460,13 +1448,12 @@ class RerankCompressor(BaseDocumentCompressor): if reranking: scores = await asyncio.to_thread(self.reranking_function, query, documents) else: - from sentence_transformers import util + from sentence_transformers import util as st_util query_embedding = await self.embedding_function(query, RAG_EMBEDDING_QUERY_PREFIX) - document_embedding = await self.embedding_function( - [doc.page_content for doc in documents], RAG_EMBEDDING_CONTENT_PREFIX - ) - scores = util.cos_sim(query_embedding, document_embedding)[0] + doc_texts = [doc.page_content for doc in documents] + document_embedding = await self.embedding_function(doc_texts, RAG_EMBEDDING_CONTENT_PREFIX) + scores = st_util.cos_sim(query_embedding, document_embedding)[0] if scores is not None: docs_with_scores = list( diff --git a/backend/open_webui/retrieval/web/brave.py b/backend/open_webui/retrieval/web/brave.py index 9e663c2684..e06d4594fb 100644 --- a/backend/open_webui/retrieval/web/brave.py +++ b/backend/open_webui/retrieval/web/brave.py @@ -1,14 +1,14 @@ +from __future__ import annotations + import logging import time -from typing import Optional import requests from open_webui.retrieval.web.main import SearchResult, get_filtered_results log = logging.getLogger(__name__) - -def search_brave(api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None) -> list[SearchResult]: +def search_brave(api_key: str, query: str, count: int, filter_list: list[str | None] = None) -> list[SearchResult]: """Search using Brave's Search API and return the results as a list of SearchResult objects. Args: diff --git a/backend/open_webui/retrieval/web/duckduckgo.py b/backend/open_webui/retrieval/web/duckduckgo.py index 5b2f076227..5eb3e73a10 100644 --- a/backend/open_webui/retrieval/web/duckduckgo.py +++ b/backend/open_webui/retrieval/web/duckduckgo.py @@ -1,20 +1,20 @@ +from __future__ import annotations + import logging import urllib.request -from typing import Optional -from open_webui.retrieval.web.main import SearchResult, get_filtered_results from ddgs import DDGS from ddgs.exceptions import RatelimitException +from open_webui.retrieval.web.main import SearchResult, get_filtered_results log = logging.getLogger(__name__) - def search_duckduckgo( query: str, count: int, - filter_list: Optional[list[str]] = None, - concurrent_requests: Optional[int] = None, - backend: Optional[str] = 'auto', + filter_list: list[str | None] = None, + concurrent_requests: int | None = None, + backend: str | None = 'auto', ) -> list[SearchResult]: """ Search using DuckDuckGo's Search API and return the results as a list of SearchResult objects. diff --git a/backend/open_webui/retrieval/web/google_pse.py b/backend/open_webui/retrieval/web/google_pse.py index bb0a852658..5391e2ba8b 100644 --- a/backend/open_webui/retrieval/web/google_pse.py +++ b/backend/open_webui/retrieval/web/google_pse.py @@ -1,19 +1,19 @@ +from __future__ import annotations + import logging -from typing import Optional import requests from open_webui.retrieval.web.main import SearchResult, get_filtered_results log = logging.getLogger(__name__) - def search_google_pse( api_key: str, search_engine_id: str, query: str, count: int, - filter_list: Optional[list[str]] = None, - referer: Optional[str] = None, + filter_list: list[str | None] = None, + referer: str | None = None, ) -> list[SearchResult]: """Search using Google's Programmable Search Engine API and return the results as a list of SearchResult objects. Handles pagination for counts greater than 10. @@ -23,7 +23,7 @@ def search_google_pse( search_engine_id (str): A Programmable Search Engine ID query (str): The query to search for count (int): The number of results to return (max 100, as PSE max results per query is 10 and max page is 10) - filter_list (Optional[list[str]], optional): A list of keywords to filter out from results. Defaults to None. + filter_list (list[str | None], optional): A list of keywords to filter out from results. Defaults to None. Returns: list[SearchResult]: A list of SearchResult objects. diff --git a/backend/open_webui/retrieval/web/jina_search.py b/backend/open_webui/retrieval/web/jina_search.py index b3266c47d0..7d6585ea3d 100644 --- a/backend/open_webui/retrieval/web/jina_search.py +++ b/backend/open_webui/retrieval/web/jina_search.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import logging import requests @@ -6,7 +8,6 @@ from yarl import URL log = logging.getLogger(__name__) - def search_jina(api_key: str, query: str, count: int, base_url: str = '') -> list[SearchResult]: """ Search using Jina's Search API and return the results as a list of SearchResult objects. @@ -17,9 +18,9 @@ def search_jina(api_key: str, query: str, count: int, base_url: str = '') -> lis base_url (str): Optional custom base URL for the Jina API Returns: - list[SearchResult]: A list of search results + A list of SearchResult objects. """ - jina_search_endpoint = base_url if base_url else 'https://s.jina.ai/' + jina_search_endpoint = base_url or 'https://s.jina.ai/' headers = { 'Accept': 'application/json', diff --git a/backend/open_webui/retrieval/web/main.py b/backend/open_webui/retrieval/web/main.py index 3a8fed52dd..e0dd30c01d 100644 --- a/backend/open_webui/retrieval/web/main.py +++ b/backend/open_webui/retrieval/web/main.py @@ -1,13 +1,11 @@ -import validators +from __future__ import annotations -from typing import Optional from urllib.parse import urlparse -from pydantic import BaseModel - +import validators from open_webui.retrieval.web.utils import resolve_hostname from open_webui.utils.misc import is_string_allowed - +from pydantic import BaseModel def get_filtered_results(results, filter_list): if not filter_list: @@ -39,8 +37,7 @@ def get_filtered_results(results, filter_list): return filtered_results - class SearchResult(BaseModel): link: str - title: Optional[str] - snippet: Optional[str] + title: str | None + snippet: str | None diff --git a/backend/open_webui/retrieval/web/mojeek.py b/backend/open_webui/retrieval/web/mojeek.py index a094ef6fc8..495180a828 100644 --- a/backend/open_webui/retrieval/web/mojeek.py +++ b/backend/open_webui/retrieval/web/mojeek.py @@ -1,13 +1,13 @@ +from __future__ import annotations + import logging -from typing import Optional import requests from open_webui.retrieval.web.main import SearchResult, get_filtered_results log = logging.getLogger(__name__) - -def search_mojeek(api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None) -> list[SearchResult]: +def search_mojeek(api_key: str, query: str, count: int, filter_list: list[str | None] = None) -> list[SearchResult]: """Search using Mojeek's Search API and return the results as a list of SearchResult objects. Args: diff --git a/backend/open_webui/retrieval/web/searxng.py b/backend/open_webui/retrieval/web/searxng.py index 2b7bd04895..f0b8b7ab51 100644 --- a/backend/open_webui/retrieval/web/searxng.py +++ b/backend/open_webui/retrieval/web/searxng.py @@ -1,17 +1,17 @@ +from __future__ import annotations + import logging -from typing import Optional import requests from open_webui.retrieval.web.main import SearchResult, get_filtered_results log = logging.getLogger(__name__) - -def search_searxng( +def search_searxng( # noqa: PLR0913 query_url: str, query: str, count: int, - filter_list: Optional[list[str]] = None, + filter_list: list[str | None] = None, **kwargs, ) -> list[SearchResult]: """ @@ -28,7 +28,7 @@ def search_searxng( language (str): Language filter for the search results; e.g., "all", "en-US", "es". Defaults to "all". safesearch (int): Safe search filter for safer web results; 0 = off, 1 = moderate, 2 = strict. Defaults to 1 (moderate). time_range (str): Time range for filtering results by date; e.g., "2023-04-05..today" or "all-time". Defaults to ''. - categories: (Optional[list[str]]): Specific categories within which the search should be performed, defaulting to an empty string if not provided. + categories: (list[str | None]): Specific categories within which the search should be performed, defaulting to an empty string if not provided. Returns: list[SearchResult]: A list of SearchResults sorted by relevance score in descending order. diff --git a/backend/open_webui/retrieval/web/serper.py b/backend/open_webui/retrieval/web/serper.py index 9f1a8e1b3a..b7fce3b417 100644 --- a/backend/open_webui/retrieval/web/serper.py +++ b/backend/open_webui/retrieval/web/serper.py @@ -1,14 +1,14 @@ +from __future__ import annotations + import json import logging -from typing import Optional import requests from open_webui.retrieval.web.main import SearchResult, get_filtered_results log = logging.getLogger(__name__) - -def search_serper(api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None) -> list[SearchResult]: +def search_serper(api_key: str, query: str, count: int, filter_list: list[str | None] = None) -> list[SearchResult]: """Search using serper.dev's API and return the results as a list of SearchResult objects. Args: diff --git a/backend/open_webui/retrieval/web/serply.py b/backend/open_webui/retrieval/web/serply.py index f245392b75..8a6b437eee 100644 --- a/backend/open_webui/retrieval/web/serply.py +++ b/backend/open_webui/retrieval/web/serply.py @@ -1,5 +1,6 @@ +from __future__ import annotations + import logging -from typing import Optional from urllib.parse import urlencode import requests @@ -7,7 +8,6 @@ from open_webui.retrieval.web.main import SearchResult, get_filtered_results log = logging.getLogger(__name__) - def search_serply( api_key: str, query: str, @@ -16,7 +16,7 @@ def search_serply( limit: int = 10, device_type: str = 'desktop', proxy_location: str = 'US', - filter_list: Optional[list[str]] = None, + filter_list: list[str | None] = None, ) -> list[SearchResult]: """Search using serper.dev's API and return the results as a list of SearchResult objects. diff --git a/backend/open_webui/retrieval/web/serpstack.py b/backend/open_webui/retrieval/web/serpstack.py index 28a4956645..93235fbfca 100644 --- a/backend/open_webui/retrieval/web/serpstack.py +++ b/backend/open_webui/retrieval/web/serpstack.py @@ -1,17 +1,17 @@ +from __future__ import annotations + import logging -from typing import Optional import requests from open_webui.retrieval.web.main import SearchResult, get_filtered_results log = logging.getLogger(__name__) - def search_serpstack( api_key: str, query: str, count: int, - filter_list: Optional[list[str]] = None, + filter_list: list[str | None] = None, https_enabled: bool = True, ) -> list[SearchResult]: """Search using serpstack.com's and return the results as a list of SearchResult objects. diff --git a/backend/open_webui/retrieval/web/tavily.py b/backend/open_webui/retrieval/web/tavily.py index 6b52bbb45b..bca2e20990 100644 --- a/backend/open_webui/retrieval/web/tavily.py +++ b/backend/open_webui/retrieval/web/tavily.py @@ -1,17 +1,17 @@ +from __future__ import annotations + import logging -from typing import Optional import requests from open_webui.retrieval.web.main import SearchResult, get_filtered_results log = logging.getLogger(__name__) - def search_tavily( api_key: str, query: str, count: int, - filter_list: Optional[list[str]] = None, + filter_list: list[str | None] = None, # **kwargs, ) -> list[SearchResult]: """Search using Tavily's Search API and return the results as a list of SearchResult objects. @@ -22,7 +22,7 @@ def search_tavily( count (int): The maximum number of results to return Returns: - list[SearchResult]: A list of search results + A list of SearchResult objects. """ url = 'https://api.tavily.com/search' headers = { diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index d59ff53277..bfb1003678 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -1,55 +1,48 @@ import asyncio -import random - -import socketio import logging +import random import sys import time from typing import Dict, Set -from redis import asyncio as aioredis + import pycrdt as Y - -from open_webui.models.users import Users, UserNameResponse -from open_webui.models.channels import Channels -from open_webui.models.chats import Chats -from open_webui.models.notes import Notes, NoteUpdateForm -from open_webui.utils.redis import ( - get_sentinels_from_env, - get_sentinel_url_from_env, -) - +import socketio from open_webui.config import ( CORS_ALLOW_ORIGIN, ) - from open_webui.env import ( - VERSION, ENABLE_WEBSOCKET_SUPPORT, + GLOBAL_LOG_LEVEL, + REDIS_KEY_PREFIX, + VERSION, + WEBSOCKET_EVENT_CALLER_TIMEOUT, WEBSOCKET_MANAGER, - WEBSOCKET_REDIS_URL, WEBSOCKET_REDIS_CLUSTER, WEBSOCKET_REDIS_LOCK_TIMEOUT, - WEBSOCKET_SENTINEL_PORT, - WEBSOCKET_SENTINEL_HOSTS, - REDIS_KEY_PREFIX, WEBSOCKET_REDIS_OPTIONS, - WEBSOCKET_SERVER_PING_TIMEOUT, - WEBSOCKET_SERVER_PING_INTERVAL, - WEBSOCKET_SERVER_LOGGING, + WEBSOCKET_REDIS_URL, + WEBSOCKET_SENTINEL_HOSTS, + WEBSOCKET_SENTINEL_PORT, WEBSOCKET_SERVER_ENGINEIO_LOGGING, - WEBSOCKET_EVENT_CALLER_TIMEOUT, + WEBSOCKET_SERVER_LOGGING, + WEBSOCKET_SERVER_PING_INTERVAL, + WEBSOCKET_SERVER_PING_TIMEOUT, ) -from open_webui.utils.auth import decode_token +from open_webui.models.access_grants import AccessGrants +from open_webui.models.channels import Channels +from open_webui.models.chats import Chats +from open_webui.models.notes import Notes, NoteUpdateForm +from open_webui.models.users import UserNameResponse, Users from open_webui.socket.utils import RedisDict, RedisLock, YdocManager from open_webui.tasks import create_task, stop_item_tasks -from open_webui.utils.redis import get_redis_connection from open_webui.utils.access_control import has_permission -from open_webui.models.access_grants import AccessGrants - - -from open_webui.env import ( - GLOBAL_LOG_LEVEL, +from open_webui.utils.auth import decode_token +from open_webui.utils.redis import ( + get_redis_connection, + get_sentinel_url_from_env, + get_sentinels_from_env, ) +from redis import asyncio as aioredis logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) log = logging.getLogger(__name__) @@ -371,15 +364,15 @@ async def connect(sid, environ, auth): @sio.on('user-join') async def user_join(sid, data): - auth = data['auth'] if 'auth' in data else None + auth = data.get('auth') if not auth or 'token' not in auth: return - data = decode_token(auth['token']) - if data is None or 'id' not in data: + token_data = decode_token(auth['token']) + if token_data is None or 'id' not in token_data: return - user = await Users.get_user_by_id(data['id']) + user = await Users.get_user_by_id(token_data['id']) if not user: return @@ -845,7 +838,7 @@ async def _make_channel_emitter(request_info): THROTTLE_INTERVAL = 0.15 # ~6 updates/sec async def _emit_channel_update(content: str, done: bool = False): - from open_webui.models.messages import Messages, MessageForm + from open_webui.models.messages import MessageForm, Messages update_form = MessageForm(content=content) if done: