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,