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