Coverage for open_webui/internal/db.py: 53%
252 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1from __future__ import annotations
3import logging
4import os
5import re
6import sys
7from contextlib import asynccontextmanager, contextmanager
8from datetime import datetime, timedelta, timezone
9from typing import Any, Optional
10from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
12from open_webui.env import (
13 DATABASE_ENABLE_IAM_TOKEN_AUTH,
14 DATABASE_ENABLE_SESSION_SHARING,
15 DATABASE_ENABLE_SQLITE_WAL,
16 DATABASE_POOL_MAX_OVERFLOW,
17 DATABASE_POOL_RECYCLE,
18 DATABASE_POOL_SIZE,
19 DATABASE_POOL_TIMEOUT,
20 DATABASE_SCHEMA,
21 DATABASE_SQLITE_PRAGMA_BUSY_TIMEOUT,
22 DATABASE_SQLITE_PRAGMA_CACHE_SIZE,
23 DATABASE_SQLITE_PRAGMA_JOURNAL_SIZE_LIMIT,
24 DATABASE_SQLITE_PRAGMA_MMAP_SIZE,
25 DATABASE_SQLITE_PRAGMA_SYNCHRONOUS,
26 DATABASE_SQLITE_PRAGMA_TEMP_STORE,
27 DATABASE_URL,
28 ENABLE_DB_MIGRATIONS,
29 OPEN_WEBUI_DIR,
30 USE_SLIM,
31)
32from open_webui.utils.json_codec import JSONCodec
33from sqlalchemy import Dialect, MetaData, create_engine, event, types
34from sqlalchemy.engine.url import make_url
35from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
36from sqlalchemy.ext.declarative import declarative_base
37from sqlalchemy.orm import Session, scoped_session, sessionmaker
38from sqlalchemy.pool import NullPool, QueuePool
39from sqlalchemy.sql.type_api import _T
40from typing_extensions import Self
42log = logging.getLogger(__name__)
45# ── SSL URL normalization (used by sync engine & Alembic migrations) ─
46#
47# psycopg2 (sync) needs ``sslmode=`` in the connection string (it does
48# not recognise the bare ``ssl=`` key that some ORMs emit). The helpers
49# below strip all SSL-related query params, normalise them, and
50# reattach them in the canonical libpq form.
51#
52# The **async** engine now uses psycopg (v3), which speaks libpq
53# natively, so it needs no translation at all — the DATABASE_URL is
54# passed through as-is.
55# ─────────────────────────────────────────────────────────────────────
58def _pop_first(params: dict[str, list[str]], key: str) -> str | None:
59 """Pop a single-valued query param, returning ``None`` if absent."""
60 values = params.pop(key, None)
61 return values[0] if values else None
64def _is_postgres_url(url: str) -> bool:
65 """Return True if *url* looks like a PostgreSQL connection string."""
66 return bool(url) and any(url.startswith(p) for p in ('postgresql://', 'postgresql+', 'postgres://'))
69def extract_ssl_params_from_url(url: str) -> tuple[str, dict[str, str]]:
70 """Strip SSL query-string parameters from a PostgreSQL URL.
72 Returns ``(url_without_ssl, ssl_dict)`` where *ssl_dict* maps
73 canonical libpq key names (``sslmode``, ``sslrootcert``, …) to
74 their values. Non-PostgreSQL URLs are returned unchanged with an
75 empty dict.
76 """
77 if not _is_postgres_url(url): 77 ↛ 80line 77 didn't jump to line 80 because the condition on line 77 was always true
78 return url, {}
80 parsed = urlparse(url)
81 qp = parse_qs(parsed.query, keep_blank_values=True)
83 # Prefer sslmode (libpq canonical) over the bare ``ssl`` key.
84 sslmode_val = _pop_first(qp, 'sslmode')
85 ssl_val = _pop_first(qp, 'ssl')
86 ssl_mode = sslmode_val or ssl_val
88 ssl_dict: dict[str, str] = {}
89 if ssl_mode:
90 ssl_dict['sslmode'] = ssl_mode
91 for key in ('sslrootcert', 'sslcert', 'sslkey', 'sslcrl'):
92 val = _pop_first(qp, key)
93 if val:
94 ssl_dict[key] = val
96 if not ssl_dict:
97 return url, ssl_dict
99 cleaned_query = urlencode(qp, doseq=True)
100 return urlunparse(parsed._replace(query=cleaned_query)), ssl_dict
103def reattach_ssl_params_to_url(url_without_ssl: str, ssl_dict: dict[str, str]) -> str:
104 """Re-append SSL query-string parameters to a cleaned PostgreSQL URL.
106 Used for psycopg2/libpq consumers that expect ``sslmode`` and the
107 certificate-file keys in the connection string.
108 """
109 if not ssl_dict:
110 return url_without_ssl
112 parts = [f'{k}={v}' for k, v in ssl_dict.items() if v]
113 if not parts:
114 return url_without_ssl
116 sep = '&' if '?' in url_without_ssl else '?'
117 return f'{url_without_ssl}{sep}{"&".join(parts)}'
120# Backwards-compatible aliases for external callers.
121extract_ssl_mode_from_url = extract_ssl_params_from_url
122reattach_ssl_mode_to_url = reattach_ssl_params_to_url
125class JSONField(types.TypeDecorator): # TEXT-backed JSON storage
126 """Store arbitrary Python objects as JSON-encoded TEXT.
128 Used instead of native JSON columns for portability across SQLite and
129 PostgreSQL. Values are serialized with ``JSONCodec.dumps`` on write and
130 deserialized with ``JSONCodec.loads`` on read.
131 """
133 impl = types.UnicodeText
134 cache_ok = True
136 def process_bind_param(self, value: _T | None, dialect: Dialect) -> Any:
137 return JSONCodec.dumps(value) if value is not None else None
139 def process_result_value(self, value: _T | None, dialect: Dialect) -> Any:
140 return JSONCodec.loads(value) if value is not None else None
142 def copy(self, **kwargs: Any) -> Self:
143 return JSONField(length=self.impl.length)
146if USE_SLIM: 146 ↛ 147line 146 didn't jump to line 147 because the condition on line 146 was never true
147 if make_url(DATABASE_URL).get_backend_name() not in ('sqlite', 'postgresql', 'postgres'): 147 ↛ anywhereline 147 didn't jump anywhere: it always raised an exception.
148 raise ValueError(
149 'Slim requires SQLite or PostgreSQL for DATABASE_URL. Use the standard image for other databases.'
150 )
151 if DATABASE_ENABLE_IAM_TOKEN_AUTH:
152 raise ValueError(
153 'AWS RDS IAM authentication requires the standard image. Slim supports PostgreSQL database credentials.'
154 )
157# Normalize SSL params from the URL once; the sync engine needs them
158# reattached in canonical libpq form for psycopg2.
159_url_without_ssl, _ssl_dict = extract_ssl_params_from_url(DATABASE_URL)
161# For psycopg2 (sync engine), re-append sslmode + cert-file params.
162SQLALCHEMY_DATABASE_URL = reattach_ssl_params_to_url(_url_without_ssl, _ssl_dict) if _ssl_dict else DATABASE_URL
165class RDSIAMTokenAuth:
166 _refresh_after = timedelta(minutes=14)
168 def __init__(self, database_url: str) -> None:
169 url = make_url(database_url)
170 if not url.drivername.startswith(('postgresql', 'postgres')):
171 raise ValueError('DATABASE_ENABLE_IAM_TOKEN_AUTH is only supported for PostgreSQL databases')
172 if not url.host or not url.username:
173 raise ValueError('DATABASE_ENABLE_IAM_TOKEN_AUTH requires a database host and user')
175 self.host = url.host
176 self.port = url.port or 5432
177 self.username = url.username
178 self._client = None
179 self._token: str | None = None
180 self._expires_at = datetime.min.replace(tzinfo=timezone.utc)
182 @property
183 def client(self):
184 if self._client is None:
185 import boto3
187 self._client = boto3.client('rds')
188 return self._client
190 def get_password(self) -> str:
191 now = datetime.now(timezone.utc)
192 if self._token and now < self._expires_at:
193 return self._token
195 self._token = self.client.generate_db_auth_token(
196 DBHostname=self.host,
197 Port=self.port,
198 DBUsername=self.username,
199 )
200 self._expires_at = now + self._refresh_after
201 log.info('AWS RDS IAM database token refreshed; next refresh after %s', self._expires_at.isoformat())
202 return self._token
205_rds_iam_token_auth = RDSIAMTokenAuth(SQLALCHEMY_DATABASE_URL) if DATABASE_ENABLE_IAM_TOKEN_AUTH else None
208def _set_iam_token_password(dialect, conn_rec, cargs, cparams):
209 if _rds_iam_token_auth is not None:
210 cparams['password'] = _rds_iam_token_auth.get_password()
213def enable_iam_token_auth(connectable) -> None:
214 if _rds_iam_token_auth is None: 214 ↛ 217line 214 didn't jump to line 217 because the condition on line 214 was always true
215 return
217 engine = getattr(connectable, 'sync_engine', connectable)
218 url = engine.url
219 auth = _rds_iam_token_auth
220 # The token is bound to one host/port/user pair; leave other databases on their own credentials.
221 if (url.host, url.port or 5432, url.username) != (auth.host, auth.port, auth.username):
222 log.warning(
223 'AWS RDS IAM token auth not applied to %s: the token is issued for %s@%s:%s, '
224 'so this connection uses the password from its own URL',
225 url.render_as_string(hide_password=True),
226 auth.username,
227 auth.host,
228 auth.port,
229 )
230 return
232 if not event.contains(engine, 'do_connect', _set_iam_token_password):
233 event.listen(engine, 'do_connect', _set_iam_token_password)
236def _make_async_url(url: str) -> str:
237 """Convert a sync database URL to its async driver equivalent.
239 The async engine uses psycopg (v3) which speaks libpq natively,
240 so all standard connection-string parameters (``sslmode``,
241 ``options``, ``target_session_attrs``, etc.) are passed through
242 without any translation.
243 """
244 if url.startswith('sqlite+sqlcipher://'): 244 ↛ 245line 244 didn't jump to line 245 because the condition on line 244 was never true
245 raise ValueError(
246 'sqlite+sqlcipher:// URLs are not supported with async engine. '
247 'Use standard sqlite:// or postgresql:// instead.'
248 )
249 if url.startswith('sqlite:///') or url.startswith('sqlite://'): 249 ↛ 252line 249 didn't jump to line 252 because the condition on line 249 was always true
250 return url.replace('sqlite://', 'sqlite+aiosqlite://', 1)
251 # psycopg v3 — auto-selects async mode with create_async_engine
252 if url.startswith('postgresql+psycopg2://'):
253 return url.replace('postgresql+psycopg2://', 'postgresql+psycopg://', 1)
254 if url.startswith('postgresql://'):
255 return url.replace('postgresql://', 'postgresql+psycopg://', 1)
256 if url.startswith('postgres://'):
257 return url.replace('postgres://', 'postgresql+psycopg://', 1)
258 # For other dialects, return as-is and let SQLAlchemy handle it
259 return url
262def _json_codec_kwargs(kwargs: dict) -> dict:
263 """Default an engine to JSONCodec for native ``JSON`` columns.
265 Unlike ``JSONField``, those serialize through the engine, which otherwise uses
266 stdlib ``json``. With ``ENABLE_ORJSON`` off JSONCodec is stdlib ``json`` anyway.
267 """
268 kwargs.setdefault('json_serializer', JSONCodec.dumps)
269 kwargs.setdefault('json_deserializer', JSONCodec.loads)
270 return kwargs
273def _create_engine(*args, **kwargs):
274 """``create_engine`` with the app JSON codec wired in."""
275 return create_engine(*args, **_json_codec_kwargs(kwargs))
278def _create_async_engine(*args, **kwargs):
279 """``create_async_engine`` with the app JSON codec wired in."""
280 return create_async_engine(*args, **_json_codec_kwargs(kwargs))
283# ============================================================
284# SYNC ENGINE (used only for: startup migrations, config loading,
285# Alembic, peewee migration, health checks)
286# ============================================================
288# Handle SQLCipher URLs
289if SQLALCHEMY_DATABASE_URL.startswith('sqlite+sqlcipher://'): 289 ↛ 290line 289 didn't jump to line 290 because the condition on line 289 was never true
290 database_password = os.environ.get('DATABASE_PASSWORD')
291 if not database_password or database_password.strip() == '':
292 raise ValueError('DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs')
294 # Extract database path from SQLCipher URL
295 db_path = SQLALCHEMY_DATABASE_URL.replace('sqlite+sqlcipher://', '')
297 # Create a custom creator function that uses sqlcipher3
298 def create_sqlcipher_connection():
299 import sqlcipher3
301 conn = sqlcipher3.connect(db_path, check_same_thread=False)
302 conn.execute(f"PRAGMA key = '{database_password}'")
303 return conn
305 # The dummy "sqlite://" URL would cause SQLAlchemy to auto-select
306 # SingletonThreadPool, which non-deterministically closes in-use
307 # connections when thread count exceeds pool_size, leading to segfaults
308 # in the native sqlcipher3 C library. Use NullPool by default for safety,
309 # or QueuePool if DATABASE_POOL_SIZE is explicitly configured.
310 if isinstance(DATABASE_POOL_SIZE, int) and DATABASE_POOL_SIZE > 0: 310 ↛ 323line 310 didn't jump to line 323 because the condition on line 310 was always true
311 engine = _create_engine(
312 'sqlite://',
313 creator=create_sqlcipher_connection,
314 pool_size=DATABASE_POOL_SIZE,
315 max_overflow=DATABASE_POOL_MAX_OVERFLOW,
316 pool_timeout=DATABASE_POOL_TIMEOUT,
317 pool_recycle=DATABASE_POOL_RECYCLE,
318 pool_pre_ping=True,
319 poolclass=QueuePool,
320 echo=False,
321 )
322 else:
323 engine = _create_engine(
324 'sqlite://',
325 creator=create_sqlcipher_connection,
326 poolclass=NullPool,
327 echo=False,
328 )
330 log.info('Connected to encrypted SQLite database using SQLCipher')
332elif 'sqlite' in SQLALCHEMY_DATABASE_URL: 332 ↛ 402line 332 didn't jump to line 402 because the condition on line 332 was always true
333 engine = _create_engine(SQLALCHEMY_DATABASE_URL, connect_args={'check_same_thread': False})
335 def _apply_sqlite_pragmas(dbapi_connection):
336 """Apply all configured SQLite PRAGMAs to a raw DBAPI connection."""
337 # SQLite LIKE folds ASCII only; SQLAlchemy SQLite ILIKE compiles to lower(x) LIKE lower(?).
338 compiled_patterns = {}
340 def like(pattern, value, escape=None):
341 if pattern is None or value is None:
342 return None
344 pattern = str(pattern).lower()
345 escape = str(escape).lower() if escape is not None else None
346 key = (pattern, escape)
347 compiled = compiled_patterns.get(key)
348 if compiled is False: 348 ↛ 349line 348 didn't jump to line 349 because the condition on line 348 was never true
349 return False
350 if compiled is None:
351 regex = []
352 escaped = False
353 for char in pattern:
354 if escape and not escaped and char == escape:
355 escaped = True
356 continue
357 regex.append(
358 '.*' if not escaped and char == '%' else '.' if not escaped and char == '_' else re.escape(char)
359 )
360 escaped = False
361 if escaped: 361 ↛ 362line 361 didn't jump to line 362 because the condition on line 361 was never true
362 compiled = False
363 if len(compiled_patterns) >= 512:
364 compiled_patterns.clear()
365 compiled_patterns[key] = compiled
366 return False
367 compiled = re.compile(''.join(regex), re.DOTALL)
368 if len(compiled_patterns) >= 512: 368 ↛ 369line 368 didn't jump to line 369 because the condition on line 368 was never true
369 compiled_patterns.clear()
370 compiled_patterns[key] = compiled
372 return compiled.fullmatch(str(value).lower()) is not None
374 dbapi_connection.create_function('like', 2, like, deterministic=True)
375 dbapi_connection.create_function('like', 3, like, deterministic=True)
376 cursor = dbapi_connection.cursor()
377 if DATABASE_ENABLE_SQLITE_WAL: 377 ↛ 380line 377 didn't jump to line 380 because the condition on line 377 was always true
378 cursor.execute('PRAGMA journal_mode=WAL')
379 else:
380 cursor.execute('PRAGMA journal_mode=DELETE')
382 # Each PRAGMA is skipped when its env var is empty, allowing opt-out.
383 if DATABASE_SQLITE_PRAGMA_SYNCHRONOUS: 383 ↛ 385line 383 didn't jump to line 385 because the condition on line 383 was always true
384 cursor.execute(f'PRAGMA synchronous={DATABASE_SQLITE_PRAGMA_SYNCHRONOUS}')
385 if DATABASE_SQLITE_PRAGMA_BUSY_TIMEOUT: 385 ↛ 387line 385 didn't jump to line 387 because the condition on line 385 was always true
386 cursor.execute(f'PRAGMA busy_timeout={DATABASE_SQLITE_PRAGMA_BUSY_TIMEOUT}')
387 if DATABASE_SQLITE_PRAGMA_CACHE_SIZE: 387 ↛ 389line 387 didn't jump to line 389 because the condition on line 387 was always true
388 cursor.execute(f'PRAGMA cache_size={DATABASE_SQLITE_PRAGMA_CACHE_SIZE}')
389 if DATABASE_SQLITE_PRAGMA_TEMP_STORE: 389 ↛ 391line 389 didn't jump to line 391 because the condition on line 389 was always true
390 cursor.execute(f'PRAGMA temp_store={DATABASE_SQLITE_PRAGMA_TEMP_STORE}')
391 if DATABASE_SQLITE_PRAGMA_MMAP_SIZE: 391 ↛ 393line 391 didn't jump to line 393 because the condition on line 391 was always true
392 cursor.execute(f'PRAGMA mmap_size={DATABASE_SQLITE_PRAGMA_MMAP_SIZE}')
393 if DATABASE_SQLITE_PRAGMA_JOURNAL_SIZE_LIMIT: 393 ↛ 395line 393 didn't jump to line 395 because the condition on line 393 was always true
394 cursor.execute(f'PRAGMA journal_size_limit={DATABASE_SQLITE_PRAGMA_JOURNAL_SIZE_LIMIT}')
395 cursor.close()
397 def on_connect(dbapi_connection, connection_record):
398 _apply_sqlite_pragmas(dbapi_connection)
400 event.listen(engine, 'connect', on_connect)
401else:
402 if isinstance(DATABASE_POOL_SIZE, int):
403 if DATABASE_POOL_SIZE > 0:
404 engine = _create_engine(
405 SQLALCHEMY_DATABASE_URL,
406 pool_size=DATABASE_POOL_SIZE,
407 max_overflow=DATABASE_POOL_MAX_OVERFLOW,
408 pool_timeout=DATABASE_POOL_TIMEOUT,
409 pool_recycle=DATABASE_POOL_RECYCLE,
410 pool_pre_ping=True,
411 poolclass=QueuePool,
412 )
413 else:
414 engine = _create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True, poolclass=NullPool)
415 else:
416 engine = _create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True)
418enable_iam_token_auth(engine)
421# Sync session — used ONLY for startup config loading (config.py runs at import time)
422SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine, expire_on_commit=False)
423metadata_obj = MetaData(schema=DATABASE_SCHEMA)
424Base = declarative_base(metadata=metadata_obj)
425ScopedSession = scoped_session(SessionLocal)
428def get_session():
429 """Sync session generator — used ONLY for startup/config operations."""
430 db = SessionLocal()
431 try:
432 yield db
433 finally:
434 db.close()
437get_db = contextmanager(get_session)
440# ============================================================
441# ASYNC ENGINE (used for ALL runtime database operations)
442# ============================================================
444# psycopg (v3) speaks libpq natively — the full DATABASE_URL is passed
445# through as-is. SSL params, ``options``, ``target_session_attrs``, etc.
446# all work without any stripping or translation.
447ASYNC_SQLALCHEMY_DATABASE_URL = _make_async_url(SQLALCHEMY_DATABASE_URL)
449# psycopg v3 cannot run in async mode under Windows' default
450# ProactorEventLoop — switch to SelectorEventLoop before creating
451# the async engine. This runs at import time, which is early enough
452# to cover every entry point (workers, reload, direct invocations).
453if sys.platform == 'win32' and _is_postgres_url(DATABASE_URL): 453 ↛ 454line 453 didn't jump to line 454 because the condition on line 453 was never true
454 import asyncio
456 asyncio.set_event_loop_policy(asyncio.WindowsSelectorEventLoopPolicy())
458if 'sqlite' in ASYNC_SQLALCHEMY_DATABASE_URL: 458 ↛ 475line 458 didn't jump to line 475 because the condition on line 458 was always true
459 # Generous default — async coroutines + no session sharing = high connection demand.
460 # No pool_pre_ping: a local SQLite file cannot drop connections, and the
461 # ping costs a worker-thread hop plus a SELECT 1 on every checkout.
462 _sqlite_pool_size = DATABASE_POOL_SIZE if isinstance(DATABASE_POOL_SIZE, int) and DATABASE_POOL_SIZE > 0 else 512
463 async_engine = _create_async_engine(
464 ASYNC_SQLALCHEMY_DATABASE_URL,
465 connect_args={'check_same_thread': False},
466 pool_size=_sqlite_pool_size,
467 pool_timeout=DATABASE_POOL_TIMEOUT,
468 pool_recycle=DATABASE_POOL_RECYCLE,
469 )
471 @event.listens_for(async_engine.sync_engine, 'connect')
472 def _set_sqlite_pragmas(dbapi_connection, connection_record):
473 _apply_sqlite_pragmas(dbapi_connection)
474else:
475 if isinstance(DATABASE_POOL_SIZE, int):
476 if DATABASE_POOL_SIZE > 0:
477 async_engine = _create_async_engine(
478 ASYNC_SQLALCHEMY_DATABASE_URL,
479 pool_size=DATABASE_POOL_SIZE,
480 max_overflow=DATABASE_POOL_MAX_OVERFLOW,
481 pool_timeout=DATABASE_POOL_TIMEOUT,
482 pool_recycle=DATABASE_POOL_RECYCLE,
483 pool_pre_ping=True,
484 )
485 else:
486 async_engine = _create_async_engine(
487 ASYNC_SQLALCHEMY_DATABASE_URL,
488 pool_pre_ping=True,
489 poolclass=NullPool,
490 )
491 else:
492 async_engine = _create_async_engine(
493 ASYNC_SQLALCHEMY_DATABASE_URL,
494 pool_pre_ping=True,
495 )
497enable_iam_token_auth(async_engine)
500AsyncSessionLocal = async_sessionmaker(
501 bind=async_engine,
502 class_=AsyncSession,
503 autocommit=False,
504 autoflush=False,
505 expire_on_commit=False,
506)
509async def get_async_session():
510 """Async session generator for FastAPI Depends()."""
511 async with AsyncSessionLocal() as db:
512 try:
513 yield db
514 finally:
515 await db.close()
518@asynccontextmanager
519async def get_async_db():
520 """Async context manager for use outside of FastAPI dependency injection."""
521 async with AsyncSessionLocal() as db:
522 try:
523 yield db
524 finally:
525 await db.close()
528@asynccontextmanager
529async def get_async_db_context(db: AsyncSession | None = None):
530 """Async context manager that reuses an existing session if provided and session sharing is enabled."""
531 if isinstance(db, AsyncSession) and DATABASE_ENABLE_SESSION_SHARING: 531 ↛ 532line 531 didn't jump to line 532 because the condition on line 531 was never true
532 yield db
533 else:
534 async with get_async_db() as session:
535 yield session