Spaces:
Sleeping
Sleeping
| 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 | |
| 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 | |