From 1aa408c25427847579802ae29c31443bf0f839ee Mon Sep 17 00:00:00 2001 From: An Long Date: Sat, 26 Sep 2026 23:20:08 +0900 Subject: [PATCH] Add ConnectionPool.close() and context manager support --- NEWS.rst | 10 ++++++++++ happybase/pool.py | 32 ++++++++++++++++++++++++++++++++ tests/test_api.py | 11 +++++++++++ 3 files changed, 53 insertions(+) diff --git a/NEWS.rst b/NEWS.rst index 04302a1..2c86998 100644 --- a/NEWS.rst +++ b/NEWS.rst @@ -4,6 +4,16 @@ Version history .. py:currentmodule:: happybase +HappyBase 1.4.0 +--------------- + +Release date: *not yet released* + +* Add :py:meth:`ConnectionPool.close` and support using the pool as a + context manager + (`issue 229 `_) + + HappyBase 1.3.0 --------------- diff --git a/happybase/pool.py b/happybase/pool.py index 5a9bd90..e9f8d75 100644 --- a/happybase/pool.py +++ b/happybase/pool.py @@ -64,6 +64,7 @@ def __init__(self, size, **kwargs): self._lock = threading.Lock() self._queue = queue.LifoQueue(maxsize=size) self._thread_connections = threading.local() + self._closed = False connection_kwargs = kwargs connection_kwargs['autoconnect'] = False @@ -80,6 +81,8 @@ def __init__(self, size, **kwargs): def _acquire_connection(self, timeout=None): """Acquire a connection from the pool.""" + if self._closed: + raise RuntimeError("Connection pool is closed") try: return self._queue.get(True, timeout) except queue.Empty: @@ -89,8 +92,37 @@ def _acquire_connection(self, timeout=None): def _return_connection(self, connection): """Return a connection to the pool.""" + with self._lock: + if self._closed: + connection.close() + return self._queue.put(connection) + def close(self): + """ + Close all connections in the pool. + + Connections that are in use are closed when they are returned. + + .. versionadded:: 1.4 + """ + with self._lock: + if self._closed: + return + self._closed = True + while True: + try: + connection = self._queue.get_nowait() + except queue.Empty: + break + connection.close() + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + self.close() + @contextlib.contextmanager def connection(self, timeout=None): """ diff --git a/tests/test_api.py b/tests/test_api.py index da1d2fe..84d9378 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -573,6 +573,17 @@ def run(): t.join() +def test_pool_close(): + with ConnectionPool(size=1, **connection_kwargs) as pool: + with pool.connection() as connection: + pass + assert not connection.transport.is_open() + + with pytest.raises(RuntimeError): + with pool.connection(): + pass + + if __name__ == '__main__': import logging import sys