Source code for sqlspec.adapters.pymysql.pool

"""PyMySQL database configuration with thread-local connections."""

import contextlib
import logging
import threading
import time
from contextlib import contextmanager
from typing import TYPE_CHECKING, Any, cast

from sqlspec.adapters.pymysql._typing import PyMysqlConnect, PyMysqlConnection, PyMysqlServerStatus
from sqlspec.utils.logging import POOL_LOGGER_NAME, get_logger, log_with_context
from sqlspec.utils.uuids import uuid4

if TYPE_CHECKING:
    from collections.abc import Callable, Generator

__all__ = ("PyMysqlConnectionPool",)


logger = get_logger(POOL_LOGGER_NAME)
_ADAPTER_NAME = "pymysql"
_pymysql_connect: "PyMysqlConnect" = cast("PyMysqlConnect", PyMysqlConnect)


[docs] class PyMysqlConnectionPool: """Thread-local connection manager for PyMySQL.""" __slots__ = ( "_connection_factory", "_connection_parameters", "_connection_registry", "_generation", "_health_check_interval", "_on_connection_create", "_pool_id", "_recycle_seconds", "_registry_lock", "_thread_local", )
[docs] def __init__( self, connection_parameters: "dict[str, Any]", recycle_seconds: int = 86400, health_check_interval: float = 30.0, on_connection_create: "Callable[[PyMysqlConnection], None] | None" = None, connection_factory: "Callable[[], PyMysqlConnection] | None" = None, ) -> None: """Initialize the thread-local connection manager. Args: connection_parameters: PyMySQL connection parameters recycle_seconds: Connection recycle time in seconds (default 24h) health_check_interval: Seconds of idle time before running health check on_connection_create: Callback executed when connection is created connection_factory: Optional factory for custom connection creation """ self._connection_parameters = connection_parameters self._connection_factory = connection_factory self._thread_local = threading.local() self._connection_registry: set[PyMysqlConnection] = set() self._generation = 0 self._registry_lock = threading.Lock() self._recycle_seconds = recycle_seconds self._health_check_interval = health_check_interval self._on_connection_create = on_connection_create self._pool_id = str(uuid4())[:8]
@property def _database_name(self) -> str: """Get sanitized database name for logging.""" return str(self._connection_parameters.get("database", "unknown")) def _create_connection(self) -> PyMysqlConnection: connection = self.new_connection() with self._registry_lock: self._connection_registry.add(connection) return connection
[docs] def new_connection(self) -> PyMysqlConnection: """Open a standalone connection configured like a pooled one. The result is owned by the caller: it is not thread-local and is not tracked for pool shutdown. Returns: PyMysqlConnection: A newly opened, fully configured connection. """ if self._connection_factory is not None: connection = self._connection_factory() else: connection = _pymysql_connect(**self._connection_parameters) if self._on_connection_create is not None: self._on_connection_create(connection) return connection
def _is_connection_alive(self, connection: PyMysqlConnection) -> bool: try: connection.ping(reconnect=False) except Exception: return False return True def _get_thread_connection(self) -> PyMysqlConnection: thread_state = self._thread_local.__dict__ if thread_state.get("generation") != self._generation: stale = thread_state.pop("connection", None) if stale is not None: self._retire_connection(cast("PyMysqlConnection", stale)) thread_state.pop("created_at", None) thread_state.pop("last_used", None) self._thread_local.generation = self._generation if "connection" not in thread_state: self._thread_local.connection = self._create_connection() self._thread_local.created_at = time.time() self._thread_local.last_used = time.time() return cast("PyMysqlConnection", self._thread_local.connection) if self._recycle_seconds > 0 and time.time() - self._thread_local.created_at > self._recycle_seconds: log_with_context( logger, logging.DEBUG, "pool.connection.recycle", adapter=_ADAPTER_NAME, pool_id=self._pool_id, database=self._database_name, recycle_seconds=self._recycle_seconds, reason="exceeded_recycle_time", ) self._retire_connection(cast("PyMysqlConnection", self._thread_local.connection)) self._thread_local.connection = self._create_connection() self._thread_local.created_at = time.time() self._thread_local.last_used = time.time() return cast("PyMysqlConnection", self._thread_local.connection) idle_time = time.time() - thread_state.get("last_used", 0) if idle_time > self._health_check_interval and not self._is_connection_alive(self._thread_local.connection): log_with_context( logger, logging.DEBUG, "pool.connection.recycle", adapter=_ADAPTER_NAME, pool_id=self._pool_id, database=self._database_name, idle_seconds=round(idle_time, 1), reason="failed_health_check", ) self._retire_connection(cast("PyMysqlConnection", self._thread_local.connection)) self._thread_local.connection = self._create_connection() self._thread_local.created_at = time.time() self._thread_local.last_used = time.time() return cast("PyMysqlConnection", self._thread_local.connection) def _retire_connection(self, connection: PyMysqlConnection) -> None: """Close a pool-owned connection and drop it from the shutdown registry.""" with self._registry_lock: self._connection_registry.discard(connection) with contextlib.suppress(Exception): connection.close() def _close_thread_connection(self) -> None: thread_state = self._thread_local.__dict__ if "connection" in thread_state: self._retire_connection(cast("PyMysqlConnection", self._thread_local.connection)) del self._thread_local.connection if "created_at" in thread_state: del self._thread_local.created_at if "last_used" in thread_state: del self._thread_local.last_used
[docs] @contextmanager def get_connection(self) -> "Generator[PyMysqlConnection, None, None]": """Get a thread-local connection. Yields: A thread-local database connection. """ connection = self._get_thread_connection() try: yield connection except Exception: with contextlib.suppress(Exception): self._close_thread_connection() raise else: self.release(connection)
[docs] def close(self) -> None: """Close every connection this pool opened, on any thread.""" self._close_thread_connection() with self._registry_lock: orphaned = list(self._connection_registry) self._connection_registry.clear() self._generation += 1 for connection in orphaned: with contextlib.suppress(Exception): connection.close()
def acquire(self) -> PyMysqlConnection: return self._get_thread_connection()
[docs] def release(self, connection: PyMysqlConnection) -> None: """Release connection back to the pool, sanitizing transactions.""" if bool(getattr(connection, "server_status", 0) & PyMysqlServerStatus.SERVER_STATUS_IN_TRANS): try: connection.rollback() except Exception: if getattr(self._thread_local, "connection", None) is connection: self._close_thread_connection() else: self._retire_connection(connection) raise
[docs] def size(self) -> int: """Report total active connections managed by this pool.""" with self._registry_lock: return len(self._connection_registry)
def checked_out(self) -> int: return 0