import os from contextlib import contextmanager from queue import Empty, Full, Queue from threading import Lock from typing import Iterator, Optional from urllib.parse import parse_qs, unquote, urlparse import pymysql from pymysql.connections import Connection from dotenv import load_dotenv load_dotenv() _pool: Optional["MySQLConnectionPool"] = None _pool_lock = Lock() def _get_env(name: str, default: Optional[str] = None) -> Optional[str]: value = os.getenv(name) if value is None: return default trimmed = value.strip() return trimmed or default def _str_to_bool(value: str) -> bool: return value.lower() in {"1", "true", "yes", "on"} def _parse_database_url(database_url: str) -> dict: parsed = urlparse(database_url) scheme = parsed.scheme.lower() if scheme not in {"mysql", "mysql+pymysql"}: raise RuntimeError( "DATABASE_URL must use the mysql scheme when PyMySQL is in use" ) database = parsed.path.lstrip("/") if not database: raise RuntimeError("DATABASE_URL must include the database name") connect_kwargs = { "host": parsed.hostname or "localhost", "port": parsed.port or 3306, "user": unquote(parsed.username or ""), "password": unquote(parsed.password or ""), "database": database, "charset": _get_env("DB_CHARSET", "utf8mb4"), "autocommit": True, } connect_timeout = _get_env("DB_CONNECT_TIMEOUT") if connect_timeout: connect_kwargs["connect_timeout"] = float(connect_timeout) ssl_root_cert = _get_env("SSL_ROOT_CERT") if ssl_root_cert: connect_kwargs["ssl"] = {"ca": ssl_root_cert} query_params = parse_qs(parsed.query, keep_blank_values=True) for key, values in query_params.items(): if not values: continue value = values[-1] if key == "autocommit": connect_kwargs["autocommit"] = _str_to_bool(value) elif key == "charset": connect_kwargs["charset"] = value elif key == "connect_timeout": connect_kwargs["connect_timeout"] = float(value) else: connect_kwargs[key] = value return {key: val for key, val in connect_kwargs.items() if val not in {None, ""}} def _build_connect_kwargs_from_env() -> Optional[dict]: host = _get_env("DB_HOST") database = _get_env("DB_NAME") or _get_env("DB_DATABASE") if not host or not database: return None connect_kwargs = { "host": host, "port": int(_get_env("DB_PORT", "3306")), "user": _get_env("DB_USER", ""), "password": _get_env("DB_PASSWORD", ""), "database": database, "charset": _get_env("DB_CHARSET", "utf8mb4"), "autocommit": _str_to_bool(_get_env("DB_AUTOCOMMIT", "true")), } connect_timeout = _get_env("DB_CONNECT_TIMEOUT") if connect_timeout: connect_kwargs["connect_timeout"] = float(connect_timeout) ssl_root_cert = _get_env("SSL_ROOT_CERT") if ssl_root_cert: connect_kwargs["ssl"] = {"ca": ssl_root_cert} return {key: val for key, val in connect_kwargs.items() if val not in {None, ""}} class MySQLConnectionPool: """Simple thread-safe connection pool for PyMySQL.""" def __init__(self, connect_kwargs: dict, min_size: int, max_size: int) -> None: if min_size < 0: raise ValueError("min_size must be non-negative") if max_size < 1: raise ValueError("max_size must be at least 1") if min_size > max_size: raise ValueError("min_size cannot exceed max_size") self._connect_kwargs = connect_kwargs self._available: Queue[Connection] = Queue(maxsize=max_size) self._lock = Lock() self._max_size = max_size self._total_created = 0 self._closed = False for _ in range(min_size): conn = self._create_connection() self._available.put(conn) self._total_created += 1 def _create_connection(self) -> Connection: if self._closed: raise RuntimeError("Connection pool is closed") return pymysql.connect(**self._connect_kwargs) def connection(self): return _PooledConnectionContext(self) def acquire(self) -> Connection: while True: try: conn = self._available.get_nowait() except Empty: with self._lock: if self._total_created < self._max_size: conn = self._create_connection() self._total_created += 1 return conn conn = self._available.get() if conn.open: conn.ping(reconnect=True) conn.autocommit(True) return conn self._discard_connection(conn) def release(self, conn: Connection) -> None: if self._closed: self._discard_connection(conn) return if not conn.open: self._discard_connection(conn) return try: self._available.put_nowait(conn) except Full: self._discard_connection(conn) def close(self) -> None: self._closed = True while True: try: conn = self._available.get_nowait() except Empty: break self._discard_connection(conn) def _discard_connection(self, conn: Connection) -> None: try: conn.close() finally: with self._lock: if self._total_created > 0: self._total_created -= 1 class _PooledConnectionContext: def __init__(self, pool: MySQLConnectionPool) -> None: self._pool = pool self._conn: Optional[Connection] = None def __enter__(self) -> Connection: self._conn = self._pool.acquire() return self._conn def __exit__(self, exc_type, exc, tb) -> None: if self._conn is not None: self._pool.release(self._conn) self._conn = None def _ensure_pool() -> MySQLConnectionPool: global _pool if _pool is not None: return _pool with _pool_lock: if _pool is not None: return _pool database_url = _get_env("DATABASE_URL") if database_url: connect_kwargs = _parse_database_url(database_url) else: connect_kwargs = _build_connect_kwargs_from_env() if connect_kwargs is None: raise RuntimeError( "DATABASE_URL or DB_HOST and DB_NAME must be set in the environment" ) min_size = int(_get_env("DB_POOL_MIN_SIZE", "1")) max_size = int(_get_env("DB_POOL_MAX_SIZE", "5")) pool = MySQLConnectionPool( connect_kwargs=connect_kwargs, min_size=min_size, max_size=max_size, ) _pool = pool return pool @contextmanager def get_connection() -> Iterator[Connection]: pool = _ensure_pool() with pool.connection() as conn: yield conn def close_pool() -> None: global _pool with _pool_lock: if _pool is not None: _pool.close() _pool = None