refac: modernize imports, standardize type hints and docstrings
This commit is contained in:
@@ -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.'
|
||||
|
||||
|
||||
@@ -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 ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user