From dc605ceba0c752be68da607c3816b71c5ca032a3 Mon Sep 17 00:00:00 2001 From: Pavel Mosein Date: Thu, 16 Apr 2026 09:25:41 +0300 Subject: [PATCH 1/3] refactor(core): split base.py into pool_state, pool_manager, health monitor MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - PoolState (hasql/pool_state.py): manages master/replica sets, waiting, pool factory delegation — replaces the stub from PR2 - BasePoolManager (hasql/pool_manager.py): thin orchestrator; all pool-state queries exposed as public proxy properties/methods - PoolHealthMonitor (hasql/health.py): background health checks — takes pool_state + scalars instead of full manager reference (no transitive private access) - hasql/base.py: reduced to 23-line re-export shim for backward compat - hasql/utils.py: genericize Stopwatch[KeyT], fix Dsn.with_() scheme preservation, defensive copy for Dsn.params Stack 3/4 Co-Authored-By: Claude Opus 4.6 (1M context) --- hasql/base.py | 725 +------------------------------- hasql/health.py | 176 ++++++++ hasql/pool_manager.py | 381 +++++++++++++++++ hasql/pool_state.py | 328 ++++++++++++++- hasql/utils.py | 54 +-- tests/conftest.py | 10 +- tests/mocks/pool_manager.py | 48 ++- tests/test_backward_compat.py | 66 +++ tests/test_balancer_policy.py | 242 +++++++++++ tests/test_base_pool_manager.py | 392 ++++++++++++----- tests/test_timeout_handling.py | 83 ++-- tests/test_trouble.py | 33 +- 12 files changed, 1620 insertions(+), 918 deletions(-) create mode 100644 hasql/health.py create mode 100644 hasql/pool_manager.py create mode 100644 tests/test_backward_compat.py diff --git a/hasql/base.py b/hasql/base.py index 0e96559..d2d9b92 100644 --- a/hasql/base.py +++ b/hasql/base.py @@ -1,709 +1,30 @@ -import asyncio -import logging -from abc import ABC, abstractmethod -from collections import defaultdict -from itertools import chain -from types import MappingProxyType -from typing import ( - Any, - AsyncContextManager, - DefaultDict, - Dict, - List, - Optional, - Sequence, - Set, - Union, +from .abc import PoolDriver +from .acquire import AcquireContext, PoolAcquireContext, TimeoutAcquireContext +from .balancer_policy import AbstractBalancerPolicy +from .exceptions import ( + HasqlError, + NoAvailablePoolError, + PoolManagerClosedError, + PoolManagerClosingError, + UnexpectedDatabaseResponseError, ) - -from .metrics import CalculateMetrics, DriverMetrics, Metrics -from .utils import Dsn, Stopwatch, split_dsn - -logger = logging.getLogger(__name__) - - -class TimeoutAcquireContext: - __slots__ = ("_context", "_timeout") - - def __init__(self, context, timeout: float): - self._context = context - self._timeout = timeout - - async def __aenter__(self): - return await asyncio.wait_for( - self._context.__aenter__(), - timeout=self._timeout, - ) - - async def __aexit__(self, *exc): - # TODO: consider adding a bounded timeout here. Currently if the - # underlying driver hangs during connection release this will block - # indefinitely. A timeout risks leaking the connection (not returned - # to pool), so this needs careful design. - await self._context.__aexit__(*exc) - - def __await__(self): - return asyncio.wait_for( - self._context.__aenter__(), - timeout=self._timeout, - ).__await__() - - -DEFAULT_REFRESH_DELAY: int = 1 -DEFAULT_REFRESH_TIMEOUT: int = 30 -DEFAULT_ACQUIRE_TIMEOUT: float = 1.0 -DEFAULT_MASTER_AS_REPLICA_WEIGHT: float = 0.0 -DEFAULT_STOPWATCH_WINDOW_SIZE: int = 128 - - -class AbstractBalancerPolicy(ABC): - def __init__(self, pool_manager: "BasePoolManager"): - raise NotImplementedError - - @abstractmethod - async def get_pool( - self, - read_only: bool, - fallback_master: bool = False, - master_as_replica_weight: Optional[float] = None, - ) -> Any: - raise NotImplementedError - - -class PoolAcquireContext(AsyncContextManager): - def __init__( - self, - pool_manager: "BasePoolManager", - read_only: bool, - master_as_replica_weight: Optional[float], - timeout: float, - metrics: CalculateMetrics, - fallback_master: bool = False, - **kwargs, - ): - self.pool_manager = pool_manager - self.read_only = read_only - self.fallback_master = fallback_master - self.master_as_replica_weight = master_as_replica_weight - self.timeout = timeout - self.kwargs = kwargs - self.pool: Any = None - self.context: Any = None - self.metrics = metrics - - def _deadline(self) -> float: - return asyncio.get_running_loop().time() + self.timeout - - def _remaining_timeout(self, deadline: float) -> float: - remaining_timeout = deadline - asyncio.get_running_loop().time() - if remaining_timeout <= 0: - raise asyncio.TimeoutError - return remaining_timeout - - # TODO: make _get_pool a clean function (return pool without mutating - # self.pool) and extract the shared preamble between __aenter__ / - # acquire_from_pool_connection into a helper. - async def _get_pool(self, deadline: float): - async def get_pool() -> Any: - with self.metrics.with_get_pool(): - return await self.pool_manager.balancer.get_pool( - read_only=self.read_only, - fallback_master=self.fallback_master, - master_as_replica_weight=self.master_as_replica_weight, - ) - - self.pool = await asyncio.wait_for( - get_pool(), - timeout=self._remaining_timeout(deadline), - ) - return self.pool - - def _acquire_kwargs(self, deadline: float) -> dict: - return self.pool_manager._prepare_acquire_kwargs( - self.kwargs, - timeout=self._remaining_timeout(deadline), - ) - - async def _resolve_pool_and_kwargs(self): - deadline = self._deadline() - await self._get_pool(deadline) - return self._acquire_kwargs(deadline) - - async def acquire_from_pool_connection(self): - acquire_kwargs = await self._resolve_pool_and_kwargs() - - with self.metrics.with_acquire(self.pool_manager.host(self.pool)): - self.conn = await self.pool_manager.acquire_from_pool( - self.pool, - **acquire_kwargs, - ) - - self.metrics.add_connection(self.pool_manager.host(self.pool)) - self.pool_manager.register_connection(self.conn, self.pool) - return self.conn - - async def __aenter__(self): - acquire_kwargs = await self._resolve_pool_and_kwargs() - - with self.metrics.with_acquire(self.pool_manager.host(self.pool)): - self.context = self.pool_manager.acquire_from_pool( - self.pool, - **acquire_kwargs, - ) - self.conn = await self.context.__aenter__() - - self.metrics.add_connection(self.pool_manager.host(self.pool)) - return self.conn - - async def __aexit__(self, *exc): - self.metrics.remove_connection(self.pool_manager.host(self.pool)) - await self.context.__aexit__(*exc) - del self.conn - - def __await__(self): - return self.acquire_from_pool_connection().__await__() - - -class BasePoolManager(ABC): - _dsn_ready_event: DefaultDict[Dsn, asyncio.Event] - _dsn_check_cond: DefaultDict[Dsn, asyncio.Condition] - _master_pool_set: Set[Any] - _replica_pool_set: Set[Any] - _unmanaged_connections: Dict[Any, Any] - - def __init__( - self, - dsn: str, - acquire_timeout: Union[float, int] = DEFAULT_ACQUIRE_TIMEOUT, - refresh_delay: Union[float, int] = DEFAULT_REFRESH_DELAY, - refresh_timeout: Union[float, int] = DEFAULT_REFRESH_TIMEOUT, - fallback_master: bool = False, - master_as_replica_weight: float = DEFAULT_MASTER_AS_REPLICA_WEIGHT, - balancer_policy: type = AbstractBalancerPolicy, - stopwatch_window_size: int = DEFAULT_STOPWATCH_WINDOW_SIZE, - pool_factory_kwargs: Optional[dict] = None, - ): - if not issubclass(balancer_policy, AbstractBalancerPolicy): - raise ValueError( - "balancer_policy must be a class BaseBalancerPolicy heir", - ) - - if balancer_policy is AbstractBalancerPolicy: - # Avoid circular import - from .balancer_policy.greedy import GreedyBalancerPolicy - - balancer_policy = GreedyBalancerPolicy - - if pool_factory_kwargs is None: - pool_factory_kwargs = {} - self._pool_factory_kwargs = MappingProxyType( - self._prepare_pool_factory_kwargs(pool_factory_kwargs), - ) - self._dsn: List[Dsn] = split_dsn(dsn) - self._dsn_ready_event = defaultdict(asyncio.Event) - self._dsn_check_cond = defaultdict(asyncio.Condition) - self._pools = [None] * len(self._dsn) - self._acquire_timeout = acquire_timeout - self._refresh_delay = refresh_delay - self._refresh_timeout = refresh_timeout - self._fallback_master = fallback_master - self._master_as_replica_weight = master_as_replica_weight - self._balancer = balancer_policy(self) - self._master_pool_set = set() - self._replica_pool_set = set() - self._master_cond = asyncio.Condition() - self._replica_cond = asyncio.Condition() - self._unmanaged_connections = {} - self._stopwatch = Stopwatch(window_size=stopwatch_window_size) - self._refresh_role_tasks = [ - asyncio.create_task(self._check_pool_task(index)) - for index in range(len(self._dsn)) - ] - self._closing = False - self._closed = False - self._metrics = CalculateMetrics() - - @property - def dsn(self) -> List[Dsn]: - return self._dsn - - @property - def refresh_delay(self): - return self._refresh_delay - - @property - def refresh_timeout(self): - return self._refresh_timeout - - @property - def pool_factory_kwargs(self): - return self._pool_factory_kwargs - - @property - def master_pool_count(self): - return len(self._master_pool_set) - - @property - def replica_pool_count(self): - return len(self._replica_pool_set) - - @property - def available_pool_count(self): - return self.master_pool_count + self.replica_pool_count - - @property - def balancer(self) -> AbstractBalancerPolicy: - return self._balancer - - @property - def closing(self) -> bool: - return self._closing - - @property - def closed(self) -> bool: - return self._closed - - @property - def pools(self) -> Sequence[Any]: - return tuple(self._pools) - - @abstractmethod - def get_pool_freesize(self, pool): - pass - - @abstractmethod - def acquire_from_pool(self, pool, **kwargs): - pass - - @abstractmethod - async def release_to_pool(self, connection, pool, **kwargs): - pass - - @abstractmethod - async def _is_master(self, connection): - pass - - @abstractmethod - async def _pool_factory(self, dsn: Dsn): - pass - - @abstractmethod - async def _close(self, pool): - pass - - @abstractmethod - async def _terminate(self, pool) -> None: - pass - - @abstractmethod - def is_connection_closed(self, connection): - pass - - @abstractmethod - def host(self, pool: Any): - pass - - @abstractmethod - def _driver_metrics(self) -> Sequence[DriverMetrics]: - pass - - def metrics(self) -> Metrics: - return Metrics( - drivers=self._driver_metrics(), - hasql=self._metrics.metrics(), - ) - - def _prepare_acquire_kwargs( - self, - kwargs: dict, - timeout: float, - ) -> dict: - return dict(kwargs) - - def acquire( - self, - read_only: bool = False, - fallback_master: Optional[bool] = None, - master_as_replica_weight: Optional[float] = None, - timeout: Optional[float] = None, - **kwargs, - ): - if fallback_master is None: - fallback_master = self._fallback_master - - if not read_only and master_as_replica_weight is not None: - raise ValueError( - "Field master_as_replica_weight is used only when " - "read_only is True", - ) - if master_as_replica_weight is not None and not ( - 0.0 <= master_as_replica_weight <= 1 - ): - raise ValueError( - "Field master_as_replica_weight must belong " - "to the segment [0; 1]", - ) - - if read_only: - if master_as_replica_weight is None: - master_as_replica_weight = self._master_as_replica_weight - - if timeout is None: - timeout = self._acquire_timeout - - ctx = PoolAcquireContext( - pool_manager=self, - read_only=read_only, - fallback_master=fallback_master, - master_as_replica_weight=master_as_replica_weight, - timeout=timeout, - metrics=self._metrics, - **kwargs, - ) - - return ctx - - def acquire_master( - self, - timeout: Optional[float] = None, - **kwargs, - ): - return self.acquire(read_only=False, timeout=timeout, **kwargs) - - def acquire_replica( - self, - fallback_master: Optional[bool] = None, - master_as_replica_weight: Optional[float] = None, - timeout: Optional[float] = None, - **kwargs, - ): - return self.acquire( - read_only=True, - fallback_master=fallback_master, - master_as_replica_weight=master_as_replica_weight, - timeout=timeout, - **kwargs, - ) - - async def release(self, connection, **kwargs): - if connection not in self._unmanaged_connections: - raise ValueError( - "Pool.release() received invalid connection: " - f"{connection!r} is not a member of this pool", - ) - - pool = self._unmanaged_connections.pop(connection) - self._metrics.remove_connection(self.host(pool)) - await self.release_to_pool(connection, pool, **kwargs) - - async def close(self): - self._closing = True - await self._clear() - await asyncio.gather( - *[self._close(pool) for pool in self._pools if pool is not None], - return_exceptions=True, - ) - self._closing = False - self._closed = True - - async def terminate(self): - self._closing = True - await self._clear() - for pool in self._pools: - if pool is None: - continue - await self._terminate(pool) - self._closing = False - self._closed = True - - async def wait_next_pool_check(self, timeout: int = 10): - tasks = [self._wait_checking_pool(dsn) for dsn in self._dsn] - await asyncio.wait_for(asyncio.gather(*tasks), timeout=timeout) - - async def _wait_checking_pool(self, dsn: Dsn): - async with self._dsn_check_cond[dsn]: - for _ in range(2): - await self._dsn_check_cond[dsn].wait() - - async def ready( - self, - masters_count: Optional[int] = None, - replicas_count: Optional[int] = None, - timeout: int = 10, - ): - - if (masters_count is not None and replicas_count is None) or ( - masters_count is None and replicas_count is not None - ): - raise ValueError( - "Arguments master_count and replicas_count " - "should both be either None or not None", - ) - - if masters_count is not None and masters_count < 0: - raise ValueError("masters_count shouldn't be negative") - if replicas_count is not None and replicas_count < 0: - raise ValueError("replicas_count shouldn't be negative") - - if masters_count is None and replicas_count is None: - await asyncio.wait_for(self.wait_all_ready(), timeout=timeout) - return - - assert isinstance(masters_count, int) - assert isinstance(replicas_count, int) - - await asyncio.wait_for( - asyncio.gather( - self.wait_masters_ready(masters_count), - self.wait_replicas_ready(replicas_count), - ), - timeout=timeout, - ) - - async def wait_all_ready(self): - for dsn in self._dsn: - await self._dsn_ready_event[dsn].wait() - - async def wait_masters_ready(self, masters_count: int): - def predicate(): - return self.master_pool_count >= masters_count - - async with self._master_cond: - await self._master_cond.wait_for(predicate) - - async def wait_replicas_ready(self, replicas_count: int): - def predicate(): - return self.replica_pool_count >= replicas_count - - async with self._replica_cond: - await self._replica_cond.wait_for(predicate) - - async def get_master_pools(self) -> List: - if not self._master_pool_set: - async with self._master_cond: - await self._master_cond.wait() - return list(self._master_pool_set) - - async def get_replica_pools(self, fallback_master: bool = False) -> List: - if not self._replica_pool_set: - if fallback_master: - return await self.get_master_pools() - async with self._replica_cond: - await self._replica_cond.wait() - return list(self._replica_pool_set) - - def pool_is_master(self, pool) -> bool: - return pool in self._master_pool_set - - def pool_is_replica(self, pool) -> bool: - return pool in self._replica_pool_set - - def register_connection(self, connection, pool): - self._unmanaged_connections[connection] = pool - - def get_last_response_time(self, pool) -> Optional[float]: - return self._stopwatch.get_time(pool) - - def _prepare_pool_factory_kwargs(self, kwargs: dict) -> dict: - return kwargs - - async def _clear(self): - self._balancer = None - if self._refresh_role_tasks is not None: - for refresh_role_task in self._refresh_role_tasks: - refresh_role_task.cancel() - - await asyncio.gather( - *self._refresh_role_tasks, - return_exceptions=True, - ) - - self._refresh_role_tasks = None - - release_tasks = [] - for connection in self._unmanaged_connections: - release_tasks.append(self.release(connection)) - - await asyncio.gather(*release_tasks, return_exceptions=True) - - self._unmanaged_connections.clear() - self._master_pool_set.clear() - self._replica_pool_set.clear() - - async def _check_pool_task(self, index: int): - logger.debug("Starting pool task") - dsn = self._dsn[index] - censored_dsn = str(dsn.with_(password="******")) - pool = await self._wait_creating_pool(dsn) - self._pools[index] = pool - - logger.debug("Setting dsn=%r event", censored_dsn) - sys_connection = None - while not self._closing: - try: - # Не использовать async with self.acquire_from_pool(pool) - # из-за большого таймаута - logger.debug( - "Acquiring connection for checking dsn=%r", - censored_dsn, - ) - sys_connection = await asyncio.wait_for( - self.acquire_from_pool(pool), - timeout=self._refresh_timeout, - ) - - logger.debug("Checking dsn=%r", censored_dsn) - await self._periodic_pool_check(pool, dsn, sys_connection) - except asyncio.TimeoutError: - logger.warning( - "Creating system connection failed for dsn=%r", - censored_dsn, - ) - self._remove_pool_from_master_set(pool, dsn) - self._remove_pool_from_replica_set(pool, dsn) - except asyncio.CancelledError as cancelled_error: - if self._closing: - raise cancelled_error from None - logger.warning( - "Cancelled error for dsn=%r", - censored_dsn, - exc_info=True, - ) - self._remove_pool_from_master_set(pool, dsn) - self._remove_pool_from_replica_set(pool, dsn) - except Exception: - logger.warning( - "Database is not available with exception for dsn=%r", - censored_dsn, - exc_info=True, - ) - self._remove_pool_from_master_set(pool, dsn) - self._remove_pool_from_replica_set(pool, dsn) - finally: - if sys_connection is not None: - try: - await self.release_to_pool(sys_connection, pool) - except asyncio.CancelledError as cancelled_error: - if self._closing: - raise cancelled_error from None - logger.warning( - "Release connection to pool with " - "Cancelled error for dsn=%r", - censored_dsn, - exc_info=True, - ) - except Exception: - logger.warning( - "Release connection to pool with " - "exception for dsn=%r", - censored_dsn, - exc_info=True, - ) - sys_connection = None - await self._notify_about_pool_has_checked(dsn) - - await asyncio.sleep(self._refresh_delay) - - async def _wait_creating_pool(self, dsn: Dsn): - while not self._closing: - try: - return await asyncio.wait_for( - self._pool_factory(dsn), - timeout=self._refresh_timeout, - ) - except Exception: - logger.warning( - "Creating pool failed with exception for dsn=%s", - dsn.with_(password="******"), - exc_info=True, - ) - await asyncio.sleep(self._refresh_delay) - - async def _periodic_pool_check(self, pool, dsn: Dsn, sys_connection): - while not self._closing: - try: - await asyncio.wait_for( - self._refresh_pool_role(pool, dsn, sys_connection), - timeout=self._refresh_timeout, - ) - await self._notify_about_pool_has_checked(dsn) - except asyncio.TimeoutError: - logger.warning( - "Periodic pool check failed for dsn=%s", - dsn.with_(password="******"), - ) - self._remove_pool_from_master_set(pool, dsn) - self._remove_pool_from_replica_set(pool, dsn) - await self._notify_about_pool_has_checked(dsn) - - await asyncio.sleep(self._refresh_delay) - - async def _notify_about_pool_has_checked(self, dsn: Dsn): - async with self._dsn_check_cond[dsn]: - self._dsn_check_cond[dsn].notify_all() - - async def _add_pool_to_master_set(self, pool, dsn: Dsn): - if pool in self._master_pool_set: - return - self._master_pool_set.add(pool) - logger.debug( - "Pool %s has been added to master set", - dsn.with_(password="******"), - ) - async with self._master_cond: - self._master_cond.notify_all() - - async def _add_pool_to_replica_set(self, pool, dsn: Dsn): - if pool in self._replica_pool_set: - return - self._replica_pool_set.add(pool) - logger.debug( - "Pool %s has been added to replica set", - dsn.with_(password="******"), - ) - async with self._replica_cond: - self._replica_cond.notify_all() - - def _remove_pool_from_master_set(self, pool, dsn: Dsn): - if pool in self._master_pool_set: - self._master_pool_set.remove(pool) - logger.debug( - "Pool %s has been removed from master set", - dsn.with_(password="******"), - ) - - def _remove_pool_from_replica_set(self, pool, dsn: Dsn): - if pool in self._replica_pool_set: - self._replica_pool_set.remove(pool) - logger.debug( - "Pool %s has been removed from replica set", - dsn.with_(password="******"), - ) - - async def _refresh_pool_role(self, pool, dsn: Dsn, sys_connection): - with self._stopwatch(pool): - is_master = await self._is_master(sys_connection) - if is_master: - await self._add_pool_to_master_set(pool, dsn) - self._remove_pool_from_replica_set(pool, dsn) - else: - await self._add_pool_to_replica_set(pool, dsn) - self._remove_pool_from_master_set(pool, dsn) - self._dsn_ready_event[dsn].set() - - def __iter__(self): - return chain(iter(self._master_pool_set), iter(self._replica_pool_set)) - - async def __aenter__(self): - await self.ready() - return self - - async def __aexit__(self, exc_type, exc_val, exc_tb): - await self.close() - +from .pool_manager import BasePoolManager, ConnT, PoolT +from .pool_state import PoolState, PoolStateProvider __all__ = ( + "AcquireContext", "BasePoolManager", "AbstractBalancerPolicy", + "HasqlError", + "NoAvailablePoolError", + "PoolDriver", + "PoolManagerClosedError", + "PoolManagerClosingError", + "PoolState", + "PoolStateProvider", "TimeoutAcquireContext", + "UnexpectedDatabaseResponseError", + "PoolAcquireContext", + "PoolT", + "ConnT", ) diff --git a/hasql/health.py b/hasql/health.py new file mode 100644 index 0000000..ec00e9b --- /dev/null +++ b/hasql/health.py @@ -0,0 +1,176 @@ +import asyncio +import logging +from collections.abc import Callable +from typing import Generic, TypeVar + +from .exceptions import PoolManagerClosingError +from .pool_state import PoolState +from .utils import Dsn + +logger = logging.getLogger(__name__) + +PoolT = TypeVar("PoolT") +ConnT = TypeVar("ConnT") + + +class PoolHealthMonitor(Generic[PoolT, ConnT]): + """Background health monitor that checks pool roles periodically.""" + + def __init__( + self, + pool_state: PoolState[PoolT, ConnT], + refresh_delay: float, + refresh_timeout: float, + closing_getter: Callable[[], bool], + ): + self._pool_state = pool_state + self._refresh_delay = refresh_delay + self._refresh_timeout = refresh_timeout + self._closing = closing_getter + self._tasks: list[asyncio.Task] | None = [ + asyncio.create_task(self._check_pool_task(index)) + for index in range(len(pool_state.dsn)) + ] + + @property + def tasks(self) -> list[asyncio.Task] | None: + return self._tasks + + async def stop(self): + if self._tasks is not None: + for task in self._tasks: + task.cancel() + + # Tasks may finish with CancelledError or + # PoolManagerClosingError — both are expected. + await asyncio.gather( + *self._tasks, + return_exceptions=True, + ) + + self._tasks = None + + async def _periodic_pool_check( + self, + pool: PoolT, + dsn: Dsn, + sys_connection: ConnT, + ): + while not self._closing(): + try: + await asyncio.wait_for( + self._pool_state.refresh_pool_role( + pool, dsn, sys_connection, + ), + timeout=self._refresh_timeout, + ) + await self._pool_state.notify_pool_checked(dsn) + except asyncio.TimeoutError: + logger.warning( + "Periodic pool check failed for dsn=%s", + dsn.with_(password="******"), + ) + self._pool_state.remove_pool_from_all_sets(pool, dsn) + await self._pool_state.notify_pool_checked(dsn) + + await asyncio.sleep(self._refresh_delay) + + async def _check_pool_task(self, index: int): + logger.debug("Starting pool task") + pool_state = self._pool_state + dsn = pool_state.dsn[index] + censored_dsn = str(dsn.with_(password="******")) + pool = await self._wait_creating_pool(dsn) + pool_state.set_pool(index, pool) + + logger.debug("Setting dsn=%r event", censored_dsn) + sys_connection: ConnT | None = None + while not self._closing(): + try: + # Don't use async with — we need a custom timeout + logger.debug( + "Acquiring connection for checking dsn=%r", + censored_dsn, + ) + sys_connection = await asyncio.wait_for( + pool_state.acquire_from_pool(pool), + timeout=self._refresh_timeout, + ) + + logger.debug("Checking dsn=%r", censored_dsn) + if sys_connection is None: + continue + await self._periodic_pool_check(pool, dsn, sys_connection) + except asyncio.TimeoutError: + logger.warning( + "Creating system connection failed for dsn=%r", + censored_dsn, + ) + pool_state.remove_pool_from_all_sets(pool, dsn) + except asyncio.CancelledError as cancelled_error: + if self._closing(): + raise cancelled_error from None + logger.warning( + "Cancelled error for dsn=%r", + censored_dsn, + exc_info=True, + ) + pool_state.remove_pool_from_all_sets(pool, dsn) + except Exception: + logger.warning( + "Database is not available with exception for dsn=%r", + censored_dsn, + exc_info=True, + ) + pool_state.remove_pool_from_all_sets(pool, dsn) + finally: + if sys_connection is not None: + await self._safe_release_connection( + sys_connection, pool, censored_dsn, + ) + sys_connection = None + await pool_state.notify_pool_checked(dsn) + + await asyncio.sleep(self._refresh_delay) + + async def _safe_release_connection( + self, connection: ConnT, pool: PoolT, censored_dsn: str, + ): + try: + await self._pool_state.release_to_pool(connection, pool) + except asyncio.CancelledError as cancelled_error: + if self._closing(): + raise cancelled_error from None + logger.warning( + "Release connection to pool with " + "Cancelled error for dsn=%r", + censored_dsn, + exc_info=True, + ) + except Exception: + logger.warning( + "Release connection to pool with " + "exception for dsn=%r", + censored_dsn, + exc_info=True, + ) + + async def _wait_creating_pool(self, dsn: Dsn) -> PoolT: + pool_state = self._pool_state + while not self._closing(): + try: + return await asyncio.wait_for( + pool_state.pool_factory(dsn), + timeout=self._refresh_timeout, + ) + except Exception: + logger.warning( + "Creating pool failed with exception for dsn=%s", + dsn.with_(password="******"), + exc_info=True, + ) + await asyncio.sleep(self._refresh_delay) + raise PoolManagerClosingError("Pool manager is closing") + + +__all__ = ("PoolHealthMonitor",) diff --git a/hasql/pool_manager.py b/hasql/pool_manager.py new file mode 100644 index 0000000..5fce561 --- /dev/null +++ b/hasql/pool_manager.py @@ -0,0 +1,381 @@ +import asyncio +import logging +from collections.abc import Sequence +from typing import ( + Generic, + TypeVar, +) + +from .abc import PoolDriver +from .acquire import PoolAcquireContext +from .balancer_policy.base import AbstractBalancerPolicy +from .balancer_policy.greedy import GreedyBalancerPolicy +from .exceptions import PoolManagerClosedError +from .constants import ( + DEFAULT_ACQUIRE_TIMEOUT, + DEFAULT_MASTER_AS_REPLICA_WEIGHT, + DEFAULT_REFRESH_DELAY, + DEFAULT_REFRESH_TIMEOUT, + DEFAULT_STOPWATCH_WINDOW_SIZE, +) +from .health import PoolHealthMonitor +from .metrics import ( + CalculateMetrics, + HasqlGauges, + Metrics, + PoolMetrics, + PoolRole, +) +from .pool_state import PoolState +from .utils import Dsn, split_dsn + +logger = logging.getLogger(__name__) + +PoolT = TypeVar("PoolT") +ConnT = TypeVar("ConnT") + + +class BasePoolManager(Generic[PoolT, ConnT]): + _unmanaged_connections: dict[ConnT, PoolT] + + def __init__( + self, + dsn: str, + *, + driver: PoolDriver[PoolT, ConnT], + acquire_timeout: float = DEFAULT_ACQUIRE_TIMEOUT, + refresh_delay: float = DEFAULT_REFRESH_DELAY, + refresh_timeout: float = DEFAULT_REFRESH_TIMEOUT, + fallback_master: bool = False, + master_as_replica_weight: float = DEFAULT_MASTER_AS_REPLICA_WEIGHT, + balancer_policy: type[AbstractBalancerPolicy] = GreedyBalancerPolicy, + stopwatch_window_size: int = DEFAULT_STOPWATCH_WINDOW_SIZE, + pool_factory_kwargs: dict | None = None, + ): + if not issubclass(balancer_policy, AbstractBalancerPolicy): + raise ValueError( + "balancer_policy must be a subclass of AbstractBalancerPolicy", + ) + + self._pool_state: PoolState[PoolT, ConnT] = PoolState( + dsn_list=split_dsn(dsn), + driver=driver, + stopwatch_window_size=stopwatch_window_size, + pool_factory_kwargs=pool_factory_kwargs, + ) + + self._balancer: AbstractBalancerPolicy[PoolT] | None = ( + balancer_policy(self._pool_state) + ) + + self._acquire_timeout = acquire_timeout + self._refresh_delay = refresh_delay + self._refresh_timeout = refresh_timeout + self._fallback_master = fallback_master + self._master_as_replica_weight = master_as_replica_weight + self._unmanaged_connections: dict[ConnT, PoolT] = {} + self._metrics = CalculateMetrics() + self._closing = False + self._closed = False + + self._health: PoolHealthMonitor[PoolT, ConnT] = PoolHealthMonitor( + pool_state=self._pool_state, + refresh_delay=refresh_delay, + refresh_timeout=refresh_timeout, + closing_getter=lambda: self._closing, + ) + + # --- Public pool-state proxy properties --- + + @property + def dsn(self) -> Sequence[Dsn]: + return self._pool_state.dsn + + @property + def master_pool_count(self) -> int: + return self._pool_state.master_pool_count + + @property + def replica_pool_count(self) -> int: + return self._pool_state.replica_pool_count + + @property + def available_pool_count(self) -> int: + return self._pool_state.available_pool_count + + @property + def pools(self) -> Sequence[PoolT | None]: + return self._pool_state.pools + + @property + def closing(self) -> bool: + return self._closing + + @property + def closed(self) -> bool: + return self._closed + + @property + def balancer(self) -> AbstractBalancerPolicy[PoolT] | None: + return self._balancer + + @property + def refresh_delay(self) -> float: + return self._refresh_delay + + @property + def refresh_timeout(self) -> float: + return self._refresh_timeout + + # --- Public pool-state proxy methods --- + + def pool_is_master(self, pool: PoolT) -> bool: + return self._pool_state.pool_is_master(pool) + + def pool_is_replica(self, pool: PoolT) -> bool: + return self._pool_state.pool_is_replica(pool) + + def get_pool_freesize(self, pool: PoolT) -> int: + return self._pool_state.get_pool_freesize(pool) + + def get_last_response_time(self, pool: PoolT) -> float | None: + return self._pool_state.get_last_response_time(pool) + + async def get_master_pools(self) -> list[PoolT]: + return await self._pool_state.get_master_pools() + + async def get_replica_pools( + self, fallback_master: bool = False, + ) -> list[PoolT]: + return await self._pool_state.get_replica_pools( + fallback_master=fallback_master, + ) + + async def wait_next_pool_check(self, timeout: int = 10) -> None: + await self._pool_state.wait_next_pool_check(timeout) + + async def wait_all_ready(self) -> None: + await self._pool_state.wait_all_ready() + + async def wait_masters_ready(self, masters_count: int) -> None: + await self._pool_state.wait_masters_ready(masters_count) + + async def wait_replicas_ready(self, replicas_count: int) -> None: + await self._pool_state.wait_replicas_ready(replicas_count) + + async def ready( + self, + masters_count: int | None = None, + replicas_count: int | None = None, + timeout: int = 10, + ) -> None: + await self._pool_state.ready( + masters_count=masters_count, + replicas_count=replicas_count, + timeout=timeout, + ) + + # --- Metrics --- + + def metrics(self) -> Metrics: + pool_state = self._pool_state + pool_metrics = [] + for pool in pool_state.pools: + if pool is None: + continue + stats = pool_state.pool_stats(pool) + + role: PoolRole | None + if pool_state.pool_is_master(pool): + role = PoolRole.MASTER + elif pool_state.pool_is_replica(pool): + role = PoolRole.REPLICA + else: + role = None + + in_flight = sum( + 1 for p in self._unmanaged_connections.values() if p is pool + ) + + pool_metrics.append(PoolMetrics( + host=pool_state.host(pool), + role=role, + healthy=role is not None, + min=stats.min, + max=stats.max, + idle=stats.idle, + used=stats.used, + response_time=pool_state.get_last_response_time(pool), + in_flight=in_flight, + extra=stats.extra, + )) + + gauges = HasqlGauges( + master_count=pool_state.master_pool_count, + replica_count=pool_state.replica_pool_count, + available_count=pool_state.available_pool_count, + active_connections=len(self._unmanaged_connections), + closing=self._closing, + closed=self._closed, + unavailable_count=( + len([p for p in pool_state.pools if p is not None]) + - pool_state.available_pool_count + ), + ) + + return Metrics( + pools=pool_metrics, + hasql=self._metrics.metrics(), + gauges=gauges, + ) + + # --- Acquire API --- + + def acquire( + self, + read_only: bool = False, + fallback_master: bool | None = None, + master_as_replica_weight: float | None = None, + timeout: float | None = None, + **kwargs, + ) -> "PoolAcquireContext[PoolT, ConnT]": + if self._closed or self._closing: + raise PoolManagerClosedError("Pool manager is closed") + + if fallback_master is None: + fallback_master = self._fallback_master + + if not read_only and master_as_replica_weight is not None: + raise ValueError( + "Field master_as_replica_weight is used only when " + "read_only is True", + ) + if master_as_replica_weight is not None and not ( + 0.0 <= master_as_replica_weight <= 1 + ): + raise ValueError( + "Field master_as_replica_weight must belong " + "to the segment [0; 1]", + ) + + if read_only: + if master_as_replica_weight is None: + master_as_replica_weight = self._master_as_replica_weight + + if timeout is None: + timeout = self._acquire_timeout + + if self._balancer is None: + raise PoolManagerClosedError("Pool manager is closed") + + return PoolAcquireContext( + pool_state=self._pool_state, + balancer=self._balancer, + register_connection=self._register_connection, + unregister_connection=self._unregister_connection, + read_only=read_only, + fallback_master=fallback_master, + master_as_replica_weight=master_as_replica_weight, + timeout=timeout, + metrics=self._metrics, + **kwargs, + ) + + def acquire_master( + self, + timeout: float | None = None, + **kwargs, + ) -> "PoolAcquireContext[PoolT, ConnT]": + return self.acquire(read_only=False, timeout=timeout, **kwargs) + + def acquire_replica( + self, + fallback_master: bool | None = None, + master_as_replica_weight: float | None = None, + timeout: float | None = None, + **kwargs, + ) -> "PoolAcquireContext[PoolT, ConnT]": + return self.acquire( + read_only=True, + fallback_master=fallback_master, + master_as_replica_weight=master_as_replica_weight, + timeout=timeout, + **kwargs, + ) + + # --- Connection tracking (internal) --- + + def _register_connection(self, connection: ConnT, pool: PoolT): + self._unmanaged_connections[connection] = pool + + def _unregister_connection(self, connection: ConnT) -> None: + self._unmanaged_connections.pop(connection, None) + + # --- Release (public, for await-pattern users) --- + + async def release(self, connection: ConnT, **kwargs) -> None: + pool = self._unmanaged_connections.pop(connection, None) + if pool is None: + return + self._metrics.remove_connection(self._pool_state.host(pool)) + await self._pool_state.release_to_pool(connection, pool, **kwargs) + + # --- Lifecycle --- + + async def close(self): + self._closing = True + await self._clear() + pool_state = self._pool_state + await asyncio.gather( + *[ + pool_state.close_pool(pool) + for pool in pool_state.pools + if pool is not None + ], + return_exceptions=True, + ) + self._closing = False + self._closed = True + + async def terminate(self): + self._closing = True + await self._clear() + pool_state = self._pool_state + await asyncio.gather( + *[ + pool_state.terminate_pool(pool) + for pool in pool_state.pools + if pool is not None + ], + return_exceptions=True, + ) + self._closing = False + self._closed = True + + async def _clear(self): + self._balancer = None + await self._health.stop() + + snapshot = list(self._unmanaged_connections.items()) + self._unmanaged_connections.clear() + + release_tasks = [] + for connection, pool in snapshot: + self._metrics.remove_connection(self._pool_state.host(pool)) + release_tasks.append( + self._pool_state.release_to_pool(connection, pool), + ) + + await asyncio.gather(*release_tasks, return_exceptions=True) + + self._pool_state.clear_sets() + + async def __aenter__(self): + await self._pool_state.ready() + return self + + async def __aexit__(self, exc_type, exc_val, exc_tb): + await self.close() + + +__all__ = ("BasePoolManager",) diff --git a/hasql/pool_state.py b/hasql/pool_state.py index 1c4d5a6..67a650a 100644 --- a/hasql/pool_state.py +++ b/hasql/pool_state.py @@ -1,28 +1,336 @@ -# Minimal stub for PR2 — full PoolState implementation is in PR3. -from typing import Protocol, TypeVar, runtime_checkable +import asyncio +import logging +from collections import defaultdict +from collections.abc import Sequence +from itertools import chain +from types import MappingProxyType +from typing import ( + Generic, + Protocol, + TypeVar, + runtime_checkable, +) + +from .abc import PoolDriver +from .acquire import AcquireContext +from .metrics import PoolStats +from .utils import Dsn, Stopwatch + +logger = logging.getLogger(__name__) PoolT = TypeVar("PoolT") +ConnT = TypeVar("ConnT") @runtime_checkable class PoolStateProvider(Protocol[PoolT]): - """Protocol for read access to pool sets — used by balancer policies.""" - @property def master_pool_count(self) -> int: ... - @property def replica_pool_count(self) -> int: ... - async def get_master_pools(self) -> list[PoolT]: ... - async def get_replica_pools( self, fallback_master: bool = False, ) -> list[PoolT]: ... - def get_pool_freesize(self, pool: PoolT) -> int: ... - def get_last_response_time(self, pool: PoolT) -> float | None: ... -__all__ = ("PoolStateProvider",) +class PoolState(Generic[PoolT, ConnT]): + """Owns the driver and all pool state: master/replica sets, + pool lifecycle, connection operations, waiting, and readiness.""" + + _dsn_ready_event: defaultdict[Dsn, asyncio.Event] + _dsn_check_cond: defaultdict[Dsn, asyncio.Condition] + _master_pool_set: set[PoolT] + _replica_pool_set: set[PoolT] + + def __init__( + self, + dsn_list: list[Dsn], + driver: PoolDriver[PoolT, ConnT], + stopwatch_window_size: int, + pool_factory_kwargs: dict | None = None, + ): + self._driver = driver + self._pool_factory_kwargs: MappingProxyType = MappingProxyType( + self._driver.prepare_pool_factory_kwargs( + dict(pool_factory_kwargs) + if pool_factory_kwargs is not None + else {}, + ), + ) + self._dsn = list(dsn_list) + self._pools: list[PoolT | None] = [None] * len(dsn_list) + self._dsn_ready_event = defaultdict(asyncio.Event) + self._dsn_check_cond = defaultdict(asyncio.Condition) + self._master_pool_set: set[PoolT] = set() + self._replica_pool_set: set[PoolT] = set() + self._master_cond = asyncio.Condition() + self._replica_cond = asyncio.Condition() + self._stopwatch: Stopwatch[PoolT] = Stopwatch( + window_size=stopwatch_window_size, + ) + + # --- Properties --- + + @property + def driver(self) -> PoolDriver[PoolT, ConnT]: + return self._driver + + @property + def dsn(self) -> Sequence[Dsn]: + return tuple(self._dsn) + + @property + def pools(self) -> Sequence[PoolT | None]: + return tuple(self._pools) + + @property + def pool_factory_kwargs(self) -> MappingProxyType: + return self._pool_factory_kwargs + + @property + def master_pool_count(self) -> int: + return len(self._master_pool_set) + + @property + def replica_pool_count(self) -> int: + return len(self._replica_pool_set) + + @property + def available_pool_count(self) -> int: + return self.master_pool_count + self.replica_pool_count + + # --- Pool state queries --- + + def pool_is_master(self, pool: PoolT) -> bool: + return pool in self._master_pool_set + + def pool_is_replica(self, pool: PoolT) -> bool: + return pool in self._replica_pool_set + + def get_pool_freesize(self, pool: PoolT) -> int: + return self._driver.get_pool_freesize(pool) + + def get_last_response_time(self, pool: PoolT) -> float | None: + return self._stopwatch.get_time(pool) + + # --- Driver operations --- + + def acquire_from_pool( + self, + pool: PoolT, + *, + timeout: float | None = None, + **kwargs, + ) -> AcquireContext[ConnT]: + return self._driver.acquire_from_pool(pool, timeout=timeout, **kwargs) + + async def release_to_pool( + self, + connection: ConnT, + pool: PoolT, + **kwargs, + ): + await self._driver.release_to_pool(connection, pool, **kwargs) + + async def pool_factory(self, dsn: Dsn) -> PoolT: + return await self._driver.pool_factory(dsn, **self._pool_factory_kwargs) + + async def close_pool(self, pool: PoolT): + await self._driver.close_pool(pool) + + async def terminate_pool(self, pool: PoolT): + await self._driver.terminate_pool(pool) + + def is_connection_closed(self, connection: ConnT) -> bool: + return self._driver.is_connection_closed(connection) + + def host(self, pool: PoolT) -> str: + return self._driver.host(pool) + + def pool_stats(self, pool: PoolT) -> PoolStats: + return self._driver.pool_stats(pool) + + # --- Pool retrieval (async, waits for availability) --- + + async def get_master_pools(self) -> list[PoolT]: + await self.wait_for_master_pools() + return list(self._master_pool_set) + + async def get_replica_pools( + self, + fallback_master: bool = False, + ) -> list[PoolT]: + if not self._replica_pool_set and fallback_master: + return await self.get_master_pools() + await self.wait_for_replica_pools() + return list(self._replica_pool_set) + + # --- Pool waiting --- + + async def wait_for_master_pools(self) -> None: + async with self._master_cond: + await self._master_cond.wait_for( + lambda: bool(self._master_pool_set), + ) + + async def wait_for_replica_pools( + self, + fallback_master: bool = False, + ) -> None: + if not self._replica_pool_set and fallback_master: + await self.wait_for_master_pools() + return + async with self._replica_cond: + await self._replica_cond.wait_for( + lambda: bool(self._replica_pool_set), + ) + + async def wait_masters_ready(self, masters_count: int): + def predicate(): + return self.master_pool_count >= masters_count + + async with self._master_cond: + await self._master_cond.wait_for(predicate) + + async def wait_replicas_ready(self, replicas_count: int): + def predicate(): + return self.replica_pool_count >= replicas_count + + async with self._replica_cond: + await self._replica_cond.wait_for(predicate) + + async def wait_all_ready(self): + for dsn in self._dsn: + await self._dsn_ready_event[dsn].wait() + + async def ready( + self, + masters_count: int | None = None, + replicas_count: int | None = None, + timeout: int = 10, + ): + if (masters_count is not None and replicas_count is None) or ( + masters_count is None and replicas_count is not None + ): + raise ValueError( + "Arguments masters_count and replicas_count " + "should both be either None or not None", + ) + + if masters_count is not None and masters_count < 0: + raise ValueError("masters_count shouldn't be negative") + if replicas_count is not None and replicas_count < 0: + raise ValueError("replicas_count shouldn't be negative") + + if masters_count is None: + await asyncio.wait_for(self.wait_all_ready(), timeout=timeout) + return + + if replicas_count is None: + raise ValueError( + "Arguments masters_count and replicas_count " + "should both be either None or not None", + ) + + await asyncio.wait_for( + asyncio.gather( + self.wait_masters_ready(masters_count), + self.wait_replicas_ready(replicas_count), + ), + timeout=timeout, + ) + + async def wait_next_pool_check(self, timeout: int = 10): + tasks = [self._wait_checking_pool(dsn) for dsn in self._dsn] + await asyncio.wait_for(asyncio.gather(*tasks), timeout=timeout) + + async def _wait_checking_pool(self, dsn: Dsn): + async with self._dsn_check_cond[dsn]: + for _ in range(2): + await self._dsn_check_cond[dsn].wait() + + # --- Pool role refresh --- + + async def refresh_pool_role( + self, pool: PoolT, dsn: Dsn, sys_connection: ConnT, + ): + with self._stopwatch(pool): + is_master = await self._driver.is_master(sys_connection) + if is_master: + await self._add_pool_to_master_set(pool, dsn) + self._remove_pool_from_replica_set(pool, dsn) + else: + await self._add_pool_to_replica_set(pool, dsn) + self._remove_pool_from_master_set(pool, dsn) + self._dsn_ready_event[dsn].set() + + def remove_pool_from_all_sets(self, pool: PoolT, dsn: Dsn): + self._remove_pool_from_master_set(pool, dsn) + self._remove_pool_from_replica_set(pool, dsn) + + def clear_sets(self): + self._master_pool_set.clear() + self._replica_pool_set.clear() + + # --- Pool registry --- + + def set_pool(self, index: int, pool: PoolT): + self._pools[index] = pool + + async def notify_pool_checked(self, dsn: Dsn): + async with self._dsn_check_cond[dsn]: + self._dsn_check_cond[dsn].notify_all() + + # --- Pool state mutations --- + + async def _add_pool_to_master_set(self, pool: PoolT, dsn: Dsn): + if pool in self._master_pool_set: + return + self._master_pool_set.add(pool) + logger.debug( + "Pool %s has been added to master set", + dsn.with_(password="******"), + ) + async with self._master_cond: + self._master_cond.notify_all() + + async def _add_pool_to_replica_set(self, pool: PoolT, dsn: Dsn): + if pool in self._replica_pool_set: + return + self._replica_pool_set.add(pool) + logger.debug( + "Pool %s has been added to replica set", + dsn.with_(password="******"), + ) + async with self._replica_cond: + self._replica_cond.notify_all() + + def _remove_pool_from_master_set(self, pool: PoolT, dsn: Dsn): + if pool in self._master_pool_set: + self._master_pool_set.remove(pool) + logger.debug( + "Pool %s has been removed from master set", + dsn.with_(password="******"), + ) + + def _remove_pool_from_replica_set(self, pool: PoolT, dsn: Dsn): + if pool in self._replica_pool_set: + self._replica_pool_set.remove(pool) + logger.debug( + "Pool %s has been removed from replica set", + dsn.with_(password="******"), + ) + + # --- Iteration --- + + def __iter__(self): + return chain( + iter(self._master_pool_set), + iter(self._replica_pool_set), + ) + + +__all__ = ("PoolState", "PoolStateProvider") diff --git a/hasql/utils.py b/hasql/utils.py index 491fc15..def81ab 100644 --- a/hasql/utils.py +++ b/hasql/utils.py @@ -3,11 +3,9 @@ import statistics import time from collections import defaultdict, deque +from collections.abc import Generator, Iterable from contextlib import contextmanager -from typing import ( - Any, DefaultDict, Deque, Dict, Generator, Iterable, List, Optional, Tuple, - Union, -) +from typing import Any, Generic, TypeVar from urllib.parse import unquote, urlencode @@ -34,9 +32,9 @@ class Dsn: def __init__( self, netloc: str, - user: Optional[str] = None, - password: Optional[str] = None, - dbname: Optional[str] = None, + user: str | None = None, + password: str | None = None, + dbname: str | None = None, scheme: str = "postgresql", **kwargs: Any, ): @@ -84,7 +82,7 @@ def parse(cls, dsn: str) -> "Dsn": return cls._parse_connection_string(dsn) @classmethod - def _parse_connection_string_params(cls, conn_str: str) -> Dict[str, str]: + def _parse_connection_string_params(cls, conn_str: str) -> dict[str, str]: """Parse key=value pairs from connection string.""" params = {} current_key = None @@ -199,12 +197,13 @@ def _compile_dsn(self) -> str: def with_( self, - netloc: Optional[str] = None, - user: Optional[str] = None, - password: Optional[str] = None, - dbname: Optional[str] = None, + netloc: str | None = None, + user: str | None = None, + password: str | None = None, + dbname: str | None = None, ) -> "Dsn": params = { + "scheme": self._scheme, "netloc": netloc if netloc is not None else self._netloc, "user": user if user is not None else self._user, "password": password if password is not None else self._password, @@ -227,20 +226,20 @@ def netloc(self) -> str: return self._netloc @property - def user(self) -> Optional[str]: + def user(self) -> str | None: return self._user @property - def password(self) -> Optional[str]: + def password(self) -> str | None: return self._password @property - def dbname(self) -> Optional[str]: + def dbname(self) -> str | None: return self._dbname @property - def params(self) -> Dict[str, str]: - return self._kwargs + def params(self) -> dict[str, str]: + return dict(self._kwargs) @property def scheme(self) -> str: @@ -251,13 +250,13 @@ def compiled_dsn(self) -> str: return self._compiled_dsn -def split_dsn(dsn: Union[Dsn, str], default_port: int = 5432) -> List[Dsn]: +def split_dsn(dsn: Dsn | str, default_port: int = 5432) -> list[Dsn]: if not isinstance(dsn, Dsn): dsn = Dsn.parse(dsn) - host_port_pairs: List[Tuple[str, Optional[int]]] = [] + host_port_pairs: list[tuple[str, int | None]] = [] port_count = 0 - port: Optional[int] + port: int | None for host in dsn.netloc.split(","): if ":" in host: host, port_str = host.rsplit(":", 1) @@ -268,7 +267,7 @@ def split_dsn(dsn: Union[Dsn, str], default_port: int = 5432) -> List[Dsn]: port = None host_port_pairs.append((host, port)) - def deduplicate(dsns: Iterable[Dsn]) -> List[Dsn]: + def deduplicate(dsns: Iterable[Dsn]) -> list[Dsn]: cache = set() result = [] for dsn in dsns: @@ -297,14 +296,17 @@ def deduplicate(dsns: Iterable[Dsn]) -> List[Dsn]: ) -class Stopwatch: +KeyT = TypeVar("KeyT") + + +class Stopwatch(Generic[KeyT]): def __init__(self, window_size: int): - self._times: DefaultDict[Any, Deque] = defaultdict( + self._times: defaultdict[KeyT, deque[float]] = defaultdict( lambda: deque(maxlen=window_size), ) - self._cache: Dict[Any, Optional[int]] = {} + self._cache: dict[KeyT, float | None] = {} - def get_time(self, obj: Any) -> Optional[float]: + def get_time(self, obj: KeyT) -> float | None: if obj not in self._times: return None if self._cache.get(obj) is None: @@ -312,7 +314,7 @@ def get_time(self, obj: Any) -> Optional[float]: return self._cache[obj] @contextmanager - def __call__(self, obj: Any) -> Generator[None, None, None]: + def __call__(self, obj: KeyT) -> Generator[None, None, None]: start_at = time.monotonic() yield self._times[obj].append(time.monotonic() - start_at) diff --git a/tests/conftest.py b/tests/conftest.py index df738b7..2ab7eba 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -43,7 +43,7 @@ def pg_dsn() -> str: @asynccontextmanager async def setup_aiopg(pg_dsn): - from hasql.aiopg import PoolManager + from hasql.driver.aiopg import PoolManager pool = PoolManager(dsn=pg_dsn, fallback_master=True) yield pool @@ -52,7 +52,7 @@ async def setup_aiopg(pg_dsn): @asynccontextmanager async def setup_aiopgsa(pg_dsn): - from hasql.aiopg_sa import PoolManager + from hasql.driver.aiopg_sa import PoolManager pool = PoolManager(dsn=pg_dsn, fallback_master=True) yield pool @@ -61,7 +61,7 @@ async def setup_aiopgsa(pg_dsn): @asynccontextmanager async def setup_asyncpg(pg_dsn): - from hasql.asyncpg import PoolManager + from hasql.driver.asyncpg import PoolManager pool = PoolManager(dsn=pg_dsn, fallback_master=True) yield pool @@ -70,7 +70,7 @@ async def setup_asyncpg(pg_dsn): @asynccontextmanager async def setup_asyncsqlalchemy(pg_dsn): - from hasql.asyncsqlalchemy import PoolManager + from hasql.driver.asyncsqlalchemy import PoolManager pool = PoolManager(dsn=pg_dsn, fallback_master=True) yield pool @@ -79,7 +79,7 @@ async def setup_asyncsqlalchemy(pg_dsn): @asynccontextmanager async def setup_psycopg3(pg_dsn): - from hasql.psycopg3 import PoolManager + from hasql.driver.psycopg3 import PoolManager pool = PoolManager(dsn=pg_dsn, fallback_master=True) yield pool diff --git a/tests/mocks/pool_manager.py b/tests/mocks/pool_manager.py index 68f41c9..eef5e98 100644 --- a/tests/mocks/pool_manager.py +++ b/tests/mocks/pool_manager.py @@ -1,10 +1,10 @@ import asyncio -from typing import Any, Sequence import mock +from hasql.abc import PoolDriver from hasql.base import BasePoolManager -from hasql.metrics import DriverMetrics +from hasql.metrics import PoolStats from hasql.utils import Dsn @@ -26,6 +26,11 @@ async def is_master(self): await asyncio.sleep(100) return self._pool.is_master + async def fetch_scalar(self, query): + if not self._pool.is_running: + raise ConnectionRefusedError + return self._pool.fetch_scalar_result + async def close(self): self._is_closed = True @@ -60,6 +65,7 @@ def __init__(self, dsn: str, maxsize: int = 10): self.is_master = dsn == "postgresql://test:test@master:5432/test" self.is_running = True self.is_behind_firewall = False + self.fetch_scalar_result = None self.used = set() self.free = asyncio.LifoQueue() self.connections = [TestConnection(self) for _ in range(maxsize)] @@ -99,42 +105,52 @@ def terminate(self): conn.terminate() -class TestPoolManager(BasePoolManager): +class TestDriver(PoolDriver[TestPool, TestConnection]): def get_pool_freesize(self, pool: TestPool): return pool.freesize - def acquire_from_pool(self, pool: TestPool, **kwargs): + def acquire_from_pool(self, pool: TestPool, *, timeout=None, **kwargs): return pool.acquire(**kwargs) async def release_to_pool( - self, connection: TestConnection, pool: TestPool, **kwargs + self, connection: TestConnection, pool: TestPool, **kwargs, ): await pool.release(connection, **kwargs) - async def _is_master(self, connection: TestConnection): + async def is_master(self, connection: TestConnection): return await connection.is_master() - async def _pool_factory(self, dsn: Dsn): + async def fetch_scalar(self, connection: TestConnection, query: str): + return await connection.fetch_scalar(query) + + async def pool_factory(self, dsn: Dsn, **kwargs): return TestPool(str(dsn)) - async def _close(self, pool: TestPool): + async def close_pool(self, pool: TestPool): await pool.close() - async def _terminate(self, pool: TestPool): + async def terminate_pool(self, pool: TestPool): loop = asyncio.get_running_loop() await loop.run_in_executor(None, pool.terminate) def is_connection_closed(self, connection: TestConnection): return connection.is_closed - def metrics(self) -> Sequence[DriverMetrics]: - return [] - - def host(self, pool: Any): + def host(self, pool: TestPool): return "test-host:5432" - def _driver_metrics(self) -> Sequence[DriverMetrics]: - return [] + def pool_stats(self, pool: TestPool) -> PoolStats: + return PoolStats( + min=0, + max=len(pool.connections), + idle=pool.freesize, + used=len(pool.used), + ) + + +class TestPoolManager(BasePoolManager[TestPool, TestConnection]): + def __init__(self, dsn, **kwargs): + super().__init__(dsn, driver=TestDriver(), **kwargs) -__all__ = ("TestPoolManager",) +__all__ = ("TestDriver", "TestPoolManager") diff --git a/tests/test_backward_compat.py b/tests/test_backward_compat.py new file mode 100644 index 0000000..d553e92 --- /dev/null +++ b/tests/test_backward_compat.py @@ -0,0 +1,66 @@ +"""Test that all old import paths still work after the split.""" + + +def test_base_exports_pool_manager(): + from hasql.base import BasePoolManager + from hasql.pool_manager import BasePoolManager as Direct + + assert BasePoolManager is Direct + + +def test_base_exports_abstract_balancer_policy(): + from hasql.balancer_policy import AbstractBalancerPolicy as Direct + from hasql.base import AbstractBalancerPolicy + + assert AbstractBalancerPolicy is Direct + + +def test_base_exports_timeout_acquire_context(): + from hasql.acquire import TimeoutAcquireContext as Direct + from hasql.base import TimeoutAcquireContext + + assert TimeoutAcquireContext is Direct + + +def test_base_exports_pool_acquire_context(): + from hasql.acquire import PoolAcquireContext as Direct + from hasql.base import PoolAcquireContext + + assert PoolAcquireContext is Direct + + +def test_base_exports_acquire_context(): + from hasql.acquire import AcquireContext as Direct + from hasql.base import AcquireContext + + assert AcquireContext is Direct + + +def test_base_exports_type_vars(): + from hasql.base import ConnT, PoolT + from hasql.pool_manager import ConnT as DirectConnT + from hasql.pool_manager import PoolT as DirectPoolT + + assert PoolT is DirectPoolT + assert ConnT is DirectConnT + + +def test_base_exports_pool_driver(): + from hasql.abc import PoolDriver as Direct + from hasql.base import PoolDriver + + assert PoolDriver is Direct + + +def test_base_exports_pool_state(): + from hasql.base import PoolState + from hasql.pool_state import PoolState as Direct + + assert PoolState is Direct + + +def test_base_exports_pool_state_provider(): + from hasql.base import PoolStateProvider + from hasql.pool_state import PoolStateProvider as Direct + + assert PoolStateProvider is Direct diff --git a/tests/test_balancer_policy.py b/tests/test_balancer_policy.py index ebe56a5..cfe4ae5 100644 --- a/tests/test_balancer_policy.py +++ b/tests/test_balancer_policy.py @@ -9,6 +9,7 @@ RoundRobinBalancerPolicy, ) from tests.mocks import TestPoolManager +from tests.mocks.pool_manager import TestPool balancer_policies = pytest.mark.parametrize( "balancer_policy", @@ -102,3 +103,244 @@ async def test_dont_acquire_master_as_replica( with pytest.raises(asyncio.TimeoutError): async with pool_manager.acquire_replica(master_as_replica_weight=0.0): pass + + +@balancer_policies +async def test_master_as_replica_weight_zero_always_false( + make_pool_manager, + balancer_policy, +): + """weight=0 should never choose master as replica, even when rand=0.""" + pool_manager = await make_pool_manager(balancer_policy, replicas_count=2) + async with timeout(1): + await pool_manager._pool_state.ready() + + # With weight=0, should never get master when requesting replica + for _ in range(20): + pool = await pool_manager._balancer.get_pool( + read_only=True, + master_as_replica_weight=0.0, + ) + assert pool is None or pool_manager._pool_state.pool_is_replica(pool) + + +@balancer_policies +async def test_get_pool_write_with_master_as_replica_weight_raises( + make_pool_manager, + balancer_policy, +): + pool_manager = await make_pool_manager(balancer_policy) + async with timeout(1): + await pool_manager._pool_state.ready() + with pytest.raises(ValueError, match="master_as_replica_weight"): + await pool_manager._balancer.get_pool( + read_only=False, + master_as_replica_weight=0.5, + ) + + +def test_random_weighted_compute_weights_equal_times(): + # Equal response times should produce equal weights + weights = RandomWeightedBalancerPolicy._compute_weights([0.1, 0.1, 0.1]) + assert len(weights) == 3 + assert weights[0] == pytest.approx(weights[1]) + assert weights[1] == pytest.approx(weights[2]) + + +def test_random_weighted_compute_weights_none_times(): + # None times treated as 0 — all equal + weights = RandomWeightedBalancerPolicy._compute_weights([None, None]) + assert len(weights) == 2 + assert weights[0] == pytest.approx(weights[1]) + + +@pytest.mark.parametrize("n", [1, 2, 3, 5, 10]) +def test_random_weighted_compute_weights_all_none_produces_uniform(n): + # When all response times are None (e.g. at startup before any health + # checks), weights should be uniform and all positive. + weights = RandomWeightedBalancerPolicy._compute_weights([None] * n) + assert len(weights) == n + assert all(w > 0 for w in weights) + assert all(w == pytest.approx(weights[0]) for w in weights) + + +def test_random_weighted_compute_weights_all_zero_produces_uniform(): + # Explicit zero times should also produce uniform positive weights + weights = RandomWeightedBalancerPolicy._compute_weights([0, 0, 0]) + assert len(weights) == 3 + assert all(w > 0 for w in weights) + assert all(w == pytest.approx(weights[0]) for w in weights) + + +def test_random_weighted_compute_weights_favors_faster(): + # Faster pool (lower time) should get higher weight + weights = RandomWeightedBalancerPolicy._compute_weights([0.1, 0.9]) + assert weights[0] > weights[1] + + +async def test_round_robin_master_as_replica(make_pool_manager): + pool_manager = await make_pool_manager( + RoundRobinBalancerPolicy, + replicas_count=0, + ) + async with timeout(1): + await pool_manager._pool_state.ready() + + async with pool_manager.acquire_replica( + master_as_replica_weight=1.0, + ) as conn: + assert await conn.is_master() + + +async def test_round_robin_waits_for_master_when_not_ready( + make_pool_manager, +): + pool_manager = await make_pool_manager( + RoundRobinBalancerPolicy, + replicas_count=0, + ) + async with timeout(2): + await pool_manager._pool_state.ready() + + # Shut down the master so master_pool_count drops to 0 + ps = pool_manager._pool_state + master_pool: TestPool = (await ps.get_master_pools())[0] + master_pool.shutdown() + + # Wait for the health monitor to detect the shutdown + await ps.wait_next_pool_check() + assert ps.master_pool_count == 0 + + # Bring it back after a short delay so the wait resolves + async def bring_master_back(): + await asyncio.sleep(0.15) + master_pool.startup() + master_pool.set_master(True) + + asyncio.ensure_future(bring_master_back()) + + # Use explicit timeout to override the short default acquire_timeout + async with timeout(2): + async with pool_manager.acquire_master(timeout=2) as conn: + assert await conn.is_master() + + +async def test_round_robin_waits_for_replica_when_not_ready( + make_pool_manager, +): + pool_manager = await make_pool_manager( + RoundRobinBalancerPolicy, + replicas_count=2, + ) + async with timeout(2): + await pool_manager._pool_state.ready() + + # Shut down all replicas so replica_pool_count drops to 0 + ps = pool_manager._pool_state + replica_pools = [ + pool + for pool in ps.pools + if pool is not None and ps.pool_is_replica(pool) + ] + for rp in replica_pools: + rp.shutdown() + + # Wait for health monitor to detect shutdowns + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.replica_pool_count == 0 + + # Bring replicas back after a short delay + async def bring_replicas_back(): + await asyncio.sleep(0.15) + for rp in replica_pools: + rp.startup() + + asyncio.ensure_future(bring_replicas_back()) + + # Use explicit timeout to override the short default acquire_timeout + async with timeout(2): + async with pool_manager.acquire_replica(timeout=2) as conn: + assert not await conn.is_master() + + +async def test_round_robin_fallback_master_waits_when_master_not_ready( + make_pool_manager, +): + pool_manager = await make_pool_manager( + RoundRobinBalancerPolicy, + replicas_count=0, + ) + async with timeout(2): + await pool_manager._pool_state.ready() + + # Shut down the master so both master and replica counts are 0 + ps = pool_manager._pool_state + master_pool: TestPool = (await ps.get_master_pools())[0] + master_pool.shutdown() + + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.master_pool_count == 0 + assert pool_manager._pool_state.replica_pool_count == 0 + + # Bring master back after a short delay + async def bring_master_back(): + await asyncio.sleep(0.15) + master_pool.startup() + master_pool.set_master(True) + + asyncio.ensure_future(bring_master_back()) + + # acquire_replica with fallback_master=True should wait for master + # Use explicit timeout to override the short default acquire_timeout + async with timeout(2): + async with pool_manager.acquire_replica( + fallback_master=True, + timeout=2, + ) as conn: + assert await conn.is_master() + + +async def test_round_robin_master_with_fallback_and_no_replicas( + make_pool_manager, +): + pool_manager = await make_pool_manager( + RoundRobinBalancerPolicy, + replicas_count=0, + ) + async with timeout(1): + await pool_manager._pool_state.ready() + + assert pool_manager._pool_state.replica_pool_count == 0 + + # Acquiring master should work even when fallback_master + # is set and there are no replicas + async with timeout(1): + async with pool_manager.acquire_master() as conn: + assert await conn.is_master() + + +@balancer_policies +async def test_master_as_replica_weight_nonzero_with_no_replicas( + make_pool_manager, + balancer_policy, +): + """weight > 0 with replica_count=0 must still return master. + + _get_candidates has a guard: the choose_master_as_replica branch that + adds masters directly is skipped when replica_pool_count == 0. + Masters must still be reachable via get_replica_pools(fallback_master=True), + which is activated because get_pool() sets fallback_master=choose_master_as_replica. + """ + pool_manager = await make_pool_manager(balancer_policy, replicas_count=0) + async with timeout(1): + await pool_manager.ready() + + assert pool_manager.replica_pool_count == 0 + + # weight=1.0 makes choose_master_as_replica deterministically True + pool = await pool_manager._balancer.get_pool( + read_only=True, + master_as_replica_weight=1.0, + ) + assert pool is not None + assert pool_manager.pool_is_master(pool) diff --git a/tests/test_base_pool_manager.py b/tests/test_base_pool_manager.py index b60baf7..353be6c 100644 --- a/tests/test_base_pool_manager.py +++ b/tests/test_base_pool_manager.py @@ -8,6 +8,7 @@ from async_timeout import timeout as timeout_context from hasql.base import BasePoolManager +from hasql.exceptions import PoolManagerClosedError from tests.mocks import TestPoolManager @@ -26,44 +27,45 @@ async def pool_manager(dsn): def pool_is_master(pool_manager: BasePoolManager, pool): - assert pool_manager.pool_is_master(pool) - assert not pool_manager.pool_is_replica(pool) + assert pool_manager._pool_state.pool_is_master(pool) + assert not pool_manager._pool_state.pool_is_replica(pool) def pool_is_replica(pool_manager: BasePoolManager, pool): - assert pool_manager.pool_is_replica(pool) - assert not pool_manager.pool_is_master(pool) + assert pool_manager._pool_state.pool_is_replica(pool) + assert not pool_manager._pool_state.pool_is_master(pool) async def test_wait_next_pool_check(pool_manager: BasePoolManager): - await pool_manager.ready() - master_pool = await pool_manager.balancer.get_pool(read_only=False) + await pool_manager._pool_state.ready() + master_pool = await pool_manager._balancer.get_pool(read_only=False) master_pool.shutdown() - assert pool_manager.master_pool_count == 1 - await pool_manager.wait_next_pool_check() - assert pool_manager.master_pool_count == 0 + assert pool_manager._pool_state.master_pool_count == 1 + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.master_pool_count == 0 async def test_ready_all_hosts(pool_manager: BasePoolManager): - await pool_manager.ready() - assert len(pool_manager.dsn) == pool_manager.available_pool_count + await pool_manager._pool_state.ready() + ps = pool_manager._pool_state + assert len(ps.dsn) == ps.available_pool_count async def test_ready_min_count_hosts(pool_manager: BasePoolManager): - await pool_manager.ready() - replica_pools = await pool_manager.get_replica_pools() + await pool_manager._pool_state.ready() + replica_pools = await pool_manager._pool_state.get_replica_pools() for replica_pool in replica_pools: replica_pool.shutdown() - master_pool = await pool_manager.balancer.get_pool(read_only=False) + master_pool = await pool_manager._balancer.get_pool(read_only=False) master_pool.shutdown() - await pool_manager.wait_next_pool_check() - assert pool_manager.master_pool_count == 0 - assert pool_manager.replica_pool_count == 0 + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.master_pool_count == 0 + assert pool_manager._pool_state.replica_pool_count == 0 master_pool.startup() master_pool.set_master(True) - await pool_manager.ready(masters_count=1, replicas_count=0) - assert pool_manager.master_pool_count == 1 - assert pool_manager.replica_pool_count == 0 + await pool_manager._pool_state.ready(masters_count=1, replicas_count=0) + assert pool_manager._pool_state.master_pool_count == 1 + assert pool_manager._pool_state.replica_pool_count == 0 @pytest.mark.parametrize( @@ -81,96 +83,97 @@ async def test_ready_with_invalid_arguments( replicas_count: Optional[int], ): with pytest.raises(ValueError): - await pool_manager.ready(masters_count, replicas_count) + await pool_manager._pool_state.ready(masters_count, replicas_count) async def test_wait_db_restart(pool_manager: BasePoolManager): - await pool_manager.ready() - master_pool = await pool_manager.balancer.get_pool(read_only=False) - assert pool_manager.pool_is_master(master_pool) + await pool_manager._pool_state.ready() + master_pool = await pool_manager._balancer.get_pool(read_only=False) + assert pool_manager._pool_state.pool_is_master(master_pool) master_pool.shutdown() - await pool_manager.wait_next_pool_check() - assert pool_manager.master_pool_count == 0 + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.master_pool_count == 0 master_pool.startup() - await pool_manager.wait_next_pool_check() - assert pool_manager.master_pool_count == 0 - assert pool_manager.pool_is_replica(master_pool) + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.master_pool_count == 0 + assert pool_manager._pool_state.pool_is_replica(master_pool) async def test_master_shutdown(pool_manager: BasePoolManager): - await pool_manager.ready() - master_pool = await pool_manager.balancer.get_pool(read_only=False) - assert pool_manager.pool_is_master(master_pool) + await pool_manager._pool_state.ready() + master_pool = await pool_manager._balancer.get_pool(read_only=False) + assert pool_manager._pool_state.pool_is_master(master_pool) master_pool.shutdown() - await pool_manager.wait_next_pool_check() - assert pool_manager.master_pool_count == 0 + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.master_pool_count == 0 async def test_replica_shutdown(pool_manager: BasePoolManager): - await pool_manager.ready() - replica_pool = await pool_manager.balancer.get_pool(read_only=True) - assert pool_manager.pool_is_replica(replica_pool) - assert pool_manager.replica_pool_count == 2 + await pool_manager._pool_state.ready() + replica_pool = await pool_manager._balancer.get_pool(read_only=True) + assert pool_manager._pool_state.pool_is_replica(replica_pool) + assert pool_manager._pool_state.replica_pool_count == 2 replica_pool.shutdown() - await pool_manager.wait_next_pool_check() - assert pool_manager.replica_pool_count == 1 + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.replica_pool_count == 1 async def test_change_master(pool_manager: BasePoolManager): - await pool_manager.ready() - master_pool = await pool_manager.balancer.get_pool(read_only=False) - replica_pool = await pool_manager.balancer.get_pool(read_only=True) + await pool_manager._pool_state.ready() + master_pool = await pool_manager._balancer.get_pool(read_only=False) + replica_pool = await pool_manager._balancer.get_pool(read_only=True) pool_is_master(pool_manager, master_pool) pool_is_replica(pool_manager, replica_pool) master_pool.set_master(False) replica_pool.set_master(True) - await pool_manager.wait_next_pool_check() + await pool_manager._pool_state.wait_next_pool_check() pool_is_master(pool_manager, replica_pool) pool_is_replica(pool_manager, master_pool) async def test_define_roles(pool_manager: BasePoolManager): - await pool_manager.ready() - master_pool = await pool_manager.balancer.get_pool(read_only=False) - replica_pool = await pool_manager.balancer.get_pool(read_only=True) + await pool_manager._pool_state.ready() + master_pool = await pool_manager._balancer.get_pool(read_only=False) + replica_pool = await pool_manager._balancer.get_pool(read_only=True) pool_is_master(pool_manager, master_pool) pool_is_replica(pool_manager, replica_pool) async def test_acquire_master_and_release(pool_manager: BasePoolManager): - await pool_manager.ready() - master_pool = await pool_manager.balancer.get_pool(read_only=False) - init_freesize = pool_manager.get_pool_freesize(master_pool) - connection = await pool_manager.acquire_master() - assert pool_manager.get_pool_freesize(master_pool) + 1 == init_freesize - assert connection in master_pool.used - await pool_manager.release(connection) + await pool_manager._pool_state.ready() + ps = pool_manager._pool_state + master_pool = await pool_manager._balancer.get_pool(read_only=False) + init_freesize = ps.get_pool_freesize(master_pool) + async with pool_manager.acquire_master() as connection: + assert ps.get_pool_freesize(master_pool) + 1 == init_freesize + assert connection in master_pool.used assert connection not in master_pool.used - assert pool_manager.get_pool_freesize(master_pool) == init_freesize + assert ps.get_pool_freesize(master_pool) == init_freesize async def test_acquire_with_context(pool_manager: BasePoolManager): - await pool_manager.ready() - master_pool = await pool_manager.balancer.get_pool(read_only=False) - init_freesize = pool_manager.get_pool_freesize(master_pool) + await pool_manager._pool_state.ready() + ps = pool_manager._pool_state + master_pool = await pool_manager._balancer.get_pool(read_only=False) + init_freesize = ps.get_pool_freesize(master_pool) async with pool_manager.acquire_master() as connection: - assert pool_manager.get_pool_freesize(master_pool) + 1 == init_freesize + assert ps.get_pool_freesize(master_pool) + 1 == init_freesize assert connection in master_pool.used assert connection not in master_pool.used - assert pool_manager.get_pool_freesize(master_pool) == init_freesize + assert ps.get_pool_freesize(master_pool) == init_freesize async def test_acquire_replica_with_fallback_master_is_true( pool_manager: BasePoolManager, ): - await pool_manager.ready() - master_pool = await pool_manager.balancer.get_pool(read_only=False) - replica_pools = await pool_manager.get_replica_pools() + await pool_manager._pool_state.ready() + master_pool = await pool_manager._balancer.get_pool(read_only=False) + replica_pools = await pool_manager._pool_state.get_replica_pools() for replica_pool in replica_pools: - assert pool_manager.pool_is_replica(replica_pool) + assert pool_manager._pool_state.pool_is_replica(replica_pool) replica_pool.shutdown() - await pool_manager.wait_next_pool_check() - assert pool_manager.replica_pool_count == 0 + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.replica_pool_count == 0 async with timeout_context(1): async with pool_manager.acquire_replica( fallback_master=True, @@ -181,79 +184,66 @@ async def test_acquire_replica_with_fallback_master_is_true( async def test_acquire_replica_with_fallback_master_is_false( pool_manager: BasePoolManager, ): - await pool_manager.ready() - replica_pools = await pool_manager.get_replica_pools() + await pool_manager._pool_state.ready() + replica_pools = await pool_manager._pool_state.get_replica_pools() for replica_pool in replica_pools: - assert pool_manager.pool_is_replica(replica_pool) + assert pool_manager._pool_state.pool_is_replica(replica_pool) replica_pool.shutdown() - await pool_manager.wait_next_pool_check() - assert pool_manager.replica_pool_count == 0 + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.replica_pool_count == 0 with pytest.raises(asyncio.TimeoutError): async with timeout_context(1): await pool_manager.acquire_replica(fallback_master=False) async def test_close(pool_manager: BasePoolManager): - await pool_manager.ready() - assert pool_manager.master_pool_count > 0 - assert pool_manager.replica_pool_count > 0 + await pool_manager._pool_state.ready() + assert pool_manager._pool_state.master_pool_count > 0 + assert pool_manager._pool_state.replica_pool_count > 0 await pool_manager.close() - assert pool_manager.master_pool_count == 0 - assert pool_manager.replica_pool_count == 0 - for pool in pool_manager: + assert pool_manager._pool_state.master_pool_count == 0 + assert pool_manager._pool_state.replica_pool_count == 0 + for pool in pool_manager._pool_state: assert pool is not None assert all( - pool_manager.is_connection_closed(conn) for conn in pool.connections + pool_manager._pool_state.is_connection_closed(conn) + for conn in pool.connections ) assert all(conn.close.call_count == 1 for conn in pool.connections) -async def test_terminate(pool_manager: BasePoolManager): - await pool_manager.ready() - assert pool_manager.master_pool_count > 0 - assert pool_manager.replica_pool_count > 0 - await pool_manager.terminate() - assert pool_manager.master_pool_count == 0 - assert pool_manager.replica_pool_count == 0 - for pool in pool_manager: - assert pool is not None - assert all( - pool_manager.is_connection_closed(conn) for conn in pool.connections - ) - assert all(conn.terminate.call_count == 1 for conn in pool.connections) - - async def test_master_behind_firewall(pool_manager: BasePoolManager): - await pool_manager.ready() - assert pool_manager.master_pool_count == 1 - master_pool = (await pool_manager.get_master_pools())[0] + await pool_manager._pool_state.ready() + assert pool_manager._pool_state.master_pool_count == 1 + master_pool = (await pool_manager._pool_state.get_master_pools())[0] master_pool.behind_firewall(True) - await pool_manager.wait_next_pool_check() - assert pool_manager.master_pool_count == 0 + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.master_pool_count == 0 master_pool.behind_firewall(False) - await pool_manager.wait_next_pool_check() - assert pool_manager.master_pool_count == 1 + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.master_pool_count == 1 async def test_replica_behind_firewall(pool_manager: BasePoolManager): - await pool_manager.ready() + await pool_manager._pool_state.ready() replica_pool_count = 2 - assert pool_manager.replica_pool_count == replica_pool_count - replica_pools = await pool_manager.get_replica_pools() + assert pool_manager._pool_state.replica_pool_count == replica_pool_count + replica_pools = await pool_manager._pool_state.get_replica_pools() for replica_pool in replica_pools: + ps = pool_manager._pool_state replica_pool.behind_firewall(True) - await pool_manager.wait_next_pool_check() - assert pool_manager.replica_pool_count == replica_pool_count - 1 + await ps.wait_next_pool_check() + assert ps.replica_pool_count == replica_pool_count - 1 replica_pool.behind_firewall(False) - await pool_manager.wait_next_pool_check() - assert pool_manager.replica_pool_count == replica_pool_count + await ps.wait_next_pool_check() + assert ps.replica_pool_count == replica_pool_count async def test_check_pool_canceled_error_while_releasing_connection( pool_manager: BasePoolManager ): - await pool_manager.ready() - master_pool = await pool_manager.balancer.get_pool(read_only=False) + await pool_manager._pool_state.ready() + master_pool = await pool_manager._balancer.get_pool(read_only=False) with ExitStack() as stack: for conn in master_pool.connections: @@ -268,5 +258,187 @@ async def test_check_pool_canceled_error_while_releasing_connection( ) ) await asyncio.sleep(1) - for task in pool_manager._refresh_role_tasks: + for task in pool_manager._health.tasks: assert not task.done() + + +def test_invalid_balancer_policy(): + with pytest.raises(ValueError, match="balancer_policy"): + TestPoolManager( + dsn="postgresql://test:test@master/test", + balancer_policy=str, + ) + + +async def test_acquire_master_as_replica_weight_write_raises( + pool_manager: BasePoolManager, +): + await pool_manager._pool_state.ready() + with pytest.raises(ValueError, match="master_as_replica_weight"): + pool_manager.acquire(read_only=False, master_as_replica_weight=0.5) + + +@pytest.mark.parametrize("weight", [-0.1, 1.1, 2.0]) +async def test_acquire_master_as_replica_weight_out_of_range( + pool_manager: BasePoolManager, + weight: float, +): + await pool_manager._pool_state.ready() + with pytest.raises(ValueError, match="segment"): + pool_manager.acquire(read_only=True, master_as_replica_weight=weight) + + +async def test_metrics_after_acquire(pool_manager: BasePoolManager): + await pool_manager._pool_state.ready() + async with pool_manager.acquire_master() as _conn: + from hasql.metrics import Metrics + m = pool_manager.metrics() + assert isinstance(m, Metrics) + assert m.hasql.pool == 1 + assert m.hasql.add_connections.get("test-host:5432") == 1 + assert len(m.pools) == 3 + assert m.gauges.master_count == 1 + assert m.gauges.replica_count == 2 + assert m.gauges.active_connections == 1 + assert m.gauges.closing is False + assert m.gauges.closed is False + master = [p for p in m.pools if p.role == "master"][0] + assert master.in_flight == 1 + assert master.healthy is True + + +async def test_aenter_aexit(dsn): + async with TestPoolManager( + dsn, refresh_timeout=0.2, refresh_delay=0.1, + ) as pm: + assert pm._pool_state.master_pool_count > 0 + assert pm._closed + + +async def test_acquire_raises_pool_manager_closed_error_when_closed( + pool_manager: BasePoolManager, +): + await pool_manager._pool_state.ready() + await pool_manager.close() + with pytest.raises(PoolManagerClosedError): + pool_manager.acquire() + + +async def test_acquire_raises_pool_manager_closed_error_when_closing( + pool_manager: BasePoolManager, +): + await pool_manager._pool_state.ready() + pool_manager._closing = True + with pytest.raises(PoolManagerClosedError): + pool_manager.acquire() + + +async def test_pool_manager_closed_error_message_contains_closed( + pool_manager: BasePoolManager, +): + await pool_manager._pool_state.ready() + await pool_manager.close() + with pytest.raises(PoolManagerClosedError, match="closed"): + pool_manager.acquire() + + +async def test_close_releases_unmanaged_connections( + pool_manager: BasePoolManager, +): + await pool_manager._pool_state.ready() + conn = await pool_manager.acquire_master() + assert conn in pool_manager._unmanaged_connections + await pool_manager.close() + assert pool_manager._closed + assert len(pool_manager._unmanaged_connections) == 0 + + +async def test_check_pool_task_cancelled_error_non_closing(): + """CancelledError during _is_master when not closing removes pool.""" + pool_manager = TestPoolManager( + "postgresql://test:test@master/test", + refresh_timeout=0.2, + refresh_delay=0.05, + ) + try: + await pool_manager._pool_state.ready() + assert pool_manager._pool_state.master_pool_count == 1 + + with patch.object( + pool_manager._pool_state.driver, + 'is_master', + AsyncMock(side_effect=asyncio.CancelledError()), + ): + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.master_pool_count == 0 + + # Recovers after the patch is removed + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.master_pool_count == 1 + finally: + await pool_manager.close() + + +async def test_wait_creating_pool_retries_on_failure(): + """_wait_creating_pool retries when pool_factory raises.""" + from hasql.pool_state import PoolState + + call_count = 0 + original_pool_factory = PoolState.pool_factory + + async def failing_factory(self, dsn): + nonlocal call_count + call_count += 1 + if call_count < 3: + raise ConnectionError("cannot connect") + return await original_pool_factory(self, dsn) + + with patch.object( + PoolState, 'pool_factory', failing_factory, + ): + pm = TestPoolManager( + "postgresql://test:test@master/test", + refresh_timeout=0.2, + refresh_delay=0.05, + ) + try: + await pm._pool_state.ready(timeout=5) + assert call_count >= 3 + assert pm._pool_state.master_pool_count == 1 + finally: + await pm.close() + + +async def test_check_pool_task_release_exception(): + """Exception during release_to_pool in _check_pool_task is handled.""" + pool_manager = TestPoolManager( + "postgresql://test:test@master/test", + refresh_timeout=0.2, + refresh_delay=0.05, + ) + try: + await pool_manager._pool_state.ready() + assert pool_manager._pool_state.master_pool_count == 1 + + master_pool = (await pool_manager._pool_state.get_master_pools())[0] + + with ExitStack() as stack: + for conn in master_pool.connections: + stack.enter_context( + patch.object( + conn, 'is_master', + AsyncMock(side_effect=Exception("db error")), + ) + ) + stack.enter_context( + patch.object( + master_pool, 'release', + AsyncMock(side_effect=RuntimeError("release failed")), + ) + ) + await asyncio.sleep(0.5) + # Tasks should still be running despite release errors + for task in pool_manager._health.tasks: + assert not task.done() + finally: + await pool_manager.close() diff --git a/tests/test_timeout_handling.py b/tests/test_timeout_handling.py index 7177bc4..d4adf3a 100644 --- a/tests/test_timeout_handling.py +++ b/tests/test_timeout_handling.py @@ -6,7 +6,7 @@ from hasql.base import PoolAcquireContext from hasql.metrics import CalculateMetrics from tests.mocks import TestPoolManager -from tests.mocks.pool_manager import TestPool +from tests.mocks.pool_manager import TestDriver, TestPool class DelayedBalancer: @@ -43,22 +43,13 @@ def __await__(self): ).__await__() -class RecordingPoolManager: - def __init__(self, pool_delay: float, acquire_delay: float): - self.pool = object() - self.balancer = DelayedBalancer(self.pool, delay=pool_delay) - self.acquire_delay = acquire_delay - self.acquire_kwargs = None +class _PoolStateProxy: + def __init__(self, parent): + self._parent = parent - def _prepare_acquire_kwargs(self, kwargs: dict, timeout): - prepared_kwargs = dict(kwargs) - prepared_kwargs["timeout"] = timeout - return prepared_kwargs - - def acquire_from_pool(self, pool, **kwargs): - self.acquire_kwargs = kwargs - timeout = kwargs.get("timeout") - slow = SlowAcquire(self.acquire_delay) + def acquire_from_pool(self, pool, *, timeout=None, **kwargs): + self._parent.acquire_timeout = timeout + slow = SlowAcquire(self._parent.acquire_delay) if timeout is not None: return _TimeoutSlowAcquire(slow, timeout) return slow @@ -66,22 +57,37 @@ def acquire_from_pool(self, pool, **kwargs): def host(self, pool): return "test-host:5432" - def register_connection(self, connection, pool): + +class RecordingPoolManager: + def __init__(self, pool_delay: float, acquire_delay: float): + self.pool = object() + self._balancer = DelayedBalancer(self.pool, delay=pool_delay) + self.acquire_delay = acquire_delay + self.acquire_timeout = None + self._pool_state = _PoolStateProxy(self) + + def _register_connection(self, connection, pool): pass -class OneConnectionPoolManager(TestPoolManager): - async def _pool_factory(self, dsn): +class OneConnectionTestDriver(TestDriver): + async def pool_factory(self, dsn, **kwargs): return TestPool(str(dsn), maxsize=1) +class OneConnectionPoolManager(TestPoolManager): + def __init__(self, dsn, **kwargs): + super().__init__(dsn, **kwargs) + self._pool_state._driver = OneConnectionTestDriver() + + class ReacquiringOneConnectionPoolManager(OneConnectionPoolManager): async def _periodic_pool_check(self, pool, dsn, sys_connection): await asyncio.wait_for( - self._refresh_pool_role(pool, dsn, sys_connection), + self._pool_state.refresh_pool_role(pool, dsn, sys_connection), timeout=self._refresh_timeout, ) - await self._notify_about_pool_has_checked(dsn) + await self._pool_state.notify_pool_checked(dsn) async def wait_until(predicate, timeout: float = 1.0): @@ -93,9 +99,12 @@ async def wait_until(predicate, timeout: float = 1.0): async def test_acquire_timeout_uses_shared_budget(): - pool_manager = RecordingPoolManager(pool_delay=0.05, acquire_delay=1.0) + recording = RecordingPoolManager(pool_delay=0.05, acquire_delay=1.0) context = PoolAcquireContext( - pool_manager=pool_manager, + pool_state=recording._pool_state, + balancer=recording._balancer, + register_connection=recording._register_connection, + unregister_connection=lambda conn: None, read_only=False, fallback_master=False, master_as_replica_weight=None, @@ -109,8 +118,8 @@ async def test_acquire_timeout_uses_shared_budget(): elapsed = asyncio.get_running_loop().time() - start assert elapsed < 0.2 - assert pool_manager.acquire_kwargs is not None - assert 0 < pool_manager.acquire_kwargs["timeout"] < 0.1 + assert recording.acquire_timeout is not None + assert 0 < recording.acquire_timeout < 0.1 async def test_refresh_timeout_removes_pool_from_available_set(): @@ -121,13 +130,17 @@ async def test_refresh_timeout_removes_pool_from_available_set(): acquire_timeout=0.5, ) try: - await pool_manager.ready() + await pool_manager._pool_state.ready() connection = await pool_manager.acquire_master() - await wait_until(lambda: pool_manager.master_pool_count == 0) + ps = pool_manager._pool_state + await wait_until(lambda: ps.master_pool_count == 0) - await pool_manager.release(connection) - await wait_until(lambda: pool_manager.master_pool_count == 1) + # Manually release: unregister + return to pool + pool = pool_manager._unmanaged_connections.pop(connection, None) + if pool is not None: + await ps.release_to_pool(connection, pool) + await wait_until(lambda: ps.master_pool_count == 1) finally: await pool_manager.close() @@ -139,17 +152,17 @@ async def test_close_preserves_cancellation_during_sys_connection_release(): refresh_delay=0.05, ) try: - await pool_manager.ready() - refresh_tasks = list(pool_manager._refresh_role_tasks) - pool_manager.release_to_pool = AsyncMock( + await pool_manager._pool_state.ready() + refresh_tasks = list(pool_manager._health.tasks) + pool_manager._pool_state.release_to_pool = AsyncMock( side_effect=asyncio.CancelledError(), ) await pool_manager.close() - assert pool_manager.closed - assert pool_manager.release_to_pool.await_count > 0 + assert pool_manager._closed + assert pool_manager._pool_state.release_to_pool.await_count > 0 assert all(task.done() for task in refresh_tasks) finally: - if not pool_manager.closed: + if not pool_manager._closed: await pool_manager.close() diff --git a/tests/test_trouble.py b/tests/test_trouble.py index ec05b64..3a5fc31 100644 --- a/tests/test_trouble.py +++ b/tests/test_trouble.py @@ -30,25 +30,30 @@ async def test_unavailable_db(pool_manager_factory, localhost, db_server_port): pass +_AIOSQA = "hasql.driver.asyncsqlalchemy.AsyncSqlAlchemyDriver" + + @pytest.mark.parametrize( - "pool_manager_factory,name", + "pool_manager_factory,driver_class", [ - (setup_aiopg, "aiopg"), - (setup_aiopgsa, "aiopg_sa"), - (setup_asyncpg, "asyncpg"), - (setup_asyncsqlalchemy, "asyncsqlalchemy"), - (setup_psycopg3, "psycopg3"), + (setup_aiopg, "hasql.driver.aiopg.AiopgDriver"), + (setup_aiopgsa, "hasql.driver.aiopg_sa.AiopgSaDriver"), + (setup_asyncpg, "hasql.driver.asyncpg.AsyncpgDriver"), + (setup_asyncsqlalchemy, _AIOSQA), + (setup_psycopg3, "hasql.driver.psycopg3.Psycopg3Driver"), ], ) -async def test_catch_cancelled_error(pool_manager_factory, pg_dsn, name): +async def test_catch_cancelled_error( + pool_manager_factory, pg_dsn, driver_class, +): async with pool_manager_factory(pg_dsn) as pool_manager: - await pool_manager.ready() - assert pool_manager.available_pool_count > 0 + await pool_manager._pool_state.ready() + assert pool_manager._pool_state.available_pool_count > 0 with mock.patch( - f"hasql.{name}.PoolManager._is_master", + f"{driver_class}.is_master", side_effect=asyncio.CancelledError(), ): - await pool_manager.wait_next_pool_check() - assert pool_manager.available_pool_count == 0 - await pool_manager.wait_next_pool_check() - assert pool_manager.available_pool_count > 0 + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.available_pool_count == 0 + await pool_manager._pool_state.wait_next_pool_check() + assert pool_manager._pool_state.available_pool_count > 0 From 06a500c8608be48fa9eda1ca56af44548b411d85 Mon Sep 17 00:00:00 2001 From: Pavel Mosein Date: Thu, 16 Apr 2026 10:00:46 +0300 Subject: [PATCH 2/3] refactor(metrics): update Metrics to use pools/gauges fields Replace Metrics(drivers=..., hasql=...) with Metrics(pools=..., hasql=..., gauges=...) now that base.py uses the new pool_manager.py. Backward-compat drivers property converts PoolMetrics back to DriverMetrics. Co-Authored-By: Claude Opus 4.6 (1M context) --- hasql/metrics.py | 19 ++++++++++++++++++- 1 file changed, 18 insertions(+), 1 deletion(-) diff --git a/hasql/metrics.py b/hasql/metrics.py index c46b0bc..fe0cd85 100644 --- a/hasql/metrics.py +++ b/hasql/metrics.py @@ -1,4 +1,5 @@ import time +import warnings from collections import defaultdict from contextlib import contextmanager from dataclasses import dataclass, field @@ -120,8 +121,24 @@ class HasqlGauges: @dataclass(frozen=True) class Metrics: - drivers: Sequence[DriverMetrics] + pools: Sequence[PoolMetrics] hasql: HasqlMetrics + gauges: HasqlGauges + + @property + def drivers(self) -> Sequence[DriverMetrics]: + """Backward-compatible accessor. Deprecated.""" + warnings.warn( + "Metrics.drivers is deprecated, use Metrics.pools instead", + DeprecationWarning, + stacklevel=2, + ) + return [ + DriverMetrics( + min=p.min, max=p.max, idle=p.idle, used=p.used, host=p.host, + ) + for p in self.pools + ] __all__ = ( From 9338dbf6e1d1e8c07b40073e3852e795b62be89b Mon Sep 17 00:00:00 2001 From: Pavel Mosein Date: Thu, 16 Apr 2026 10:25:01 +0300 Subject: [PATCH 3/3] fix: restore updated test_metrics.py for refactored pool_manager Co-Authored-By: Claude Opus 4.6 (1M context) --- tests/test_metrics.py | 44 +++++++++++++++++++++++++------------------ 1 file changed, 26 insertions(+), 18 deletions(-) diff --git a/tests/test_metrics.py b/tests/test_metrics.py index 96c2fcb..d16245b 100644 --- a/tests/test_metrics.py +++ b/tests/test_metrics.py @@ -27,7 +27,9 @@ async def test_hasql_context_metrics(pool_manager_factory, pg_dsn): metrics = pool_manager.metrics().hasql assert metrics == HasqlMetrics( pool=1, - acquire={pool_manager.host(pool_manager.pools[0]): 1}, + acquire={pool_manager._pool_state.host( + pool_manager._pool_state.pools[0], + ): 1}, pool_time=mock.ANY, acquire_time=mock.ANY, add_connections=mock.ANY, @@ -39,7 +41,9 @@ async def test_hasql_context_metrics(pool_manager_factory, pg_dsn): metrics = pool_manager.metrics().hasql assert metrics == HasqlMetrics( pool=1, - acquire={pool_manager.host(pool_manager.pools[0]): 1}, + acquire={pool_manager._pool_state.host( + pool_manager._pool_state.pools[0], + ): 1}, pool_time=mock.ANY, acquire_time=mock.ANY, add_connections=mock.ANY, @@ -61,25 +65,27 @@ async def test_hasql_context_metrics(pool_manager_factory, pg_dsn): ) async def test_hasql_metrics(pool_manager_factory, pg_dsn): async with pool_manager_factory(pg_dsn) as pool_manager: - _conn = await pool_manager.acquire_master() - metrics = pool_manager.metrics().hasql - assert metrics == HasqlMetrics( - pool=1, - acquire={pool_manager.host(pool_manager.pools[0]): 1}, - pool_time=mock.ANY, - acquire_time=mock.ANY, - add_connections=mock.ANY, - remove_connections=mock.ANY, - ) - assert list(metrics.add_connections.values()) == [1] - assert metrics.remove_connections == {} - - await pool_manager.release(connection=_conn) + async with pool_manager.acquire_master() as _conn: + metrics = pool_manager.metrics().hasql + assert metrics == HasqlMetrics( + pool=1, + acquire={pool_manager._pool_state.host( + pool_manager._pool_state.pools[0], + ): 1}, + pool_time=mock.ANY, + acquire_time=mock.ANY, + add_connections=mock.ANY, + remove_connections=mock.ANY, + ) + assert list(metrics.add_connections.values()) == [1] + assert metrics.remove_connections == {} metrics = pool_manager.metrics().hasql assert metrics == HasqlMetrics( pool=1, - acquire={pool_manager.host(pool_manager.pools[0]): 1}, + acquire={pool_manager._pool_state.host( + pool_manager._pool_state.pools[0], + ): 1}, pool_time=mock.ANY, acquire_time=mock.ANY, add_connections=mock.ANY, @@ -107,7 +113,9 @@ async def test_hasql_close_metrics(pool_manager_factory, pg_dsn): metrics = pool_manager.metrics().hasql assert metrics == HasqlMetrics( pool=1, - acquire={pool_manager.host(pool_manager.pools[0]): 1}, + acquire={pool_manager._pool_state.host( + pool_manager._pool_state.pools[0], + ): 1}, pool_time=mock.ANY, acquire_time=mock.ANY, add_connections=mock.ANY,