refac: modernize imports, standardize type hints and docstrings

This commit is contained in:
Timothy Jaeryang Baek
2026-05-12 06:30:38 +09:00
parent 998c86a52b
commit a59c967d7e
16 changed files with 118 additions and 142 deletions
+3 -3
View File
@@ -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.'
+1 -2
View File
@@ -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 ─────────────────────────────────────────────────
+4 -5
View File
@@ -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:
+37 -50
View File
@@ -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(
+3 -3
View File
@@ -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:
@@ -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.
@@ -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.
@@ -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',
+5 -8
View File
@@ -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
+3 -3
View File
@@ -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:
+5 -5
View File
@@ -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.
+3 -3
View File
@@ -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:
+3 -3
View File
@@ -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.
@@ -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.
+4 -4
View File
@@ -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 = {
+29 -36
View File
@@ -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: