This commit is contained in:
Timothy Jaeryang Baek
2026-04-21 14:58:28 +09:00
parent a27916d1db
commit c4aac0415c
2 changed files with 120 additions and 5 deletions
+115 -5
View File
@@ -1,8 +1,10 @@
import os
import json
import logging
import ssl as _stdlib_ssl
from contextlib import asynccontextmanager, contextmanager
from typing import Any, Optional
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
from open_webui.internal.wrappers import register_connection
from open_webui.env import (
@@ -35,6 +37,96 @@ from typing_extensions import Self
log = logging.getLogger(__name__)
def extract_ssl_mode_from_url(url: str) -> tuple[str, str | None]:
"""Strip SSL query-string parameters from a PostgreSQL URL.
asyncpg and psycopg2 use different query-string keys for SSL
(``ssl`` vs ``sslmode``). This helper removes **both** from the
URL so that each driver can receive the correct parameter through
its own mechanism (query-string re-injection for psycopg2,
``connect_args`` for asyncpg).
Returns
-------
(url_without_ssl, ssl_mode)
*url_without_ssl* is the original URL with ``ssl`` / ``sslmode``
query parameters removed. *ssl_mode* is the extracted mode
string (e.g. ``'require'``), or ``None`` if neither parameter
was present.
Non-PostgreSQL URLs are returned unchanged with ``ssl_mode=None``.
"""
if not url or not any(
url.startswith(prefix)
for prefix in ('postgresql://', 'postgresql+', 'postgres://')
):
return url, None
parsed = urlparse(url)
query_params = parse_qs(parsed.query, keep_blank_values=True)
# Prefer sslmode (libpq canonical) over the asyncpg-only ssl key.
ssl_mode: str | None = None
for key in ('sslmode', 'ssl'):
values = query_params.pop(key, None)
if values and ssl_mode is None:
ssl_mode = values[0]
if ssl_mode is None:
# Nothing to strip — return the URL untouched.
return url, None
# Rebuild the query string without the SSL keys.
remaining_query = urlencode(query_params, doseq=True)
url_without_ssl = urlunparse(parsed._replace(query=remaining_query))
return url_without_ssl, ssl_mode
def build_asyncpg_ssl_args(ssl_mode: str | None) -> dict:
"""Convert a libpq-style SSL mode value to asyncpg ``connect_args``.
Returns a dict suitable for unpacking into
``create_async_engine(..., connect_args=...)``.
"""
if ssl_mode is None:
return {}
mode = ssl_mode.lower()
if mode == 'disable':
return {'connect_args': {'ssl': False}}
if mode in ('allow', 'prefer'):
# asyncpg has no direct equivalent — omit to let it try without.
return {}
if mode == 'require':
# SSL required but no certificate verification (matches libpq).
ctx = _stdlib_ssl.create_default_context()
ctx.check_hostname = False
ctx.verify_mode = _stdlib_ssl.CERT_NONE
return {'connect_args': {'ssl': ctx}}
if mode in ('verify-ca', 'verify-full'):
# Full verification — use the system trust store.
ctx = _stdlib_ssl.create_default_context()
if mode == 'verify-ca':
ctx.check_hostname = False
return {'connect_args': {'ssl': ctx}}
# Unknown value — pass through as-is and let asyncpg decide.
return {'connect_args': {'ssl': ssl_mode}}
def reattach_ssl_mode_to_url(url_without_ssl: str, ssl_mode: str | None) -> str:
"""Re-append ``sslmode=<value>`` to a cleaned PostgreSQL URL.
Used for psycopg2 / libpq consumers that expect the canonical
``sslmode`` query-string key.
"""
if ssl_mode is None:
return url_without_ssl
separator = '&' if '?' in url_without_ssl else '?'
return f'{url_without_ssl}{separator}sslmode={ssl_mode}'
class JSONField(types.TypeDecorator):
impl = types.Text
cache_ok = True
@@ -60,10 +152,14 @@ class JSONField(types.TypeDecorator):
# Workaround to handle the peewee migration
# This is required to ensure the peewee migration is handled before the alembic migration
def handle_peewee_migration(DATABASE_URL):
# db = None
db = None
try:
# Normalize SSL params so psycopg2 always sees `sslmode=` (never `ssl=`).
url_without_ssl, ssl_mode = extract_ssl_mode_from_url(DATABASE_URL)
normalized_url = reattach_ssl_mode_to_url(url_without_ssl, ssl_mode)
# Replace the postgresql:// with postgres:// to handle the peewee migration
db = register_connection(DATABASE_URL.replace('postgresql://', 'postgres://'))
db = register_connection(normalized_url.replace('postgresql://', 'postgres://'))
migrate_dir = OPEN_WEBUI_DIR / 'internal' / 'migrations'
router = Router(db, logger=log, migrate_dir=migrate_dir)
router.run()
@@ -79,14 +175,20 @@ def handle_peewee_migration(DATABASE_URL):
db.close()
# Assert if db connection has been closed
assert db.is_closed(), 'Database connection is still open.'
if db is not None:
assert db.is_closed(), 'Database connection is still open.'
if ENABLE_DB_MIGRATIONS:
handle_peewee_migration(DATABASE_URL)
SQLALCHEMY_DATABASE_URL = DATABASE_URL
# Normalize SSL params from the URL once; each engine branch re-injects
# the driver-appropriate form.
DATABASE_URL_WITHOUT_SSL, DATABASE_SSL_MODE = extract_ssl_mode_from_url(DATABASE_URL)
# For psycopg2 (sync engine), re-append sslmode=<value>.
SQLALCHEMY_DATABASE_URL = reattach_ssl_mode_to_url(DATABASE_URL_WITHOUT_SSL, DATABASE_SSL_MODE) if DATABASE_SSL_MODE else DATABASE_URL
def _make_async_url(url: str) -> str:
@@ -229,7 +331,8 @@ get_db = contextmanager(get_session)
# ASYNC ENGINE (used for ALL runtime database operations)
# ============================================================
ASYNC_SQLALCHEMY_DATABASE_URL = _make_async_url(SQLALCHEMY_DATABASE_URL)
# Use the SSL-stripped URL for asyncpg — SSL is injected via connect_args.
ASYNC_SQLALCHEMY_DATABASE_URL = _make_async_url(DATABASE_URL_WITHOUT_SSL if DATABASE_SSL_MODE else SQLALCHEMY_DATABASE_URL)
if 'sqlite' in ASYNC_SQLALCHEMY_DATABASE_URL:
# Generous default — async coroutines + no session sharing = high connection demand.
@@ -251,6 +354,10 @@ if 'sqlite' in ASYNC_SQLALCHEMY_DATABASE_URL:
def _set_sqlite_pragmas(dbapi_connection, connection_record):
_apply_sqlite_pragmas(dbapi_connection)
else:
# Inject asyncpg-compatible SSL connect_args when the user specified
# sslmode/ssl in DATABASE_URL.
asyncpg_ssl_args = build_asyncpg_ssl_args(DATABASE_SSL_MODE)
if isinstance(DATABASE_POOL_SIZE, int):
if DATABASE_POOL_SIZE > 0:
async_engine = create_async_engine(
@@ -260,17 +367,20 @@ else:
pool_timeout=DATABASE_POOL_TIMEOUT,
pool_recycle=DATABASE_POOL_RECYCLE,
pool_pre_ping=True,
**asyncpg_ssl_args,
)
else:
async_engine = create_async_engine(
ASYNC_SQLALCHEMY_DATABASE_URL,
pool_pre_ping=True,
poolclass=NullPool,
**asyncpg_ssl_args,
)
else:
async_engine = create_async_engine(
ASYNC_SQLALCHEMY_DATABASE_URL,
pool_pre_ping=True,
**asyncpg_ssl_args,
)
+5
View File
@@ -5,6 +5,7 @@ from alembic import context
from open_webui.models.auths import Auth
from open_webui.models.calendar import Calendar, CalendarEvent, CalendarEventAttendee # noqa: F401
from open_webui.env import DATABASE_URL, DATABASE_PASSWORD, LOG_FORMAT
from open_webui.internal.db import extract_ssl_mode_from_url, reattach_ssl_mode_to_url
from sqlalchemy import engine_from_config, pool, create_engine
# this is the Alembic Config object, which provides
@@ -36,6 +37,10 @@ target_metadata = Auth.metadata
DB_URL = DATABASE_URL
# Normalize SSL query params for psycopg2 (Alembic uses psycopg2, not asyncpg).
url_without_ssl, ssl_mode = extract_ssl_mode_from_url(DB_URL)
DB_URL = reattach_ssl_mode_to_url(url_without_ssl, ssl_mode) if ssl_mode else DB_URL
if DB_URL:
config.set_main_option('sqlalchemy.url', DB_URL.replace('%', '%%'))