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

1from __future__ import annotations 

2 

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 

11 

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 

41 

42log = logging.getLogger(__name__) 

43 

44 

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# ───────────────────────────────────────────────────────────────────── 

56 

57 

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 

62 

63 

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://')) 

67 

68 

69def extract_ssl_params_from_url(url: str) -> tuple[str, dict[str, str]]: 

70 """Strip SSL query-string parameters from a PostgreSQL URL. 

71 

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, {} 

79 

80 parsed = urlparse(url) 

81 qp = parse_qs(parsed.query, keep_blank_values=True) 

82 

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 

87 

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 

95 

96 if not ssl_dict: 

97 return url, ssl_dict 

98 

99 cleaned_query = urlencode(qp, doseq=True) 

100 return urlunparse(parsed._replace(query=cleaned_query)), ssl_dict 

101 

102 

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. 

105 

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 

111 

112 parts = [f'{k}={v}' for k, v in ssl_dict.items() if v] 

113 if not parts: 

114 return url_without_ssl 

115 

116 sep = '&' if '?' in url_without_ssl else '?' 

117 return f'{url_without_ssl}{sep}{"&".join(parts)}' 

118 

119 

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 

123 

124 

125class JSONField(types.TypeDecorator): # TEXT-backed JSON storage 

126 """Store arbitrary Python objects as JSON-encoded TEXT. 

127 

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 """ 

132 

133 impl = types.UnicodeText 

134 cache_ok = True 

135 

136 def process_bind_param(self, value: _T | None, dialect: Dialect) -> Any: 

137 return JSONCodec.dumps(value) if value is not None else None 

138 

139 def process_result_value(self, value: _T | None, dialect: Dialect) -> Any: 

140 return JSONCodec.loads(value) if value is not None else None 

141 

142 def copy(self, **kwargs: Any) -> Self: 

143 return JSONField(length=self.impl.length) 

144 

145 

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 ) 

155 

156 

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) 

160 

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 

163 

164 

165class RDSIAMTokenAuth: 

166 _refresh_after = timedelta(minutes=14) 

167 

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') 

174 

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) 

181 

182 @property 

183 def client(self): 

184 if self._client is None: 

185 import boto3 

186 

187 self._client = boto3.client('rds') 

188 return self._client 

189 

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 

194 

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 

203 

204 

205_rds_iam_token_auth = RDSIAMTokenAuth(SQLALCHEMY_DATABASE_URL) if DATABASE_ENABLE_IAM_TOKEN_AUTH else None 

206 

207 

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() 

211 

212 

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 

216 

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 

231 

232 if not event.contains(engine, 'do_connect', _set_iam_token_password): 

233 event.listen(engine, 'do_connect', _set_iam_token_password) 

234 

235 

236def _make_async_url(url: str) -> str: 

237 """Convert a sync database URL to its async driver equivalent. 

238 

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 

260 

261 

262def _json_codec_kwargs(kwargs: dict) -> dict: 

263 """Default an engine to JSONCodec for native ``JSON`` columns. 

264 

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 

271 

272 

273def _create_engine(*args, **kwargs): 

274 """``create_engine`` with the app JSON codec wired in.""" 

275 return create_engine(*args, **_json_codec_kwargs(kwargs)) 

276 

277 

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)) 

281 

282 

283# ============================================================ 

284# SYNC ENGINE (used only for: startup migrations, config loading, 

285# Alembic, peewee migration, health checks) 

286# ============================================================ 

287 

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') 

293 

294 # Extract database path from SQLCipher URL 

295 db_path = SQLALCHEMY_DATABASE_URL.replace('sqlite+sqlcipher://', '') 

296 

297 # Create a custom creator function that uses sqlcipher3 

298 def create_sqlcipher_connection(): 

299 import sqlcipher3 

300 

301 conn = sqlcipher3.connect(db_path, check_same_thread=False) 

302 conn.execute(f"PRAGMA key = '{database_password}'") 

303 return conn 

304 

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 ) 

329 

330 log.info('Connected to encrypted SQLite database using SQLCipher') 

331 

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}) 

334 

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 = {} 

339 

340 def like(pattern, value, escape=None): 

341 if pattern is None or value is None: 

342 return None 

343 

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 

371 

372 return compiled.fullmatch(str(value).lower()) is not None 

373 

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') 

381 

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() 

396 

397 def on_connect(dbapi_connection, connection_record): 

398 _apply_sqlite_pragmas(dbapi_connection) 

399 

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) 

417 

418enable_iam_token_auth(engine) 

419 

420 

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) 

426 

427 

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() 

435 

436 

437get_db = contextmanager(get_session) 

438 

439 

440# ============================================================ 

441# ASYNC ENGINE (used for ALL runtime database operations) 

442# ============================================================ 

443 

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) 

448 

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 

455 

456 asyncio.set_event_loop_policy(asyncio.WindowsSelectorEventLoopPolicy()) 

457 

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 ) 

470 

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 ) 

496 

497enable_iam_token_auth(async_engine) 

498 

499 

500AsyncSessionLocal = async_sessionmaker( 

501 bind=async_engine, 

502 class_=AsyncSession, 

503 autocommit=False, 

504 autoflush=False, 

505 expire_on_commit=False, 

506) 

507 

508 

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() 

516 

517 

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() 

526 

527 

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